File size: 24,169 Bytes
1944112
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
/**
 * Multi-step decoding: N forward steps per GPU->CPU sync.
 *
 * Why this exists: decode here is not compute-bound, it is *sync*-bound. Firefox
 * resolves `onSubmittedWorkDone()` / `mapAsync()` only on a 100 ms poll tick
 * (AI.md, "The 10 tok/s ceiling"), and stock WebLLM needs exactly one sync per
 * token β€” it reads the sampled token id back to JS before it can build the next
 * step's input. One token per tick = 9.6 tok/s, of which ~7 ms is real compute.
 *
 * The fix is the one vLLM ships as `--num-scheduler-steps`: run K steps before
 * paying the per-batch cost once. What makes it possible here without touching
 * the compiled model is that WebLLM's sampling path is *already* on the GPU β€”
 * `softmax_with_temperature`, `argsort_probs` and `sample_with_top_p` hand back
 * an int32[1] device tensor, and `Tensor.copyFrom(Tensor)` is a device-to-device
 * copy. So the sampled id feeds straight back into `embed` without ever becoming
 * a JS number:
 *
 *   embed -> decode -> penalties -> softmax -> argsort -> sample -> embed -> ...
 *
 * Each step stages its id into its own CPU tensor, and the burst ends with
 * **one** `device.sync()`. tvmjs queues GPU->CPU copies into `pendingGPUToCPUCopy`
 * and only awaits them in `sync()`, so K readbacks still cost one tick.
 *
 * Two things follow from the 100 ms grid, and they are why `steps` is a dial:
 *
 *  - The win is quantized, not linear. A burst costs `ceil(K * perStepMs / 100)`
 *    ticks, so throughput is a sawtooth and the good values of K are the ones
 *    landing just under a boundary. On a 0.8B at ~7.3 ms/step that is K=13
 *    (~130 tok/s); K=14 already spills into a second tick and halves it.
 *  - The best K shrinks as the model grows, because `perStepMs` grows. A model
 *    at 25 ms/step wants K=4, not K=15.
 *
 * Cost of the trick: the sampler cannot see its own output mid-burst. Repetition
 * and presence/frequency penalties use the token history as it stood when the
 * burst started, and stop conditions are only checked after the readback, so a
 * burst can overshoot a stop token and must then be rewound. Both are the same
 * trade vLLM makes. Anything needing per-token CPU feedback (grammar-constrained
 * JSON, logprobs, a logit processor) falls back to single-step, where behaviour
 * is identical to stock WebLLM.
 *
 * The other cost is that all of this drives ~30 undocumented tvmjs internals. A
 * WebLLM upgrade that renames one does not break generation β€” it turns the fast
 * path off and takes the throughput with it, silently. `PIPELINE_CONTRACT` below
 * is that surface written down and checked against the live pipeline before the
 * first burst, so the failure announces itself instead of being measured months
 * later.
 */

/** vLLM's documented sweet spot, and the value this extension ships. */
export const DEFAULT_DECODE_STEPS = 15;

/**
 * Above this, the lookahead thrown away at a stop token outweighs the tick it
 * saves, and the transient logits/argsort buffers stop being free.
 */
export const MAX_DECODE_STEPS = 32;

export const clampSteps = (n) => Math.max(1, Math.min(MAX_DECODE_STEPS, Math.round(Number(n)) || 1));

// -------------------------------------------------- the pipeline contract ----

/**
 * Every tvmjs pipeline internal a burst drives, and how each must behave.
 *
 * None of these are documented, none are part of WebLLM's public surface, and
 * nothing upstream promises they will keep their names. The contract test checks
 * them against the *bundle* on every `npm test`; this checks them against the
 * *live object*, which is a different question β€” a member can survive in the
 * bundle and still not be on the pipeline handed to us, if upstream moves it to
 * a subclass, a different pipeline type, or behind a factory.
 *
 * Split three ways because presence alone is not the failure that hurts:
 *
 *  - **`calls`** must be callable. A rename here throws, which is the *good*
 *    case β€” it is loud.
 *  - **`numbers`** are read arithmetically or incremented in place. This is the
 *    silent one: `pipeline.filledKVCacheLength += 1` on a member that no longer
 *    exists creates a new property, nothing throws, and the KV cache accounting
 *    quietly drifts. A missing `contextWindowSize` makes `burstSize` NaN.
 *  - **`reads`** need only exist.
 *
 * `logitProcessor` is deliberately optional: `burstSize` tests it for
 * `undefined`, so absent is the normal case, not a broken one.
 *
 * The list is not maintained by hand β€” `webllm-contract.test.mjs` derives the
 * set this file actually reaches for from its own source and asserts it matches
 * this declaration exactly, so adding a `pipeline.newThing` without declaring it
 * fails the build.
 */
