// Encoder pass: embedding → ENC_LAYERS post-LN transformer layers → fused // cross-K/V projections for every decoder layer. Everything is recorded into ONE // command encoder / ONE compute pass (WebGPU orders dispatches that touch the // same storage buffers within a pass) and submitted once. // // Per layer (post-LN Marian): // qkvOut = x @ qkv.w + qkv.b [B·S, 3·HD] fused q|k|v // attnOut = attention(qkvOut) reads the fused buffer // directly via strides // y = attnOut @ out.w + out.b // x' = LN1(y + x) // ffnTmp = SiLU(x' @ fc1.w + fc1.b) // y = ffnTmp @ fc2.w + fc2.b // x'' = LN2(y + x') // LN output must not alias its residual input (read + read_write on one // buffer is a WebGPU usage conflict), so the hidden state ping-pongs // a → b → a within each layer and ends every layer back in `a`. import { createSplitArena } from './arena.js'; import { dispatchGemm, dispatchAttention, dispatchAddLn, dispatchEmbed, dispatchScatterRows } from './pipelines.js'; import { D_MODEL, HEADS, HEAD_DIM, FFN, ENC_LAYERS, DEC_LAYERS, assertModelActive, } from './constants.js'; // ctx = {device, dtype?}; activation dtype follows weights.dtype (the kernels' // {{T}} must match the storage type of the weight tensors). // Returns {encOut, crossKV: [GPUBuffer×2], lensBuf, B, S, arena}. The returned // arena owns the retained buffers; encoder-only scratch is destroyed right // after a normal submit. `retainEncOut=false` skips the inspectable final // hidden output used only by encoder debug/readback gates. // // recordInto: profiling hook — an already-begun compute pass (or a proxy // implementing setPipeline/setBindGroup/dispatchWorkgroups, see // profile.js makeProfilingPass). When given, dispatches are recorded into it // and NOTHING is submitted: the caller owns pass.end()/submit and must // destroy the returned `scratch` buffers after its submit. // // gemmOverrides: tuning hook — merged over the default GEMM flags (e.g. // {bkk: 32} to change the tile K-slice, {tiled: false} for the naive-kernel // control in sweep tests). attnOverrides: same for the self-attention site // (e.g. {block: false} for the unblocked control, {qb, jb} tile shapes). // Production callers leave both unset. // // packed: encoder row-packing — drop the pad rows (file23k batch fill: 80.8% // at b128) so every GEMM runs at T = Σ lens rows instead of B·S and // attention early-exits whole query blocks past len. crossKV and encOut are // scattered back to the padded [B·S, ·] layout at the end of the pass, so // the returned contract (and the decoder, compaction, every test) is // unchanged — bit-exact on valid rows, zeros on pad rows (enc_pack_equiv; // unpacked pad rows hold computed garbage nothing reads). 'auto' (default) // packs except under attnOverrides — the attention sweeps A/B the padded // kernel variants (incl. block:false, which has no packed path). // splitSubmits: instead of ONE submit, cut the pass at the embed and every // layer boundary (8 submits total). Dispatch order and results are identical // (queue order sequences chunks exactly like one pass); only the submit // boundaries move, so no single submit keeps a slow GPU busy for seconds — // the Android-watchdog regime where the batch with the largest B·S dies with // VK_ERROR_DEVICE_LOST at its encoder submit. Ignored under recordInto. export async function runEncoder(ctx, weights, { ids, lens, B, S }, { recordInto = null, gemmOverrides = null, attnOverrides = null, packed = 'auto', sg = false, attnQbAlign8 = false, retainEncOut = true, splitSubmits = false } = {}) { assertModelActive(weights.model, 'runEncoder weights'); const HD = HEADS * HEAD_DIM; // == D_MODEL (enforced by applyModelConfig) const QKV_N = 3 * HD; // fused q|k|v const CROSS_KV_N = 2 * HD; // fused k|v const { device } = ctx; const t = weights.dtype; const eb = t === 'f16' ? 2 : 4; const padRows = B * S; let total = 0; for (let i = 0; i < B; i++) total += lens[i]; const usePacked = (packed === 'auto' ? !attnOverrides : !!packed) && total < padRows && S <= 0xffff; const rows = usePacked ? total : padRows; const flags = { t }; // Subgroup add_ln (translateBatch passes tuned sg through): perf knob on // Metal, CORRECTNESS requirement on Adreno-class devices where the tree // reduction miscompiles (m1_encoder_parity failure, 2026-07 Android round). const lnFlags = sg ? { ...flags, sg: true } : flags; // Encoder GEMMs are large-M (rows = B·S ≥ 32) — route them to the tiled // kernel (enc_profile baseline: GEMMs = ~88% of encoder time on the naive // kernel). tm8 (8×4 register subtile, gemm_tiled2.wgsl): tiled_sweep // 2026-07-06 b64 medians v1 84.5 / v2 63.0 / v2+tm8 58.9ms. Non-GEMM sites // (attention/add_ln/embed) keep plain `flags`. const gemmFlags = { ...flags, tiled: true, tm8: true, ...(gemmOverrides ?? {}) }; // Self-attention: blocked kernel (QB queries share staged K/V tiles — // attn_sweep 2026-07-06 picks the default QB/JB in pipelines.js). const attnFlags = { ...flags, block: true, packed: usePacked, ...(attnQbAlign8 ? { qbAlign8: true } : {}), ...(attnOverrides ?? {}) }; const arena = createSplitArena(device); const scratch = []; try { const retained = arena.retained; const scratchArena = arena.scratch; const upload = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST; const act = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC; const idsBuf = scratchArena.buf(rows * 4, upload, 'enc ids'); const lensBuf = retained.buf(B * 4, upload, 'enc lens'); let startsBuf = null; if (usePacked) { // Packed rows: sequence bi occupies rows starts[bi] .. starts[bi]+len; // each ids word carries its own position (embed.wgsl PACKED). const words = new Uint32Array(rows); const starts = new Uint32Array(B); let r = 0; for (let bi = 0; bi < B; bi++) { starts[bi] = r; for (let m = 0; m < lens[bi]; m++, r++) words[r] = (m << 16) | ids[bi * S + m]; } startsBuf = scratchArena.buf(B * 4, upload, 'enc starts'); device.queue.writeBuffer(idsBuf, 0, words); device.queue.writeBuffer(startsBuf, 0, starts); } else { device.queue.writeBuffer(idsBuf, 0, ids); } device.queue.writeBuffer(lensBuf, 0, lens); // In the unpacked inspectable path `a` is also encOut, so retain it. The // production path and every packed path keep `a` scratch-only. const aArena = retainEncOut && !usePacked ? retained : scratchArena; const a = aArena.buf(rows * D_MODEL * eb, act, 'enc hidden a'); const b = scratchArena.buf(rows * D_MODEL * eb, act, 'enc hidden b'); const y = scratchArena.buf(rows * D_MODEL * eb, act, 'enc sublayer y'); const qkvOut = scratchArena.buf(rows * QKV_N * eb, act, 'enc qkv out'); const attnOut = scratchArena.buf(rows * HD * eb, act, 'enc attn out'); const ffnTmp = scratchArena.buf(rows * FFN * eb, act, 'enc ffn tmp'); // crossKV keeps the padded [B·S, 2·HD] contract; packed runs project into // packed temporaries and scatter into these at the end of the pass. const crossKV = [ retained.buf(padRows * CROSS_KV_N * eb, act, 'crossKV dec.0'), retained.buf(padRows * CROSS_KV_N * eb, act, 'crossKV dec.1'), ]; const crossKVP = usePacked ? [ scratchArena.buf(rows * CROSS_KV_N * eb, act, 'crossKV packed dec.0'), scratchArena.buf(rows * CROSS_KV_N * eb, act, 'crossKV packed dec.1'), ] : crossKV; const W = (name) => weights.bindingFor(name); let encoder = recordInto ? null : device.createCommandEncoder({ label: 'encoder pass' }); let pass = recordInto ?? encoder.beginComputePass({ label: 'encoder pass' }); const rec = ({ scratch: s }) => scratch.push(...s); // splitSubmits: close the current chunk and open the next. Buffers written // by a submitted chunk are queue-retained; later chunks read them in queue // order, so results are bit-identical to the single-submit pass. const split = !recordInto && splitSubmits; let chunkIdx = 0; const cut = () => { if (!split) return; pass.end(); device.queue.submit([encoder.finish()]); chunkIdx++; encoder = device.createCommandEncoder({ label: `encoder chunk ${chunkIdx}` }); pass = encoder.beginComputePass({ label: `encoder chunk ${chunkIdx}` }); }; // 1. Embedding: x = shared[id]·EMBED_SCALE + pos_embed[pos] → a. rec(dispatchEmbed(device, pass, { ids: idsBuf, table: W('shared.weight'), posEmbed: W('pos_embed'), y: a, mode: 'src', nRows: rows, batch: B, s: S, packed: usePacked, flags, })); cut(); // 2. Transformer layers. Hidden state: in `a` at layer start and layer end. for (let l = 0; l < ENC_LAYERS; l++) { const p = (name) => `enc.${l}.${name}`; // Self-attention block: a → b rec(dispatchGemm(device, pass, { x: a, w: W(p('qkv.weight')), b: W(p('qkv.bias')), y: qkvOut, M: rows, K: D_MODEL, N: QKV_N, flags: gemmFlags, })); rec(dispatchAttention(device, pass, { q: qkvOut, k: qkvOut, v: qkvOut, lens: lensBuf, y: attnOut, starts: startsBuf, B, M: S, L: S, lenMode: 1, qStride: QKV_N, qOff: 0, kvStride: QKV_N, kOff: HD, vOff: 2 * HD, flags: attnFlags, })); rec(dispatchGemm(device, pass, { x: attnOut, w: W(p('out.weight')), b: W(p('out.bias')), y, M: rows, K: HD, N: D_MODEL, flags: gemmFlags, })); rec(dispatchAddLn(device, pass, { x: y, r: a, gamma: W(p('ln1.weight')), beta: W(p('ln1.bias')), y: b, rows, flags: lnFlags, })); // FFN block: b → a rec(dispatchGemm(device, pass, { x: b, w: W(p('fc1.weight')), b: W(p('fc1.bias')), y: ffnTmp, M: rows, K: D_MODEL, N: FFN, flags: { ...gemmFlags, silu: true }, })); rec(dispatchGemm(device, pass, { x: ffnTmp, w: W(p('fc2.weight')), b: W(p('fc2.bias')), y, M: rows, K: FFN, N: D_MODEL, flags: gemmFlags, })); rec(dispatchAddLn(device, pass, { x: y, r: b, gamma: W(p('ln2.weight')), beta: W(p('ln2.bias')), y: a, rows, flags: lnFlags, })); cut(); } // 3. Fused cross-attention K/V for both decoder layers (attention reads // these directly with kvStride 896, kOff 0, vOff 448). for (let l = 0; l < DEC_LAYERS; l++) { rec(dispatchGemm(device, pass, { x: a, w: W(`dec.${l}.cross_kv.weight`), b: W(`dec.${l}.cross_kv.bias`), y: crossKVP[l], M: rows, K: D_MODEL, N: CROSS_KV_N, flags: gemmFlags, })); } // 4. Packed runs: scatter crossKV and the final hidden state back to the // padded layout (valid rows copied, pad rows stay zero — see // scatter_rows.wgsl) so downstream consumers see the unpacked contract. let encOut = retainEncOut && !usePacked ? a : null; if (usePacked) { for (let l = 0; l < DEC_LAYERS; l++) { rec(dispatchScatterRows(device, pass, { x: crossKVP[l], y: crossKV[l], starts: startsBuf, lens: lensBuf, B, S, N: CROSS_KV_N, flags, })); } if (retainEncOut) { encOut = retained.buf(padRows * D_MODEL * eb, act, 'enc out padded'); rec(dispatchScatterRows(device, pass, { x: a, y: encOut, starts: startsBuf, lens: lensBuf, B, S, N: D_MODEL, flags, })); } } if (recordInto) { return { encOut, crossKV, lensBuf, B, S, arena, scratch, packed: usePacked }; } pass.end(); device.queue.submit([encoder.finish()]); // Capture the encoder-only queue boundary before any caller can submit // decode work. The promise is used both for honest submitted-byte accounting // and for the existing non-blocking encoder timing in translateBatch. const submittedDone = device.queue.onSubmittedWorkDone(); for (const buf of scratch) buf.destroy(); // safe post-submit // Submitted commands retain their resources until execution completes. // Retiring JS ownership here lets WebKit release the large encoder-only // backing allocations as soon as that submit drains, instead of after the // entire decode batch. arena.retireScratch(submittedDone); return { encOut, crossKV, lensBuf, B, S, arena, packed: usePacked, submittedDone }; } catch (err) { // A failure before ownership is returned must retire both dispatch-local // uniforms and every partially allocated arena buffer. Cleanup is best // effort so the original encoder error remains the one callers see. for (const buf of scratch) { try { buf.destroy(); } catch {} } try { arena.destroy(); } catch {} throw err; } }