You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
[Triton/Gluon] [gfx942] Support fp8 KV caches and fp8 dots in sparse_mla_fwd on gfx942 - #6199
#5721 enabled sparse_mla_fwd on gfx942 for bf16 caches only. The kernel decoded fp8 as OCP e4m3, while gfx942's native fp8, which is also what vLLM's fp8 KV cache stores there, is e4m3fnuz. This PR adds the fp8 cache on gfx942 under both dot precisions. For GLM-5.3-Flash on MI325X, it turns --kv-cache-dtype fp8 from a crash into a speedup (end-to-end numbers below).
Technical details
Two commits:
Read the fp8_scalar cache as e4m3fnuz on gfx942. An FP8_FNUZ constexpr picks the fp8 type the dequant reads, and the wrapper sets it from the arch. gfx942 now takes the per-tensor fp8_scalar cache under bf16 dots. A typed cache in the other arch's encoding is refused instead of misread.
Run fp8 dots on gfx942.
The staging, the in-kernel Q quantization (now to 240, fnuz's max) and P use the arch's fp8 type.
The MFMA layouts take version 3 for fnuz operands, because only CDNA3 intrinsics accept them.
P is quantized as p * 128, and the epilogue divides that back out along with the V scale.
gfx942 publishes a _sparse_mla_fp8 launch config at BLOCK_K 64, since one-byte tiles fit twice the bf16 tile's rows.
The LDS budget check models one-byte tiles. It matches the compiled kernels: 34816 B rope-free and 39936 B with 64 rope.
gfx950 keeps its numerics: the OCP type, MFMA version 4, and a P scale of 1.
Results
MI325X, at the GLM-5.3-Flash TP4 shape (16 heads, rope-free 512, top-k 2048). Kernel time per call:
bf16 cache
fp8 cache, bf16 dots
fp8 cache, fp8 dots
Decode, 2 tokens
33.2 µs
30.7 µs
23.7 µs
Decode, 16 tokens
37.4 µs
32.7 µs
25.4 µs
Prefill, 16K tokens
10.55 ms
9.35 ms
5.78 ms
For fp8 dots, BLOCK_K 32 and 8 warps were both slower.
Accuracy. Relative L2 error against f32 attention over the same fp8 cache is 0.2-0.3% with bf16 dots. With fp8 dots it is 3.4% on N(0,1) inputs and 4.5% on peaky attention, almost all of it from quantizing Q. Without the P scale, a long softmax tail flushes to zero in fnuz. In a test with one key scoring 9 above 2047 others that together carry a fifth of the output (split-K off), the error is 14.5% unscaled and 1.7% scaled.
End to end. vLLM serving GLM-5.3-Flash at TP4, 131K tokens in and 1K out, concurrency 2-16. The attend was routed to this kernel by a vLLM build carrying vllm-project/vllm#57134 and a small Gluon dispatch helper that upstream vLLM does not have.
bf16 KV
fp8 KV, bf16 dots
fp8 KV, fp8 dots
Output tok/s (geomean)
155.1
157.1
165.7
TPOT
28.1 ms
27.5 ms
26.4 ms
gsm8k, 5-shot strict (±0.005)
0.974
0.970
0.973
Needle in a haystack, 4K-128K
60/60
60/60
60/60
KV cache capacity
12.2M tokens
23.1M
23.1M
Test plan
op_tests/triton_tests/attention/test_sparse_mla.py on gfx942 (MI325X): 77 passed, 1 skipped (fp8_dsv32_mla, which is gfx950-only). The fp8-dot cases now run on gfx942, and new tests cover the foreign-encoding rejection and the fp8-tile LDS budget.
I have no gfx950 hardware, so these changes are not run there. gfx950 keeps its fp8 type, MFMA version and a P scale of 1, but it still needs CI coverage.
All standard extended tests (excludes ci:atom_full)
Only add ci:atom_full for FlyDSL or Triton upgrades.
Add labels via the sidebar or gh pr edit 6199 --add-label <label>
One backend per PR:
A PR changes one kernel backend: [Triton/Gluon] (Triton and Gluon count as one), [HIP], [ASM], [CK], [OPUS] or [FlyDSL]. If the title ends up with two backend tags, split the PR -- as stacked pull requests when one part cannot merge without the other.
PR title tags & labels:
Component tags ([Triton/Gluon], [HIP], [CK], [ASM], ...) are added to the PR title and as PR labels automatically from the changed files and re-synced on every push — change-type tags like [fix]/[Perf], op tags like [MLA], and human labels (ci:*) are left untouched. Add the no-auto-title label to stop the title rewrites; labels stay in sync either way.
The existing fp8 accuracy matrix uses small random Q/K values, so its per-tile softmax probabilities stay close to one and it does not exercise the long-tail underflow that this new P_SCALE is intended to prevent. A regression to P_SCALE = 1 could therefore pass the suite while substantially changing peaky-attention output. Add the described one-dominant-key/2047-tail case (or an equivalent focused case) so the scaling and epilogue compensation are asserted.
🧠 Review effort: Balanced
Give feedback about Copilot approvals in this survey to enter a drawing for a $150 gift card.
The kernel decoded every fp8 byte as OCP e4m3, so gfx942, whose native fp8 (and
vLLM's fp8 KV cache) is e4m3fnuz, took bf16 caches only. Add an FP8_FNUZ
constexpr that picks the fp8 element type the dequant helpers read, and set it
from the arch in the wrapper. gfx942 now takes the per-tensor fp8_scalar cache
under bf16 dots; fp8 q, fp8 dots and fp8_dsv32_mla stay gfx950-only.
A typed cache in the other arch's encoding is refused instead of misread; a
uint8 view is taken to be the arch's own fp8.
gfx942 (MI325X), GLM-5.3-Flash shape (16 heads, rope-free 512, top-k 2048),
fp8 vs bf16 cache: decode 31.0 vs 33.2 us at 2 tokens, 33.1 vs 35.5 us at 16;
16K-token prefill 9.34 vs 10.54 ms. test_sparse_mla.py: 61 passed, 13 skipped
(fp8 dots, fp8_dsv32_mla).
Co-authored-by: Cursor <cursoragent@cursor.com>
dot_precision="fp8" was gfx950-only because the kernel fed the matrix core OCP
e4m3. The fp8 staging, the in-kernel Q quantization (now to the arch's fp8
range, 240 on fnuz) and P now use the arch's fp8 type, and the MFMA layouts take
version 3 for fnuz operands, which only have CDNA3 intrinsics.
P is quantized as p * 128 (exact; p <= 1 stays under 240) and the epilogue
divides it out with the V-side scale. Unscaled, softmax tails below fnuz's
smallest subnormal flush to zero: with one key scoring 9 above 2047 others that
carry a fifth of the output, rel-L2 error is 14.5% unscaled and 1.7% scaled
(split-K off). gfx950 keeps its numerics (P scale 1, OCP, version 4).
gfx942 publishes a _sparse_mla_fp8 config at BLOCK_K 64: one-byte tiles fit
twice the bf16 tile's rows. The LDS model takes the element size, and matches
the compiled kernels (bf16 34304/34816/38912/39424 B, fp8 34816/39936 B).
MI325X, GLM-5.3-Flash shape (16 heads, rope-free 512, top-k 2048), fp8 cache
under fp8 vs bf16 dots, and the bf16 cache: decode at 2 tokens 23.7 / 30.7 /
33.2 us, at 16 tokens 25.4 / 32.7 / 37.4 us; 16K-token prefill 5.78 / 9.35 /
10.55 ms. BLOCK_K 32 and 8 warps were slower.
Rel-L2 vs f32 attention over the same fp8 cache, fp8 dots: 3.4% on N(0,1)
inputs, 4.5% on peaky ones (bf16 dots: 0.2-0.3%).
test_sparse_mla.py on gfx942: 77 passed, 1 skipped (fp8_dsv32_mla).
Co-authored-by: Cursor <cursoragent@cursor.com>
…ench fp8 dots on gfx942
_sparse_mla_repr left out FP8_MFMA and FP8_FNUZ, so bf16 and fp8 dots over the
same fp8_scalar cache got the same trace name whenever they ran at the same
BLOCK_K, which gfx950 does for prefill-sized grids and with the async path off.
Add both keys.
bench_sparse_mla.py gated its fp8-dot series on FP8_ARCHS and pre-quantized Q,
so gfx942 skipped the path with a stale note that the kernel reads OCP e4m3.
Gate on FP8_SCALAR_ARCHS, pre-quantize Q only on FP8_ARCHS (gfx942 quantizes
bf16 Q inside the kernel), and count Q at its own element size in the
bandwidth metric.
MI325X decode at 1/8/64 sequences (16 heads, context 8192, top-k 2048): bf16
dots 39.5/43.6/103.8 us, fp8 dots 25.0/27.0/72.3 us. test_sparse_mla.py: 77
passed, 1 skipped.
Co-authored-by: Cursor <cursoragent@cursor.com>
… softmax tail
_classify_flat took any one-byte dtype other than the other arch's e4m3, and
the kernel then read the bytes as this arch's e4m3, so a typed e5m2 (or int8,
bool) cache was decoded with the wrong encoding. Accept only the native e4m3
dtype or a uint8 view of it, as the sparse_mla_fwd docstring says.
The fp8 accuracy cases use mild random inputs, where every p stays near 1, and
they all pass with P_SCALE forced to 1. test_fp8_dots_keep_the_softmax_tail
gives each query a first key scoring 9 above 2047 others that carry a fifth of
the output at p ~ 1e-4, with split-K off. On MI325X the fp8-dot output is 1.8%
off its f32 reference (max-rel); with P_SCALE = 1 the tail's share comes out
zero and the error is 100%. It runs on gfx942, the arch with a P_SCALE.
test_sparse_mla.py on gfx942: 82 passed, 1 skipped.
Co-authored-by: Cursor <cursoragent@cursor.com>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Motivation
#5721 enabled
sparse_mla_fwdon gfx942 for bf16 caches only. The kernel decoded fp8 as OCP e4m3, while gfx942's native fp8, which is also what vLLM's fp8 KV cache stores there, is e4m3fnuz. This PR adds the fp8 cache on gfx942 under both dot precisions. For GLM-5.3-Flash on MI325X, it turns--kv-cache-dtype fp8from a crash into a speedup (end-to-end numbers below).Technical details
Two commits:
FP8_FNUZconstexpr picks the fp8 type the dequant reads, and the wrapper sets it from the arch. gfx942 now takes the per-tensorfp8_scalarcache under bf16 dots. A typed cache in the other arch's encoding is refused instead of misread.p * 128, and the epilogue divides that back out along with the V scale._sparse_mla_fp8launch config at BLOCK_K 64, since one-byte tiles fit twice the bf16 tile's rows.gfx950 keeps its numerics: the OCP type, MFMA version 4, and a P scale of 1.
Results
MI325X, at the GLM-5.3-Flash TP4 shape (16 heads, rope-free 512, top-k 2048). Kernel time per call:
For fp8 dots, BLOCK_K 32 and 8 warps were both slower.
Accuracy. Relative L2 error against f32 attention over the same fp8 cache is 0.2-0.3% with bf16 dots. With fp8 dots it is 3.4% on N(0,1) inputs and 4.5% on peaky attention, almost all of it from quantizing Q. Without the P scale, a long softmax tail flushes to zero in fnuz. In a test with one key scoring 9 above 2047 others that together carry a fifth of the output (split-K off), the error is 14.5% unscaled and 1.7% scaled.
End to end. vLLM serving GLM-5.3-Flash at TP4, 131K tokens in and 1K out, concurrency 2-16. The attend was routed to this kernel by a vLLM build carrying vllm-project/vllm#57134 and a small Gluon dispatch helper that upstream vLLM does not have.
Test plan
op_tests/triton_tests/attention/test_sparse_mla.pyon gfx942 (MI325X): 77 passed, 1 skipped (fp8_dsv32_mla, which is gfx950-only). The fp8-dot cases now run on gfx942, and new tests cover the foreign-encoding rejection and the fp8-tile LDS budget.I have no gfx950 hardware, so these changes are not run there. gfx950 keeps its fp8 type, MFMA version and a P scale of 1, but it still needs CI coverage.