From 2d8d2951ee7d93cf2adad7b89b56c0adc1dbd0c5 Mon Sep 17 00:00:00 2001 From: Egon Ruiter Date: Thu, 16 Jul 2026 15:37:20 +0200 Subject: [PATCH] fix(prefetch): address second round of copilot review comments - Fix ENV VARS header: document PILOT=0-3, SMOOTH, CONF_LIMIT; remove stale REBAL entry - Fix per-layer EMA: apply routing momentum to all layers (not just layer 0) with correct offset - Fix in-flight slot race in expert_get: LRU eviction now skips slots with eid==-1 (being loaded) - Fix in-flight slot race in pilot_realload: same fix, prevents concurrent writes into active slot - Fix idx[] buffer overflow: clamp max_cand to 128 before E in pilot_prefetch --- c/olmoe.c | 30 ++++++++++++++++-------------- 1 file changed, 16 insertions(+), 14 deletions(-) diff --git a/c/olmoe.c b/c/olmoe.c index 36fc0b0..51e32ae 100644 --- a/c/olmoe.c +++ b/c/olmoe.c @@ -7,12 +7,13 @@ * Matmul multi-thread con OpenMP (niente BLAS). * * ENV VARS: - * PILOT=0/1/2 : 0=no prefetch, 1=1-layer lookahead, 2=2-layer lookahead [IMPROVEMENT 1] - * HOT=N : pin top-N hot experts per layer permanently (never evict) [IMPROVEMENT 2] - * WARMUP=N : tokens before hot pinning activates (default 5) [IMPROVEMENT 2] - * REBAL=N : rebalance cache per-layer every N tokens (0=off) [IMPROVEMENT 3] - * WIDE=N : prefetch top-K*N candidates (default 1, try 2 or 3) [IMPROVEMENT 4] - * (expert queue is sorted by eid for SSD locality) [IMPROVEMENT 5] + * PILOT=0/1/2/3 : 0=no prefetch, 1=1-layer lookahead, 2=2-layer, 3=3-layer lookahead + * HOT=N : pin top-N hot experts per layer permanently (never evict) + * WARMUP=N : tokens before hot pinning activates (default 5) + * WIDE=N : prefetch top-K*N candidates (default 1, try 2 or 3) + * SMOOTH=F : EMA coefficient for routing momentum (default 0.3, range 0.0-0.95) + * CONF_LIMIT=F : cumulative gate probability threshold for prefetch cutoff (default 0.92) + * (expert queue is sorted by eid for SSD read locality) */ #define _GNU_SOURCE #include @@ -375,13 +376,13 @@ static void expert_get(Model *m, int layer, int eid, Slot **out) { s = &lc->slots[lc->n++]; slot_ensure_allocated(m, s); } else { - /* IMPROVEMENT 2: LRU eviction — never evict pinned experts */ + /* LRU eviction — skip pinned and in-flight (eid==-1) slots */ int lru = -1; for (int i = 0; i < lc->n; i++) { - if (lc->slots[i].pinned) continue; + if (lc->slots[i].pinned || lc->slots[i].eid < 0) continue; if (lru < 0 || lc->slots[i].used < lc->slots[lru].used) lru = i; } - if (lru < 0) lru = 0; /* all pinned: fallback evict oldest */ + if (lru < 0) lru = 0; /* all pinned/in-flight: fallback evict oldest */ s = &lc->slots[lru]; s->pinned = 0; } @@ -539,8 +540,8 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out) { float *g = falloc(I), *u = falloc(I), *hh = falloc(D); for (int s = 0; s < S; s++) { float *pr = logits + (int64_t)s*E; - if (layer == 0 && m->momentum_logits) { - float *ema = m->momentum_logits; + if (m->momentum_logits && m->pilot_smooth > 0.f) { + float *ema = m->momentum_logits + (int64_t)layer * E; int is_zero = 1; for (int e = 0; e < E; e++) { if (ema[e] != 0.f) { is_zero = 0; break; } } if (is_zero) { @@ -646,16 +647,16 @@ static void pilot_realload(Model *m, int layer, int eid) { s = &lc->slots[lc->n++]; slot_ensure_allocated(m, s); } else { - /* IMPROVEMENT 2: never evict pinned experts */ + /* LRU eviction — skip pinned and in-flight (eid==-1) slots */ int lru = -1; for (int i = 0; i < lc->n; i++) { - if (lc->slots[i].pinned) continue; + if (lc->slots[i].pinned || lc->slots[i].eid < 0) continue; if (lru < 0 || lc->slots[i].used < lc->slots[lru].used) lru = i; } if (lru < 0) { m->is_queued[layer * c->n_experts + eid] = 0; pthread_mutex_unlock(&g_pilot_mx); - return; /* all pinned, skip */ + return; /* all pinned/in-flight, skip */ } s = &lc->slots[lru]; s->pinned = 0; } @@ -745,6 +746,7 @@ static void pilot_prefetch(Model *m, int lnext, const float *x, int S) { int min_cand = c->topk; int max_cand = c->topk * g_wide; if (max_cand < min_cand) max_cand = min_cand; + if (max_cand > 128) max_cand = 128; /* idx[] buffer bound */ if (max_cand > E) max_cand = E; for (int kk = 0; kk < max_cand; kk++) {