export const PIPELINE_CONTRACT = {
  calls: [
    "embed",
    "fKVCacheBeginForward",
    "fKVCacheEndForward",
    "fapplyLogitBias",
    "fapplyPenalty",
    "fargsortProbs",
    "fsampleWithTopP",
    "fsoftmaxWithTemperature",
    "getActiveKVStates",
    "invokeDecode",
    "processNextToken",
    "resetChat",
    "stopped",
  ],
  numbers: [
    "contextWindowSize",
    "curRoundDecodingTotalTime",
    "curRoundDecodingTotalTokens",
    "decodingTotalTime",
    "decodingTotalTokens",
    "filledKVCacheLength",
    "fullVocabSize",
    "slidingWindowSize",
  ],
  reads: [
    "appearedTokensFreq",
    "config",
    "device",
    "outputIds",
    "params",
    "sampleIndices",
    "sampleIndicesDevice",
    "topPDevice",
    "tvm",
  ],
  optional: ["logitProcessor"],
};

/**
 * What this pipeline is missing, as sentences a reader can act on.
 * Empty means a burst is safe to run.
 */
export function missingPipelineMembers(pipeline) {
  if (!pipeline || typeof pipeline !== "object") return ["the pipeline itself is not an object"];
  const missing = [];
  for (const name of PIPELINE_CONTRACT.calls) {
    if (typeof pipeline[name] !== "function") missing.push(`${name}() is not a function`);
  }
  for (const name of PIPELINE_CONTRACT.numbers) {
    if (typeof pipeline[name] !== "number") missing.push(`${name} is not a number`);
  }
  for (const name of PIPELINE_CONTRACT.reads) {
    if (pipeline[name] === undefined) missing.push(`${name} is missing`);
  }
  return missing;
}

/**
 * Replaces `engine.decode` with a burst-and-drain version.
 *
 * `decode` stays a one-token call β€” the caller's loop still checks
 * `pipeline.stopped()` between tokens and still emits one chunk per token β€” but
 * only one call in K actually touches the GPU. The rest drain a buffer.
 *
 * @param {object} engine an MLCEngine; in this project the one inside the worker
 * @param {object} [options]
 * @param {number} [options.steps] forward steps per sync; 1 disables the path
 * @param {(info: {steps: number, tokens: number, ms: number}) => void} [options.onBurst]
 * @param {(info: {missing: string[]}) => void} [options.onFallback] fired once
 *   per pipeline that fails the contract, before it is routed to stock decoding
 * @returns {{setSteps: (n: number) => void, readonly steps: number,
 *            readonly fallbacks: number}}
 */
export function installMultiStepDecoding(
  engine,
  { steps = DEFAULT_DECODE_STEPS, onBurst, onFallback } = {},
) {
  const config = { steps: clampSteps(steps), fallbacks: 0 };
  const lookahead = new WeakMap();
  const baseDecode = engine.decode.bind(engine);
  const basePrefill = engine.prefill.bind(engine);

  const stateFor = (pipeline) => {
    let state = lookahead.get(pipeline);
    if (!state) lookahead.set(pipeline, (state = { queue: [] }));
    return state;
  };

  /** Contract verdict per pipeline; the check runs once, the answer is reused. */
  const supported = new WeakMap();

  /**
   * Whether this pipeline may be burst, decided once and remembered.
   *
   * Checked at first decode rather than at install time because there is no
   * pipeline yet when this function runs β€” the engine gets one per `reload()`,
   * and hands it to us as an argument. So the guard lives at the first place a
   * pipeline is ever seen.
   *
   * Failing here means an upgrade moved something and multi-step decoding is
   * gone. Stock decoding still produces correct tokens, so the danger is not a
   * crash but silence: ~18.4 -> ~9.7 tok/s with nothing in the log to explain
   * it. Hence one loud report, and a `fallbacks` count the worker can surface.
   */
  const canBurst = (pipeline) => {
    const known = supported.get(pipeline);
    if (known !== undefined) return known;

    const missing = missingPipelineMembers(pipeline);
    supported.set(pipeline, missing.length === 0);
    if (missing.length > 0) {
      config.fallbacks += 1;
      console.error(
        "[everything-webgpu] multi-step decoding disabled β€” falling back to stock " +
          "single-step decode. Generation stays correct, throughput roughly halves.\n" +
          `  The pipeline is missing ${missing.length} of the internals a burst drives:\n` +
          missing.map((line) => `    - ${line}`).join("\n") +
          "\n  This is what a WebLLM upgrade looks like from here. `npm test` " +
          "(webllm-contract) says whether the names are gone from the bundle too.",
      );
      onFallback?.({ missing });
    }
    return missing.length === 0;
  };

  // A round can end with tokens still buffered β€” a stop token mid-burst, or an
  // interrupt that breaks the caller's loop. Those tokens are already in the KV
  // cache, so they must come back out before the next round reuses it.
  engine.prefill = async (input, pipeline, chatConfig, genConfig) => {
    discardLookahead(pipeline, stateFor(pipeline));
    return basePrefill(input, pipeline, chatConfig, genConfig);
  };

  engine.decode = async (pipeline, genConfig) => {
    const state = stateFor(pipeline);

    if (state.queue.length === 0) {
      // Before the first burst on this pipeline, not before every one: the
      // verdict is cached, so a healthy pipeline pays one property scan for the
      // whole conversation.
      if (!canBurst(pipeline)) return baseDecode(pipeline, genConfig);

      const burst = burstSize(pipeline, genConfig, config.steps);
      if (burst <= 1) return baseDecode(pipeline, genConfig);

      const probe = {};
      const tstart = performance.now();
      state.queue = await sampleBurst(pipeline, genConfig, burst, probe);
      const ms = performance.now() - tstart;

      // One burst is one wall-clock cost; its tokens are counted as they drain,
      // so a rewound overshoot never inflates the reported rate.
      pipeline.decodingTotalTime += ms / 1e3;
      pipeline.curRoundDecodingTotalTime += ms / 1e3;
      onBurst?.({ steps: burst, tokens: state.queue.length, ms, ...probe });
    }

    const token = state.queue.shift();
    pipeline.decodingTotalTokens += 1;
    pipeline.curRoundDecodingTotalTokens += 1;
    pipeline.processNextToken(token, genConfig);

    // The burst ran past a stop token; nothing after it was ever emitted.
    if (pipeline.stopped() && state.queue.length > 0) discardLookahead(pipeline, state);
  };

  return {
    setSteps: (n) => void (config.steps = clampSteps(n)),
    get steps() {
      return config.steps;
    },
    /** Pipelines that failed the contract. Non-zero means the fast path is off. */
    get fallbacks() {
      return config.fallbacks;
    },
  };
}

// ------------------------------------------------------------- burst size ---

/**
 * How many steps may run before the next stop condition *has* to be checked.
 *
 * `max_tokens` and the context window are countable, so they are clamped rather
 * than overshot β€” which leaves stop tokens as the only reason a burst is ever
 * rewound. Returns 1 when multi-step cannot be used at all, routing the caller
 * to stock single-step decoding.
 */
export function burstSize(pipeline, genConfig, steps) {
  if (steps <= 1) return 1;

  // Per-token CPU feedback: the next step's logits depend on a JS-side decision
  // about this step's token, so there is nothing to overlap.
  const format = genConfig?.response_format?.type;
  if (format === "json_object" || format === "grammar" || format === "structural_tag") return 1;
  if (genConfig?.logprobs) return 1;
  if (pipeline.logitProcessor !== undefined) return 1;

  const maxTokens = genConfig?.max_tokens;
  const untilMax = maxTokens ? maxTokens - pipeline.outputIds.length : Infinity;
  const untilContextEnd =
    pipeline.slidingWindowSize === -1
      ? pipeline.contextWindowSize - pipeline.filledKVCacheLength
      : Infinity;

  return Math.max(1, Math.min(steps, untilMax, untilContextEnd));
}

