vp2vi / engine /encoder.js
DanVP's picture
feat: publish vp2vi WebGPU browser app
c971a45 verified
Raw
History Blame Contribute Delete
12.6 kB
// 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;
}
}