File size: 50,736 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 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 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 1004 1005 1006 1007 1008 1009 1010 1011 1012 1013 1014 1015 1016 1017 1018 1019 1020 1021 1022 1023 1024 1025 1026 1027 1028 1029 1030 1031 1032 1033 1034 1035 1036 1037 1038 1039 1040 1041 1042 1043 1044 1045 1046 1047 1048 1049 1050 1051 1052 1053 1054 1055 1056 1057 1058 1059 1060 1061 1062 1063 | // Decoder: per-run persistent state + a single decode step recorded into the
// caller's compute pass. One step = 25 dispatches:
// embed DECODE (1)
// per layer ×2 (11 each):
// x = LN1(x + SelfAttn(x)) gemm self_qkv (epilogue scatters k|v into the
// caches — no separate kv_append dispatch)
// → attention → gemm self_out → add_ln
// x = LN2(x + CrossAttn(x)) gemm cross_q → attention (K/V = encoder crossKV) → gemm cross_out → add_ln
// x = LN3(x + FFN(x)) gemm fc1 (SiLU) → gemm fc2 → add_ln
// LM head gemm x @ shared.weightᵀ → f32 logits (1) — final_logits_bias is
// NOT added here; the argmax kernel adds it (matching HF ordering)
// argmax_penalty → token ring (1)
// (tiled lm_head default: the two fuse — gemm emits argmax partials, no
// logits store, and argmax_reduce writes the ring; same dispatch count)
//
// LN output must not alias its residual input (read + read_write usage
// conflict — see encoder.js), so the hidden state ping-pongs hidden[0] ↔
// hidden[1]. Three LNs per layer × 2 layers = 6 swaps: a step both starts and
// ends with the hidden state in hidden[0].
import { createArena } from './arena.js';
import { deviceSupportsImmediates } from './device.js';
import {
dispatchGemm, dispatchAttention, dispatchAddLn, dispatchEmbed,
dispatchArgmaxPenalty, dispatchArgmaxReduce, dispatchGemmRowLn,
dispatchGemmReduce, dispatchDecoderMega, decodeMegaSharedBytes, splitKParts,
dispatchCompactGather, purgeBindGroupsForBuffers,
} from './pipelines.js';
import {
D_MODEL, HEADS, HEAD_DIM, FFN, VOCAB, DECODE_CAP, BITMASK_WORDS, DEC_LAYERS,
DECODER_START, assertModelActive,
} from './constants.js';
// Per-run persistent decode state: K/V caches, token ring, done flags,
// repetition bitmask, logits, and the step's activation buffers. K/V use a
// separate replaceable arena so a planner-sized generation can grow without
// rebuilding the rest of the decode state. state.destroy() frees both owners.
//
// ctx = {device}; S is unused for sizing (kept for symmetry/debug) — cross
// K/V stay in encRun's buffers. maxSteps bounds the token ring.
//
// lmHead: 'auto' (default) | 'gemv' | 'tiled' | 'q8' — which kernel serves the
// [B,448]×[448,24000] logits projection. The tuned GEMV re-sweeps the 21.5MB
// embedding matrix every ceil(B/MT=8) rows, so past B≈32 the tiled kernel
// (one W sweep total) wins; below that the GEMV's latency shape wins.
// 'auto' picks tiled at B ≥ LM_HEAD_TILED_MIN_B (threshold measured by the
// lm_head_sweep debug test); the explicit values exist for that sweep.
// lmHeadFlags: extra GEMM flags merged at the lm_head site when tiled (sweep
// hook for kernel-geometry A/Bs, e.g. {tm8: true} or {tiledV: 1}).
//
// tiledProj: 'auto' (default) | array of projection site kinds routed to the
// tiled kernel instead of the GEMV — subset of ['self_out', 'cross_q',
// 'cross_out', 'fc1', 'fc2'] (self_qkv stays GEMV: its storeKV cache-scatter
// epilogue is a GEMV-only feature). 'auto' = PROJ_TILED_KINDS at
// B ≥ PROJ_TILED_MIN_B, else none (measured by proj_sweep).
//
// lmHeadFuse: 'auto' (default) | 'on' | 'off' — fuse the greedy argmax into
// the tiled lm_head's epilogue (gemm_tiled2 IF_ARGMAX + argmax_reduce): the
// [B, 24000] f32 logits are never materialized. Bit-identical token picks
// (same f32 ops, same tie order — argmax_fuse gate), so 'auto' fuses whenever
// the tiled v2 kernel serves lm_head; 'off' is the A/B control. The GEMV
// path (B < 16) keeps the unfused argmax_penalty scan.
// fuseLn: 'auto' (default) | 'on' | 'off' — fuse the three LN-terminated
// projections (self_out+ln1, cross_out+ln2, fc2+ln3) into single
// gemm_row_ln dispatches (25 → 19 per step). One workgroup per row caps the
// fused GEMM at M workgroups, so 'auto' fuses only at B ≤ FUSE_LN_MAX_B
// (measured by fuse_ln_sweep); numerics shift at f32-ULP level (accumulation
// order), gated by m3/goldens like every routing change. When ffnMode is
// 'q8' the fc2 site stays unfused (the fused kernel reads float weights and
// would silently unquantize it); self_out/cross_out still fuse.
// proj: 'auto' | 'f16' | 'wt' — W layout for the four attention-side decode
// projections (self_qkv incl. its storeKV cache scatter, self_out, cross_q,
// cross_out). 'f16' = the original [K,N] tensors (GEMV NWT: the whole W is
// re-read per batch ROW — 46.6% of the b128 step, prod_profile 2026-07-06);
// 'wt' = transposed copies (loadWeights {projWT: true}, +4.8MB) on the GEMV
// WT path (one W-tile per MT=8 rows). 'auto' picks 'wt' whenever the copies
// were loaded (range confirmed by proj_wt_sweep).
// ffnSplitK: 'auto' (default) | 0 | sk | {fc1, fc2} — split-K for the TILED
// fc1/fc2 sites (the B ≥ PROJ_TILED_MIN_B route), which are workgroup-starved
// (fc2 at b128: N=448 → 7×2 = 14 workgroups). sk partitions K over grid.z
// and a gemm_reduce dispatch folds the partials + bias/SiLU (two dispatches
// per site instead of one). Only applies where the site actually runs tiled
// non-q8; numerics shift at the f32 re-association seams (gated like every
// routing change). 'auto' = FFN_SPLITK_AUTO at the tiled sites.
// projSplitK: 'auto' (default) | 0 | sk | {qkv, out} — the same split-K
// medicine for the four attention-side projections (post-split-K profile:
// self_qkv ×2 = 24.6% of the b64 step, the three N=448 sites ×2 = 24.1%).
// Routes tiled(+splitK) on the WT copies instead of the GEMV WT sweep;
// self_qkv's K/V-cache scatter moves into the gemm_reduce epilogue
// (bit-identical to Y's slices, the kv_append contract). Needs projMode
// 'wt'. 'auto' = PROJ_SPLITK_AUTO per site group from its
// PROJ_SPLITK_MIN_B threshold (qkv ≥ 64, out ≥ 128 — set by the E2E A/B,
// not the steady-state sweep; see the consts' comment).
// decodeMega: 'auto' (default) | 'on' | 'off' — the small-batch decode-step
// MEGAKERNEL (decoder_mega.wgsl): one dispatch per LAYER computes a whole
// row's layer in one workgroup (embed folds into layer 0), so a step becomes
// mega L0 → mega L1 → lm_head → argmax = 4 dispatches instead of 19. The
// one-workgroup-per-row shape serializes each row's weight streaming through
// one SM, so it only wins where the step is dispatch-overhead-bound —
// 'auto' engages at B ≤ DECODE_MEGA_MAX_B, and needs f16 weights with the
// projWT+ffnWT transposed copies (q8 FFN has no mega path). Not bit-exact
// vs the unfused chain (accumulation order differs at every site) — gated
// by m3/goldens + mega_equiv like every routing change; pinned across
// compaction like every knob.
export function normalizeKvCapacity(value, maxSteps) {
if (!Number.isInteger(maxSteps) || maxSteps < 1 || maxSteps > DECODE_CAP) {
throw new Error(`decode maxSteps must be an integer in [1, ${DECODE_CAP}], got ${maxSteps}`);
}
// Missing preserves the legacy allocation exactly. Explicit planner/forced
// values are clamped to work the run can actually produce.
if (value === null || value === undefined) return DECODE_CAP;
if (!Number.isInteger(value) || value < 1) {
throw new Error(`kvCapacity must be a positive integer, got ${value}`);
}
return Math.min(value, maxSteps, DECODE_CAP);
}
export function nextKvCapacity({ current, required, maxSteps, groupSteps = 8 }) {
for (const [name, value] of Object.entries({ current, required, maxSteps, groupSteps })) {
if (!Number.isInteger(value) || value < 1) {
throw new Error(`KV grow ${name} must be a positive integer, got ${value}`);
}
}
if (current > DECODE_CAP || maxSteps > DECODE_CAP) {
throw new Error(`KV grow capacity exceeds decode cap ${DECODE_CAP}`);
}
if (required > maxSteps) {
throw new Error(`KV grow requires ${required} steps but run maxSteps=${maxSteps}`);
}
if (required <= current) return current;
const wanted = Math.max(required, current * 2);
const aligned = Math.ceil(wanted / groupSteps) * groupSteps;
return Math.min(maxSteps, DECODE_CAP, aligned);
}
function allocateKvGeneration(device, { B, capacity, HD, eb, usage }) {
const arena = createArena(device);
const kvCacheK = [];
const kvCacheV = [];
try {
for (let l = 0; l < DEC_LAYERS; l++) {
kvCacheK.push(arena.buf(B * capacity * HD * eb, usage, `dec kvK.${l} cap${capacity}`));
kvCacheV.push(arena.buf(B * capacity * HD * eb, usage, `dec kvV.${l} cap${capacity}`));
}
} catch (err) {
arena.destroy();
throw err;
}
return { arena, kvCacheK, kvCacheV };
}
export function createDecodeState(ctx, weights, { B, S, maxSteps, kvCapacity = null, lmHead = 'auto', lmHeadFlags = null, lmHeadFuse = 'auto', tiledProj = 'auto', ffn = 'auto', ffnFlags = null, fuseLn = 'auto', proj = 'auto', ffnSplitK = 'auto', projSplitK = 'auto', decodeMega = 'auto', sg = 'auto', immediates = 'auto', inPlaceCompact = false }) {
assertModelActive(weights.model, 'createDecodeState weights');
const HD = HEADS * HEAD_DIM; // == D_MODEL (enforced by applyModelConfig)
const QKV_N = 3 * HD; // fused q|k|v
const { device } = ctx;
const eb = weights.dtype === 'f16' ? 2 : 4;
const normalizedKvCapacity = normalizeKvCapacity(kvCapacity, maxSteps);
// Large-batch guard: every decode-side buffer is bound whole, so each must
// fit maxStorageBufferBindingSize. The two candidates that grow with B are
// the per-layer KV caches at their possible full-cap grow size and the f32 logits
// [B, 24000] (~25.7 MB and ~12.3 MB at B=128/f16) — assert both up front
// with a clear error instead of an opaque createBuffer validation failure.
const bindingLimit = ctx?.limits?.maxStorageBufferBindingSize ?? 134217728;
const biggest = [
[`kv cache max [B, ${DECODE_CAP}, ${HD}]`, B * DECODE_CAP * HD * eb],
['f32 logits [B, 24000]', B * VOCAB * 4],
['token ring [maxSteps, B]', maxSteps * B * 4],
];
for (const [what, bytes] of biggest) {
if (bytes > bindingLimit) {
throw new Error(
`createDecodeState: ${what} = ${bytes} bytes exceeds ` +
`maxStorageBufferBindingSize=${bindingLimit} at B=${B} — reduce batch`,
);
}
}
const arena = createArena(device);
const act = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC;
const rw = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST | GPUBufferUsage.COPY_SRC;
// Self-attention K/V caches, [B, kvCapacity, H, D] each, per layer.
// COPY_DST: compactDecodeState copies live-row prefixes into a fresh state.
const kvUsage = act | GPUBufferUsage.COPY_DST;
const kv = allocateKvGeneration(device, {
B, capacity: normalizedKvCapacity, HD, eb, usage: kvUsage,
});
try {
const { kvCacheK, kvCacheV } = kv;
// tokenRing/done rely on WebGPU zero-initialization of fresh buffers.
const tokenRing = arena.buf(maxSteps * B * 4, rw, 'dec token ring');
const done = arena.buf(B * 4, rw, 'dec done');
// Repetition bitmask: the DECODER_START bit pre-set per row — HF counts
// decoder_start in the penalized input_ids too.
const bitmask = arena.buf(B * BITMASK_WORDS * 4, rw, 'dec bitmask');
const maskInit = new Uint32Array(B * BITMASK_WORDS);
for (let b = 0; b < B; b++) {
maskInit[b * BITMASK_WORDS + (DECODER_START >> 5)] = 1 << (DECODER_START & 31);
}
device.queue.writeBuffer(bitmask, 0, maskInit);
// Step activations. hidden[0]/hidden[1] ping-pong across add_ln outputs.
const hidden = [
arena.buf(B * D_MODEL * eb, act, 'dec hidden a'),
arena.buf(B * D_MODEL * eb, act, 'dec hidden b'),
];
const y = arena.buf(B * D_MODEL * eb, act, 'dec sublayer y');
const attnOut = arena.buf(B * HD * eb, act, 'dec attn out');
const qkvOut = arena.buf(B * QKV_N * eb, act, 'dec qkv out');
const ffnTmp = arena.buf(B * FFN * eb, act, 'dec ffn tmp');
const crossQOut = arena.buf(B * HD * eb, act, 'dec cross q');
// 'q8': the tiled v2 kernel's W8A16 path over the int8-quantized embedding
// matrix (weights loaded with lmHeadQ8: true). Halves lm_head W traffic —
// lm_head_sweep step wall vs f16 tiled: b16 780→639µs, b64 1406→1358µs.
// NOT token-exact vs f16 (quantization error; golden 29/30 — q8_lmhead
// gate), so 'auto' only picks it when the caller opted into quantization
// at loadWeights — parity suites load without it and keep f16 semantics.
const lmHeadQ8 = lmHead === 'q8'
|| (lmHead === 'auto' && B >= LM_HEAD_TILED_MIN_B && weights.tensors.has('lm_head.q8'));
if (lmHeadQ8 && !weights.tensors.has('lm_head.q8')) {
throw new Error("lmHead 'q8' needs weights loaded with lmHeadQ8: true");
}
const lmHeadTiled =
lmHeadQ8 || lmHead === 'tiled' || (lmHead === 'auto' && B >= LM_HEAD_TILED_MIN_B);
// ffn: 'auto' | 'f16' | 'wt' | 'q8' — decode fc1/fc2 weight path.
// 'f16' the original [K,N] tensors (GEMV NWT / tiled WNT). NWT re-streams
// the whole W once per batch ROW — B×6.4MB per step.
// 'wt' f16 transposed copies (loadWeights {ffnWT}): GEMV WT sweeps one
// W tile for MT=8 rows (traffic ÷8), full f16 precision.
// 'q8' int8 [N,K]-packed (loadWeights {ffnQ8}): WT traffic halved again,
// BUT quality is below the golden bar (q8_full_golden 27/30 vs
// need 29) — explicit opt-in only, never chosen by 'auto'.
const ffnMode = (() => {
if (ffn === 'q8') {
if (!weights.tensors.has('dec.0.fc1.q8')) throw new Error("ffn 'q8' needs weights loaded with ffnQ8: true");
return 'q8';
}
if (ffn === 'wt') {
if (!weights.tensors.has('dec.0.fc1.wt')) throw new Error("ffn 'wt' needs weights loaded with ffnWT: true");
return 'wt';
}
if (ffn === 'f16') return 'nwt';
// auto: 'wt' below the tiled threshold (ffn_q8_sweep: step b64 1400→1209µs,
// b16 770→667); at B ≥ 128 the tiled WNT staging beats WT (1788 vs 1856) —
// keep the f16 [K,N] tensors there.
return (B < PROJ_TILED_MIN_B && weights.tensors.has('dec.0.fc1.wt')) ? 'wt' : 'nwt';
})();
// Desired split-K per FFN site (before tiled-eligibility). Split-K widens
// the tiled kernel's winning range: with sk8 the tiled route beats the
// production GEMV from B ≥ FFN_SPLITK_MIN_B (ffn_splitk_sweep: b32 668→600,
// b64 977→779, b128 1443→1267µs), so 'auto' tiledProj flips fc1/fc2 tiled
// there too — but ONLY the split-K kinds (plain tiled still loses below
// B=128, and the q8 kernel has no split-K path).
const wantSK = ffnSplitK === 'auto'
? { fc1: FFN_SPLITK_AUTO.fc1, fc2: FFN_SPLITK_AUTO.fc2 }
: typeof ffnSplitK === 'number' || !ffnSplitK
? { fc1: ffnSplitK || 0, fc2: ffnSplitK || 0 }
: { fc1: ffnSplitK.fc1 || 0, fc2: ffnSplitK.fc2 || 0 };
const skKinds = ffnMode === 'q8' ? [] : PROJ_TILED_KINDS.filter((k) => wantSK[k] > 0);
const tiledProjKinds = new Set(
tiledProj === 'auto'
? (B >= PROJ_TILED_MIN_B ? PROJ_TILED_KINDS
: B >= FFN_SPLITK_MIN_B ? skKinds : [])
: tiledProj,
);
// Fused lm_head argmax: tiled v2 only (the v1 fallback and the GEMV have no
// fused epilogue). Token picks are bit-identical to the unfused path, so
// 'auto' means "whenever eligible"; 'off' is the sweep control.
if (lmHeadFuse === 'on' && !lmHeadTiled) {
throw new Error("lmHeadFuse 'on' needs a tiled lm_head (B >= 16 or lmHead 'tiled'/'q8')");
}
const lmHeadFused = lmHeadFuse !== 'off' && lmHeadTiled && (lmHeadFlags?.tiledV ?? 2) === 2;
// Static lm_head shortlist (loadWeights {lmHeadIds}): the q8 tensor holds
// only the emittable rows, so every lm_head consumer below sizes by lmN and
// the argmax epilogues translate local→vocab via lm_head.idmap. Rides the
// q8 route only — GEMV/f16-tiled states keep the full vocab.
const lmShortMeta = weights.lmHeadShort ?? null;
const lmShort = !!(lmShortMeta && lmHeadQ8);
const lmN = lmShort ? lmShortMeta.ids.length : VOCAB;
const lmMaskWords = Math.ceil(lmN / 32);
// Partial (val, idx) pairs per [row, column-tile] — NT must match the BN
// the lm_head dispatch will use.
const lmNT = Math.ceil(lmN / (lmHeadFlags?.bn ?? 64));
const logits = lmHeadFused ? null : arena.buf(B * lmN * 4, act, 'dec logits'); // always f32
const argmaxPartials = lmHeadFused ? arena.buf(B * lmNT * 8, act, 'dec argmax partials') : null;
// Local-space twin of the repetition bitmask: the fused epilogue and the
// shortlisted argmax_penalty read/write THIS one (contiguous local bits —
// the quad trick in gemm_tiled2 stays valid); the vocab-space `bitmask`
// above stays maintained in parallel so a full-vocab state inheriting these
// rows (routing is pinned, but belt-and-braces) reads correct history.
// Same DECODER_START pre-set, at its local index.
let bitmaskL = null;
if (lmShort) {
const startL = lmShortMeta.ids.indexOf(DECODER_START);
if (startL < 0) throw new Error('lm_head shortlist must contain decoderStart');
bitmaskL = arena.buf(B * lmMaskWords * 4, rw, 'dec bitmaskL');
const initL = new Uint32Array(B * lmMaskWords);
for (let b = 0; b < B; b++) {
initL[b * lmMaskWords + (startL >> 5)] = 1 << (startL & 31);
}
device.queue.writeBuffer(bitmaskL, 0, initL);
}
const lnFused = fuseLn === 'on' || (fuseLn === 'auto' && B <= FUSE_LN_MAX_B);
// Effective split-K per FFN site: 0 wherever the site doesn't run the
// tiled non-q8 kernel (splitK is a tiled-v2 float-W feature). Resolved
// here — not in encodeDecodeStep — so compaction can pin the exact values.
const eligibleSK = (kind) => tiledProjKinds.has(kind) && ffnMode !== 'q8';
const ffnSK = {
fc1: eligibleSK('fc1') ? wantSK.fc1 : 0,
fc2: eligibleSK('fc2') ? wantSK.fc2 : 0,
};
const projMode = (() => {
if (proj === 'wt') {
if (!weights.tensors.has('dec.0.self_qkv.wt')) throw new Error("proj 'wt' needs weights loaded with projWT: true");
return 'wt';
}
if (proj === 'f16') return 'nwt';
return weights.tensors.has('dec.0.self_qkv.wt') ? 'wt' : 'nwt';
})();
// Effective split-K for the attention-side projections: qkv = self_qkv
// (N=1344, cache scatter in the reduce), out = self_out/cross_q/cross_out
// (N=448). Needs the WT copies (tiled wt layout) — 0 when projMode 'nwt'.
const wantPSK = projSplitK === 'auto'
? {
qkv: B >= PROJ_SPLITK_MIN_B.qkv ? PROJ_SPLITK_AUTO.qkv : 0,
out: B >= PROJ_SPLITK_MIN_B.out ? PROJ_SPLITK_AUTO.out : 0,
}
: typeof projSplitK === 'number' || !projSplitK
? { qkv: projSplitK || 0, out: projSplitK || 0 }
: { qkv: projSplitK.qkv || 0, out: projSplitK.out || 0 };
const projSK = projMode === 'wt' ? wantPSK : { qkv: 0, out: 0 };
// Shared raw-partials scratch for every split-K site, sized for the
// largest active [nz, B, N] f32 layout (the sites run sequentially in the
// step's dependency chain, so one buffer serves them all).
const skBK = ffnFlags?.bkk ?? 16;
const partsBytes = Math.max(
ffnSK.fc1 ? splitKParts(D_MODEL, ffnSK.fc1, skBK).nz * B * FFN * 4 : 0,
ffnSK.fc2 ? splitKParts(FFN, ffnSK.fc2, skBK).nz * B * D_MODEL * 4 : 0,
projSK.qkv ? splitKParts(D_MODEL, projSK.qkv).nz * B * QKV_N * 4 : 0,
projSK.out ? splitKParts(HD, projSK.out).nz * B * D_MODEL * 4 : 0,
);
const skParts = partsBytes ? arena.buf(partsBytes, act, 'dec splitK parts') : null;
// Megakernel eligibility: f16 weights, a float FFN (mega reads the
// original [K,N] .weight tensors — always present — but has no int8 path,
// so a q8-FFN config must keep the unfused chain rather than silently
// dropping the quantization), and a model whose dims fit the kernel's
// 16KB workgroup-shared budget (Moxhi-30 13,184B, Hachimi-60 16,256B —
// both inside; a larger member falls back to the split pipeline).
const megaOk = weights.dtype === 'f16' && ffnMode !== 'q8'
&& decodeMegaSharedBytes() <= 16384;
if (decodeMega === 'on' && !megaOk) {
throw new Error(
"decodeMega 'on' needs f16 weights, ffn != 'q8', and dims inside the 16KB shared budget");
}
const mega = decodeMega === 'on'
|| (decodeMega === 'auto' && megaOk && B <= DECODE_MEGA_MAX_B);
// In-place compaction uses persistent plumbing so the ownership boundary
// itself allocates no GPUBuffer. Rebuild/default states pay no memory for
// the experiment. One aligned parameter slot serves each possible row-major
// buffer (K/V/cross per layer plus global/local repetition masks).
const compactMap = inPlaceCompact
? arena.buf(B * 4, GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, 'dec compact map')
: null;
const compactParams = inPlaceCompact
? arena.buf(
(DEC_LAYERS * 3 + 2) * COMPACT_PARAM_STRIDE,
GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST,
'dec compact params',
)
: null;
// Subgroup reductions ('sg') at the WT-GEMV and LN(-fused) sites. Needs
// the feature AND a TK-slice that can't straddle a subgroup (TK=16
// default → subgroupMinSize ≥ 16; Intel can report 8). 'auto' stays OFF
// until the sg sweeps pick per-device defaults; 'on' is the sweep arm.
const sgOk = !!ctx.hasSubgroups && (ctx.subgroupMinSize ?? 0) >= 16;
if (sg === 'on' && !sgOk) {
throw new Error("sg 'on' needs the subgroups feature and subgroupMinSize ≥ 16");
}
const sgOn = sg === 'on';
if (!['auto', 'on', 'off'].includes(immediates)) {
throw new Error(`immediates must be 'auto', 'on', or 'off', got ${immediates}`);
}
const immediateAvailable = !!ctx.hasImmediates || deviceSupportsImmediates(device);
if (immediates === 'on' && !immediateAvailable) {
throw new Error("immediates 'on' needs WGSL immediate_address_space support");
}
const immediateOn = immediates === 'on' || (immediates === 'auto' && immediateAvailable);
const state = {
B, S, maxSteps, kvCapacity: normalizedKvCapacity, arena, kvArena: kv.arena,
lmHeadTiled, lmHeadQ8, lmHeadFused, lmHeadFlags, lmNT,
lmShort, lmN, lmMaskWords, bitmaskL,
tiledProjKinds, ffnMode, ffnFlags, lnFused, projMode, ffnSK, projSK, skParts,
mega, sg: sgOn, immediate: immediateOn, inPlaceCompact,
kvCacheK, kvCacheV, tokenRing, done, bitmask, logits, argmaxPartials,
hidden, y, attnOut, qkvOut, ffnTmp, crossQOut, compactMap, compactParams,
destroy() {
state.kvArena?.destroy();
state.kvArena = null;
arena.destroy();
},
};
return state;
} catch (err) {
// A later persistent-state allocation can still fail at high B. Keep the
// split ownership from leaking the already-created KV generation.
kv.arena.destroy();
arena.destroy();
throw err;
}
}
// Replace only the self-attention K/V generation. Copies are submitted after
// any already-queued decode groups, and later groups are submitted after this
// copy, so queue order supplies the synchronization without a CPU wait.
export function growDecodeKV(
ctx, weights, state,
{ requiredCapacity, submittedSteps, groupSteps = 8 } = {},
) {
assertModelActive(weights.model, 'growDecodeKV weights');
if (!state?.kvArena || !Array.isArray(state.kvCacheK) || !Array.isArray(state.kvCacheV)) {
throw new Error('growDecodeKV: state has no owned KV generation');
}
if (!Number.isInteger(submittedSteps) || submittedSteps < 0
|| submittedSteps > state.kvCapacity) {
throw new Error(
`growDecodeKV: submittedSteps=${submittedSteps} outside [0, ${state.kvCapacity}]`,
);
}
const nextCapacity = nextKvCapacity({
current: state.kvCapacity,
required: requiredCapacity,
maxSteps: state.maxSteps,
groupSteps,
});
if (nextCapacity === state.kvCapacity) {
return {
grown: false,
oldCapacity: state.kvCapacity,
newCapacity: state.kvCapacity,
copiedSteps: submittedSteps,
bindGroupsPurged: 0,
};
}
const { device } = ctx;
const HD = HEADS * HEAD_DIM;
const eb = weights.dtype === 'f16' ? 2 : 4;
const kvUsage = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST;
const next = allocateKvGeneration(device, {
B: state.B, capacity: nextCapacity, HD, eb, usage: kvUsage,
});
const oldCapacity = state.kvCapacity;
const oldK = state.kvCacheK;
const oldV = state.kvCacheV;
const oldArena = state.kvArena;
const prefixBytes = submittedSteps * HD * eb;
let submitted = false;
try {
if (prefixBytes > 0) {
const encoder = device.createCommandEncoder({
label: `grow decode KV ${oldCapacity} -> ${nextCapacity}`,
});
const oldRowBytes = oldCapacity * HD * eb;
const newRowBytes = nextCapacity * HD * eb;
for (let b = 0; b < state.B; b++) {
for (let l = 0; l < DEC_LAYERS; l++) {
encoder.copyBufferToBuffer(
oldK[l], b * oldRowBytes, next.kvCacheK[l], b * newRowBytes, prefixBytes,
);
encoder.copyBufferToBuffer(
oldV[l], b * oldRowBytes, next.kvCacheV[l], b * newRowBytes, prefixBytes,
);
}
}
device.queue.submit([encoder.finish()]);
submitted = true;
}
} catch (err) {
next.arena.destroy();
throw err;
}
let drained = Promise.resolve();
if (submitted && typeof device.queue.onSubmittedWorkDone === 'function') {
try { drained = device.queue.onSubmittedWorkDone(); } catch { /* device loss surfaces elsewhere */ }
}
// Removing JS cache ownership is safe for already-submitted command buffers
// and must happen before GPUBuffer.destroy() on the retired generation.
const bindGroupsPurged = purgeBindGroupsForBuffers(device, [...oldK, ...oldV]);
state.kvCacheK = next.kvCacheK;
state.kvCacheV = next.kvCacheV;
state.kvArena = next.arena;
state.kvCapacity = nextCapacity;
oldArena.destroyDeferred(drained);
return {
grown: true,
oldCapacity,
newCapacity: nextCapacity,
copiedSteps: submittedSteps,
bindGroupsPurged,
};
}
// Uniform dynamic offsets must satisfy minUniformBufferOffsetAlignment. WebGPU
// guarantees that 256-byte slots meet the default/supported alignment limit.
const COMPACT_PARAM_STRIDE = 256;
// Measured crossover (lm_head_sweep, RTX 5070 Ti, medians of 7 interleaved):
// step wall gemv→tiled at B=16: 833→791µs, B=32: 1273→1028µs, B=64:
// 2137→1461µs — tiled wins from B=16 up; below that stays on the GEMV's
// latency-optimized shape (unmeasured territory, and b1–b8 is where GEMV was
// tuned). Details in notes-m4-tuning.md.
const LM_HEAD_TILED_MIN_B = 16;
// Which projection kinds go tiled in 'auto' mode, and from what batch —
// measured by proj_sweep (medians of 7 interleaved, step wall): tiled LOSES
// at these sites for B ≤ 64 (b64: gemv-all 1438µs, ffn-tiled 1701, all-5
// 1938 — N=448–1792 tiles yield only 7–28 workgroups, the GPU is starved;
// unlike lm_head's N=24000 → 375). Only the FFN pair at B=128 wins
// (2539→2347µs, −7.6%); the out/cross_q sites lose everywhere. Details in
// notes-m4-tuning.md.
const PROJ_TILED_KINDS = ['fc1', 'fc2'];
const PROJ_TILED_MIN_B = 128;
// Split-K factors 'auto' uses at the tiled FFN sites (0 = off), and the
// batch where the split-K tiled route starts beating the GEMV one. The
// starved shapes (fc2: N=448 → 14 workgroups, fc1: 56) gain nz× workgroups;
// ffn_splitk_sweep (medians of 7 interleaved, step wall): sk8 b32 668→600µs
// (−10%), b64 977→779 (−20%), b128 1443→1267 (−12%); b16's −3.7% sits
// inside run noise (716µs outliers) — GEMV keeps it.
const FFN_SPLITK_AUTO = { fc1: 8, fc2: 8 };
const FFN_SPLITK_MIN_B = 32;
// Split-K factors for the attention-side projections in 'auto' (0 = off),
// and the batch each site GROUP engages from. The steady-state step sweep
// (proj_splitk_sweep: sk8 both groups b64 −8.5%, b128 −25.1%) is NOT the
// decider here — split-K is pinned across compaction, so an e2e run spends
// most steps at live B far below the group size, where the N=448 trio tanks
// (+13.8% at b32, +32% at b16). proj_splitk_e2e_ab (full translateBatch,
// interleaved): qkv-only b64 −8.8% / b128 −10.5%; adding the out-trio at
// b64 flips to +13.3% but reaches −18.1% at b128 — hence the split
// thresholds. qkv never loses at any measured live B (b16 −0.4%).
const PROJ_SPLITK_AUTO = { qkv: 8, out: 8 };
const PROJ_SPLITK_MIN_B = { qkv: 64, out: 128 };
// Largest batch where the decode megakernel wins (mega_sweep + mega_e2e_ab):
// one workgroup per row means B workgroups for the whole layer — pure
// latency shape. Below the threshold the step is dispatch-overhead-bound
// (19 dispatches × ~15µs fixed cost) and collapsing a layer's 8 dispatches
// into 1 wins; past it the starved GEMVs lose more than the overhead saved.
const DECODE_MEGA_MAX_B = 4;
// Largest batch where the fused projection+LN kernel wins (fuse_ln_sweep,
// medians of 7 interleaved: b1 608→486µs −20%, b4 −7.8%, b8 +21% LOSES,
// b32 +37%): its one-workgroup-per-row shape starves the GPU as B grows, but
// below the threshold the step is dispatch-overhead-bound and 6 fewer
// dispatches win. 0 would disable 'auto' fusing entirely.
const FUSE_LN_MAX_B = 4;
// EOS row compaction: build a fresh, smaller decode context holding only the
// live rows of a running decode, so finished rows stop consuming GEMM rows
// and attention workgroups. Called between step groups (t0 = the next step
// index; the KV caches hold positions 0..t0-1).
//
// prev {state, crossKV, lensBuf, S} — the CURRENT decode view (crossKV/
// lensBuf are runEncoder's on the first compaction, a previous
// compaction's after that). The caller destroys the old pieces
// AFTER the returned copies have been submitted (this function
// submits them itself — old buffers are queue-retained).
// liveIdx current-row indices to keep, in order (new row i = old liveIdx[i])
// lens Uint32Array[newB] — source lengths of the kept rows
// lastTok Uint32Array[newB] — step t0-1 tokens of the kept rows (from the
// group readback); embed at step t0 reads ring[(t0-1)·B + b], and
// ring history is NOT copied (the CPU already collected it)
//
// Kernel routing (lmHead/tiledProj + lmHeadFlags) is PINNED from the old
// state rather than re-derived from the smaller B: per-row math is identical
// in every kernel here, so a compacted run must produce token-exact output
// vs the uncompacted run (compact_equiv gate) — re-routing at the new B could
// legally flip near-tie argmaxes and would make that equivalence untestable.
//
// Returns {state, crossKV, lensBuf, S, arena} — `arena` owns the new
// crossKV/lens buffers; destroy it alongside state.
export function compactDecodeState(ctx, weights, prev, { liveIdx, t0, cap, lens, lastTok }) {
const HD = HEADS * HEAD_DIM;
const CROSS_KV_STRIDE = 2 * HD; // fused k|v (encoder crossKV layout)
const { device } = ctx;
const { state, crossKV, S } = prev;
const eb = weights.dtype === 'f16' ? 2 : 4;
const newB = liveIdx.length;
const next = createDecodeState(ctx, weights, {
B: newB, S, maxSteps: cap, kvCapacity: state.kvCapacity,
lmHead: state.lmHeadQ8 ? 'q8' : state.lmHeadTiled ? 'tiled' : 'gemv',
lmHeadFlags: state.lmHeadFlags,
lmHeadFuse: state.lmHeadFused ? 'on' : 'off',
tiledProj: [...state.tiledProjKinds],
ffn: { q8: 'q8', wt: 'wt', nwt: 'f16' }[state.ffnMode],
ffnFlags: state.ffnFlags,
fuseLn: state.lnFused ? 'on' : 'off',
proj: state.projMode === 'wt' ? 'wt' : 'f16',
ffnSplitK: { ...state.ffnSK },
projSplitK: { ...state.projSK },
decodeMega: state.mega ? 'on' : 'off',
sg: state.sg ? 'on' : 'off',
immediates: state.immediate ? 'on' : 'off',
});
const arena = createArena(device);
const cross = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST;
const newCrossKV = [
arena.buf(newB * S * CROSS_KV_STRIDE * eb, cross, 'compact crossKV dec.0'),
arena.buf(newB * S * CROSS_KV_STRIDE * eb, cross, 'compact crossKV dec.1'),
];
const newLens = arena.buf(newB * 4, GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, 'compact lens');
device.queue.writeBuffer(newLens, 0, lens);
// Step t0-1 tokens land where the next step's embed will read them.
device.queue.writeBuffer(next.tokenRing, (t0 - 1) * newB * 4, lastTok);
const enc = device.createCommandEncoder({ label: `compact B ${state.B} -> ${newB}` });
const oldKvRow = state.kvCapacity * HD * eb;
const newKvRow = next.kvCapacity * HD * eb;
const kvPrefix = Math.min(t0, state.kvCapacity, next.kvCapacity) * HD * eb;
const crossRow = S * CROSS_KV_STRIDE * eb;
const maskRow = BITMASK_WORDS * 4;
for (let i = 0; i < newB; i++) {
const o = liveIdx[i];
for (let l = 0; l < DEC_LAYERS; l++) {
enc.copyBufferToBuffer(
state.kvCacheK[l], o * oldKvRow, next.kvCacheK[l], i * newKvRow, kvPrefix,
);
enc.copyBufferToBuffer(
state.kvCacheV[l], o * oldKvRow, next.kvCacheV[l], i * newKvRow, kvPrefix,
);
enc.copyBufferToBuffer(crossKV[l], o * crossRow, newCrossKV[l], i * crossRow, crossRow);
}
enc.copyBufferToBuffer(state.bitmask, o * maskRow, next.bitmask, i * maskRow, maskRow);
if (state.bitmaskL && next.bitmaskL) {
const maskRowL = state.lmMaskWords * 4;
enc.copyBufferToBuffer(state.bitmaskL, o * maskRowL, next.bitmaskL, i * maskRowL, maskRowL);
}
}
device.queue.submit([enc.finish()]);
return { state: next, crossKV: newCrossKV, lensBuf: newLens, S, arena };
}
// The same-buffer gather is safe only for the stable live-row ordering used by
// the decode loop. Strictly ascending integer indices imply liveIdx[i] >= i,
// so every row moves down and an earlier destination cannot be a later source.
export function validateLiveIdx(liveIdx, oldB) {
if (!Array.isArray(liveIdx)) throw new Error('compact in place: liveIdx must be an array');
const newB = liveIdx.length;
if (newB < 1 || newB > oldB) {
throw new Error(`compact in place: bad live count ${newB} (B=${oldB})`);
}
for (let i = 0; i < newB; i++) {
const value = liveIdx[i];
if (!Number.isInteger(value) || value < 0 || value >= oldB
|| (i > 0 && value <= liveIdx[i - 1])) {
throw new Error(
`compact in place: liveIdx must ascend within [0, ${oldB}) — [${liveIdx}]`,
);
}
}
}
// Pure row geometry for the in-place gather. Strides and copy lengths are u32
// counts because the kernel aliases native f16/f32 data as raw storage words.
export function inPlaceGatherPlan({
t0, S, eb, kvCapacity = DECODE_CAP, lmMaskWords = 0, lmShort = false,
}) {
const HD = HEADS * HEAD_DIM;
const kvStride = (kvCapacity * HD * eb) / 4;
const kvCopy = (Math.min(t0, kvCapacity) * HD * eb) / 4;
const crossStride = (S * 2 * HD * eb) / 4;
const plan = [];
for (let l = 0; l < DEC_LAYERS; l++) {
plan.push({ key: `kvK.${l}`, strideU32: kvStride, copyU32: kvCopy });
plan.push({ key: `kvV.${l}`, strideU32: kvStride, copyU32: kvCopy });
plan.push({ key: `crossKV.${l}`, strideU32: crossStride, copyU32: crossStride });
}
plan.push({ key: 'bitmask', strideU32: BITMASK_WORDS, copyU32: BITMASK_WORDS });
if (lmShort) {
plan.push({ key: 'bitmaskL', strideU32: lmMaskWords, copyU32: lmMaskWords });
}
for (const item of plan) {
if (!Number.isInteger(item.strideU32) || item.strideU32 < 1
|| !Number.isInteger(item.copyU32) || item.copyU32 < 1
|| item.copyU32 > item.strideU32) {
throw new Error(`compact in place: invalid u32 row shape at ${item.key}`);
}
}
return plan;
}
// Gather live rows downward inside the current buffers, preserving every
// GPUBuffer identity and the pinned kernel route. No buffer is allocated or
// destroyed at this boundary, so pooled/immediate bind groups stay valid and
// WebKit sees no old+new resource generation overlap.
export function compactDecodeStateInPlace(
ctx, weights, prev, { liveIdx, t0, lens, lastTok },
) {
const { device } = ctx;
const { state, crossKV, lensBuf, S } = prev;
const oldB = state.B;
const newB = liveIdx.length;
validateLiveIdx(liveIdx, oldB);
if (!state.inPlaceCompact || !state.compactMap || !state.compactParams) {
throw new Error('compact in place: state was not created for in-place compaction');
}
if (!Number.isInteger(t0) || t0 < 1 || t0 > state.maxSteps) {
throw new Error(`compact in place: bad t0=${t0} for maxSteps=${state.maxSteps}`);
}
if (lens?.length !== newB || lastTok?.length !== newB) {
throw new Error(
`compact in place: CPU row data mismatch live=${newB} lens=${lens?.length} token=${lastTok?.length}`,
);
}
const eb = weights.dtype === 'f16' ? 2 : 4;
const plan = inPlaceGatherPlan({
t0, S, eb, kvCapacity: state.kvCapacity,
lmMaskWords: state.lmMaskWords, lmShort: !!state.bitmaskL,
});
const targets = { bitmask: state.bitmask };
for (let l = 0; l < DEC_LAYERS; l++) {
targets[`kvK.${l}`] = state.kvCacheK[l];
targets[`kvV.${l}`] = state.kvCacheV[l];
targets[`crossKV.${l}`] = crossKV[l];
}
if (state.bitmaskL) targets.bitmaskL = state.bitmaskL;
// When finished rows already form a suffix, all live data is in its final
// prefix. Stride-dependent CPU rewrites below are still required.
if (liveIdx.some((source, i) => source !== i)) {
device.queue.writeBuffer(state.compactMap, 0, Uint32Array.from(liveIdx));
const slotU32 = COMPACT_PARAM_STRIDE / 4;
const slab = new Uint32Array(slotU32 * plan.length);
plan.forEach((item, index) => {
slab.set([newB, item.strideU32, item.copyU32, 0], index * slotU32);
});
device.queue.writeBuffer(state.compactParams, 0, slab);
const encoder = device.createCommandEncoder({
label: `compact in place B ${oldB} -> ${newB}`,
});
const pass = encoder.beginComputePass({ label: 'compact gather' });
plan.forEach((item, index) => {
dispatchCompactGather(device, pass, {
data: targets[item.key],
map: state.compactMap,
params: {
buffer: state.compactParams,
offset: index * COMPACT_PARAM_STRIDE,
size: 16,
},
rowStrideU32: item.strideU32,
copyLenU32: item.copyU32,
});
});
pass.end();
device.queue.submit([encoder.finish()]);
}
device.queue.writeBuffer(lensBuf, 0, lens);
device.queue.writeBuffer(state.done, 0, new Uint32Array(newB));
// Embed at step t0 reads only the immediately preceding token row. All
// earlier rows were already collected on CPU; all later rows use newB.
device.queue.writeBuffer(state.tokenRing, (t0 - 1) * newB * 4, lastTok);
state.B = newB;
return prev;
}
// Record one decode step (step index t) into the caller's compute pass — no
// submit here. encRun is runEncoder's result (crossKV/lensBuf/S must match
// state.B's batch). Per-dispatch scratch uniforms are collected and returned;
// the caller destroys them after submitting. (Measured at 4.8% of decode wall
// in M4a — persistent uniforms were never needed; targets met without them.)
//
// Returns {dispatches, scratch}. Logits land in state.logits (f32, bias NOT
// included), the picked token in state.tokenRing[t·B + b].
export function encodeDecodeStep(ctx, weights, encRun, state, t, pass) {
const HD = HEADS * HEAD_DIM; // == D_MODEL
const QKV_N = 3 * HD; // fused q|k|v
const CROSS_KV_STRIDE = 2 * HD; // fused k|v (encoder crossKV layout)
const { device } = ctx;
const { B } = state;
const flags = { t: weights.dtype, immediate: state.immediate };
// Decode projections are tiny-M GEMVs (M = B): the plain gemm kernel is
// occupancy/latency-bound there (Task 17 profile) — route them through the
// GEMV kernel. At large B the GEMV's per-8-row W re-sweep loses to the
// tiled kernel; state.tiledProjKinds (proj_sweep-measured) flips sites.
// Non-gemm kernels and the add_ln/embed sites keep `flags`.
// Subgroup reductions (state.sg): applied at the GEMV, LN(-fused) and
// attention sites — the other kernels have no SG variant and adding the
// flag there would just fork their pipeline-cache entries. On devices
// where autotune detects the Adreno tree-reduce miscompile, sg is FORCED
// on: the sg variants are the only reduction shape that driver compiles
// correctly (2026-07 Android probe rounds).
const sgFlag = state.sg ? { sg: true } : {};
const gemvFlags = { ...flags, gemv: true, ...sgFlag };
const tiledFlags = { ...flags, tiled: true };
const projFlags = (kind) => (state.tiledProjKinds.has(kind) ? tiledFlags : gemvFlags);
const W = (name) => weights.bindingFor(name);
const scratch = [];
let dispatches = 0;
const rec = ({ scratch: s }) => { scratch.push(...s); dispatches++; };
// Hidden-state ping-pong: cur holds the current hidden state; each add_ln
// writes into the OTHER buffer and swaps.
let cur = 0;
const addLn = (prefix) => {
rec(dispatchAddLn(device, pass, {
x: state.y, r: state.hidden[cur],
gamma: W(`${prefix}.weight`), beta: W(`${prefix}.bias`),
y: state.hidden[1 - cur], rows: B, flags: { ...flags, ...sgFlag },
}));
cur = 1 - cur;
};
// Fused projection + residual + LN (state.lnFused): one gemm_row_ln
// dispatch replaces the gemm→state.y + addLn pair at an LN-terminated site.
const projLnFused = (x, wPrefix, lnPrefix, K) => {
rec(dispatchGemmRowLn(device, pass, {
x, w: W(`${wPrefix}.weight`), b: W(`${wPrefix}.bias`),
r: state.hidden[cur],
gamma: W(`${lnPrefix}.weight`), beta: W(`${lnPrefix}.bias`),
y: state.hidden[1 - cur], M: B, K, N: D_MODEL, flags: { ...flags, ...sgFlag },
}));
cur = 1 - cur;
};
// Megakernel route (state.mega): one dispatch per layer replaces the embed
// + 8-dispatch layer chain — the step is mega L0 (embed folded) → mega L1
// → lm_head → argmax. The final hidden state lands in hidden[0] (cur = 0).
if (state.mega) {
for (let l = 0; l < DEC_LAYERS; l++) {
rec(dispatchDecoderMega(device, pass, {
weights, layer: l, embed: l === 0,
ring: state.tokenRing, kCache: state.kvCacheK[l], vCache: state.kvCacheV[l],
crossKV: encRun.crossKV[l], lens: encRun.lensBuf, x: state.hidden[0],
B, t, S: encRun.S, kvCapacity: state.kvCapacity,
flags: { ...flags, ...sgFlag },
}));
}
} else {
// Embedding: t=0 → DECODER_START, else tokenRing[(t-1)·B + b]; pos = t.
rec(dispatchEmbed(device, pass, {
ids: state.tokenRing, table: W('shared.weight'), posEmbed: W('pos_embed'),
y: state.hidden[cur], mode: 'decode', nRows: B, step: t, batch: B, flags,
}));
for (let l = 0; l < DEC_LAYERS; l++) {
const p = (name) => `dec.${l}.${name}`;
// Attention-side projections: 'wt' swaps in the transposed copies on the
// GEMV WT path (one W-tile per MT=8 rows instead of NWT's per-row full-W
// re-read); the storeKV epilogue exists on both layouts. With split-K
// (state.projSK, wt-only) the site becomes a tiled GEMM storing raw
// partials plus a gemm_reduce fold — self_qkv's cache scatter rides the
// reduce epilogue instead of the GEMV one.
const projDispatch = (kind, { x, y, b, K, N, storeKV = null }) => {
const sk = kind === 'self_qkv' ? state.projSK.qkv : state.projSK.out;
if (sk) {
rec(dispatchGemm(device, pass, {
x, w: W(p(`${kind}.wt`)), y: null, M: B, K, N,
splitK: { parts: state.skParts, sk },
flags: { ...flags, tiled: true, tm8: true, wt: true },
}));
rec(dispatchGemmReduce(device, pass, {
parts: state.skParts, b, y, M: B, N, nz: splitKParts(K, sk).nz,
storeKV, flags,
}));
return;
}
const args = state.projMode === 'wt'
? { w: W(p(`${kind}.wt`)), flags: { ...flags, gemv: true, wt: true, ...sgFlag } }
: { w: W(p(`${kind}.weight`)), flags: projFlags(kind) };
rec(dispatchGemm(device, pass, { x, w: args.w, b, y, M: B, K, N, storeKV, flags: args.flags }));
};
// Self-attention block: x = LN1(x + SelfAttn(x)).
projDispatch('self_qkv', {
x: state.hidden[cur], b: W(p('self_qkv.bias')), y: state.qkvOut,
K: D_MODEL, N: QKV_N,
storeKV: {
kCache: state.kvCacheK[l], vCache: state.kvCacheV[l],
t, Lmax: state.kvCapacity,
},
});
rec(dispatchAttention(device, pass, {
q: state.qkvOut, k: state.kvCacheK[l], v: state.kvCacheV[l], y: state.attnOut,
B, M: 1, L: state.kvCapacity, lenMode: 0, step: t,
qStride: QKV_N, flags: { ...flags, ...sgFlag }, // K/V: default compact [B, L, H, D] strides
}));
if (state.lnFused) {
projLnFused(state.attnOut, p('self_out'), p('ln1'), HD);
} else {
projDispatch('self_out', {
x: state.attnOut, b: W(p('self_out.bias')), y: state.y,
K: HD, N: D_MODEL,
});
addLn(p('ln1'));
}
// Cross-attention block: x = LN2(x + CrossAttn(x)). K/V read the fused
// encoder crossKV [B·S, 2·HD] (k at 0, v at HD) via strides.
projDispatch('cross_q', {
x: state.hidden[cur], b: W(p('cross_q.bias')), y: state.crossQOut,
K: D_MODEL, N: HD,
});
rec(dispatchAttention(device, pass, {
q: state.crossQOut, k: encRun.crossKV[l], v: encRun.crossKV[l],
lens: encRun.lensBuf, y: state.attnOut,
B, M: 1, L: encRun.S, lenMode: 1,
kvStride: CROSS_KV_STRIDE, kOff: 0, vOff: HD, flags: { ...flags, ...sgFlag },
}));
if (state.lnFused) {
projLnFused(state.attnOut, p('cross_out'), p('ln2'), HD);
} else {
projDispatch('cross_out', {
x: state.attnOut, b: W(p('cross_out.bias')), y: state.y,
K: HD, N: D_MODEL,
});
addLn(p('ln2'));
}
// FFN block: x = LN3(x + FFN(x)). ffnMode 'wt'/'q8' swap in the [N,K]
// tensors (f16 transposed / int8 packed) and run the WT layout on either
// kernel (GEMV WT below the tiled threshold, tiled above) — one W-tile
// per MT=8 rows instead of NWT's per-row full-W re-stream.
const wtKernel = (kind) => (state.tiledProjKinds.has(kind)
? { tiled: true, tm8: true } : { gemv: true, ...sgFlag });
const ffnArgs = (kind) => {
if (state.ffnMode === 'q8') {
return {
w: W(p(`${kind}.q8`)), scales: W(p(`${kind}.scales`)),
flags: { ...flags, wt: true, wq8: true, ...wtKernel(kind), ...state.ffnFlags },
};
}
if (state.ffnMode === 'wt') {
return { w: W(p(`${kind}.wt`)), flags: { ...flags, wt: true, ...wtKernel(kind), ...state.ffnFlags } };
}
return { w: W(p(`${kind}.weight`)), flags: projFlags(kind) };
};
// Split-K FFN site (state.ffnSK, tiled route only): the GEMM stores raw
// f32 partials over grid.z K-partitions and gemm_reduce folds them with
// the bias/SiLU epilogue — 14 starved workgroups become 14·nz.
const fcSplit = (kind, { x, y, b, K, N, silu = false }) => {
const sk = state.ffnSK[kind];
const args = ffnArgs(kind);
if (!sk) {
rec(dispatchGemm(device, pass, {
x, w: args.w, scales: args.scales, b, y, M: B, K, N,
flags: silu ? { ...args.flags, silu: true } : args.flags,
}));
return;
}
rec(dispatchGemm(device, pass, {
x, w: args.w, y: null, M: B, K, N,
splitK: { parts: state.skParts, sk }, flags: args.flags,
}));
rec(dispatchGemmReduce(device, pass, {
parts: state.skParts, b, y, M: B, N,
nz: splitKParts(K, sk, args.flags.bkk ?? 16).nz,
flags: { ...flags, silu },
}));
};
fcSplit('fc1', {
x: state.hidden[cur], b: W(p('fc1.bias')), y: state.ffnTmp,
K: D_MODEL, N: FFN, silu: true,
});
// fc2 stays unfused in 'q8' mode — the fused kernel reads float weights.
if (state.lnFused && state.ffnMode !== 'q8') {
projLnFused(state.ffnTmp, p('fc2'), p('ln3'), FFN);
} else {
fcSplit('fc2', {
x: state.ffnTmp, b: W(p('fc2.bias')), y: state.y,
K: FFN, N: D_MODEL,
});
addLn(p('ln3'));
}
}
} // end non-mega route
// LM head: logits = x @ shared.weightᵀ ([24000,448] row-major → wt), f32
// out, NO bias — final_logits_bias is added inside the argmax kernel.
// Kernel per state.lmHeadTiled: GEMV re-sweeps the 21.5MB W per 8 rows —
// batch-linear past B≈32 — while the tiled kernel reads W exactly once.
// Fused (state.lmHeadFused): the tiled epilogue applies bias + penalty and
// emits per-tile argmax partials instead of storing logits; argmax_reduce
// finishes the row (same dispatch count, ~37MB/step less traffic at B=128).
// Shortlist (state.lmShort): the q8 tensor holds only emittable rows, so N
// shrinks to lmN, the bias/seen bindings switch to their local-space twins,
// and the argmax epilogues translate the winner back to a vocab id.
const lmShortArgs = state.lmShort ? {
n: state.lmN, maskWords: state.lmMaskWords,
idmap: W('lm_head.idmap'), gmask: state.bitmask,
} : null;
rec(dispatchGemm(device, pass, {
x: state.hidden[cur],
w: state.lmHeadQ8 ? W('lm_head.q8') : W('shared.weight'),
...(state.lmHeadQ8 ? { scales: W('lm_head.scales') } : {}),
y: state.logits,
...(state.lmHeadFused ? {
fusedArgmax: state.lmShort ? {
partials: state.argmaxPartials, lbias: W('lm_head.sbias'),
seen: state.bitmaskL, maskWords: state.lmMaskWords,
} : {
partials: state.argmaxPartials, lbias: W('final_logits_bias'), seen: state.bitmask,
},
} : {}),
M: B, K: D_MODEL, N: state.lmN,
flags: {
...flags, outT: 'f32', wt: true,
...(state.lmHeadTiled
// q8 defaults to the 8×4 subtile (b16 665→639µs; a wash above).
? { tiled: true, ...(state.lmHeadQ8 ? { wq8: true, tm8: true } : {}), ...state.lmHeadFlags }
: { gemv: true, ...sgFlag }),
},
}));
if (state.lmHeadFused) {
rec(dispatchArgmaxReduce(device, pass, {
partials: state.argmaxPartials,
bitmask: state.lmShort ? state.bitmaskL : state.bitmask,
done: state.done, tokens: state.tokenRing, B, t, NT: state.lmNT,
short: lmShortArgs,
flags: { immediate: state.immediate },
}));
} else {
rec(dispatchArgmaxPenalty(device, pass, {
logits: state.logits,
bias: W(state.lmShort ? 'lm_head.sbias' : 'final_logits_bias'),
bitmask: state.lmShort ? state.bitmaskL : state.bitmask,
done: state.done, tokens: state.tokenRing,
B, t,
short: lmShortArgs,
flags: { immediate: state.immediate },
}));
}
return { dispatches, scratch };
}
|