Skip to content

[Triton/Gluon] [gfx942] Support fp8 KV caches and fp8 dots in sparse_mla_fwd on gfx942 - #6199

Open
jin-amd wants to merge 4 commits into
ROCm:mainfrom
jin-amd:gfx942-sparse-mla-fp8
Open

jin-amd wants to merge 4 commits into
ROCm:mainfrom
jin-amd:gfx942-sparse-mla-fp8

Conversation

@jin-amd

@jin-amd jin-amd commented Oct 7, 2026

Copy link
Copy Markdown
Contributor

Motivation

#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:

  1. 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.
  2. 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.

@jin-amd
jin-amd requested review from a team and a balanced review from Copilot October 7, 2026 07:57
@github-actions

github-actions Bot commented Oct 7, 2026

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every PR:

  • ✅ Pre-checks (submodule verification, code formatting)
  • ✅ Aiter op tests (gfx942 + gfx950)
  • ✅ Triton tests on MI35X (only when aiter/ops/triton/** or related paths are changed)

Extended tests (opt-in via labels):

Label Tests
ci:gfx1250-ffm-triton Run the five-shard gfx1250 FFM Triton test suite
ci:triton-300x Run an additional Triton test job on MI300X in PRs (added automatically when gfx942 configs change); main branch always runs both MI35X and MI300X
ci:triton-355 Run the full Triton test suite on MI35X, not only the tests the change affects
multigpu Aiter multi-GPU tests on the 8-GPU runner
ci:sglang SGLang integration tests: DeepSeek-R1-MXFP4 accuracy, Qwen 3.5 accuracy
ci:atom ATOM benchmark: DeepSeek-R1-0528, GPT-OSS-120B
ci:atom_full ATOM accuracy suite for PR and main models from ATOM models_accuracy.json
ci:vllm vLLM benchmark: GPT-OSS-120B, DeepSeek-R1-0528, Kimi-K2.5
ci:all 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.

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot review overview

🟡 Changes recommended

The benchmark still skips the new gfx942 path, and kernel trace representations omit its fp8 encoding specialization.

Review effort: Balanced
Findings: 1 Medium severity · 1 Low severity

Open (2)
What changed in this PR

Adds gfx942-native fp8 cache decoding and fp8 dot support to sparse MLA while preserving gfx950 behavior.

Changes:

  • Selects architecture-native fp8 encoding and MFMA layouts.
  • Adds gfx942 fp8 launch configuration and LDS validation.
  • Expands correctness and encoding-rejection tests.
File Description
aiter/​ops/​triton/​attention/​sparse_mla.py Enables and dispatches gfx942 fp8 paths.
aiter/​ops/​triton/​_gluon_kernels/​gfx950/​attention/​sparse_mla.py Implements fnuz decoding, quantization, and scaling.
aiter/​ops/​triton/​configs/​gfx942/​gluon/​attention/​sparse_mla/​DEFAULT.json Adds the gfx942 fp8 launch configuration.
op_tests/​triton_tests/​attention/​test_sparse_mla.py Covers fp8 execution, validation, and LDS limits.

💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment thread aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py
Comment thread aiter/ops/triton/attention/sparse_mla.py
@zufayu
zufayu requested review from a team and Dewei-Wang-sh October 8, 2026 02:08
Copilot AI balanced review requested due to automatic review settings October 8, 2026 08:25

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

Flat-cache validation permits unsupported byte dtypes, and the FP8 probability-scaling invariant lacks focused regression coverage.

1 open finding
2 resolved since last review
Previously missed (1)

In code that hasn't changed since last review

Medium severity Add peaky-attention regression coverage for FP8 probability scaling

aiter/​ops/​triton/​_gluon_kernels/​gfx950/​attention/​sparse_mla.py:341

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.

Comment thread aiter/ops/triton/attention/sparse_mla.py Outdated
Copilot AI balanced review requested due to automatic review settings October 8, 2026 08:45

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔵 Needs a closer look

The architecture-sensitive numerical kernel changes were validated on gfx942 but still lack gfx950 execution coverage.

0 open findings

1 resolved since last review

🧠 Review effort: Balanced


Give feedback about Copilot approvals in this survey to enter a drawing for a $150 gift card.

jin-amd and others added 4 commits October 8, 2026 10:13
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>
@jin-amd
jin-amd force-pushed the gfx942-sparse-mla-fp8 branch from b140a94 to 013f408 Compare October 8, 2026 10:23
Copilot AI balanced review requested due to automatic review settings October 8, 2026 10:23

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔵 Needs a closer look

The architecture-specific kernel changes require final human review and gfx950 CI validation.

0 open findings

1 resolved since last review

🧠 Review effort: Balanced


Give feedback about Copilot approvals in this survey to enter a drawing for a $150 gift card.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants