Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 22 additions & 2 deletions c/colibri.c
Original file line number Diff line number Diff line change
Expand Up @@ -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;i<V;i++){ float v=lo[i];
for(int k=0;k<5;k++) if(v>val[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<V)?(double)lo[expected]:0.0,
got, (got>=0&&got<V)?(double)lo[got]:0.0);
for(int k=0;k<5&&idx[k]>=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;s<S;s++) embed_row(m, ids[s], x+(int64_t)s*D);
Expand All @@ -4818,6 +4837,7 @@ static void forward_all(Model *m, const int *ids, int S, int *pred){
matmul_qt(lo, row, &m->lm_head, 1);
int best=0; float bv=lo[0]; for(int i=1;i<c->vocab;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);
}
Expand Down Expand Up @@ -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<nfull;i++){
if(pred[i]==tf[i]) ok++;
else fprintf(stderr,"[ORACLE] mismatch pos=%d expected=%d got=%d\n",i,tf[i],pred[i]);
Expand Down
10 changes: 9 additions & 1 deletion docs/metal.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,15 @@ the PCIe copy tax that keeps CUDA's streaming experts on the CPU — so colibrì
has an opt-in Metal backend that runs the **routed-expert SwiGLU (batched,
zero-copy from the RAM slabs)**, the **fused decode attention** (full MLA layer
in one command buffer, S≤4), and **prefill's large GEMMs** on the GPU.
Token-exact vs the CPU path.
Decode is token-exact vs the CPU path. Prefill's large GEMMs run on the GPU in a
different accumulation order, so on **near-tie logits** they can occasionally pick a
different top token than the CPU — a floating-point ordering difference, not a kernel
bug (`make metal-test` passes the GEMM at ~3e-6 against a 1e-4 tolerance; see
[#622](https://github.com/JustVugg/colibri/issues/622)). It is invisible in normal use
but can surface in teacher-forced oracle comparisons on pathological 4-bit toy
containers. Set `COLI_METAL_GEMM_MIN=100000` to keep every GEMM on the CPU for
bit-exact prefill (`DEBUG_LOGITS=1` on a `TF=1` run dumps the top-5 logits and the
top1–top2 margin at each mismatch, so you can see how close the tie was).

```bash
cd c
Expand Down
Loading