const WORKGROUP = 64; const TOP_K = 16; const OUTPUT_STRIDE = TOP_K + 2; const EMPTY_ID = 0xffffffff; const pipelines = new WeakMap(); const shader = /* wgsl */ ` struct Params { rows: u32, vocab: u32, maskWords: u32, _pad: u32, } @group(0) @binding(0) var logits: array; @group(0) @binding(1) var bias: array; @group(0) @binding(2) var maskBits: array; @group(0) @binding(3) var maskModes: array; @group(0) @binding(4) var output: array; @group(0) @binding(5) var params: Params; var sharedValues: array; var sharedIds: array; fn better(av: f32, ai: u32, bv: f32, bi: u32) -> bool { return av > bv || (av == bv && ai < bi); } fn insert16( values: ptr>, ids: ptr>, value: f32, id: u32, ) { if (!better(value, id, (*values)[15], (*ids)[15])) { return; } var pos = 15u; loop { if (pos == 0u || !better(value, id, (*values)[pos - 1u], (*ids)[pos - 1u])) { break; } (*values)[pos] = (*values)[pos - 1u]; (*ids)[pos] = (*ids)[pos - 1u]; pos -= 1u; } (*values)[pos] = value; (*ids)[pos] = id; } @compute @workgroup_size(64) fn main( @builtin(workgroup_id) group: vec3, @builtin(local_invocation_id) local: vec3, ) { let row = group.x; let lane = local.x; let mode = maskModes[row]; var localValues: array; var localIds: array; for (var k = 0u; k < 16u; k += 1u) { localValues[k] = -3.402823466e+38; localIds[k] = 0xffffffffu; } var rawBest = -3.402823466e+38; var rawId = 0u; var maskedBest = -3.402823466e+38; var maskedId = 0u; for (var token = lane; token < params.vocab; token += 64u) { let raw = logits[row * params.vocab + token] + bias[token]; if (better(raw, token, rawBest, rawId)) { rawBest = raw; rawId = token; } let word = maskBits[row * params.maskWords + (token >> 5u)]; let marked = (word & (1u << (token & 31u))) != 0u; let blocked = (mode == 1u && marked) || (mode == 2u && !marked); // Must match decode-guard.js NEG = -1e30 and its valid cutoff NEG/2. let value = select(raw, raw + -1.0e30, blocked); if (better(value, token, maskedBest, maskedId)) { maskedBest = value; maskedId = token; } if (value > -5.0e29) { insert16(&localValues, &localIds, value, token); } } for (var k = 0u; k < 16u; k += 1u) { let at = lane * 16u + k; sharedValues[at] = localValues[k]; sharedIds[at] = localIds[k]; } workgroupBarrier(); if (lane == 0u) { for (var k = 0u; k < 16u; k += 1u) { localValues[k] = -3.402823466e+38; localIds[k] = 0xffffffffu; } for (var at = 0u; at < 1024u; at += 1u) { let id = sharedIds[at]; if (id != 0xffffffffu) { insert16(&localValues, &localIds, sharedValues[at], id); } } for (var k = 0u; k < 16u; k += 1u) { output[row * 18u + k] = localIds[k]; } } workgroupBarrier(); sharedValues[lane] = maskedBest; sharedIds[lane] = maskedId; workgroupBarrier(); if (lane == 0u) { var bestValue = sharedValues[0]; var bestId = sharedIds[0]; for (var at = 1u; at < 64u; at += 1u) { if (better(sharedValues[at], sharedIds[at], bestValue, bestId)) { bestValue = sharedValues[at]; bestId = sharedIds[at]; } } output[row * 18u + 16u] = bestId; } workgroupBarrier(); sharedValues[lane] = rawBest; sharedIds[lane] = rawId; workgroupBarrier(); if (lane == 0u) { var bestValue = sharedValues[0]; var bestId = sharedIds[0]; for (var at = 1u; at < 64u; at += 1u) { if (better(sharedValues[at], sharedIds[at], bestValue, bestId)) { bestValue = sharedValues[at]; bestId = sharedIds[at]; } } output[row * 18u + 17u] = bestId; } } `; // Speculative-window variant (design addendum 2026-08-11 ยง3). Same ranking // core; three additions wired around it: // - the mask applies ONLY at each row's frontier step (t == frontier[row]), // free-run steps rank unmasked; // - the kernel writes the token ring itself (overwriting the vendor // argmax_penalty write earlier in the same encoder): t < frontier // restores the accepted history (histTok), t == frontier commits // maskedTop, t > frontier free-runs rawTop; // - output has one 18-u32 slot per window step so the whole window reads // back in a single mapAsync. const windowShader = /* wgsl */ ` struct Params { rows: u32, vocab: u32, maskWords: u32, t: u32, slot: u32, _pad0: u32, _pad1: u32, _pad2: u32, } @group(0) @binding(0) var logits: array; @group(0) @binding(1) var bias: array; @group(0) @binding(2) var maskBits: array; @group(0) @binding(3) var maskModes: array; @group(0) @binding(4) var frontier: array; @group(0) @binding(5) var histTok: array; @group(0) @binding(6) var tokenRing: array; @group(0) @binding(7) var output: array; @group(0) @binding(8) var params: Params; var sharedValues: array; var sharedIds: array; fn better(av: f32, ai: u32, bv: f32, bi: u32) -> bool { return av > bv || (av == bv && ai < bi); } fn insert16( values: ptr>, ids: ptr>, value: f32, id: u32, ) { if (!better(value, id, (*values)[15], (*ids)[15])) { return; } var pos = 15u; loop { if (pos == 0u || !better(value, id, (*values)[pos - 1u], (*ids)[pos - 1u])) { break; } (*values)[pos] = (*values)[pos - 1u]; (*ids)[pos] = (*ids)[pos - 1u]; pos -= 1u; } (*values)[pos] = value; (*ids)[pos] = id; } @compute @workgroup_size(64) fn main( @builtin(workgroup_id) group: vec3, @builtin(local_invocation_id) local: vec3, ) { let row = group.x; let lane = local.x; let front = frontier[row]; let atFrontier = params.t == front; let mode = select(0u, maskModes[row], atFrontier); let outBase = (params.slot * params.rows + row) * 18u; var localValues: array; var localIds: array; for (var k = 0u; k < 16u; k += 1u) { localValues[k] = -3.402823466e+38; localIds[k] = 0xffffffffu; } var rawBest = -3.402823466e+38; var rawId = 0u; var maskedBest = -3.402823466e+38; var maskedId = 0u; for (var token = lane; token < params.vocab; token += 64u) { let raw = logits[row * params.vocab + token] + bias[token]; if (better(raw, token, rawBest, rawId)) { rawBest = raw; rawId = token; } let word = maskBits[row * params.maskWords + (token >> 5u)]; let marked = (word & (1u << (token & 31u))) != 0u; let blocked = (mode == 1u && marked) || (mode == 2u && !marked); // Must match decode-guard.js NEG = -1e30 and its valid cutoff NEG/2. let value = select(raw, raw + -1.0e30, blocked); if (better(value, token, maskedBest, maskedId)) { maskedBest = value; maskedId = token; } if (value > -5.0e29) { insert16(&localValues, &localIds, value, token); } } for (var k = 0u; k < 16u; k += 1u) { let at = lane * 16u + k; sharedValues[at] = localValues[k]; sharedIds[at] = localIds[k]; } workgroupBarrier(); // Round-2: pairwise tree merge of 64 sorted-desc top-16 lists (6 levels) // replaces the serial 1,024-entry lane-0 merge. Each active lane selects // the top-16 of two sorted lists with the SAME comparator, so the result // is order-exact. Inside the 16-pick loop ai+bi == k <= 15, so neither // cursor can run past its list. Empty slots (-3.4e38 / 0xffffffff) sort // last naturally. Barriers stay in uniform control flow. var mergedV: array; var mergedI: array; for (var stride = 1u; stride < 64u; stride = stride << 1u) { let partner = lane + stride; let laneMerges = (lane & (2u * stride - 1u)) == 0u && partner < 64u; if (laneMerges) { var ai = 0u; var bi = 0u; for (var k = 0u; k < 16u; k += 1u) { let av = sharedValues[lane * 16u + ai]; let aid = sharedIds[lane * 16u + ai]; let bv = sharedValues[partner * 16u + bi]; let bid = sharedIds[partner * 16u + bi]; if (better(av, aid, bv, bid)) { mergedV[k] = av; mergedI[k] = aid; ai += 1u; } else { mergedV[k] = bv; mergedI[k] = bid; bi += 1u; } } } workgroupBarrier(); if (laneMerges) { for (var k = 0u; k < 16u; k += 1u) { sharedValues[lane * 16u + k] = mergedV[k]; sharedIds[lane * 16u + k] = mergedI[k]; } } workgroupBarrier(); } if (lane == 0u) { for (var k = 0u; k < 16u; k += 1u) { output[outBase + k] = sharedIds[k]; } } workgroupBarrier(); sharedValues[lane] = maskedBest; sharedIds[lane] = maskedId; workgroupBarrier(); if (lane == 0u) { var bestValue = sharedValues[0]; var bestId = sharedIds[0]; for (var at = 1u; at < 64u; at += 1u) { if (better(sharedValues[at], sharedIds[at], bestValue, bestId)) { bestValue = sharedValues[at]; bestId = sharedIds[at]; } } output[outBase + 16u] = bestId; } workgroupBarrier(); sharedValues[lane] = rawBest; sharedIds[lane] = rawId; workgroupBarrier(); if (lane == 0u) { var bestValue = sharedValues[0]; var bestId = sharedIds[0]; for (var at = 1u; at < 64u; at += 1u) { if (better(sharedValues[at], sharedIds[at], bestValue, bestId)) { bestValue = sharedValues[at]; bestId = sharedIds[at]; } } output[outBase + 17u] = bestId; var pick: u32; if (params.t < front) { pick = histTok[params.slot * params.rows + row]; } else if (atFrontier) { pick = output[outBase + 16u]; } else { pick = output[outBase + 17u]; } tokenRing[params.t * params.rows + row] = pick; } } `; const windowPipelines = new WeakMap(); async function pipelineFor(device) { let hit = pipelines.get(device); if (!hit) { hit = device.createComputePipelineAsync({ label: 'guarded top-16', layout: 'auto', compute: { module: device.createShaderModule({ label: 'guarded top-16 shader', code: shader }), entryPoint: 'main', }, }); pipelines.set(device, hit); } return hit; } async function windowPipelineFor(device) { let hit = windowPipelines.get(device); if (!hit) { hit = device.createComputePipelineAsync({ label: 'guarded window top-16', layout: 'auto', compute: { module: device.createShaderModule({ label: 'guarded window top-16 shader', code: windowShader }), entryPoint: 'main', }, }); windowPipelines.set(device, hit); } return hit; } export function prepareGuardedTopK(device) { return pipelineFor(device); } export async function createGuardedTopK(device, weights, logits, rows, vocab) { const pipeline = await pipelineFor(device); const maskWords = Math.ceil(vocab / 32); const masksCpu = new Uint32Array(rows * maskWords); const modesCpu = new Uint32Array(rows); const maskBits = device.createBuffer({ label: 'guarded top-16 mask bits', size: masksCpu.byteLength, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, }); const maskModes = device.createBuffer({ label: 'guarded top-16 mask modes', size: modesCpu.byteLength, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, }); const output = device.createBuffer({ label: 'guarded top-16 output', size: rows * OUTPUT_STRIDE * 4, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC, }); const staging = device.createBuffer({ label: 'guarded top-16 readback', size: rows * OUTPUT_STRIDE * 4, usage: GPUBufferUsage.MAP_READ | GPUBufferUsage.COPY_DST, }); const params = device.createBuffer({ label: 'guarded top-16 params', size: 16, usage: GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST, }); device.queue.writeBuffer(params, 0, new Uint32Array([rows, vocab, maskWords, 0])); const bias = weights.bindingFor('final_logits_bias'); const bindGroup = device.createBindGroup({ label: 'guarded top-16 bindings', layout: pipeline.getBindGroupLayout(0), entries: [ { binding: 0, resource: { buffer: logits } }, { binding: 1, resource: bias }, { binding: 2, resource: { buffer: maskBits } }, { binding: 3, resource: { buffer: maskModes } }, { binding: 4, resource: { buffer: output } }, { binding: 5, resource: { buffer: params } }, ], }); return { outputStride: OUTPUT_STRIDE, emptyId: EMPTY_ID, prepare(requests) { masksCpu.fill(0); modesCpu.fill(0); requests.forEach((request, row) => { if (request.mode !== 'ban' && request.mode !== 'allow') return; modesCpu[row] = request.mode === 'ban' ? 1 : 2; const base = row * maskWords; for (const id of request.ids) { if (id >= 0 && id < vocab) masksCpu[base + (id >> 5)] |= 1 << (id & 31); } }); device.queue.writeBuffer(maskBits, 0, masksCpu); device.queue.writeBuffer(maskModes, 0, modesCpu); }, record(commandEncoder) { const pass = commandEncoder.beginComputePass({ label: 'guarded top-16' }); pass.setPipeline(pipeline); pass.setBindGroup(0, bindGroup); pass.dispatchWorkgroups(rows); pass.end(); commandEncoder.copyBufferToBuffer(output, 0, staging, 0, rows * OUTPUT_STRIDE * 4); }, async read() { await staging.mapAsync(GPUMapMode.READ); const result = new Uint32Array(staging.getMappedRange().slice(0)); staging.unmap(); return result; }, destroy() { maskBits.destroy(); maskModes.destroy(); output.destroy(); staging.destroy(); params.destroy(); }, }; } // Window variant: kMax slots, per-slot params/bind groups, kernel-side ring // writes. prepare() uploads the per-row mask (applied only at each row's // frontier step), the frontier itself, and the accepted-history restore // tokens; record(cmd, k) appends slot k's ranking pass; copyOut + read() // fetch the whole window in one staging round-trip. export async function createGuardedTopKWindow(device, weights, logits, tokenRing, rows, vocab, kMax) { const pipeline = await windowPipelineFor(device); const maskWords = Math.ceil(vocab / 32); const masksCpu = new Uint32Array(rows * maskWords); const modesCpu = new Uint32Array(rows); const maskBits = device.createBuffer({ label: 'guarded window mask bits', size: masksCpu.byteLength, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, }); const maskModes = device.createBuffer({ label: 'guarded window mask modes', size: modesCpu.byteLength, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, }); const frontier = device.createBuffer({ label: 'guarded window frontier', size: rows * 4, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, }); const histTok = device.createBuffer({ label: 'guarded window history tokens', size: kMax * rows * 4, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, }); const output = device.createBuffer({ label: 'guarded window output', size: kMax * rows * OUTPUT_STRIDE * 4, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC, }); const staging = device.createBuffer({ label: 'guarded window readback', size: kMax * rows * OUTPUT_STRIDE * 4, usage: GPUBufferUsage.MAP_READ | GPUBufferUsage.COPY_DST, }); const bias = weights.bindingFor('final_logits_bias'); const paramBufs = []; const bindGroups = []; for (let k = 0; k < kMax; k += 1) { const params = device.createBuffer({ label: `guarded window params ${k}`, size: 32, usage: GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST, }); paramBufs.push(params); bindGroups.push(device.createBindGroup({ label: `guarded window bindings ${k}`, layout: pipeline.getBindGroupLayout(0), entries: [ { binding: 0, resource: { buffer: logits } }, { binding: 1, resource: bias }, { binding: 2, resource: { buffer: maskBits } }, { binding: 3, resource: { buffer: maskModes } }, { binding: 4, resource: { buffer: frontier } }, { binding: 5, resource: { buffer: histTok } }, { binding: 6, resource: { buffer: tokenRing } }, { binding: 7, resource: { buffer: output } }, { binding: 8, resource: { buffer: params } }, ], })); } return { outputStride: OUTPUT_STRIDE, emptyId: EMPTY_ID, prepare(requests, frontierCpu, histCpu, t0, kUsed) { masksCpu.fill(0); modesCpu.fill(0); requests.forEach((request, row) => { if (request.mode !== 'ban' && request.mode !== 'allow') return; modesCpu[row] = request.mode === 'ban' ? 1 : 2; const base = row * maskWords; for (const id of request.ids) { if (id >= 0 && id < vocab) masksCpu[base + (id >> 5)] |= 1 << (id & 31); } }); device.queue.writeBuffer(maskBits, 0, masksCpu); device.queue.writeBuffer(maskModes, 0, modesCpu); device.queue.writeBuffer(frontier, 0, frontierCpu); device.queue.writeBuffer(histTok, 0, histCpu); for (let k = 0; k < kUsed; k += 1) { device.queue.writeBuffer(paramBufs[k], 0, new Uint32Array([rows, vocab, maskWords, t0 + k, k, 0, 0, 0])); } }, record(commandEncoder, k) { const pass = commandEncoder.beginComputePass({ label: `guarded window top-16 slot ${k}` }); pass.setPipeline(pipeline); pass.setBindGroup(0, bindGroups[k]); pass.dispatchWorkgroups(rows); pass.end(); }, copyOut(commandEncoder, kUsed) { commandEncoder.copyBufferToBuffer(output, 0, staging, 0, kUsed * rows * OUTPUT_STRIDE * 4); }, async read(kUsed) { const bytes = kUsed * rows * OUTPUT_STRIDE * 4; await staging.mapAsync(GPUMapMode.READ, 0, bytes); const result = new Uint32Array(staging.getMappedRange(0, bytes).slice(0)); staging.unmap(); return result; }, destroy() { maskBits.destroy(); maskModes.destroy(); frontier.destroy(); histTok.destroy(); output.destroy(); staging.destroy(); for (const params of paramBufs) params.destroy(); }, }; }