// Full GPU-resident greedy decode: tokenize → encoder (one submit) → decode // loop in groups of K=8 steps per submit. Within a group all 25·K dispatches // are recorded into one compute pass — the argmax kernel writes each step's // token into the ring and the next step's embed reads it, so the CPU never // sits between steps. After each group ONE readback (ring slice + done flags) // tells the CPU which tokens were produced and whether every row has hit eos. // // Decode-ahead (default): group g+1 is encoded and submitted BEFORE awaiting // group g's readback, so the GPU never drains while the CPU maps/collects. // GPU-side correctness is free — queue order guarantees group g's argmax // writes the ring before group g+1's embed reads it. The cost is that EOS // early-exit lags one group: when group g's flags show all rows done, group // g+1 has already been submitted (≤ (inFlight-1)·GROUP_STEPS wasted steps per // run; done rows just produce PAD, which never enters a sequence). inFlight // staging buffers rotate: a buffer is copied into, mapped, read, unmapped — // and only then reused (N groups in flight ⇒ N buffers is exactly enough). // Pass {decodeAhead: false} for the sequential A/B path. // // Adaptive submit budget (submitBudgetMs, default 1000): group readback // waits over the budget halve the steps-per-submit for the rest of the // batch — slow mobile GPUs converge to submits the OS watchdog tolerates // instead of dying with VK_ERROR_DEVICE_LOST; fast GPUs never trip it. // // Timing semantics (spike-grade, documented not perfect): // encoderMs wall time from the runEncoder() call until // queue.onSubmittedWorkDone() resolves for its submit — // includes tokenized-batch upload + command recording overhead. // decodeMs wall time of the whole decode loop (GPU work + per-group // mapAsync readbacks + CPU collection). // cpuEncodeMs Σ per-group CPU time recording commands + submitting. // awaitMs Σ per-group time awaiting mapAsync + copying the readback. import { tokenizeBatch } from './tokenizer.js'; import { runEncoder } from './encoder.js'; import { createDecodeState, compactDecodeState, compactDecodeStateInPlace, encodeDecodeStep, growDecodeKV, } from './decoder.js'; import { createUniformParamPool, getDispatchStats, shouldUseUniformParamPool, } from './pipelines.js'; import { maxNewTokensFor, maxBatchForLimits } from './shapes.js'; import { EOS, DECODER_START, DECODE_CAP } from './constants.js'; const GROUP_STEPS = 8; // decode steps recorded per submit const SUBMIT_BUDGET_MS = 1000; // adaptive per-submit ceiling (submitBudgetMs) // One-way ratchet for the adaptive submit budget: halve steps-per-submit // whenever a group's readback wait exceeded the budget. With the pipeline // primed the CPU parks on mapAsync for almost exactly the GPU tail of that // group, so the wait is a lower bound on the group's GPU time — overshooting // means the submit kept the GPU busy past the budget, the regime where // Android compositor fences start missing and the driver eventually kills // the context (measured on Adreno 710: 3000ms fence misses, then // vkQueueSubmit VK_ERROR_DEVICE_LOST once submits reach ~4-5s; ~1.5s // submits survive a 286k-char run). Never grows back within a batch: file // mode sorts long rows last so pressure only rises, a false shrink costs a // few % submit overhead, a missed shrink costs the device. export function nextGroupSteps(cur, groupAwaitMs, budgetMs) { if (!budgetMs || cur <= 1 || !(groupAwaitMs > budgetMs)) return cur; return Math.max(1, cur >> 1); } // translateBatch(ctx, weights, tok, sources, // {maxNewTokens?, onProgress?, onPrimed?, decodeAhead?, // inFlight?, compact?, overlapEnc?}) // ctx {device, limits?, ...} from initDevice() // sources array of source strings (one batch) // onPrimed called ONCE, right after the encoder and the first decode // groups are submitted and the loop is about to park on the GPU — // the spot where caller CPU work (pre-tokenizing the next batch) // overlaps GPU execution instead of delaying submits // overlapEnc (default true) skip the queue drain between encoder and // decode: queue order already sequences crossKV before its // readers, so the first decode groups are recorded while the GPU // still encodes. encoderMs is then measured via a non-blocking // onSubmittedWorkDone.then() and overlaps decodeMs by a few ms // (stage sums read slightly high; wall is what drops). false // restores the drained baseline (A/B arm). // inFlight decode-ahead depth: 2 (default, double-buffered staging) or 3 // (triple-buffered). Measured on 64 file23k chunks at B=64 // (m5_inflight_ab, 2026-07-06): median 28737 vs 28724 tok/s — a // wash; the readback gap is already hidden at depth 2, so 2 stays // the default. // compact EOS row compaction (default true): when ≥ max(8, B·compactFrac) // rows of the current batch have emitted eos, stop submitting, // drain the in-flight groups, and rebuild the decode context with // only the live rows (compactDecodeState) — finished rows // otherwise keep burning GEMM rows and attention workgroups until // the whole group hits eos. Kernel routing is pinned across // compactions, so output is token-exact vs {compact: false} // (compact_equiv gate). // compactFrac dead fraction of the CURRENT batch that triggers a compaction // (default 0.25 — the original B>>2). Lower = more compactions: // each costs a drain of the in-flight groups plus the live-row // copies, each saves dead-row GEMM/attention work for every // remaining step. Compaction count and step timings shift but // tokens stay exact at ANY value. MEASURED INSENSITIVE on sorted // file batches (dec_compact_sweep 2026-07-09: 0.25→0.03 identical // wall/compactions — uniform rows die in one synchronized wave, so // every threshold fires at the same group boundary). // groupSteps decode steps recorded per submit (default 8). Smaller halves // the EOS-detection latency (dead rows compute until their group's // readback lands) at the cost of more submits; token-exact at any // value — compaction timing shifts, row math doesn't // (dec_group_sweep A/Bs this). // submitBudgetMs adaptive submit ceiling (default 1000, null/0 disables): // whenever one group's readback wait exceeds this, steps-per-submit // halve for the REST of the batch (8→4→2→1, never back up). Fast // GPUs never trip it (group waits are tens of ms); slow mobile // GPUs converge within a few groups to submits the OS watchdog // tolerates instead of VK_ERROR_DEVICE_LOST. Token-exact like any // groupSteps value. Callers can seed the next batch with this // batch's landing point via metrics.groupStepsFinal. // Returns: // rows [{ids, text, forcedEos, steps}] — ids = [0, ...tokens] trimmed at // the first eos inclusive; rows that never emitted eos within the // cap get one appended (forcedEos: true — HF max_length semantics) // metrics {tokenizeMs, encoderMs, decodeMs, detokMs, cpuEncodeMs, awaitMs, // submits, steps, tokensGenerated, tokPerSec} // tokenizeMs/detokMs bracket the CPU tokenizer calls — the two // stages the GPU timings can't see (the app-stage breakdown needs // them to locate wall-clock loss on large files) // tokensGenerated counts argmax-produced tokens across ALL rows // (incl. an emitted eos, excl. a force-appended one); tokPerSec is // total tokens over decodeMs. export async function translateBatch(ctx, weights, tok, sources, { maxNewTokens, onProgress, onPrimed = null, decodeAhead = true, inFlight = 2, compact = true, compactFrac = 0.25, inPlaceCompact = false, groupSteps = GROUP_STEPS, submitBudgetMs = SUBMIT_BUDGET_MS, overlapEnc = true, kvCapacity = null, // runEncoder passthrough (row-packing A/B; encSplitSubmits cuts the // encoder into per-layer submits — the watchdog guard for the encoder // side, see runEncoder splitSubmits) encPacked = 'auto', encSplitSubmits = false, // createDecodeState passthrough ('q8' / fusion / layout A/B tests, and the // per-batch options tunedOptions(tuned, B) resolves after an autotune run) lmHead = 'auto', ffn = 'auto', lmHeadFuse = 'auto', fuseLn = 'auto', proj = 'auto', ffnSplitK = 'auto', projSplitK = 'auto', decodeMega = 'auto', sg = 'auto', encAttnSafe = false, immediates = 'auto', uniformPool = false, tiledProj = 'auto', } = {}) { const { device } = ctx; if (typeof inPlaceCompact !== 'boolean') { throw new Error(`translateBatch: inPlaceCompact must be boolean, got ${inPlaceCompact}`); } const dispatchBefore = getDispatchStats(device); if (inFlight !== 2 && inFlight !== 3) { throw new Error(`translateBatch: inFlight must be 2 or 3, got ${inFlight}`); } if (!Number.isInteger(groupSteps) || groupSteps < 1) { throw new Error(`translateBatch: groupSteps must be a positive integer, got ${groupSteps}`); } const tTok0 = performance.now(); const batch = await tokenizeBatch(tok, sources); const tokenizeMs = performance.now() - tTok0; const { B } = batch; const srcTruncated = new Set(batch.truncated ?? []); // Memory guard: the encoder ffnTmp [B·S, 1792] binding is the ceiling. No // silent sub-batching — the bench/app layer owns batch-size policy, the // engine stays explicit. const maxB = maxBatchForLimits(ctx, batch.S, weights.dtype === 'f16' ? 2 : 4); if (B > maxB) { throw new Error( `translateBatch: batch ${B} at S=${batch.S} exceeds ` + `maxStorageBufferBindingSize=${ctx?.limits?.maxStorageBufferBindingSize ?? 'default 134217728'} ` + `(encoder ffnTmp [B·S, 1792]); split into batches of ≤ ${maxB}`, ); } const cap = Math.min( DECODE_CAP, maxNewTokens ?? Math.max(...sources.map((s) => maxNewTokensFor(s.length))), ); const tEnc0 = performance.now(); const encRun = await runEncoder(ctx, weights, batch, { packed: encPacked, sg: sg === 'on', attnQbAlign8: encAttnSafe, retainEncOut: false, splitSubmits: encSplitSubmits, }); let cur = null; const stagings = []; let paramPool = null; let cleaned = false; const cleanup = () => { if (cleaned) return; cleaned = true; // Cleanup is best effort so an allocation/decode failure is never masked // by a secondary destroy error. Every owner is idempotent. for (const staging of stagings) { try { staging.destroy(); } catch {} } try { paramPool?.destroy(); } catch {} try { cur?.state?.destroy(); } catch {} try { cur?.arena?.destroy(); } catch {} try { encRun.arena.destroy(); } catch {} }; try { // Don't drain the queue between encoder and decode (overlapEnc, default): // queue order already guarantees the decode dispatches see the finished // crossKV, so the CPU can record/submit the first decode groups WHILE the // GPU is still encoding — the old blocking await left the GPU idle for // exactly the CPU-side recording time of those groups every batch. // encoderMs still brackets submit → GPU-done via a non-blocking then() // (±one macrotask); with overlap on, the decode loop starts inside that // window, so encoderMs and decodeMs overlap by up to a few ms — the stage // sums in app_stage/translateText read slightly high, wall clock is what // dropped. {overlapEnc: false} restores the drained A/B baseline. let encoderMs = 0; const encDone = (encRun.submittedDone ?? device.queue.onSubmittedWorkDone()).then( () => { encoderMs = performance.now() - tEnc0; }, () => {}, // measurement only — a lost device surfaces via the decode loop ); if (!overlapEnc) await encDone; // The mutable decode context: state + the cross-attention view. Replaced // wholesale by each compaction; curMap maps current row -> original row. cur = { state: createDecodeState(ctx, weights, { B, S: encRun.S, maxSteps: cap, lmHead, ffn, lmHeadFuse, fuseLn, proj, tiledProj, ffnSplitK, projSplitK, decodeMega, sg, immediates, inPlaceCompact, kvCapacity, }), crossKV: encRun.crossKV, lensBuf: encRun.lensBuf, S: encRun.S, arena: null, }; let curB = B; let curMap = Array.from({ length: B }, (_, i) => i); const rowTokens = Array.from({ length: B }, () => []); const rowDone = new Array(B).fill(false); // CPU mirror: saw eos in the ring let submits = 0; let steps = 0; let compactions = 0; const initialKvCapacity = cur.state.kvCapacity; const kvCapacitySequence = [initialKvCapacity]; let kvGrows = 0; let kvGrowsWithPending = 0; let kvGrowBindGroupsPurged = 0; let cpuEncodeMs = 0; let awaitMs = 0; // Adaptive submit budget state: curGroupSteps only ever shrinks (see // nextGroupSteps); groupIndex rotates staging/pool banks by SUBMIT order, // which g/groupSteps no longer encodes once the group size changes. let curGroupSteps = groupSteps; let groupIndex = 0; let submitShrinks = 0; let maxGroupAwaitMs = 0; // inFlight staging buffers, rotated by submit index. Sized for a full // group at the STARTING groupSteps (the adaptive path only shrinks); // shorter groups copy less and read accordingly. const stagingBytes = (groupSteps * B + B) * 4; for (let i = 0; i < inFlight; i++) { stagings.push(device.createBuffer({ label: `decode group staging ${i}`, size: stagingBytes, usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.MAP_READ, })); } // Safari/WebKit currently lacks immediate_address_space. The selected pool // keeps ordinary uniform bindings but gives each in-flight group a stable // buffer bank, so parameter uploads collapse to one write and bind groups // become reusable. Immediate-mode states keep their existing path. paramPool = shouldUseUniformParamPool(uniformPool, cur.state) ? createUniformParamPool(device, { banks: inFlight }) : null; // Record + submit one group [g, tEnd). Returns the in-flight descriptor — // it carries the batch geometry (B, row map) it was submitted with, since a // compaction may swap the context before its readback is processed. const encodeGroup = (g, pendingGroups = 0) => { const t0 = performance.now(); const tEnd = Math.min(g + curGroupSteps, cap); if (tEnd > cur.state.kvCapacity) { const grow = growDecodeKV(ctx, weights, cur.state, { requiredCapacity: tEnd, submittedSteps: g, groupSteps: curGroupSteps, }); if (grow.grown) { kvGrows++; if (pendingGroups > 0) kvGrowsWithPending++; kvGrowBindGroupsPurged += grow.bindGroupsPurged; kvCapacitySequence.push(grow.newCapacity); } } const bank = groupIndex % stagings.length; groupIndex++; const staging = stagings[bank]; if (staging.mapState !== 'unmapped') { throw new Error(`decode staging reused while ${staging.mapState}`); } const scratch = []; let poolFrameActive = false; try { if (paramPool) { paramPool.begin(bank); poolFrameActive = true; } const encoder = device.createCommandEncoder({ label: `decode ${g}..${tEnd}` }); const pass = encoder.beginComputePass({ label: `decode ${g}..${tEnd}` }); for (let t = g; t < tEnd; t++) { scratch.push(...encodeDecodeStep(ctx, weights, cur, cur.state, t, pass).scratch); } pass.end(); if (paramPool) { paramPool.flush(); poolFrameActive = false; } // One readback per group: ring slice [g·B, tEnd·B) + done[B]. const nTok = (tEnd - g) * curB; encoder.copyBufferToBuffer(cur.state.tokenRing, g * curB * 4, staging, 0, nTok * 4); encoder.copyBufferToBuffer(cur.state.done, 0, staging, nTok * 4, curB * 4); device.queue.submit([encoder.finish()]); submits++; for (const buf of scratch) buf.destroy(); // safe post-submit cpuEncodeMs += performance.now() - t0; return { g, tEnd, nTok, staging, B: curB, map: curMap, poolBank: paramPool ? bank : null }; } catch (err) { if (poolFrameActive) paramPool.abort(); for (const buf of scratch) buf.destroy(); throw err; } }; // Compact away the finished rows: pays when a decent slice of the batch is // dead AND there are steps left. The max(8, ·) floor keeps small batches // (b1–b8 gates, latency runs) permanently on the no-compaction path. const compactWanted = () => { if (!compact || steps === 0) return false; const live = curMap.reduce((n, orig) => n + (rowDone[orig] ? 0 : 1), 0); return live > 0 && curB - live >= Math.max(8, Math.ceil(curB * compactFrac)); }; const doCompact = () => { const liveIdx = []; for (let i = 0; i < curB; i++) if (!rowDone[curMap[i]]) liveIdx.push(i); const newMap = liveIdx.map((i) => curMap[i]); const lens = new Uint32Array(newMap.map((orig) => batch.lens[orig])); const lastTok = new Uint32Array(newMap.map((orig) => rowTokens[orig].at(-1))); if (inPlaceCompact) { compactDecodeStateInPlace(ctx, weights, cur, { liveIdx, t0: steps, lens, lastTok, }); } else { const prev = cur; // pending is empty here, so every pool bank has completed its readback // and been released. Drop bind groups for the old resource generation // BEFORE allocating the compacted state: on WebKit those bind groups // keep the old KV/crossKV/arena backing memory alive after destroy(). paramPool?.invalidateBindings(); cur = compactDecodeState(ctx, weights, prev, { liveIdx, t0: steps, cap, lens, lastTok, }); // Old buffers are queue-retained by the just-submitted copies; destroy() // only blocks future submissions. First compaction: the encoder arena // (encOut + old crossKV/lens) is no longer referenced either. prev.state.destroy(); if (prev.arena) prev.arena.destroy(); else encRun.arena.destroy(); } curB = liveIdx.length; curMap = newMap; compactions++; }; const tDec0 = performance.now(); try { const depth = decodeAhead ? inFlight : 1; // groups in flight (submitted, unread) const pending = []; // submitted-but-unread groups, oldest first let nextG = 0; let allDone = false; for (;;) { // A wanted compaction stalls new submits so the in-flight groups (built // against the OLD layout) drain first — one pipeline bubble per compact. const wantCompact = !allDone && nextG < cap && compactWanted(); if (wantCompact && pending.length === 0) { doCompact(); continue; } while (!allDone && !wantCompact && nextG < cap && pending.length < depth) { const grp = encodeGroup(nextG, pending.length); pending.push(grp); nextG = grp.tEnd; // curGroupSteps may shrink between submits } // The pipeline is primed: encoder + the first decode groups are all // submitted, and the next await parks on the GPU. This is the one spot // where a caller can burn CPU for free (e.g. pre-tokenizing the NEXT // batch into the SPM LRU) — earlier would delay these submits, later // (onProgress) the GPU is already half done. if (onPrimed) { const cb = onPrimed; onPrimed = null; cb(); } const grp = pending.shift(); if (!grp) break; if (allDone) continue; // submitted before all-done was seen; results are // PAD-only for done rows — safe to ignore unread. const tA0 = performance.now(); await grp.staging.mapAsync(GPUMapMode.READ, 0, (grp.nTok + grp.B) * 4); const data = new Uint32Array(grp.staging.getMappedRange(0, (grp.nTok + grp.B) * 4).slice(0)); grp.staging.unmap(); // staging is now free for group g+2 if (grp.poolBank !== null) paramPool.release(grp.poolBank); const groupAwait = performance.now() - tA0; awaitMs += groupAwait; if (groupAwait > maxGroupAwaitMs) maxGroupAwaitMs = groupAwait; // Adaptive submit budget: a long wait here means the GPU chewed on one // submit past the budget — shrink the groups still to be submitted. const shrunk = nextGroupSteps(curGroupSteps, groupAwait, submitBudgetMs); if (shrunk !== curGroupSteps) { curGroupSteps = shrunk; submitShrinks++; } // Collect per row (map current -> original), stopping at its first eos — // done rows produce PAD afterwards, which must NOT enter the sequence. for (let t = grp.g; t < grp.tEnd; t++) { for (let b = 0; b < grp.B; b++) { const orig = grp.map[b]; if (rowDone[orig]) continue; const id = data[(t - grp.g) * grp.B + b]; rowTokens[orig].push(id); if (id === EOS) rowDone[orig] = true; } } steps = grp.tEnd; onProgress?.({ step: grp.tEnd, cap, done: rowDone.filter(Boolean).length, B }); // GPU done flags (queue-ordered snapshot after step tEnd-1). const doneFlags = data.subarray(grp.nTok, grp.nTok + grp.B); if (doneFlags.every((d) => d === 1)) allDone = true; } } finally { cleanup(); } const decodeMs = performance.now() - tDec0; // The decode loop's readbacks are queue-ordered after the encoder submit, // so encDone has long resolved — this await only pins encoderMs before the // metrics object is built. await encDone; const tDet0 = performance.now(); const rows = []; let tokensGenerated = 0; for (let b = 0; b < B; b++) { const forcedEos = !rowDone[b]; tokensGenerated += rowTokens[b].length; // argmax-produced (incl. emitted eos) const ids = [DECODER_START, ...rowTokens[b]]; if (forcedEos) ids.push(EOS); const text = tok.decode(ids, { skip_special_tokens: true }).trim(); // srcTruncated: the ENCODER saw a cut-off source (tokenizeBatch hit // SRC_CAP) — the decode itself is fine, but `text` only translates the // prefix. Distinct from forcedEos, which is the TARGET ring cap. rows.push({ ids, text, forcedEos, steps: rowTokens[b].length, srcTruncated: srcTruncated.has(b), }); } const detokMs = performance.now() - tDet0; const dispatchAfter = getDispatchStats(device); const dispatch = {}; for (const field of [ 'uniformBuffersCreated', 'uniformPoolBuffersCreated', 'uniformPoolBuffersDestroyed', 'uniformPoolFramesBegun', 'uniformPoolFramesFlushed', 'uniformPoolBlocks', 'uniformPoolBytes', 'uniformPoolBindGroupCacheHits', 'uniformPoolWarmBindGroupLookups', 'uniformPoolWarmBindGroupCacheHits', 'uniformPoolWarmBindGroupResets', 'uniformPoolGenerationInvalidations', 'uniformPoolCachePurges', 'dummyBuffersCreated', 'bindGroupsCreated', 'bindGroupCacheHits', 'bindGroupEvictions', 'bindGroupTargetedPurgeCalls', 'bindGroupTargetedPurges', 'immediateSets', ]) { dispatch[field] = dispatchAfter[field] - dispatchBefore[field]; } dispatch.bindGroupCacheSize = dispatchAfter.bindGroupCacheSize; dispatch.bindGroupCacheLimit = dispatchAfter.bindGroupCacheLimit; return { rows, metrics: { tokenizeMs, encoderMs, decodeMs, detokMs, cpuEncodeMs, awaitMs, B, S: encRun.S, cap, submits, steps, compactions, groupStepsStart: groupSteps, groupStepsFinal: curGroupSteps, submitShrinks, maxGroupAwaitMs, initialKvCapacity, finalKvCapacity: cur.state.kvCapacity, kvGrows, kvGrowsWithPending, kvCapacitySequence, kvGrowBindGroupsPurged, srcTruncated: srcTruncated.size, compactMode: inPlaceCompact ? 'inplace' : 'realloc', tokensGenerated, dispatch, uniformPool: paramPool?.snapshot() ?? null, tokPerSec: decodeMs > 0 ? (tokensGenerated / decodeMs) * 1000 : 0, }, }; } finally { // Covers initialization failures before the decode loop's narrower // finally is entered (decode-state, partial staging, or pool setup). cleanup(); } }