From cda2e4ee03d8478addbfbc2814550f2a310a0f23 Mon Sep 17 00:00:00 2001 From: Kyle Kelley Date: Sat, 1 Aug 2026 04:29:34 -0700 Subject: [PATCH 1/2] inkling: overlap shared experts with the last GPU round Routed and shared experts are independent given the layer input: both read x, both accumulate into out, and the GPU only touches out in the _end scatter-add after the wait. So the last round submits async and the CPU computes the shared experts while the GPU runs the routed ones. Frozen probe: 1.75 -> 1.96 tok/s decode, tokens identical, oracles exact. On a GPU fault _end returns 0 and the CPU redoes the round; shared_done guards double-compute. --- c/inkling.c | 63 ++++++++++++++++++++++++++++++++++++++++++----------- 1 file changed, 50 insertions(+), 13 deletions(-) diff --git a/c/inkling.c b/c/inkling.c index c698d587b..48147e343 100644 --- a/c/inkling.c +++ b/c/inkling.c @@ -1128,6 +1128,28 @@ static void dense_mlp(Model *m, Layer *l, float *x, int S, float *out) { * (2) fill ALL missing experts in one parallel burst (the NVMe wants queue * depth — during prefill this batches the whole sequence's misses), then * (3) compute. */ +/* shared experts for all S positions: gamma inside (before down_proj is + * linear, so applied at the end). Factored out so the Metal path can run it + * on the CPU while the last routed-expert round is in flight on the GPU. */ +static void shared_experts_cpu(Model *m, Layer *l, const float *x, int S, + float *out, const float *wgt, + float *g, float *u, float *hh) { + Cfg *c = &m->c; + int D = c->hidden, K = c->topk, I = c->moe_inter, ns = c->n_shared; + for (int s = 0; s < S; s++) { + const float *xs = x + (int64_t)s*D; + float *os = out + (int64_t)s*D; + const float *w = wgt + (int64_t)s*(K+ns); + for (int j = 0; j < ns; j++) { + matmul_w(g, xs, wt_off_i(l->sh_g, (int64_t)j*I*D, D), 1, D, I); + matmul_w(u, xs, wt_off_i(l->sh_u, (int64_t)j*I*D, D), 1, D, I); + for (int i = 0; i < I; i++) g[i] = siluf(g[i]) * u[i]; + matmul_w(hh, g, wt_off_i(l->sh_d, (int64_t)j*D*I, I), 1, I, D); + for (int d = 0; d < D; d++) os[d] += w[K+j] * hh[d]; + } + } +} + static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out) { Cfg *c = &m->c; int D = c->hidden, E = c->n_experts, K = c->topk, I = c->moe_inter, ns = c->n_shared; @@ -1201,6 +1223,7 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out) { int cap = m->cache[layer].cap; if (cap < 1) cap = 1; float *g = falloc(2*I), *u = g + I, *hh = falloc(D); int q4 = m->xq && m->rb13*2 == D; /* packed int4 vs int8 container */ + int shared_done = 0; /* set when overlapped with the last GPU round */ int64_t npair = (int64_t)S*K; #ifdef COLI_METAL /* per-round scratch for the batched GPU submit: pairs grouped by expert, @@ -1279,6 +1302,28 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out) { mgp[j] = e->p13; mup[j] = e->p13 + (int64_t)I*m->rb13; mdp[j] = e->p2; mgs[j] = e->s13; mus[j] = e->s13 + I; mds[j] = e->s2; } + /* Last round: submit async and run the shared experts on the + * CPU while the GPU computes the routed ones — they are + * independent (both read x, both accumulate into out, and the + * GPU only touches out in _end's scatter-add, after the wait). + * On a GPU fault _end returns 0 and the CPU loop below redoes + * the round; shared_done stays set either way. */ + if (base + cap >= npair) { + ColiMetalMoeHandle *h = coli_metal_moe_block_begin( + nb, D, I, q4 ? 2 : 1, mgp, mup, mdp, mgs, mus, mds, + mxg, mxoff, mnr, mrows, mrw); + if (h) { + double ts = now_s(); + shared_experts_cpu(m, l, x, S, out, wgt, g, u, hh); + double sh = now_s() - ts; + m->t_shared += sh; + shared_done = 1; + int ok = coli_metal_moe_block_end(h, out); + m->t_expert += (now_s() - te) - sh; + if (ok) continue; /* round done on the GPU */ + te = now_s(); /* fault: CPU redo below */ + } + } if (coli_metal_moe_block(nb, D, I, q4 ? 2 : 1, mgp, mup, mdp, mgs, mus, mds, mxg, mxoff, mnr, mrows, mrw, out, S)) { m->t_expert += now_s() - te; @@ -1317,20 +1362,12 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out) { } m->t_expert += now_s() - te; } - /* shared experts: una volta per token, fuori dai giri (non usano la cache) */ - for (int s = 0; s < S; s++) { - const float *xs = x + (int64_t)s*D; - float *os = out + (int64_t)s*D; - float *w = wgt + (int64_t)s*(K+ns); + /* shared experts: una volta per token, fuori dai giri (non usano la cache). + * Under Metal the last routed round overlaps this on the CPU (see the + * begin/end path above) and sets shared_done. */ + if (!shared_done) { double ts = now_s(); - /* gamma inside (before down_proj is linear, so applied at the end) */ - for (int j = 0; j < ns; j++) { - matmul_w(g, xs, wt_off_i(l->sh_g, (int64_t)j*I*D, D), 1, D, I); - matmul_w(u, xs, wt_off_i(l->sh_u, (int64_t)j*I*D, D), 1, D, I); - for (int i = 0; i < I; i++) g[i] = siluf(g[i]) * u[i]; - matmul_w(hh, g, wt_off_i(l->sh_d, (int64_t)j*D*I, I), 1, I, D); - for (int d = 0; d < D; d++) os[d] += w[K+j] * hh[d]; - } + shared_experts_cpu(m, l, x, S, out, wgt, g, u, hh); m->t_shared += now_s() - ts; } free(logits); free(idx); free(keff); free(wgt); free(use); free(fill); free(fl); From b9ee0edbb050ab97a9ce236c35f578fbf0799bd8 Mon Sep 17 00:00:00 2001 From: Kyle Kelley Date: Sat, 1 Aug 2026 04:35:07 -0700 Subject: [PATCH 2/2] metal+inkling: shared experts on the GPU as a bf16 block MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds fmt=5 (raw bf16 rows, no scales) to moe_gemv — an additive branch, existing formats untouched — and moves inkling's shared experts onto it. The bf16 weights relocate into a page-aligned slab at load (the Wt handles repoint into it, so the CPU fallback reads the same bytes), and each MoE layer submits one fmt=5 block BEFORE the routed rounds: the GPU chews the shared experts while the CPU routes, fills, and packs. Every token uses every shared expert, so group j is just all S rows. Shared experts were the profile's biggest line at 8.0s per 24-token probe; on the GPU they are 0.1s. Frozen probe: 1.96 -> 2.46 tok/s decode, prefill 10.3 -> 6.4s. Audio prefill (113-token prompt) 22.4s. Greedy output is token-identical to the pre-Metal baseline through eos, and the tiny oracle stays exact (its f32 shared weights correctly skip the bf16 path and take the CPU loop). --- c/backend_metal.mm | 14 +++++++-- c/inkling.c | 71 ++++++++++++++++++++++++++++++++++++++++++++-- 2 files changed, 80 insertions(+), 5 deletions(-) diff --git a/c/backend_metal.mm b/c/backend_metal.mm index 484eb0016..f3f603bc9 100644 --- a/c/backend_metal.mm +++ b/c/backend_metal.mm @@ -200,13 +200,23 @@ kernel void moe_gemv(device const ulong* waddr [[buffer(0)]], device const ulong acc += (dot(m0*sA,xs[l*2]) + dot(m1*sB,xs[l*2+1])) * (0.5f*db); } } + } else if (fmt == 5) { // bf16 raw rows: no scales, upper half of f32 + device const ushort* w=(device const ushort*)(waddr[e])+(long)o*K; + device const ushort4* w4=(device const ushort4*)w; + for(int c=slane;c(uint(a.x)<<16), as_type(uint(a.y)<<16), + as_type(uint(a.z)<<16), as_type(uint(a.w)<<16)); + float4 w1=float4(as_type(uint(b.x)<<16), as_type(uint(b.y)<<16), + as_type(uint(b.z)<<16), as_type(uint(b.w)<<16)); + acc+=dot(w0,x4[2*c])+dot(w1,x4[2*c+1]); } + for(int i=K8*8+slane;i(uint(w[i])<<16)*xr[i]; } else { device const char* w=(device const char*)(waddr[e])+(long)o*K; device const char4* w4=(device const char4*)w; for(int c=slane;c xg_buf, id gg_buf, id uu_buf, id hh_buf) { - if (!g_dev || (fmt != 1 && fmt != 2 && fmt != 6)) return nil; + if (!g_dev || (fmt != 1 && fmt != 2 && fmt != 5 && fmt != 6)) return nil; if (fmt == 6) { /* e8 kernel assumes clean block tiling, and every FWHT tile of the * down input (CPU tiling rule, e8_rot_rows) must fit threadgroup mem */ if ((D & 255) || (Iinter & 31)) return nil; diff --git a/c/inkling.c b/c/inkling.c index 48147e343..69e130880 100644 --- a/c/inkling.c +++ b/c/inkling.c @@ -795,6 +795,28 @@ static void model_init(Model *m, const char *snap, int cap, int bits) { LDW(sh_g, "mlp.shared_experts.gate_proj"); LDW(sh_u, "mlp.shared_experts.up_proj"); LDW(sh_d, "mlp.shared_experts.down_proj"); +#ifdef COLI_METAL + /* Move the bf16 shared-expert weights into one page-aligned slab + * (gates, then ups, then downs) and repoint the Wt handles into + * it: the CPU path reads the same bytes it always did, and the + * slab registers once so the GPU fmt=5 path can resolve it. */ + if (g_metal && l->sh_g.h && l->sh_u.h && l->sh_d.h) { + int64_t I = c->moe_inter, ns = c->n_shared; + size_t one = (size_t)ns*I*D*2, pg = 16384; + size_t len = (3*one + pg - 1) / pg * pg; + void *sl = NULL; + if (!posix_memalign(&sl, pg, len)) { + memcpy((char*)sl, l->sh_g.h, one); + memcpy((char*)sl + one, l->sh_u.h, one); + memcpy((char*)sl + 2*one, l->sh_d.h, one); + free(l->sh_g.h); free(l->sh_u.h); free(l->sh_d.h); + l->sh_g.h = (uint16_t*)sl; + l->sh_u.h = (uint16_t*)((char*)sl + one); + l->sh_d.h = (uint16_t*)((char*)sl + 2*one); + coli_metal_register(sl, len); + } + } +#endif } #undef LD #undef LDW @@ -1245,6 +1267,39 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out) { mxoff = malloc((cap+1)*sizeof(int)); mnr = malloc(cap*sizeof(int)); mrows = malloc(cap*sizeof(int)); mgi = malloc(cap*sizeof(int)); mfp = malloc(cap*sizeof(int)); } + /* Shared experts as ONE bf16 (fmt=5) block, submitted BEFORE the routed + * rounds so it rides the GPU while the CPU routes, fills, and packs. + * Every token uses every shared expert, so group j is simply all S rows + * with weight w[K+j]. Falls back to the CPU loop if begin refuses. */ + ColiMetalMoeHandle *sh_h = NULL; + float *sxg = NULL, *srw = NULL; + if (mxg && ns > 0 && ns <= 8 && l->sh_g.h && l->sh_u.h && l->sh_d.h && + !(getenv("INK_METAL_SHARED") && *getenv("INK_METAL_SHARED") == '0')) { + double ts = now_s(); + sxg = falloc((int64_t)ns*S*D); srw = falloc((int64_t)ns*S); + const void *sgp[8], *sup[8], *sdp[8]; + const float *sscale[8]; + int sxoff[9], snr[8]; + int *srows = malloc((size_t)ns*S*sizeof(int)); + for (int j = 0; j < ns; j++) { + sgp[j] = l->sh_g.h + (int64_t)j*I*D; + sup[j] = l->sh_u.h + (int64_t)j*I*D; + sdp[j] = l->sh_d.h + (int64_t)j*D*I; + sscale[j] = (const float*)sgp[j]; /* fmt=5 never reads scales */ + sxoff[j] = j*S; snr[j] = S; + for (int s = 0; s < S; s++) { + memcpy(sxg + ((int64_t)j*S + s)*D, x + (int64_t)s*D, (size_t)D*sizeof(float)); + srows[j*S + s] = s; + srw[j*S + s] = wgt[(int64_t)s*(K+ns) + K + j]; + } + } + sxoff[ns] = ns*S; + sh_h = coli_metal_moe_block_begin(ns, D, I, 5, sgp, sup, sdp, + sscale, sscale, sscale, + sxg, sxoff, snr, srows, srw); + free(srows); + m->t_shared += now_s() - ts; + } #endif for (int64_t base = 0; base < npair; base += cap) { int64_t end = base + cap < npair ? base + cap : npair; @@ -1308,7 +1363,7 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out) { * GPU only touches out in _end's scatter-add, after the wait). * On a GPU fault _end returns 0 and the CPU loop below redoes * the round; shared_done stays set either way. */ - if (base + cap >= npair) { + if (base + cap >= npair && !sh_h) { ColiMetalMoeHandle *h = coli_metal_moe_block_begin( nb, D, I, q4 ? 2 : 1, mgp, mup, mdp, mgs, mus, mds, mxg, mxoff, mnr, mrows, mrw); @@ -1362,9 +1417,18 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out) { } m->t_expert += now_s() - te; } +#ifdef COLI_METAL + /* GPU shared block: wait + scatter-add. A fault falls through to CPU. */ + if (sh_h) { + double ts = now_s(); + if (coli_metal_moe_block_end(sh_h, out)) shared_done = 1; + m->t_shared += now_s() - ts; + } +#endif /* shared experts: una volta per token, fuori dai giri (non usano la cache). - * Under Metal the last routed round overlaps this on the CPU (see the - * begin/end path above) and sets shared_done. */ + * Under Metal these run on the GPU as an fmt=5 block overlapped with the + * routed rounds (or on the CPU during the last GPU round); either path + * sets shared_done. */ if (!shared_done) { double ts = now_s(); shared_experts_cpu(m, l, x, S, out, wgt, g, u, hh); @@ -1378,6 +1442,7 @@ static void moe(Model *m, Layer *l, int layer, float *x, int S, float *out) { free(mgs); free(mus); free(mds); free(mslot); free(mxoff); free(mnr); free(mrows); free(mgi); free(mfp); } + free(sxg); free(srw); #endif }