From 7ceb3a024bde764b2438bbbcf90f771616b2323c Mon Sep 17 00:00:00 2001 From: ZacharyZcR Date: Fri, 24 Jul 2026 11:47:43 +0800 Subject: [PATCH] =?UTF-8?q?serve:=20single-slot=20MTP=20speculation=20in?= =?UTF-8?q?=20run=5Fserve=5Fmux=20(KV=5FSLOTS=3D1)=20=E2=80=94=20#492/#358?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- c/colibri.c | 81 ++++++++++++++++++++++++++++++++++++++++++++--------- 1 file changed, 68 insertions(+), 13 deletions(-) diff --git a/c/colibri.c b/c/colibri.c index 82eeac21b..71c2f1afe 100644 --- a/c/colibri.c +++ b/c/colibri.c @@ -4697,7 +4697,7 @@ static void intr_install(void){ static void intr_install(void){} #endif static int spec_decode(Model *m, int *all, int kv, int n_new, int eos, float *logit, - void (*emit)(int,void*), void *ud, int *kv_out){ + void (*emit)(int,void*), void *ud, int *kv_out, float **logit_out){ Cfg *c=&m->c; int V=c->vocab; int emitted=0, done=0; int draft[64]; if(g_draft>63) g_draft=63; int carry_ban=-1; /* token rifiutato dalla verifica: escluso dal resample */ @@ -4772,7 +4772,12 @@ static int spec_decode(Model *m, int *all, int kv, int n_new, int eos, float *lo repin_pass(m); /* safe point: all device work is synchronized */ } g_spec_live = 0; /* prefill/decode successivi: gate normali / next prefill: normal gates */ - if(logit) free(logit); + /* logit_out (mux chunking): hand the continuation logits to the caller. + * NULL here means the loop exited right after emitting a token that was + * never forwarded — it sits at all[kv], and the caller must forward it + * before the next chunk. */ + if(logit_out) *logit_out=logit; + else if(logit) free(logit); if(kv_out) *kv_out=kv; return emitted; } @@ -4881,7 +4886,7 @@ static void generate(Model *m, const int *prompt, int np, int n_new, int *out){ for(int i=0;ihits+m->miss; int nsp=0; for(int i=0;in_layers;i++) if(m->L[i].sparse) nsp++; @@ -5418,6 +5423,10 @@ static void serve_ctx_free(Model *m, ServeCtx *s){ typedef struct { int active, pending, emitted, maximum, prompt_tokens, length_limited; + int spec; /* single-slot speculation (#492): decode runs through + spec_decode chunks instead of the shared batch */ + float *spec_logit; /* continuation logits between chunks; NULL = the last + emitted token sits at hist[len], not yet forwarded */ unsigned long long id; float temp, top_p; double started; @@ -5432,6 +5441,10 @@ static void mux_data(Tok *T, unsigned long long id, int token){ fflush(stdout); } +/* emit callback for the single-slot speculative path: stream straight to the mux protocol */ +typedef struct { Tok *T; unsigned long long id; } MuxEmit; +static void mux_spec_emit(int t, void *ud){ MuxEmit *e=(MuxEmit*)ud; mux_data(e->T,e->id,t); } + static void mux_done(Model *m, ServeCtx *sc, ServeReq *r){ double dt=now_s()-r->started; if(dt<1e-6) dt=1e-6; double dh=(double)(m->hits-r->hits0), dm=(double)(m->miss-r->miss0); @@ -5460,6 +5473,8 @@ static void mux_done(Model *m, ServeCtx *sc, ServeReq *r){ * share the batched forwards, so the shares describe the engine, not the * single request (same convention as the STAT hit%% above). */ if(g_prof) prof_report(m,&r->pb,dt,r->emitted,stderr); + if(r->spec_logit){ free(r->spec_logit); r->spec_logit=NULL; } + r->spec=0; r->active=0; } @@ -5478,6 +5493,8 @@ static int mux_submit(Model *m, Tok *T, ServeCtx *ctx, ServeReq *req, GrDraft *g } for(int i=0;ipb); /* a few loads: cheap enough to always track */ int room=maxctx-sc->len-1; if(r->maximum>room){r->maximum=room; r->length_limited=1;} g_temp=r->temp; g_nuc=r->top_p; + /* Single-slot speculation (#492/#358): with one KV slot there is no ragged + * batch — the decode is a single contiguous sequence, exactly the regime + * spec_decode already serves in chat/run/run_serve. Hand the request to the + * scheduler's chunked spec path instead of the pending-token cycle; grammar + * requests keep the forced-draft path (its acceptance is ~1 where it fires, + * and the two draft sources would fight over the same forward). */ + if(g_draft>0 && nctx==1 && !grd[sub.slot].on){ + if(r->maximum<=0){ free(logit); mux_done(m,sc,r); return 1; } + r->spec=1; r->spec_logit=logit; r->active=1; + return 1; + } int next=pick_tok(logit,m->c.vocab,-1); free(logit); if(r->maximum<=0 || next==eos || is_stop(next)){ mux_done(m,sc,r); return 1; } r->pending=next; r->emitted=1; r->active=1; sc->hist[sc->len]=next; m->n_emit++; @@ -5608,13 +5636,18 @@ static int mux_submit(Model *m, Tok *T, ServeCtx *ctx, ServeReq *req, GrDraft *g static void run_serve_mux(Model *m, const char *snap){ char tkp[2048]; snprintf(tkp,sizeof(tkp),"%s/tokenizer.json",snap); Tok T; tok_load(&T,tkp); int eos=tok_id_of(&T,"<|endoftext|>"); stops_arm_tok(&m->c,eos,&T); - g_draft=0; /* one scheduler owns every forward; MTP/n-gram speculation is not ragged-safe. - * Grammar-forced drafts ARE mux-safe (below): a drafting slot leaves the shared - * batch for one forward and runs the proven single-sequence verify path - * (kv_bind + step_all), exactly like prefill already does per submission. */ int maxctx=getenv("CTX")?atoi(getenv("CTX")):4096; int nctx=getenv("KV_SLOTS")?atoi(getenv("KV_SLOTS")):1; if(nctx<1||nctx>512){fprintf(stderr,"KV_SLOTS must be between 1 and 512\n");exit(2);} + /* MTP/n-gram speculation is not ragged-safe across KV slots, so multi-slot + * serve keeps one scheduler owning every forward (g_draft=0). At KV_SLOTS=1 + * there IS no ragged batch — decode is one contiguous sequence, the exact + * regime spec_decode already serves in chat/run/run_serve — so the engine's + * resolved draft setting stays live (#492/#358). Grammar-forced drafts stay + * mux-safe on every slot count, as before. */ + if(nctx>1) g_draft=0; + else if(g_draft>0) + fprintf(stderr,"[MTP] single-slot serve: speculation active (draft=%d)\n",g_draft); g_kvsave=getenv("KVSAVE")?atoi(getenv("KVSAVE")):1; KVState *initial=m->kv; free(initial->kv_start); free(initial); ServeCtx *ctx=calloc(nctx,sizeof(*ctx)); ServeReq *req=calloc(nctx,sizeof(*req)); @@ -5674,6 +5707,24 @@ static void run_serve_mux(Model *m, const char *snap){ DecodeRow rows[512]; int slots[512], S=0; for(int i=0;ispec){ + /* Single-slot speculative decode (#492): KV_SLOTS=1 has no ragged + * batch, so the whole turn runs through spec_decode in one call — + * the exact contract run_serve (non-mux) already uses. No chunk + * splicing (that would re-enter spec_decode mid-turn and double + * the boundary token); Ctrl-C interrupts through g_intr inside + * spec_decode, same as every other serve path. r->spec_logit + * holds the prefill continuation from mux_submit. */ + kv_bind(m,&sc->kv); + g_temp=r->temp; g_nuc=r->top_p; + float *lg=r->spec_logit; r->spec_logit=NULL; /* spec_decode takes ownership */ + MuxEmit ud={&T,r->id}; + int prod=spec_decode(m,sc->hist,sc->len,r->maximum-r->emitted,eos,lg, + mux_spec_emit,&ud,&sc->len,NULL); + r->emitted+=prod; + mux_done(m,sc,r); + continue; /* whole turn handled outside the shared batch */ + } /* grammar-forced drafts (greedy requests only: verification under sampling * needs rejection resampling, out of scope here). The slot leaves the shared * batch for one forward and runs the single-sequence verify path. */ @@ -5787,7 +5838,7 @@ static void run_serve(Model *m, const char *snap){ float *logit=step(m,hist+len-1,1,len-1); EmitStream es={&T,m,now_s(),0,1}; int prod=0; - if(cur>0) prod=spec_decode(m,hist,len,cur,eos,logit,emit_stream,&es,&len); + if(cur>0) prod=spec_decode(m,hist,len,cur,eos,logit,emit_stream,&es,&len,NULL); else free(logit); double tdt=now_s()-tt0; if(tdt<1e-6) tdt=1e-6; double dh=(double)(m->hits-h0), dm=(double)(m->miss-ms0); @@ -5870,7 +5921,7 @@ static void run_serve(Model *m, const char *snap){ EmitStream es={&T,m,now_s(),0,1}; int prod=0; grammar_reset(&g_grd); /* nuova risposta = nuovo documento (MORE invece continua) */ - if(cur>0) prod=spec_decode(m,hist,len,cur,eos,logit,emit_stream,&es,&len); + if(cur>0) prod=spec_decode(m,hist,len,cur,eos,logit,emit_stream,&es,&len,NULL); else free(logit); double tdt=now_s()-tt0; if(tdt<1e-6) tdt=1e-6; double dh=(double)(m->hits-h0), dm=(double)(m->miss-ms0); @@ -6786,15 +6837,19 @@ int main(int argc, char **argv){ * altrimenti "MTP active (draft=8)" mentirebbe: il messaggio e' stampato * prima della scelta del path (run_serve_mux, sotto), e con DRAFT=8 diceva * "active" per poi disabilitarlo in silenzio (#358, LordMZTE). */ - int mux_will_disable_mtp = getenv("SERVE") && getenv("SERVE_BATCH") && atoi(getenv("SERVE_BATCH")); + /* Multi-slot mux only: at KV_SLOTS=1 there is no ragged batch and the mux + * keeps the resolved draft setting live (#492 single-slot speculation). */ + int mux_slots = getenv("KV_SLOTS") ? atoi(getenv("KV_SLOTS")) : 1; + int mux_will_disable_mtp = getenv("SERVE") && getenv("SERVE_BATCH") && + atoi(getenv("SERVE_BATCH")) && mux_slots>1; int eff_draft = mux_will_disable_mtp ? 0 : g_draft; printf("loaded in %.2fs | resident dense: %.2f MB | layers=%d experts=%d | MTP %s (draft=%d)\n", now_s()-t0, m.resident_bytes/(1024.0*1024.0), m.c.n_layers, m.c.n_experts, m.has_mtp?(mux_will_disable_mtp?"DISABLED (multiplexed serve)":"ACTIVE"):"absent", eff_draft); /* anche su stderr: e' il canale che le UI (coli) mostrano all'utente */ if(mux_will_disable_mtp && m.has_mtp) - fprintf(stderr,"[MTP] disabled in multiplexed serve (SERVE_BATCH=1): speculation is not " - "ragged-safe across KV slots. Single-client interactive use (`coli chat`) keeps MTP.\n"); + fprintf(stderr,"[MTP] disabled in multiplexed serve (SERVE_BATCH=1, KV_SLOTS>1): speculation is " + "not ragged-safe across KV slots. Single-slot serve (KV_SLOTS=1) keeps MTP.\n"); else fprintf(stderr,"[MTP] %s (draft=%d)\n", m.has_mtp?"active: native speculative decoding":"absent", eff_draft); #ifdef __linux__