// ----------------------------------------------------------------- burst ----

/**
 * Runs `steps` forward+sample steps with no GPU->CPU sync between them, then
 * pays exactly one.
 *
 * @returns {Promise<number[]>} the sampled token ids, in order
 */
async function sampleBurst(pipeline, genConfig, steps, out) {
  const { tvm, device } = pipeline;
  let probe = null;
  const vocab = pipeline.fullVocabSize;
  const sampling = resolveSampling(pipeline, genConfig);

  tvm.beginScope();
  let temperatures;
  let bias;
  let penalty;
  /** The last committed token, which seeds step 0. Owned here, not by a scope. */
  let seedTokens;
  try {
    temperatures = tvm.detachFromCurrentScope(
      tvm.empty([1], "float32", device).copyFrom([Math.max(1e-6, sampling.temperature)]),
    );
    bias = makeLogitBias(pipeline, sampling);
    penalty = makePenalty(pipeline, sampling);
    // top_p lives in a tensor the pipeline owns and reuses, set up exactly as
    // `sampleTokenFromLogits` does. It is constant for the whole burst.
    const topPHost = new Float32Array(pipeline.topPDevice.shape[0]).fill(-1);
    const topP = Math.max(sampling.top_p, 1e-5);
    pipeline.sampleIndices.forEach((row) => (topPHost[row] = topP));
    pipeline.topPDevice.copyFrom(topPHost);
    seedTokens = tvm.detachFromCurrentScope(
      tvm.empty([1], "int32", device).copyFrom([pipeline.outputIds[pipeline.outputIds.length - 1]]),
    );
  } finally {
    tvm.endScope();
  }
  let tokens = seedTokens;

  /**
   * Sampled ids stay on the device for the whole loop; the host copies happen
   * after it, never interleaved with compute.
   *
   * The order is load-bearing. `flushCommands()` nulls tvmjs's
   * `pendingGPUToCPUCopy` whenever it submits an encoder, and every GPU->CPU
   * copy calls it. Interleaving copies with compute therefore made each step
   * discard the previous step's pending readback, leaving `device.sync()`
   * awaiting only the last one β€” correct in practice only because the
   * `mapAsync` promises happen to resolve in FIFO order. Doing all the copies
   * after the loop means the first flushes and starts the chain while the rest
   * find no pending encoder, so the chain accumulates intact.
   */
  const sampledIds = [];
  /** One CPU int32[1] per step. All of their reads land on the same poll tick. */
  const staged = [];

  // The decisive probe. The K-step loop below contains no `await`, so it is one
  // synchronous JS turn: everything it costs is content-process CPU β€” command
  // encoding, `createBindGroup`, IPC to the GPU process. The `await` after it is
  // everything else: GPU execution plus the wait for the next 100 ms poll tick.
  // Splitting the two says which one the budget actually goes to.
  const gpuCtx = tvm.lib?.webGPUContext;
  const dispatchesBefore = gpuCtx?.shaderSubmitCounter ?? 0;
  const flushesBefore = countFlushes(gpuCtx);
  let forwardDispatches = 0;
  const tEncodeStart = performance.now();

  try {
    for (let step = 0; step < steps; step++) {
      tvm.beginScope();
      const stepStart = gpuCtx?.shaderSubmitCounter ?? 0;
      try {
        // `tokens` is owned by `sampledIds` (or is the seed), not by this scope.
        const embeddings = pipeline.embed(tokens, pipeline.params);
        const batched = embeddings.view([1].concat(embeddings.shape));

        const states = pipeline.getActiveKVStates();
        const seqIds = tvm.makeShapeTuple([0]);
        const inputLen = tvm.makeShapeTuple([1]);
        for (const state of states) pipeline.fKVCacheBeginForward(state, seqIds, inputLen);
        const forwarded = pipeline.invokeDecode(batched);
        for (let i = states.length - 1; i >= 0; i--) pipeline.fKVCacheEndForward(states[i]);
        pipeline.filledKVCacheLength += 1;

        // Split the launch count at the forward/sample boundary. The sampling
        // tail is `argsort_probs` over the full vocab (248k here), which is a
        // multi-pass sort and belongs to the runtime, not the model β€” so it is
        // worth knowing how much of the per-token kernel budget it owns.
        forwardDispatches += (gpuCtx?.shaderSubmitCounter ?? 0) - stepStart;

        const logits = forwarded.get(0);
        if (bias) {
          pipeline.fapplyLogitBias(logits.view([1, vocab]), bias.pos2seqIds, bias.tokenIds, bias.values);
        }
        if (penalty) {
          pipeline.fapplyPenalty(
            logits.view([1, vocab]),
            penalty.seqIds,
            penalty.pos2seqIds,
            penalty.tokenIds,
            penalty.counts,
            penalty.penalties,
          );
        }

        const probs = pipeline
          .fsoftmaxWithTemperature(logits.view([1, 1, vocab]), temperatures)
          .view([1, vocab]);
        const sorted = pipeline.fargsortProbs(probs);
        const sampled = pipeline.fsampleWithTopP(
          sorted.get(0),
          sorted.get(1),
          tvm.uniform([1], 0, 1, device),
          pipeline.sampleIndicesDevice,
          pipeline.topPDevice,
        );

        tokens = tvm.detachFromCurrentScope(sampled);
        sampledIds.push(tokens);
      } finally {
        tvm.endScope();
      }
    }

    // Every readback together, after all compute: one flush, one intact chain.
    tvm.beginScope();
    try {
      for (const id of sampledIds) {
        staged.push(tvm.detachFromCurrentScope(tvm.empty([1], "int32", tvm.cpu()).copyFrom(id)));
      }
    } finally {
      tvm.endScope();
    }

    // Encoding the copies is still CPU work, so the boundary sits after them.
    const tEncoded = performance.now();

    // The one sync the whole burst pays for.
    await device.sync();

    probe = {
      encodeMs: tEncoded - tEncodeStart,
      syncMs: performance.now() - tEncoded,
      dispatches: (gpuCtx?.shaderSubmitCounter ?? 0) - dispatchesBefore,
      forwardDispatches,
      flushes: countFlushes(gpuCtx) - flushesBefore,
    };
    return staged.map((host) => host.toArray()[0]);
  } finally {
    if (probe) Object.assign(out, probe);
    for (const host of staged) host.dispose();
    for (const id of sampledIds) id.dispose();
    seedTokens?.dispose();
    temperatures.dispose();
    disposeAll(bias);
    disposeAll(penalty);
  }
}

