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] Enable sparse_mla_fwd on gfx942 - #5721
Re-land of #5539, which was auto-closed on Sep 18 when its base branch
(cagri/sparse_pa_optimizations) was deleted three seconds after #4919 merged.
Nothing was rejected there — the one review comment is addressed and folded into
this commit. Now targets main directly; the commit cherry-picks onto it with no
conflicts and no file in it has changed on main since #4919 landed.
The Gluon kernel needs no change to run there — the gl.amd.cdna4 intrinsics it
uses (async_copy, mfma, buffer_load/buffer_store) all lower fine on
gfx942 under Triton 3.7.1. The arch gate was the only thing in the way, plus two
launch-config values sized against gfx950's LDS budget.
Why the tile has to change
gfx942 (CDNA3) has 64 KB of LDS per workgroup against gfx950's 160 KB. The
bf16 KV tile alone is BLOCK_K * (kv_lora_rank + LDS_PAD) * 2 B, so the gfx950
tile of 64 asks for 68 608 B and will not launch:
OutOfResources: out of resource: shared memory, Required: 68608, Hardware limit: 65536
BLOCK_K=32 is the largest that fits.
num_warps is decoupled from BLOCK_K in the same change. It was derived as BLOCK_K // 16, so capping the tile for LDS would have dropped num_warps from 4
to 2 as a side effect rather than as a tuning decision. Keeping it at 4 is
worth 1.19× → 1.55× on the prefill shapes, so it is worth making explicit.
fp8 is rejected on gfx942
The kernel reads every fp8 byte (q, the cache, and the operands of the fp8 dots) as OCP e4m3, which is gfx950's native fp8. On gfx942 the native fp8 is fnuz: vLLM's
fp8 KV cache and aiter's quantizers both write float8_e4m3fnuz there. The encodings
share a bit layout with exponent bias 7 vs 8, so the same bytes decode 2x apart; 0x80
is -0 in OCP but NaN in fnuz, and fnuz's saturated ±240 (0x7F/0xFF) are NaN in OCP.
It first showed up as dot_precision="fp8" being silently wrong (rel-err 7.5e-1):
the CDNA3 fp8 MFMA decodes fnuz and was fed OCP. The fp8 caches have the same problem,
and a cache usually arrives as a uint8 view, so the wrapper cannot tell the encoding
from the tensor. gfx942 therefore takes bf16 q and a bf16 cache only and raises
otherwise, naming the reason. That is what the GLM-5.3-Flash path uses: vLLM keeps this
model's MLA cache in bf16 and dispatches only bf16 q/kv to the kernel. Reading fnuz
natively (dequant, plus dot_precision="fp8" on CDNA3's fnuz MFMA) is a follow-up.
Two consequences, both in this PR:
The bench's fp8-dot series is built from FP8_ARCHS. With the arch gate widened it
would otherwise reach that series on gfx942 and raise on its first point; triton.testing.perf_report does not catch it, so the run died before printing a
table. It now prints a note so the omission is visible rather than silent. (This was
frida-andersson's catch on [Triton/Gluon] [gfx942] Enable sparse_mla_fwd on gfx942 #5539.)
The tests keep main's arch-native fp8 fixtures and skip their fp8 cases on gfx942. A
GPU-free test covers the gate on every arch, and a wrapper-level test feeds gfx942's
own fnuz cache behind a uint8 view and expects the error.
gfx950 is unchanged
gfx950's JSON keeps 64 and 4, and gfx950 is in FP8_ARCHS and PACKED_ARCHS.
Validation
On MI325X (gfx942), TP4, GLM-5.3-Flash at 131 k context. This matters on that
model because its NoPE MLA has qk_rope_head_dim = 0, which has no AITER path
today, so vLLM falls back to a vendored Triton gather+dot — the single largest
kernel in the model in both phases.
Kernel bucket, from a real torch profile at concurrency 12 (rank 0):
phase
bucket
Triton (vLLM)
this kernel
speedup
decode
per step
3.099 ms
0.400 ms
7.74×
prefill
per call
10 988 µs
10 330 µs
1.06×
End-to-end, 131 k in / 1024 out, same image both arms with only the dispatch
switched:
conc
out tok/s before → after
delta
2
80.92 → 91.36
+12.90%
4
110.96 → 121.09
+9.13%
8
129.61 → 137.74
+6.27%
12
141.29 → 147.94
+4.71%
16
143.24 → 149.89
+4.64%
Won 7/7 rows, mean +7.14% output throughput, zero failed requests. gsm8k
(5-shot, full 1319 questions) is neutral: 0.9719 vs 0.9712 strict-match, against
a ±0.0046 standard error.
The gain being largest at low concurrency is the signature of this kernel rather
than noise: the operator is batch-independent, so replacing it takes a roughly
constant ~2.7 ms off every decode step, which is 16.6% of TPOT at concurrency 2
but 4.9% by concurrency 16.
Note on prefill
The 1.06× prefill figure is the honest one, measured in a profile. A
microbenchmark using randperm top-k indices reports 1.55× for the same shape,
but that flatters it — random indices give the Triton gather far worse locality
than production does (11.0 ms real vs 15.7 ms synthetic for the same shape),
while this kernel gathers whole 512 B rows and is locality-insensitive. Prefill microbenchmarks on this operator should not be trusted without a
profile.
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 5721 --add-label <label>
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 opt this PR out.
Enables sparse_mla_fwd on gfx942 with LDS-safe tiling, explicit FP8 handling, and updated tests and benchmarks.
Changes:
Adds gfx942 dispatch and launch configuration.
Rejects unsupported gfx942 FP8 matrix-core dots.
Updates cache tests and benchmark gating.
File
Summary
Review findings
aiter/ops/triton/attention/sparse_mla.py
Adds gfx942 support and launch adjustments.
Critical (1 vote): Arbitrary geometry can exceed gfx942’s LDS limit; bound it or derive tile size from the footprint. Moderate (1 vote): Some advertised FP8 cache formats still crash; route them safely or reject them explicitly. Moderate (2 votes): Add direct coverage for gfx942 FP8-dot rejection. Moderate (1 vote): Source waves_per_eu from shared configuration. Nit (3 votes): Move architecture launch parameters into the shared tuned configuration. Nit (1 vote): Document the gfx942 FP8 restriction and bf16 alternative.
[verified] On gfx942, _skip_unless_supported("fp8") skips every fp8-dot parametrization before _run_and_check, so the new _resolve_dot_precision(..., arch) rejection is never asserted. If the guard is accidentally removed or widened, this suite would still pass while exercising the silently-wrong CDNA3 fp8 path. Author must add a gfx942-only negative test that calls the public wrapper with dot_precision="fp8" and asserts the documented ValueError.
Although this new gfx942 LDS guard is the safety net for oversized geometries, the added tests only exercise the default 512/64 shape and never assert the rejection boundary. A regression that removes or miscomputes _check_lds_budget would leave the current tests green and restore the original launch-time out-of-resources failure; add parameterized coverage for a fitting and an over-budget geometry.
The rejecting branch of this new gfx942 LDS guard is not covered by the test suite. The existing GPU tests exercise supported geometries, but no test calls this helper with a footprint over 64 KB, so a regression could turn the documented invalid shapes back into an opaque OutOfResources; add a GPU-independent boundary/failure test for the 64 KB budget.
Benchmark gfx942 bf16 dots with the production fp8 cache
[verified] On gfx942 this leaves only the dots='bf16' series, but the benchmark callback still selects cache=kv (the bf16 cache) for that series at lines 203-204. The production configuration enabled here is bf16 dots over an fp8 cache, so this benchmark does not measure the new arch's dequant/gather path and its numbers can miss its cost. Author must add a labeled bf16-dot/fp8-cache series (using OCP e4m3 bytes, not the gfx942-native fnuz builder) or otherwise make the gfx942 series use the production cache format.
sparse_mla_fwd was gated to gfx950. The Gluon kernel itself needs no change to
run on gfx942 -- the gl.amd.cdna4 intrinsics it uses lower fine there -- so the
gate was the only thing in the way, plus a tile size chosen against gfx950's
LDS budget.
gfx942 has 64 KB of LDS per workgroup against gfx950's 160 KB, so BLOCK_K=64
asks for 68608 B and fails to launch:
OutOfResources: shared memory, Required: 68608, Hardware limit: 65536
BLOCK_K=32 is the largest tile that fits. num_warps is pinned at 4 rather than
following block_k // 16, so the LDS cap does not halve it as a side effect
rather than as a tuning decision; that is worth 1.19x -> 1.55x on the prefill
shapes.
dot_precision="fp8" is rejected on gfx942. It feeds the cache's own code points
to the matrix core, which needs OCP e4m3 MFMA; CDNA3 has fp8 MFMA in the fnuz
encoding only, so those code points come out wrong. It was silently wrong
(rel-err 7.5e-1) rather than faulting, so this is now an explicit error naming
the alternative. dot_precision="bf16" dequantizes the tile into LDS and works
on every arch and every cache format, fp8 caches included.
The bench's fp8-dot series is built from FP8_DOT_ARCHS for the same reason:
with the arch gate widened it would otherwise reach that series on gfx942 and
raise on its first point, and triton's perf_report does not catch it, so the
run died before printing a table.
The test's cache builders go back to an explicit OCP e4m3 rather than the
arch-native fp8 dtype. These records are OCP by definition of the format, and
the kernel's dequant reads OCP code points; on gfx942 the native dtype is
float8_e4m3fnuz, so building with it decoded to garbage (rel-err 1.0) on every
fp8-cache case. _classify_flat already rejects fnuz when a uint8 view is not
hiding the dtype.
gfx950 is unchanged: _arch_block_k -> 64, _arch_num_warps -> 4, lds_limited
False, FP8_DOT_ARCHS contains gfx950, and FP8_DTYPE is OCP e4m3 there either
way, so every value and branch resolves exactly as before.
Validated on MI325X (gfx942), TP4, GLM-5.3-Flash at 131k context, where this
path serves the rope-free NoPE MLA that has no AITER kernel today:
op_tests test_sparse_mla 32 passed, 12 skipped (fp8 dots)
bench_sparse_mla runs clean, bf16 series only
sparse-MLA decode bucket 3.099 -> 0.400 ms/step 7.74x
sparse-MLA prefill 10988 -> 10330 us/call 1.06x
end-to-end sweep +7.14% mean output tok/s, 7/7 rows won
gsm8k 0.9719 vs 0.9712 strict-match (neutral)
Co-authored-by: Cursor <cursoragent@cursor.com>
…s LDS
_check_geometry has never had an upper bound on kv_lora_rank, and gfx942 already
takes the smaller of the two tiles this wrapper selects, so a latent too wide to
fit has nowhere left to go. It surfaced as Triton's OutOfResources at launch
rather than as something a caller can act on.
Measured on MI325X by reading the required byte count back out of
OutOfResources, the gfx942 footprint is exactly
BLOCK_K * (kv_lora_rank + 8) * 2
+ BLOCK_K * (qk_rope_head_dim + 8) * 2 (separate rope only)
+ 32 * BLOCK_K
on every point at BLOCK_K 32 and 64: 67072 B at kv_lora_rank=1024 rope-free,
71680 with a 64-wide rope, 132608 at 2048. Those constants hold only for bf16
tiles with the async path off, which is every gfx942 launch, since fp8 dots are
rejected there and lds_limited forces ASYNC_LDS off.
The check raises a ValueError naming the geometry, its footprint and the budget.
gfx950 is not in the table, so it is left to the launcher exactly as before.
test_sparse_mla 32 passed, 12 skipped (unchanged)
guard boundary 512 and 512+rope still launch; 1008, 1024 and 2048 are
rejected, and the computed figures match the measured
Required: byte counts exactly
Co-authored-by: Cursor <cursoragent@cursor.com>
The gate had no test on any arch, not just gfx942: the shape matrix skips its
fp8 cases off gfx950, so the branch that rejects fp8 dots was never reached and
the branch that accepts them only incidentally.
_resolve_dot_precision's arch parameter no longer defaults. That default was the
real hazard behind the missing coverage: there is one call site, and dropping
its third argument would have silently resolved every launch as gfx950 and
re-enabled fp8 dots on gfx942 -- which a unit test on the helper cannot see,
since the helper still behaves correctly when handed an explicit arch. Required,
that same edit is a TypeError caught by 32 of the 34 tests; verified by making
it and running the suite.
test_dot_precision_arch_gate parametrizes over SUPPORTED_ARCHS, so it needs no
GPU, never skips, and covers a newly added arch automatically.
test_sparse_mla 34 passed, 12 skipped (was 32 passed, 12 skipped)
Co-authored-by: Cursor <cursoragent@cursor.com>
fp8_dsv4_mla, fp8_g64 and the SWA+top-k two-loop return through _forward_paged
before the kernel's own arch gate, so widening that gate never reached them.
They route to pa_decode_sparse, whose packed driver is gated on
DEVICE_ARCH == "gfx950"; anywhere else they land in its fallback path, which
reads a plain grouped fp8 pool rather than these records and rejects them on
dtype:
fp8_g64 AssertionError: kv_scales supplied but unified_kv is
torch.uint8, expected torch.float8_e4m3fnuz
fp8_dsv4_mla RuntimeError: unified_kv dtype mismatch: kv=torch.uint8,
q=torch.bfloat16
Neither names the arch, so on gfx942 both read as a caller mistake. The
ordering is pre-existing, since the early return precedes the assert on main
too, but listing gfx942 in SUPPORTED_ARCHS advertises formats that have no
implementation there, so the wrapper now says so itself.
Reaching the fallback path instead is not an option: it wants the arch-native
fnuz dtype and these records carry OCP code points, so it would decode them
wrong in exactly the way dot_precision="fp8" does.
Verified on gfx942 through the public wrapper:
bf16, fp8_scalar, fp8_dsv32_mla launch, unchanged
fp8_g64, fp8_dsv4_mla ValueError naming the arch
test_sparse_mla 36 passed, 12 skipped (was 34 and 12)
gfx950 is in PACKED_ARCHS, so every branch resolves there as before.
Co-authored-by: Cursor <cursoragent@cursor.com>
The GPU suite only launches the default 512/64 geometry, so a regression
in _check_lds_budget would restore OutOfResources while tests stayed
green. CPU-only cases pin the measured rope-free and separated-rope
boundaries, including the exact 64 KB point.
Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
The kernel reads every fp8 byte -- q, the cache, and the operands of the fp8
dots -- as OCP e4m3, which is gfx950's native fp8. On gfx942 the native fp8
is fnuz: vLLM's fp8 KV cache (current_platform.fp8_dtype()) and aiter's
quantizers (dtypes.fp8) both write float8_e4m3fnuz there. The encodings share
a bit layout with exponent bias 7 against 8, so a gfx942 cache read as OCP
comes out 2x too large, and fnuz's saturated +-240 (0x7F/0xFF) are NaN in OCP.
The fp8-cache tests passed on gfx942 only because their fixtures had been
switched to OCP to match the kernel, which no gfx942 producer writes. A cache
arrives as bytes, usually behind a uint8 view, so the wrapper cannot tell the
encoding from the tensor; _classify_flat catches a fnuz-typed flat pool and
nothing else.
So gfx942 now takes bf16 q and a bf16 cache only, and says why. That is what
the GLM-5.3-Flash path this PR enables uses: vLLM keeps that model's MLA
cache in bf16 and dispatches only bf16 q/kv here. FP8_DOT_ARCHS becomes
FP8_ARCHS and covers fp8 q and caches as well as the dots. The test fixtures
go back to the arch-native fp8, and the fp8 cases skip on gfx942 instead of
running on data nothing there produces. Reading fnuz natively, which would
also let the fp8 dots use CDNA3's fnuz MFMA directly, is left to a follow-up.
The >2 GB global-load test gains a bf16 case, so gfx942 keeps that path
covered now that the fp8 case skips there.
Verified on MI325X (gfx942):
test_sparse_mla 32 passed, 29 skipped (was 44 passed, 12 skipped;
every skip is an fp8 case)
native fnuz cache rejected through sparse_mla_fwd as fp8_scalar and
fp8_dsv32_mla; with the gate stubbed out the same call
returns rel-err 1.01 and the wrapper test fails
bench_sparse_mla runs, bf16 series only
gfx950 is in FP8_ARCHS, so every branch resolves there as before.
Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
_check_lds_budget kept its own copy of gfx942's LDS capacity. It now reads
arch_info._LDS_CAP_BYTES, as pa_decode_sparse and the GEMM num_stages picker
do, so the capacity has one definition.
_LDS_CAP_BYTES lists gfx950 as well, so dict membership can no longer pick
the checked arch, and the gfx942-only gate is now explicit. The footprint
model was measured for gfx942's bf16, non-async tiles; gfx950 stays with the
launcher as before, which test_lds_budget_gfx950_is_unchecked pins.
test_sparse_mla 32 passed, 29 skipped (unchanged)
Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
_async_launch_config took an lds_limited flag that sparse_mla_fwd derived
from the arch. It now checks arch_info.get_arch() itself, so the arch
switch is visible where the config is chosen, and its signature is back to
main's.
No launch changes: lds_limited was exactly arch == "gfx942".
test_sparse_mla 32 passed, 29 skipped (unchanged)
Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
"{arch}'s decodes fnuz" read as malformed. It now says the arch's native
fp8 is fnuz, the wording _check_fp8_arch already uses.
test_sparse_mla 32 passed, 29 skipped (unchanged)
Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
… arch
The dot_precision docstring still said bf16 dots work with every cache
format, and did not say fp8 dots are gfx950-only. On gfx942 the wrapper
takes a bf16 cache alone and rejects the rest, so the public docstring now
states the same scope the code enforces.
Docstring only.
Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
This PR's import edits in test_sparse_mla.py and bench_sparse_mla.py left
both blocks unsorted, which the pinned ruff 0.16.0 in pre-checks reports as
I001. Both files are clean at the merge base. Import order only.
ruff 0.16.0 check all three sparse_mla files pass
test_sparse_mla 32 passed, 29 skipped (unchanged)
bench_sparse_mla runs, bf16 series only
Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
…'s KV pad
_check_lds_budget modelled the KV tile at the kernel's default 8-element row
pad, but a prefill launch (num_queries >= _PREFILL_MIN_ROWS, bf16 dots)
passes KV_LDS_PAD=16, so the guard undercounted every prefill launch by
BLOCK_K * 8 * 2 = 512 B. The pad is now computed once and handed to both
the guard and the launch, so the two cannot drift apart again.
No power-of-two geometry changes its accept/reject result, so no launch
that used to work is rejected. What changes is that the guard now reports
the footprint the launch actually needs, and it no longer depends on the
gap between pow-2 widths to stay safe. Widths in between never reach the
launcher: the KV tile's PaddedSharedLayout asserts a pow-2 shape.
Measured on MI325X with the guard stubbed out, reading metadata.shared when
the kernel fits and OutOfResources' Required: when it does not, at BLOCK_K
32, decode (8 queries) and prefill (2048):
kv_lora_rank / rope decode prefill
512 / 0 34304 34816 fit
512 / 64 38912 39424 fit
512 / 256 51200 51712 fit
512 / 512 67584 68096 rejected
1024 / 0 67072 67584 rejected
1024 / 64 71680 72192 rejected
2048 / 0 132608 133120 rejected
The guard now reproduces every one of those through sparse_mla_fwd. The
boundary test gains the prefill pad, with exact-fit points at 1000 (decode)
and 992 (prefill). Dropping the pad from the guard again fails the five
prefill rejects.
test_sparse_mla 43 passed, 29 skipped (was 32 and 29; the new ones are
boundary cases)
Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
_ARCH_BLOCK_K and _ARCH_NUM_WARPS were per-arch launch values in Python,
which the config rules in aiter-ops-triton.instructions.md and
configs/CLAUDE.md place in JSON. They now live in
configs/gfx942/gluon/attention/sparse_mla/DEFAULT.json BLOCK_K 32, num_warps 4
configs/gfx950/gluon/attention/sparse_mla/DEFAULT.json BLOCK_K 64, num_warps 4
read through resolve_config_dir() + load_config_json(), the single
DEFAULT.json read that configs/CLAUDE.md describes for attention wrappers.
The gfx950 file carries main's hardcoded block_k = 64 and
num_warps = block_k // 16, so nothing changes there. The entry is keyed by
the kernel, _sparse_mla, so the reduce can get its own entry later.
Only the values this PR made per-arch move. The rest of the launch policy
(the BLOCK_M rule, split-K, the async tile choice, the prefill KV pad) is
main's and stays as it was. Moving it all would rework gfx950's launch path,
which is its own change.
test_launch_config_published checks from any machine that every arch in
SUPPORTED_ARCHS ships the file, so adding an arch without one fails on
every runner rather than only on that arch's. The LDS boundary test takes
gfx942's BLOCK_K from the same file.
Verified on MI325X (gfx942): the compile-time launch arguments the wrapper
passes (BLOCK_K, num_warps, waves_per_eu, ASYNC_LDS, KV_LDS_PAD and the
rest) are identical before and after over 96 shapes spanning decode and
prefill, H 8/16/64, rope 0/64, has_invalid and forced split-K.
_get_config() costs 0.65 us per call.
test_sparse_mla 45 passed, 29 skipped (was 43 and 29; +2 config checks)
Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
_async_launch_config's gfx942 branch returned early with its own copy of
main's waves_per_eu rule, so gfx942 carried a second, arch-specific
statement of a launch value. All the branch has to do is keep the async
tiles off, which are sized for gfx950's LDS, so the arch check now sits in
the condition that enables them, and gfx942 falls through to main's rule
for everything else. The switch still reads arch_info.get_arch() inside
_async_launch_config, and the signature is still main's.
The async path also needs fp8 dots, which gfx942 rejects before it gets
here, so today the check only takes effect once fnuz fp8 dots exist there.
It stays for that case.
No launch changes. Checked on MI325X:
_async_launch_config old and new agree on all 10752 combinations of
its inputs under both gfx942 and gfx950, fp8_dots
included
launch arguments identical to the PR head over the same 96 shapes
test_sparse_mla 45 passed, 29 skipped (unchanged)
Signed-off-by: Jin Tao <jin.tao@amd.com>
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.
Enables
sparse_mla_fwdon gfx942.Re-land of #5539, which was auto-closed on Sep 18 when its base branch
(
cagri/sparse_pa_optimizations) was deleted three seconds after #4919 merged.Nothing was rejected there — the one review comment is addressed and folded into
this commit. Now targets
maindirectly; the commit cherry-picks onto it with noconflicts and no file in it has changed on
mainsince #4919 landed.The Gluon kernel needs no change to run there — the
gl.amd.cdna4intrinsics ituses (
async_copy,mfma,buffer_load/buffer_store) all lower fine ongfx942 under Triton 3.7.1. The arch gate was the only thing in the way, plus two
launch-config values sized against gfx950's LDS budget.
Why the tile has to change
gfx942 (CDNA3) has 64 KB of LDS per workgroup against gfx950's 160 KB. The
bf16 KV tile alone is
BLOCK_K * (kv_lora_rank + LDS_PAD) * 2 B, so the gfx950tile of 64 asks for 68 608 B and will not launch:
BLOCK_K=32is the largest that fits.num_warpsis decoupled fromBLOCK_Kin the same change. It was derived asBLOCK_K // 16, so capping the tile for LDS would have dropped num_warps from 4to 2 as a side effect rather than as a tuning decision. Keeping it at 4 is
worth 1.19× → 1.55× on the prefill shapes, so it is worth making explicit.
fp8 is rejected on gfx942
The kernel reads every fp8 byte (q, the cache, and the operands of the fp8 dots) as
OCP e4m3, which is gfx950's native fp8. On gfx942 the native fp8 is fnuz: vLLM's
fp8 KV cache and aiter's quantizers both write
float8_e4m3fnuzthere. The encodingsshare a bit layout with exponent bias 7 vs 8, so the same bytes decode 2x apart;
0x80is -0 in OCP but NaN in fnuz, and fnuz's saturated ±240 (
0x7F/0xFF) are NaN in OCP.It first showed up as
dot_precision="fp8"being silently wrong (rel-err 7.5e-1):the CDNA3 fp8 MFMA decodes fnuz and was fed OCP. The fp8 caches have the same problem,
and a cache usually arrives as a uint8 view, so the wrapper cannot tell the encoding
from the tensor. gfx942 therefore takes bf16 q and a bf16 cache only and raises
otherwise, naming the reason. That is what the GLM-5.3-Flash path uses: vLLM keeps this
model's MLA cache in bf16 and dispatches only bf16 q/kv to the kernel. Reading fnuz
natively (dequant, plus
dot_precision="fp8"on CDNA3's fnuz MFMA) is a follow-up.Two consequences, both in this PR:
FP8_ARCHS. With the arch gate widened itwould otherwise reach that series on gfx942 and raise on its first point;
triton.testing.perf_reportdoes not catch it, so the run died before printing atable. It now prints a note so the omission is visible rather than silent. (This was
frida-andersson's catch on [Triton/Gluon] [gfx942] Enable sparse_mla_fwd on gfx942 #5539.)
GPU-free test covers the gate on every arch, and a wrapper-level test feeds gfx942's
own fnuz cache behind a uint8 view and expects the error.
gfx950 is unchanged
gfx950's JSON keeps 64 and 4, and gfx950 is in FP8_ARCHS and PACKED_ARCHS.
Validation
On MI325X (gfx942), TP4, GLM-5.3-Flash at 131 k context. This matters on that
model because its NoPE MLA has
qk_rope_head_dim = 0, which has no AITER pathtoday, so vLLM falls back to a vendored Triton gather+dot — the single largest
kernel in the model in both phases.
op_tests/triton_tests/attention/test_sparse_mla.py:Kernel bucket, from a real torch profile at concurrency 12 (rank 0):
End-to-end, 131 k in / 1024 out, same image both arms with only the dispatch
switched:
Won 7/7 rows, mean +7.14% output throughput, zero failed requests. gsm8k
(5-shot, full 1319 questions) is neutral: 0.9719 vs 0.9712 strict-match, against
a ±0.0046 standard error.
The gain being largest at low concurrency is the signature of this kernel rather
than noise: the operator is batch-independent, so replacing it takes a roughly
constant ~2.7 ms off every decode step, which is 16.6% of TPOT at concurrency 2
but 4.9% by concurrency 16.
Note on prefill
The 1.06× prefill figure is the honest one, measured in a profile. A
microbenchmark using
randpermtop-k indices reports 1.55× for the same shape,but that flatters it — random indices give the Triton gather far worse locality
than production does (11.0 ms real vs 15.7 ms synthetic for the same shape),
while this kernel gathers whole 512 B rows and is locality-insensitive.
Prefill microbenchmarks on this operator should not be trusted without a
profile.