File size: 16,960 Bytes
c971a45
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
// APE A2 WebGPU AR ship engine — guarded decode loop over the vendored webMT
// engine. Design: docs/plans/2026-08-10-ape-a2-webgpu-ar-ship-engine-design.md
// (§3.4). Zero vendor modifications:
//   - encoder: stock runEncoder over caller-tokenized ids (the Python
//     contract: add_special_tokens=False — NO trailing </s>).
//   - decode: stock encodeDecodeStep with lmHeadFuse 'off' so the f32 logits
//     buffer materializes (COPY_SRC); one step per submit, logits copied to a
//     staging buffer and mapped. The vendor argmax_penalty still runs but its
//     ring write is OVERWRITTEN each step by the CPU-guarded pick
//     (queue-ordered writeBuffer — the compactDecodeState precedent); its
//     done/bitmask side effects are never read on this path.
//   - Python-semantics note: the reference project() ADDS final_logits_bias
//     to the logits; the engine leaves the bias to the argmax kernel — so
//     this loop adds the bias on CPU after readback (readTensorF32 once).
// The per-step token decision is decode-guard.js's GuardedRow — the exact
// stack P1 proved equivalent (21,348/21,348 steps) to the frozen Python loop.

import { runEncoder } from './engine/encoder.js';
import { createDecodeState, encodeDecodeStep } from './engine/decoder.js';
import { VOCAB, PAD, SRC_CAP } from './engine/constants.js';
import { GuardedRow, EOS_ID, MAX_NEW_TOKENS, NEG, hasNewRepeat } from './decode-guard.js';
import { createGuardedLogitRanker } from './guarded-logits.js';
import { createGuardedTopK, createGuardedTopKWindow } from './guarded-topk.js';

// Read one tensor back from the uploaded weights mega-buffer as f32.
export async function readTensorF32(device, weights, name) {
  const meta = weights.tensors.get(name);
  if (!meta) throw new Error(`readTensorF32: unknown tensor ${name}`);
  if (meta.dtype !== 'f32') throw new Error(`readTensorF32: ${name} is ${meta.dtype}, not f32`);
  const bind = weights.bindingFor(name);
  const buffer = bind.buffer ?? bind;
  const offset = bind.offset ?? 0;
  const staging = device.createBuffer({
    size: meta.byteLength,
    usage: GPUBufferUsage.MAP_READ | GPUBufferUsage.COPY_DST,
  });
  const enc = device.createCommandEncoder({ label: `read ${name}` });
  enc.copyBufferToBuffer(buffer, offset, staging, 0, meta.byteLength);
  device.queue.submit([enc.finish()]);
  await staging.mapAsync(GPUMapMode.READ);
  const out = new Float32Array(staging.getMappedRange().slice(0));
  staging.unmap();
  staging.destroy();
  return out;
}

function tierHasAcceptable(guard, ids) {
  for (const token of ids) {
    if (token === EOS_ID) {
      if (guard.name.broken || guard.name.eosOk()) return true;
      continue;
    }
    const text = guard.pieceTable.pieceText[token];
    if (!text) continue;
    if (!(guard.name.broken || guard.name.trial(text))) continue;
    if (hasNewRepeat(guard.sim + text, guard.draftPatterns)) continue;
    return true;
  }
  return false;
}