/**
 * Kernel launches per `flushCommands()`, which decides whether batching tvmjs's
 * per-kernel compute passes into one pass is worth anything.
 *
 * `flushCommands()` submits the pending encoder β€” so it would also close a
 * shared pass β€” and it fires from `deviceFreeDataSpace`, the buffer copies and
 * `sync`. If TVM frees an intermediate between every op then flushes β‰ˆ kernels,
 * the pass stream is already chopped up, and there is nothing to merge. tvmjs
 * keeps no counter of its own, so wrap the method once per context.
 */
function countFlushes(gpuCtx) {
  if (!gpuCtx) return 0;
  if (gpuCtx.__ewgpuFlushCount === undefined) {
    const base = gpuCtx.flushCommands.bind(gpuCtx);
    gpuCtx.__ewgpuFlushCount = 0;
    gpuCtx.flushCommands = () => {
      gpuCtx.__ewgpuFlushCount += 1;
      base();
    };
  }
  return gpuCtx.__ewgpuFlushCount;
}

// ---------------------------------------------------------------- rewind ----

/**
 * Drops un-emitted lookahead and takes it back out of the KV cache.
 *
 * `kv_state_popn` is the clean path. If the runtime has not registered it the
 * cache cannot be trimmed, so it is thrown away instead: the next round pays a
 * full re-prefill (one sync, not one per token) rather than attending over
 * tokens the caller never saw.
 */
