Batched MoE prefill is bit-identical to a token-by-token replay of the decode path on 23 of 23 test prompts and reaches 246 to 247 tok/s at 32k context against about 67 for the replay, while two variants that measured faster still stay off by default because they reorder floating-point sums and fail that same identity check.
Prefilling a prompt through a mixture-of-experts model one token at a time means running the full decode kernel path once per token, even though every token in the prompt is already known and none of the router decisions depend on tokens that haven't been generated yet. That's wasteful in an obvious way, and on 2026-09-19 this project landed a batched version, BARO_PREFILL=1, that processes prefill in chunks instead. The chunked path passed identity on 23 of 23 test prompts and reached 246 to 247 tokens per second at 32k context, against roughly 67 for the token-by-token replay of the decode path on the same prompts. Two other variants tried in the same lane were faster still and are not in the shipped default, because they don't compute the same sums the replay does, and identity is the gate here, not speed.
Every MoE phase in the batched path (kernels/moe_rows.mojo) is written as a sibling of the existing m=1 decode kernel, sharing the same per-element dot-product helper and the same summation order, but taking many rows instead of one. The gate/up and down expert kernels sort (token, k) pairs by expert id on the device first, so that the pairs routed to the same expert sit together, compute their dot products per pair, and then sum the top-k contributions afterward in k order, which is the same order the m=1 kernel would produce them in one at a time. kernels/test_moe_rows.mojo checks 16 outputs bit-exact against the m=1 kernel directly, which is the actual gate, not a token-level comparison: a byte match against the single-row kernel is what proves the batching didn't change the arithmetic, rather than just producing plausible-looking text.
The SSM state gets the same treatment through amar_ssm_delta_rows, a one-launch delta scan written in the decode step's own per-column order and checked bit-exact over 37 rows through its 9-slot ring in kernels/test_ssm_rows.mojo. Attention inside a chunk uses amar_attn_decode applied per row, the exact decode kernel, not an approximation.
None of that would matter if the experts a chunk needs didn't fit in VRAM, and a full prefill chunk touches nearly all 256 experts of a layer at once, unlike decode's top-8. So prefill mode bypasses the 64-slot decode cache entirely: one layer's experts are staged from the pinned host store into one of two VRAM slots, 0.52 GB each, with layer L+1's copy running on a second device stream while layer L computes, ordered by device events rather than a host-side sync per row. A per-row prepare() call, which would mean a host synchronization for every row of every layer, was tried and rejected; so was in-kernel zero-copy for this path, because a full-layer chunk means PCIe traffic once per token, expert, and row rather than a bounded set of misses.
The comparison arm is a replay of the decode kernel path, one token at a time, over the same prompt: about 67 tok/s regardless of context length, because replay does the same fixed amount of per-token work whether the context is 1k or 32k tokens long.
| context | replay | batched, final build | speedup |
|---|---|---|---|
| 1k | ~67 tok/s | 524 to 531 tok/s | 7.9x |
| 8k | ~67 tok/s | 428 to 432 tok/s | 6.4x |
| 32k | ~67 tok/s | 246 to 247 tok/s | 3.7x |
These are engine-clock timings from the final build (c5b664a), two repeats, spread under 1.2%, not the identity gate itself; the identity gate that governs whether this path is trustworthy at all ran separately, 23 of 23 prompts equal over 64 tokens in both tier mode and resident mode, 12 of those 23 exercising the prefill path directly.
The predictions were frozen (70de96a) before these numbers existed: tier mode at 1k was predicted 350 to 900 tok/s and landed at 524, inside the band; at 8k predicted 330 to 850, landed at 428, inside the band; at 32k predicted 250 to 700, landed at 246 to 247, missing the low end of the band by less than 2%. The rule that every cell would beat 3x the replay held everywhere it was tested. What did not hold on the first attempt was identity itself: the initial design item, which used a dense chunked attention kernel and a dense chunked delta scan instead of the row-wise kernels described above, failed identity on 3 of 23 and 4 of 23 prompts respectively, before being replaced with the row-wise versions that check bit-exact against the m=1 kernel.
A resident-mode comparison at 8k and 32k context couldn't be run at all: the full 21 GB pack leaves only about 0.08 GB of MAX's memory pool free at a 10240-token tmax, which is an out-of-memory condition, not a slow result, and at 1k context resident mode measured 51 tok/s against tier mode's 220 in the same round, which means resident mode is not a usable speed arm on this card regardless of context length.
Two further optimizations were tried in the same lane and both measured faster in isolation. BARO_PF_ATT=wmma, a WMMA-based attention kernel for chunks above 256 tokens, read teacher-forced agreement of 61, 63, and 64 out of 64 tokens across three test runs, a mean of 97.92%, against a 99% bar that was set before any of the three runs happened; a same-arm control ran 64 of 64 in the same session, ruling out a broken harness as the explanation. BARO_PF_SSM=chunk, a dense chunked SSM scan, diverges from the decode step's own summation order in 131,837 of 151,552 outputs (about 87%) on a direct comparison, and in the identity gate itself it flips a greedy token at position 4 of a generated sequence. Both stay opt-in, off by default, documented in docs/CAPABILITIES.md as failing identity rather than as merely unverified.
The reason both fail is the same reason the row-wise kernels above were built to pass: floating-point addition is not associative, and a kernel that sums the same set of numbers in a different order, however mathematically equivalent on paper, can and does produce a different final bit pattern, and occasionally a different argmax. WMMA attention and the dense chunk scan are both faster because they process more of a chunk's work in a single wide operation instead of the decode kernel's narrower, ordered accumulation; that's also exactly what makes their sums come out in a different order.
A kernel that is faster because it sums things in a different order has changed what it computes, not just how fast it computes it, and the identity gate exists to catch that difference before a user does.
What counts: - A byte-exact comparison against the m=1 decode kernel, not a token-level comparison, for any new batched or chunked kernel, because a token-level match can hide a numeric difference that simply didn't change the argmax on the prompts tested. - Teacher-forced agreement against a bar set before the run, 99% here, for a kernel that can't pass byte-exact identity by construction, so that "close" has a number attached rather than being a feeling. - A speedup measured against the actual comparison arm being replaced, replay's ~67 tok/s here, not against some other engine's number that answers a different question.
What doesn't: - A prompt set where the failure mode never shows up. Of the 20 prompts named in the original brief for this lane, only 9 reached the minimum length to exercise a chunk boundary at all, and none crossed a chunk; the identity gate used here was extended with 128, 512, and 1024-token prompts specifically so the boundary condition would be tested, and it's the extended set, not the brief's original one, that the 23/23 result is measured against. - Calling 97.92% "basically 99%." It's a number below a preregistered bar, and the post that would round it up is the post that lets a wrong-but-close kernel through next time.
The dense chunked attention and delta scan were both built first, and both had to be discarded once identity failed, which is real engineering time spent on kernels that never shipped. The row-wise replacements that did ship pay a cost of their own: the sort-by-expert step for gate/up and down kernels is extra device work that a dense chunk kernel wouldn't need, and it's part of why 32k context prefill is limited by the exact attention kernel now (133 seconds at 32k, against 84 seconds if the rejected WMMA kernel were used instead) rather than by the expert kernels that were the original bottleneck.
BARO_PREFILL stays opt-in. The lane's own adoption rule needed identity, a 100,000-token needle test, and the full repository gate suite to pass before recommending a default flip, and the needle test only ran to 75,052 tokens rather than the planned 100,000 because of how the test prompt tokenized; a rerun at the full length, and the speed gate itself (bench/moe-prefill-speed.sh, written and ready but not run because the GPU queue was needed elsewhere that session), are both still open.
bench/moe-prefill-identity.sh runs the same 23-prompt identity check this post cites, with BARO_PREFILL=1 for the batched path and =0 for replay; a divergent token anywhere falsifies the identity claim. bench/moe-prefill-speed.sh is written and ready to reproduce the 1k/8k/32k timing table above. bench/moe-prefill-needle.sh reproduces the long-context needle test. All three are checked into bench/ in the public repo.
Every number in this post comes from one RX 7900 XTX. It says nothing about whether the sort-by-expert step, the two-slot layer staging, or the WMMA attention rejection would come out the same way on a card with a different balance of compute to memory bandwidth, or on a model with a different expert count or routing width.
AMD RX 7900 XTX (gfx1100, RDNA3), 24 GB, one card. ROCm 7.2. Mojo 1.0.0 via max[all]==26.5.0. Lane branch lane-moepf, final commit c5b664a, merged to main at ea57b60. Report exchange/lane-MOEPF-report.md; frozen predictions at 70de96a; identity, speed, and needle scripts in bench/, cited above, in amarbaro/mojo-baro.
Commenti
Ancora nessun commento.
Accedi per commentare.