// Decode one batch with the guarded ship stack.
//   ctx        initDevice() result
//   weights    loadWeights() result (f32 oracle or f16 ship build)
//   pieceTable {pieceText, digitIds, bannedBytes} (piece_table.json)
//   bias       Float32Array[VOCAB] — final_logits_bias (readTensorF32)
//   rows       [{id, draft, srcIds: number[] (NO trailing eos)}]
//   opts       {specWindow}: >= 2 selects the speculative-window executor
//              (design addendum 2026-08-11); default/1 keeps the per-step
//              path. Both must produce bit-identical outputs (receipt _v2).
// Returns {rows: [{id, outputs, sim, stats, finished, steps}], timing}.
export async function guardedDecodeBatch(ctx, weights, pieceTable, bias, rows, opts = {}) {
  const { device } = ctx;
  const B = rows.length;
  if (!B) throw new Error('guardedDecodeBatch: empty batch');
  let S = 0;
  for (const row of rows) {
    if (!row.srcIds.length || row.srcIds.length > SRC_CAP) {
      throw new Error(`guardedDecodeBatch: row ${row.id} has ${row.srcIds.length} src tokens (cap ${SRC_CAP})`);
    }
    if (row.srcIds[row.srcIds.length - 1] === EOS_ID) {
      throw new Error(`guardedDecodeBatch: row ${row.id} ends with EOS — the contract is add_special_tokens=False`);
    }
    S = Math.max(S, row.srcIds.length);
  }
  if (bias.length !== VOCAB) throw new Error('guardedDecodeBatch: bias length != VOCAB');

  const ids = new Uint32Array(B * S).fill(PAD);
  const lens = new Uint32Array(B);
  rows.forEach((row, i) => {
    ids.set(row.srcIds, i * S);
    lens[i] = row.srcIds.length;
  });

  const specWindow = Math.trunc(opts.specWindow ?? 0);
  const useWindow = specWindow >= 2;

  const t0 = performance.now();
  const encRun = await runEncoder(ctx, weights, { ids, lens, B, S });
  const state = createDecodeState(ctx, weights, { B, S, maxSteps: MAX_NEW_TOKENS, lmHeadFuse: 'off' });
  const gpuTopK = useWindow ? null : await createGuardedTopK(device, weights, state.logits, B, VOCAB);
  const win = useWindow
    ? await createGuardedTopKWindow(device, weights, state.logits, state.tokenRing, B, VOCAB, specWindow)
    : null;
  const t1 = performance.now();

  const memo = new Map();
  const guards = rows.map((row) => new GuardedRow(row.draft, pieceTable, memo));
  const stepsPerRow = new Array(B).fill(0);
  const ranker = createGuardedLogitRanker(VOCAB, NEG);
  let stepsRun = 0;
  const diag = { windows: 0, corrections: 0, deferrals: 0, rescues: 0 };

  try {
    if (useWindow) {
      // Speculative-window executor. Per-row frontier = accepted steps; the
      // GuardedRow is always exactly at its frontier (stepped forward once
      // per accepted step, never re-stepped, never rolled back).
      const FINISHED = 0xffffffff;
      const frontier = new Uint32Array(B);
      const needsFull = new Array(B).fill(false);
      while (true) {
        let t0w = Infinity;
        for (let i = 0; i < B; i += 1) {
          if (!guards[i].finished) t0w = Math.min(t0w, frontier[i]);
        }
        if (t0w === Infinity || t0w >= MAX_NEW_TOKENS) break;
        // A row whose frontier step needs full logits forces a 1-step window
        // once it defines t0w — state.logits then holds exactly that step.
        const forceSingle = guards.some((g, i) => !g.finished && needsFull[i] && frontier[i] === t0w);
        // Scheduling-only tail clamp (round-2 design §2.2): any Kp yields the
        // same outputs; this just trims window overshoot past row finishes.
        // Polish outputs track source length, so a source-token estimate with
        // slack sizes the tail window; when a row outruns its estimate the
        // clamp simply stops applying to it. Never used for token decisions.
        let maxSpan = 1;
        for (let i = 0; i < B; i += 1) {
          if (guards[i].finished) continue;
          const est = Math.ceil(1.2 * rows[i].srcIds.length) + 4;
          const remaining = est > frontier[i] ? est - frontier[i] : specWindow;
          maxSpan = Math.max(maxSpan, frontier[i] - t0w + remaining);
        }
        const Kp = forceSingle ? 1 : Math.min(specWindow, MAX_NEW_TOKENS - t0w, maxSpan);

        const requests = guards.map((guard) => (guard.finished
          ? { mode: 'none', ids: [] }
          : guard.maskRequest()));
        const frontCpu = new Uint32Array(B);
        for (let i = 0; i < B; i += 1) frontCpu[i] = guards[i].finished ? FINISHED : frontier[i];
        const hist = new Uint32Array(Kp * B).fill(EOS_ID);
        for (let i = 0; i < B; i += 1) {
          for (let k = 0; k < Kp; k += 1) {
            if (t0w + k < frontCpu[i]) {
              const tok = guards[i].outputs[t0w + k];
              if (tok !== undefined) hist[k * B + i] = tok;
            }
          }
        }
        win.prepare(requests, frontCpu, hist, t0w, Kp);

        const cmd = device.createCommandEncoder({ label: `guarded window ${t0w}+${Kp}` });
        const scratchAll = [];
        for (let k = 0; k < Kp; k += 1) {
          const pass = cmd.beginComputePass({ label: `guarded window step ${t0w + k}` });
          const { scratch } = encodeDecodeStep(ctx, weights, encRun, state, t0w + k, pass);
          pass.end();
          win.record(cmd, k);
          scratchAll.push(...scratch);
        }
        win.copyOut(cmd, Kp);
        device.queue.submit([cmd.finish()]);
        for (const buf of scratchAll) buf.destroy();
        stepsRun += Kp;
        diag.windows += 1;

        const top = await win.read(Kp);
        const corrections = [];
        const fullRows = [];
        for (let i = 0; i < B; i += 1) {
          const guard = guards[i];
          if (guard.finished) continue;
          let f = frontier[i];
          for (let j = Math.max(t0w, f); j < t0w + Kp; j += 1) {
            const base = ((j - t0w) * B + i) * 18;
            const atFrontier = j === f;
            if (atFrontier && needsFull[i]) {
              if (Kp === 1 && f === t0w) fullRows.push(i);
              break;
            }
            const visible = [];
            for (let k = 0; k < 16; k += 1) {
              const id = top[base + k];
              if (id === win.emptyId) break;
              visible.push(id);
            }
            const rawTop = top[base + 17];
            let tier;
            let maskedTop;
            if (atFrontier) {
              // The kernel applied the true mask here — per-step semantics.
              maskedTop = top[base + 16];
              if (!tierHasAcceptable(guard, visible)) {
                needsFull[i] = true;
                diag.rescues += 1;
                break;
              }
              tier = visible;
            } else {
              // Free-run step: the readback is unmasked. The permitted
              // sub-sequence of the unmasked top-16 is an exact prefix of
              // the masked descending order (unbanned/allowed values keep
              // their raw value; everything hidden ranks below them), so
              // the guard decision is exact whenever it lands inside it.
              const req = guard.maskRequest();
              let prefix = visible;
              if (req.mode === 'ban') {
                const banned = new Set(req.ids);
                prefix = visible.filter((id) => !banned.has(id));
              } else if (req.mode === 'allow') {
                const allowed = new Set(req.ids);
                prefix = visible.filter((id) => allowed.has(id));
              }
              if (!prefix.length || !tierHasAcceptable(guard, prefix)) {
                f = j; // next window applies the true mask at j
                diag.deferrals += 1;
                break;
              }
              maskedTop = prefix[0];
              tier = prefix;
            }
            const ringVal = atFrontier ? top[base + 16] : rawTop;
            const chosen = guard.step(rawTop, maskedTop, (tierIndex) => {
              if (tierIndex !== 0) throw new Error('guarded window tier escalation past preflight');
              return { ids: tier, validCount: tier.length };
            });
            stepsPerRow[i] += 1;
            f = j + 1;
            if (chosen !== ringVal) {
              // The free-run trajectory diverged: pin the true token and
              // recompute later steps next window (unaffected rows below
              // their own frontier restore identical history/KV).
              corrections.push([j, i, chosen]);
              break;
            }
            if (guard.finished) break;
          }
          frontier[i] = f;
        }

        if (fullRows.length) {
          // Rescue: tier-16 had no acceptable candidate at a frontier step
          // (P1 frequency: 0). Kp == 1, so state.logits holds this step.
          const rowBytes = VOCAB * 4;
          const rescueStaging = device.createBuffer({
            label: 'guarded window rescue logits',
            size: fullRows.length * rowBytes,
            usage: GPUBufferUsage.MAP_READ | GPUBufferUsage.COPY_DST,
          });
          const rescueCmd = device.createCommandEncoder({ label: 'guarded window rescue readback' });
          fullRows.forEach((row, index) => {
            rescueCmd.copyBufferToBuffer(state.logits, row * rowBytes, rescueStaging, index * rowBytes, rowBytes);
          });
          device.queue.submit([rescueCmd.finish()]);
          await rescueStaging.mapAsync(GPUMapMode.READ);
          const view = new Float32Array(rescueStaging.getMappedRange());
          fullRows.forEach((row, index) => {
            const guard = guards[row];
            const ranked = ranker.rank(view, index * VOCAB, bias, requests[row]);
            const chosen = guard.step(ranked.rawTop, ranked.maskedTop, ranked.getTier);
            stepsPerRow[row] += 1;
            needsFull[row] = false;
            frontier[row] = t0w + 1;
            const ringVal = top[row * 18 + 16];
            if (chosen !== ringVal) corrections.push([t0w, row, chosen]);
          });
          rescueStaging.unmap();
          rescueStaging.destroy();
        }

        for (const [t, rowIdx, token] of corrections) {
          device.queue.writeBuffer(state.tokenRing, (t * B + rowIdx) * 4, new Uint32Array([token]));
        }
        diag.corrections += corrections.length;
      }
    } else {
    for (let t = 0; t < MAX_NEW_TOKENS; t += 1) {
      const requests = guards.map((guard) => (guard.finished
        ? { mode: 'none', ids: [] }
        : guard.maskRequest()));
      gpuTopK.prepare(requests);
      const cmd = device.createCommandEncoder({ label: `guarded step ${t}` });
      const pass = cmd.beginComputePass({ label: `guarded step ${t}` });
      const { scratch } = encodeDecodeStep(ctx, weights, encRun, state, t, pass);
      pass.end();
      gpuTopK.record(cmd);
      device.queue.submit([cmd.finish()]);
      for (const buf of scratch) buf.destroy();
      stepsRun += 1;

      const top = await gpuTopK.read();
      const picks = new Uint32Array(B).fill(EOS_ID);
      const fallbackRows = [];
      const gpuTiers = new Array(B);
      for (let i = 0; i < B; i += 1) {
        if (guards[i].finished) continue;
        const base = i * gpuTopK.outputStride;
        const ids = [];
        for (let k = 0; k < 16; k += 1) {
          if (top[base + k] === gpuTopK.emptyId) break;
          ids.push(top[base + k]);
        }
        gpuTiers[i] = ids;
        if (!tierHasAcceptable(guards[i], ids)) fallbackRows.push(i);
      }

      let fallbackStaging = null;
      let fallbackView = null;
      if (fallbackRows.length) {
        const rowBytes = VOCAB * 4;
        fallbackStaging = device.createBuffer({
          label: 'guarded fallback logits',
          size: fallbackRows.length * rowBytes,
          usage: GPUBufferUsage.MAP_READ | GPUBufferUsage.COPY_DST,
        });
        const fallbackCmd = device.createCommandEncoder({ label: 'guarded fallback readback' });
        fallbackRows.forEach((row, index) => {
          fallbackCmd.copyBufferToBuffer(state.logits, row * rowBytes, fallbackStaging, index * rowBytes, rowBytes);
        });
        device.queue.submit([fallbackCmd.finish()]);
        await fallbackStaging.mapAsync(GPUMapMode.READ);
        fallbackView = new Float32Array(fallbackStaging.getMappedRange());
      }

      const fallbackIndex = new Map(fallbackRows.map((row, index) => [row, index]));
      let live = 0;
      for (let i = 0; i < B; i += 1) {
        const guard = guards[i];
        if (guard.finished) continue;
        const compactIndex = fallbackIndex.get(i);
        let chosen;
        if (compactIndex !== undefined) {
          const ranked = ranker.rank(fallbackView, compactIndex * VOCAB, bias, requests[i]);
          chosen = guard.step(ranked.rawTop, ranked.maskedTop, ranked.getTier);
        } else {
          const base = i * gpuTopK.outputStride;
          const tier = { ids: gpuTiers[i], validCount: gpuTiers[i].length };
          chosen = guard.step(top[base + 17], top[base + 16], (tierIndex) => {
            if (tierIndex !== 0) throw new Error('guarded top-16 preflight disagreement');
            return tier;
          });
        }
        picks[i] = chosen;
        stepsPerRow[i] += 1;
        if (!guard.finished) live += 1;
      }
      if (fallbackStaging) {
        fallbackStaging.unmap();
        fallbackStaging.destroy();
      }
      device.queue.writeBuffer(state.tokenRing, t * B * 4, picks);
      if (live === 0) break;
    }
    }
  } finally {
    (useWindow ? win : gpuTopK).destroy();
    state.destroy();
    encRun.arena.destroy();
  }
  const t2 = performance.now();

  return {
    rows: rows.map((row, i) => ({
      id: row.id,
      outputs: guards[i].outputs,
      sim: guards[i].sim,
      stats: guards[i].stats,
      finished: guards[i].finished,
      steps: stepsPerRow[i],
    })),
    timing: {
      encoderMs: Math.round(t1 - t0),
      decodeMs: Math.round(t2 - t1),
      stepsRun,
      spec: useWindow ? specWindow : 0,
      ...(useWindow ? diag : {}),
    },
  };
}