function discardLookahead(pipeline, state) {
  const n = state.queue.length;
  state.queue = [];
  if (n === 0) return;

  const { tvm } = pipeline;
  try {
    const popn = getPopN(pipeline);
    if (popn) {
      tvm.beginScope();
      try {
        for (const kvState of pipeline.getActiveKVStates()) {
          popn(kvState, tvm.scalar(0, "int64"), tvm.scalar(n, "int32"));
        }
      } finally {
        tvm.endScope();
      }
      pipeline.filledKVCacheLength -= n;
      return;
    }
  } catch {
    // Fall through: a trim that threw is handled the same as no trim at all.
  }
  pipeline.resetChat(/* keepStats= */ true);
}

const popNCache = new WeakMap();

function getPopN(pipeline) {
  if (popNCache.has(pipeline)) return popNCache.get(pipeline);
  let popn = null;
  const { tvm } = pipeline;
  tvm.beginScope();
  try {
    popn = tvm.detachFromCurrentScope(tvm.getGlobalFunc("vm.builtin.kv_state_popn"));
  } catch {
    popn = null;
  } finally {
    tvm.endScope();
  }
  popNCache.set(pipeline, popn);
  return popn;
}

// ------------------------------------------------------- sampling inputs ----

/**
 * The subset of `sampleTokenFromLogits`'s config resolution a burst can honour,
 * in the same precedence order: the request overrides `mlc-chat-config.json`.
 */
function resolveSampling(pipeline, genConfig) {
  const has = (v) => v !== undefined && v !== null;
  const pick = (key, fallback) => (has(genConfig?.[key]) ? genConfig[key] : fallback);
  return {
    temperature: pick("temperature", pipeline.config.temperature),
    top_p: pick("top_p", pipeline.config.top_p) ?? 1,
    repetition_penalty: pick("repetition_penalty", pipeline.config.repetition_penalty),
    frequency_penalty: pick("frequency_penalty", pipeline.config.frequency_penalty) ?? 0,
    presence_penalty: pick("presence_penalty", pipeline.config.presence_penalty) ?? 0,
    logit_bias: pick("logit_bias", undefined),
  };
}

/** Static for the whole request, so it is uploaded once and reused every step. */
function makeLogitBias(pipeline, { logit_bias }) {
  const ids = Object.keys(logit_bias ?? {});
  if (ids.length === 0) return null;
  const { tvm, device } = pipeline;
  const int32 = (values) =>
    tvm.detachFromCurrentScope(tvm.empty([values.length], "int32", device).copyFrom(values));
  return {
    pos2seqIds: int32(new Int32Array(ids.length)),
    tokenIds: int32(Int32Array.from(ids, (id) => parseInt(id, 10))),
    values: tvm.detachFromCurrentScope(
      tvm.empty([ids.length], "float32", device).copyFrom(Float32Array.from(ids, (id) => logit_bias[id])),
    ),
  };
}

/**
 * Frozen token history for the burst.
 *
 * This is the one place multi-step is not equivalent to single-step: tokens
 * sampled *within* a burst are not penalised against each other, because their
 * ids are still on the GPU. At K=15 the penalty state is at most 15 tokens
 * stale. Anything that cannot tolerate that should run with `decodeSteps: 1`.
 */
function makePenalty(pipeline, { repetition_penalty, frequency_penalty, presence_penalty }) {
  const active = frequency_penalty !== 0 || presence_penalty !== 0 || (repetition_penalty ?? 1) !== 1;
  if (!active) return null;

  const appeared = [...pipeline.appearedTokensFreq.keys()];
  if (appeared.length === 0) return null;
  const freqs = [...pipeline.appearedTokensFreq.values()];

  const { tvm, device } = pipeline;
  const int32 = (values) =>
    tvm.detachFromCurrentScope(tvm.empty([values.length], "int32", device).copyFrom(values));
  return {
    seqIds: int32(new Int32Array(1)),
    pos2seqIds: int32(new Int32Array(appeared.length)),
    tokenIds: int32(Int32Array.from(appeared)),
    counts: int32(Int32Array.from(freqs)),
    penalties: tvm.detachFromCurrentScope(
      tvm
        .empty([1, 3], "float32", device)
        .copyFrom(new Float32Array([presence_penalty, frequency_penalty, repetition_penalty ?? 1])),
    ),
  };
}

function disposeAll(inputs) {
  if (!inputs) return;
  for (const tensor of Object.values(inputs)) tensor.dispose();
}