From a9b8a476d060ba87cf7f4ec42221be15ce7093ab Mon Sep 17 00:00:00 2001 From: JustVugg Date: Sun, 26 Jul 2026 03:06:06 +0200 Subject: [PATCH] metal/oracle: scope the token-exact claim to decode + implement DEBUG_LOGITS top-5 dump (#622) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit @lBroth reported that Metal prefill GEMM (matmul_qt on GPU, S >= COLI_METAL_GEMM_MIN, default 16) is not token-exact vs CPU on NEAR-TIE logits: on 4-bit toy fixtures a teacher-forced prefill disagrees at 2/32 positions, deterministically. Decode is unaffected (S=1, and short prompts stay below the GEMM threshold). The GEMM kernel is numerically fine (make metal-test passes at ~3e-6 vs a 1e-4 tolerance); it is a GPU-vs-CPU accumulation-order difference that only flips the argmax when two logits sit inside that drift. Confirmed: COLI_METAL_GEMM_MIN=100000 (keep every GEMM on the CPU) makes it bit-identical. So the concrete defect is a docs overclaim, not a kernel bug. - docs/metal.md: "Token-exact vs the CPU path" -> decode is token-exact; prefill's GPU GEMM can diverge on near-tie logits by accumulation order (not a kernel bug, #622), with COLI_METAL_GEMM_MIN=100000 as the bit-exact escape hatch for teacher-forced comparisons. - Implement DEBUG_LOGITS: the oracle mismatch message pointed at "TF=1 DEBUG_LOGITS=1 for top-5 logit dump", but no such flag existed (a dead string, also flagged in #622). It now works: on a teacher-forcing mismatch it dumps the top-5 logits, the top1-top2 margin, and the expected/got tokens' logits, so a near-tie divergence reads as the tiny gap it is instead of a bare token mismatch. Opt-in, stderr, fires only on a mismatch — normal runs are byte-for-byte unchanged (forward_all gains a nullable ref arg; its one caller updated). Verified: oracle stays 32/32 with DEBUG_LOGITS=1 on a clean run (no dump); a forced mismatch prints e.g. "[LOGITS] pos=5 top1-top2 gap=2.449e-03 | expected=35 (-0.23799) got=34 (1.73265) | top5: 34:1.73265 197:1.73020 ...". Also relevant to #457/#587: verify Metal fmt=4 token-exactness with COLI_METAL_GEMM_MIN=100000 to isolate the new kernel from this pre-existing prefill drift. Co-Authored-By: Claude Opus 4.8 (1M context) --- c/colibri.c | 24 ++++++++++++++++++++++-- docs/metal.md | 10 +++++++++- 2 files changed, 31 insertions(+), 3 deletions(-) diff --git a/c/colibri.c b/c/colibri.c index 3b0f13713..869f9cb38 100644 --- a/c/colibri.c +++ b/c/colibri.c @@ -4805,8 +4805,27 @@ static void emit_stream(int t, void *ud){ } /* teacher-forcing: un solo forward su ids[S], argmax per posizione in pred[S] */ -static void forward_all(Model *m, const int *ids, int S, int *pred){ +/* DEBUG_LOGITS=1: on a teacher-forcing mismatch, dump the top-5 logits, the top1-top2 margin, + * and where the expected/got tokens land. Makes a near-tie divergence (e.g. the Metal prefill + * GEMM accumulation-order drift, #622) visible as the tiny gap it is, rather than a bare token + * mismatch. Stderr, opt-in, only on a mismatch — normal runs are byte-for-byte unchanged. */ +static void dump_top5_logits(int pos, const float *lo, int V, int expected, int got){ + int idx[5]; float val[5]; + for(int k=0;k<5;k++){ idx[k]=-1; val[k]=-INFINITY; } + for(int i=0;ival[k]){ + for(int j=4;j>k;j--){ val[j]=val[j-1]; idx[j]=idx[j-1]; } + val[k]=v; idx[k]=i; break; } } + double gap = (idx[1]>=0) ? (double)(val[0]-val[1]) : 0.0; + fprintf(stderr,"[LOGITS] pos=%d top1-top2 gap=%.3e | expected=%d (%.5f) got=%d (%.5f) | top5:", + pos, gap, expected, (expected>=0&&expected=0&&got=0;k++) fprintf(stderr," %d:%.5f", idx[k], (double)val[k]); + fprintf(stderr,"\n"); +} +static void forward_all(Model *m, const int *ids, int S, int *pred, const int *ref){ Cfg *c=&m->c; int D=c->hidden; + int dbg = ref && getenv("DEBUG_LOGITS"); kv_alloc(m,S); float *x=falloc((int64_t)S*D); for(int s=0;slm_head, 1); int best=0; float bv=lo[0]; for(int i=1;ivocab;i++) if(lo[i]>bv){bv=lo[i];best=i;} pred[s]=best; + if(dbg && pred[s]!=ref[s]) dump_top5_logits(s, lo, c->vocab, ref[s], pred[s]); } free(x); free(lo); free(row); } @@ -6998,7 +7018,7 @@ int main(int argc, char **argv){ if(getenv("TF")){ int *tf=read_arr(ref,"tf_pred",&(int){0}); int *pred=malloc(nfull*sizeof(int)); double tt=now_s(); - forward_all(&m, full, nfull, pred); double tdt=now_s()-tt; + forward_all(&m, full, nfull, pred, tf); double tdt=now_s()-tt; int ok=0; for(int i=0;i