// Decode-step MEGAKERNEL: one workgroup computes one batch row's ENTIRE // decoder layer — 8 dispatches collapse into 1 (small-B latency: the b1 step // is dispatch-overhead-bound, ~15µs of fixed cost per dispatch against // ~5-10µs kernels). With IF_EMBED (layer 0) the decode embedding folds in // too, so a b1 step becomes: mega L0 → mega L1 → lm_head → argmax. // // Stages (workgroupBarrier between each; all math f32; every stage boundary // value is ROUNDED THROUGH THE STORAGE TYPE first — replicating the unfused // path's activation-buffer round trips, the gemm_row_ln precedent): // S0 x ← embed (L0: ring/table/pos, matching embed.wgsl DECODE) or the // global hidden buffer X (L1) // S1 qkv = x·Wqkvᵀ + b; q stays in shared, k|v quads store to the caches // at position t (the same rounded value — the kv_append contract) // S2 self-attention over positions 0..t (phase structure and score/exp/ // fold order mirror attention.wgsl; heads loop serially) // S3 x = LN1(x + self_out(attn)) // S4 q = cross_q(x) // S5 cross-attention over lens[b] encoder positions (fused crossKV k|v) // S6 x = LN2(x + cross_out(attn)) // S7 ffn = SiLU(fc1(x)) (SiLU in f32 BEFORE the f16 round, as gemv) // S8 X ← LN3(x + fc2(ffn)) (written back to the global hidden buffer) // // All projections read the ORIGINAL [K, N] row-major `.weight` tensors with // the gemm_row_ln access pattern: thread q owns output quad q, and at each k // the threads read CONSECUTIVE quads of W's k-row — fully coalesced. (The // first version walked the transposed [N,K] copies, one row per thread — // every load touched 32 distinct lines and the whole kernel ran at ~4.5GB/s, // 4.3× SLOWER than the chain it replaced. One workgroup has only ~8 warps of // latency-hiding; coalescing is everything here.) Four independent // accumulators per k-quad keep 4 loads in flight per thread. Every tensor is // addressed inside the ONE weights buffer via compile-time vec4-element // offsets (…4 defines = byteOffset/8; manifest offsets are 256-aligned so /8 // is exact). One pipeline per layer. // // NOT bit-exact vs the unfused chain (accumulation/reduction ORDER differs // at every site) — routed like every kernel change: m3/golden gates + the // step-0 + divergence-rate equiv (mega_equiv), e2e A/B decides the batch // threshold. // // Shared budget (16KB): (2·HD4 + TMP4)·16 + SCORES_CAP·4 + WG·4 — 13,184B // for MoxhiMT-30 (448/1792), 16,256B for HachimiMT-60 (576/2304); checked at // dispatch. tmp4 is max(FFN4, HD4 + WG) quads: the attention phase-3 partial // scratch tmp4[HD4 .. HD4+WG) must fit even when the model's FFN is small // (q lives in [0..HD4) for self, out4 for cross; the fold result lands in // [0..HD4) only after all partial reads). // // Template placeholders (buildShader in pipelines.js): // ENABLE_F16, T (must be f16 — the .wt copies only exist as f16), WG (256) // ENABLE_SG + IF_SG/IF_NOSG subgroup wgMax/wgSum (flags.sg) — see below // IF_EMBED / IF_NOEMBED layer-0 embedding fold (TABLE4/POS4/EMBED_SCALE/ // DECODER_START live inside IF_EMBED) // H, D, FFN4, LMAX, SCORES_CAP, ATTN_SCALE, EPS // QKVW4 QKVB4 OUTW4 OUTB4 LN1G4 LN1B4 CQW4 CQB4 COW4 COB4 LN2G4 LN2B4 // FC1W4 FC1B4 FC2W4 FC2B4 LN3G4 LN3B4 per-tensor vec4 offsets into W // // SG mode (flags.sg): wgMax/wgSum are where this kernel's barriers live — // each tree call is 10 workgroupBarriers, and one layer makes 38 of them // (2 per attention head × 2 sides × H, 2 per LN × 3), ~400 barriers per // step per layer. That is exactly what the megakernel pays on Apple GPUs // (Metal mega_sweep: mega LOSES b1 there while winning −12% on NVIDIA). // With sg each call is subgroupMax/Add → one partial per subgroup → ONE // barrier → serial fold over ≤ WG/4 partials, ~5× fewer barriers overall. {{ENABLE_IMMEDIATE}} {{ENABLE_SG}} {{ENABLE_F16}} struct Params { B: u32, // batch rows (grid.x) t: u32, // decode step (cache position; embed pos; self len = t+1) S: u32, // encoder crossKV position capacity (padded S) _pad: u32, } {{PARAM_BINDING}}var<{{PARAM_ADDRESS}}> params: Params; @group(0) @binding(1) var W: array>; // whole weights buffer @group(0) @binding(2) var ring: array; // token ring @group(0) @binding(3) var Kc: array>; // [B, LMAX, H·D] @group(0) @binding(4) var Vc: array>; @group(0) @binding(5) var CKV: array>; // [B·S, 2·H·D] fused k|v @group(0) @binding(6) var lens: array; @group(0) @binding(7) var X: array>; // hidden [B, H·D] const H: u32 = {{H}}u; const D: u32 = {{D}}u; const D4: u32 = D / 4u; // quads per head const HD4: u32 = H * D4; // quads per d_model row const QKV4: u32 = 3u * HD4; // fused q|k|v quads const FFN4: u32 = {{FFN4}}u; // ffn quads (FFN/4) const KQ_FFN: u32 = FFN4; // fc2 K quads (K = FFN) const LMAX: u32 = {{LMAX}}u; const SCORES_MAX: u32 = {{SCORES_CAP}}u; const ATTN_SCALE: f32 = {{ATTN_SCALE}}; const WG: u32 = {{WG}}u; const JT: u32 = WG / D4; // attention phase-3 j-lanes per d-quad const TMP4: u32 = max(FFN4, HD4 + WG); // ffn row AND attn partial scratch fit var xs4: array, HD4>; // hidden state (residual base) var tmp4: array, TMP4>; // q / vbuf / ffn / attn partials var out4: array, HD4>; // stage outputs var scores: array; {{IF_NOSG}} var red: array; {{/IF_NOSG}} {{IF_SG}} // One partial per subgroup, in TWO alternating slots of NSG_CAP (WG/4 covers // the spec-minimum subgroup size 4). The alternation is what buys the single // barrier per call: call N's fold reads slot A strictly before every thread // passes call N+1's barrier (slot B), and call N+2's elect-writes to slot A // happen strictly after it — so no trailing barrier is needed to protect // reuse. sgId = tid/sgSize assumes the linear tid→subgroup layout (same bet // as add_layernorm.wgsl; the equiv gates catch a violating backend). const NSG_CAP: u32 = WG / 4u; var red: array; var sgId: u32; var nSg: u32; var redSlot: u32 = 0u; {{/IF_SG}} fn wgMax(tid: u32, v: f32) -> f32 { {{IF_SG}} let s1 = subgroupMax(v); let base = redSlot * NSG_CAP; if (subgroupElect()) { red[base + sgId] = s1; } workgroupBarrier(); var r = red[base]; for (var i = 1u; i < nSg; i = i + 1u) { r = max(r, red[base + i]); } redSlot = 1u - redSlot; return r; {{/IF_SG}} {{IF_NOSG}} red[tid] = v; workgroupBarrier(); for (var s = WG / 2u; s > 0u; s = s >> 1u) { if (tid < s) { red[tid] = max(red[tid], red[tid + s]); } workgroupBarrier(); } let r = red[0]; workgroupBarrier(); // red[0] reads done before the next reduction reuses red return r; {{/IF_NOSG}} } fn wgSum(tid: u32, v: f32) -> f32 { {{IF_SG}} let s1 = subgroupAdd(v); let base = redSlot * NSG_CAP; if (subgroupElect()) { red[base + sgId] = s1; } workgroupBarrier(); var r = red[base]; for (var i = 1u; i < nSg; i = i + 1u) { r = r + red[base + i]; } redSlot = 1u - redSlot; return r; {{/IF_SG}} {{IF_NOSG}} red[tid] = v; workgroupBarrier(); for (var s = WG / 2u; s > 0u; s = s >> 1u) { if (tid < s) { red[tid] = red[tid] + red[tid + s]; } workgroupBarrier(); } let r = red[0]; workgroupBarrier(); return r; {{/IF_NOSG}} } // One GEMV output quad, gemm_row_ln-style: thread computes outputs // 4·n4 .. 4·n4+3 from the [K, N] row-major W — at each k, threads read // consecutive quads of the k-row (coalesced across the workgroup). srcSel // picks the shared source (0 = xs4, 1 = out4, 2 = tmp4); kq = K/4 source // quads, nq = N/4 output quads (the W row stride). Four independent // accumulators keep 4 loads in flight; the fold order is fixed // (a0+a1)+(a2+a3). Returns f32 WITHOUT rounding — the caller rounds/routes. fn gemvQuad(wOff: u32, bOff: u32, n4: u32, kq: u32, nq: u32, srcSel: u32) -> vec4 { var a0 = vec4(0.0); var a1 = vec4(0.0); var a2 = vec4(0.0); var a3 = vec4(0.0); for (var k4 = 0u; k4 < kq; k4 = k4 + 1u) { var xq: vec4; if (srcSel == 0u) { xq = xs4[k4]; } else if (srcSel == 1u) { xq = out4[k4]; } else { xq = tmp4[k4]; } let kBase = wOff + (k4 << 2u) * nq + n4; a0 = fma(vec4(xq.x), vec4(W[kBase]), a0); a1 = fma(vec4(xq.y), vec4(W[kBase + nq]), a1); a2 = fma(vec4(xq.z), vec4(W[kBase + 2u * nq]), a2); a3 = fma(vec4(xq.w), vec4(W[kBase + 3u * nq]), a3); } return (a0 + a1) + (a2 + a3) + vec4(W[bOff + n4]); } @compute @workgroup_size({{WG}}) fn main(@builtin(workgroup_id) wid: vec3, @builtin(local_invocation_id) lid: vec3{{IF_SG}}, @builtin(subgroup_size) sgSize: u32{{/IF_SG}}) { // Uniform per workgroup — safe early return before the first barrier. if (wid.x >= params.B) { return; } let b = wid.x; let tid = lid.x; let t = params.t; {{IF_SG}} sgId = tid / sgSize; nSg = (WG + sgSize - 1u) / sgSize; {{/IF_SG}} // ---- S0: hidden state into xs4 ---- {{IF_EMBED}} // embed.wgsl DECODE semantics: id = DECODER_START at t=0, else the ring // token; pos = t; y = f16round(table·EMBED_SCALE + pos_embed). var id: u32 = {{DECODER_START}}u; if (t != 0u) { id = ring[(t - 1u) * params.B + b]; } for (var i = tid; i < HD4; i = i + WG) { let e = vec4(W[{{TABLE4}}u + id * HD4 + i]) * {{EMBED_SCALE}} + vec4(W[{{POS4}}u + t * HD4 + i]); xs4[i] = vec4(vec4<{{T}}>(e)); } {{/IF_EMBED}} {{IF_NOEMBED}} // Phony use: only the embed fold reads the ring, but the binding must stay // statically used or layout 'auto' drops @binding(2) and the bind group // (which always supplies it) fails validation — killing the whole submit. _ = ring[0]; for (var i = tid; i < HD4; i = i + WG) { xs4[i] = vec4(X[b * HD4 + i]); } {{/IF_NOEMBED}} workgroupBarrier(); // ---- S1: fused qkv projection; q → tmp4[0..HD4), k|v quads → caches ---- let kvBase = (b * LMAX + t) * HD4; for (var n4 = tid; n4 < QKV4; n4 = n4 + WG) { let g = vec4<{{T}}>(gemvQuad({{QKVW4}}u, {{QKVB4}}u, n4, HD4, QKV4, 0u)); if (n4 < HD4) { tmp4[n4] = vec4(g); } else if (n4 < 2u * HD4) { Kc[kvBase + n4 - HD4] = g; } else { Vc[kvBase + n4 - 2u * HD4] = g; } } workgroupBarrier(); // ---- S2: self-attention over positions 0..t (attention.wgsl phases) ---- { let len = min(t + 1u, LMAX); for (var h = 0u; h < H; h = h + 1u) { let hq = h * D4; var lm: f32 = -1e30; for (var j = tid; j < len; j = j + WG) { let koff = (b * LMAX + j) * HD4 + hq; var dot4 = vec4(0.0); for (var i = 0u; i < D4; i = i + 1u) { dot4 = dot4 + tmp4[hq + i] * vec4(Kc[koff + i]); } let sc = (dot4.x + dot4.y + dot4.z + dot4.w) * ATTN_SCALE; scores[j] = sc; lm = max(lm, sc); } let rowMax = wgMax(tid, lm); var ls: f32 = 0.0; for (var j = tid; j < len; j = j + WG) { let e = exp(scores[j] - rowMax); scores[j] = e; ls = ls + e; } let denom = wgSum(tid, ls); let dq = tid % D4; let jg = tid / D4; var acc = vec4(0.0); if (jg < JT) { for (var j = jg; j < len; j = j + JT) { acc = acc + scores[j] * vec4(Vc[(b * LMAX + j) * HD4 + hq + dq]); } } tmp4[HD4 + tid] = acc; // partial scratch; q region [0..HD4) untouched workgroupBarrier(); if (tid < D4) { var o = vec4(0.0); for (var g = 0u; g < JT; g = g + 1u) { o = o + tmp4[HD4 + g * D4 + tid]; } out4[hq + tid] = vec4(vec4<{{T}}>(o / denom)); } workgroupBarrier(); // out4 + partial reads done before the next head } } // ---- S3: x = LN1(x + self_out(attn)); vbuf = tmp4 ---- for (var n4 = tid; n4 < HD4; n4 = n4 + WG) { let g = vec4<{{T}}>(gemvQuad({{OUTW4}}u, {{OUTB4}}u, n4, HD4, HD4, 1u)); tmp4[n4] = vec4(g) + xs4[n4]; } workgroupBarrier(); { var s: f32 = 0.0; for (var i = tid; i < HD4; i = i + WG) { let v = tmp4[i]; s = s + v.x + v.y + v.z + v.w; } let mu = wgSum(tid, s) / f32(H * D); var sq: f32 = 0.0; for (var i = tid; i < HD4; i = i + WG) { let dv = tmp4[i] - vec4(mu); sq = sq + dot(dv, dv); } let inv = inverseSqrt(wgSum(tid, sq) / f32(H * D) + {{EPS}}); for (var i = tid; i < HD4; i = i + WG) { let o = vec4(W[{{LN1G4}}u + i]) * (tmp4[i] - vec4(mu)) * inv + vec4(W[{{LN1B4}}u + i]); xs4[i] = vec4(vec4<{{T}}>(o)); } } workgroupBarrier(); // ---- S4: q = cross_q(x) → out4 ---- for (var n4 = tid; n4 < HD4; n4 = n4 + WG) { out4[n4] = vec4(vec4<{{T}}>(gemvQuad({{CQW4}}u, {{CQB4}}u, n4, HD4, HD4, 0u))); } workgroupBarrier(); // ---- S5: cross-attention over lens[b] encoder positions → tmp4[0..HD4) ---- { let len = min(lens[b], SCORES_MAX); for (var h = 0u; h < H; h = h + 1u) { let hq = h * D4; var lm: f32 = -1e30; for (var j = tid; j < len; j = j + WG) { let koff = (b * params.S + j) * 2u * HD4 + hq; // k slice at offset 0 var dot4 = vec4(0.0); for (var i = 0u; i < D4; i = i + 1u) { dot4 = dot4 + out4[hq + i] * vec4(CKV[koff + i]); } let sc = (dot4.x + dot4.y + dot4.z + dot4.w) * ATTN_SCALE; scores[j] = sc; lm = max(lm, sc); } let rowMax = wgMax(tid, lm); var ls: f32 = 0.0; for (var j = tid; j < len; j = j + WG) { let e = exp(scores[j] - rowMax); scores[j] = e; ls = ls + e; } let denom = wgSum(tid, ls); let dq = tid % D4; let jg = tid / D4; var acc = vec4(0.0); if (jg < JT) { for (var j = jg; j < len; j = j + JT) { // v slice at element offset H·D within the fused k|v position acc = acc + scores[j] * vec4(CKV[(b * params.S + j) * 2u * HD4 + HD4 + hq + dq]); } } tmp4[HD4 + tid] = acc; // fold results land in [0..HD4) only afterwards workgroupBarrier(); if (tid < D4) { var o = vec4(0.0); for (var g = 0u; g < JT; g = g + 1u) { o = o + tmp4[HD4 + g * D4 + tid]; } tmp4[hq + tid] = vec4(vec4<{{T}}>(o / denom)); } workgroupBarrier(); } } // ---- S6: x = LN2(x + cross_out(attn)); vbuf = out4 ---- for (var n4 = tid; n4 < HD4; n4 = n4 + WG) { let g = vec4<{{T}}>(gemvQuad({{COW4}}u, {{COB4}}u, n4, HD4, HD4, 2u)); out4[n4] = vec4(g) + xs4[n4]; } workgroupBarrier(); { var s: f32 = 0.0; for (var i = tid; i < HD4; i = i + WG) { let v = out4[i]; s = s + v.x + v.y + v.z + v.w; } let mu = wgSum(tid, s) / f32(H * D); var sq: f32 = 0.0; for (var i = tid; i < HD4; i = i + WG) { let dv = out4[i] - vec4(mu); sq = sq + dot(dv, dv); } let inv = inverseSqrt(wgSum(tid, sq) / f32(H * D) + {{EPS}}); for (var i = tid; i < HD4; i = i + WG) { let o = vec4(W[{{LN2G4}}u + i]) * (out4[i] - vec4(mu)) * inv + vec4(W[{{LN2B4}}u + i]); xs4[i] = vec4(vec4<{{T}}>(o)); } } workgroupBarrier(); // ---- S7: ffn = SiLU(fc1(x)) → tmp4[0..FFN4) (SiLU in f32, then round) ---- for (var n4 = tid; n4 < FFN4; n4 = n4 + WG) { var v = gemvQuad({{FC1W4}}u, {{FC1B4}}u, n4, HD4, FFN4, 0u); v = v / (vec4(1.0) + exp(-v)); tmp4[n4] = vec4(vec4<{{T}}>(v)); } workgroupBarrier(); // ---- S8: X ← LN3(x + fc2(ffn)); vbuf = out4 ---- for (var n4 = tid; n4 < HD4; n4 = n4 + WG) { let g = vec4<{{T}}>(gemvQuad({{FC2W4}}u, {{FC2B4}}u, n4, KQ_FFN, HD4, 2u)); out4[n4] = vec4(g) + xs4[n4]; } workgroupBarrier(); { var s: f32 = 0.0; for (var i = tid; i < HD4; i = i + WG) { let v = out4[i]; s = s + v.x + v.y + v.z + v.w; } let mu = wgSum(tid, s) / f32(H * D); var sq: f32 = 0.0; for (var i = tid; i < HD4; i = i + WG) { let dv = out4[i] - vec4(mu); sq = sq + dot(dv, dv); } let inv = inverseSqrt(wgSum(tid, sq) / f32(H * D) + {{EPS}}); for (var i = tid; i < HD4; i = i + WG) { let o = vec4(W[{{LN3G4}}u + i]) * (out4[i] - vec4(mu)) * inv + vec4(W[{{LN3B4}}u + i]); X[b * HD4 + i] = vec4<{{T}}>(o); } } }