Skip to content

[Triton/Gluon] [gfx942] Enable sparse_mla_fwd on gfx942 - #5721

Merged
k50112113 merged 14 commits into
ROCm:mainfrom
jin-amd:gfx942-sparse-mla-main
Oct 6, 2026
Merged

k50112113 merged 14 commits into
ROCm:mainfrom
jin-amd:gfx942-sparse-mla-main

Conversation

@jin-amd

@jin-amd jin-amd commented Sep 21, 2026 •

Copy link
Copy Markdown
Contributor

Enables sparse_mla_fwd on 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 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.

op_tests/triton_tests/attention/test_sparse_mla.py:

The test count should read 45 passed, 29 skipped.

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.

@jin-amd
jin-amd requested review from a team and a lite review from Copilot September 21, 2026 07:28
@github-actions

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; main branch always runs both MI35X and MI300X
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 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.

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

Unresolved critical and moderate launch-safety, compatibility, configuration, and test-coverage findings remain.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 1 High severity · 1 Medium severity · 1 Low severity

Open (3)
What changed in this PR

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.
op_tests/​triton_tests/​attention/​test_sparse_mla.py Extends gfx942 correctness coverage and skips unsupported FP8 cases. Moderate (1 vote): Add a gfx942-only public-wrapper test asserting FP8 rejection while retaining numerical FP8 skips.
op_tests/​op_benchmarks/​triton/​bench_sparse_mla.py Gates benchmark series by architecture. No final comments.

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

Comment thread aiter/ops/triton/attention/sparse_mla.py Outdated
Comment thread aiter/ops/triton/attention/sparse_mla.py Outdated
Comment thread aiter/ops/triton/attention/sparse_mla.py Outdated
Copilot AI review requested due to automatic review settings September 21, 2026 09:39

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

🔵 Needs a closer look

Resolve launch-config integration, add the gfx942 FP8 rejection test, and cover the production FP8-cache benchmark.

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

Open (2)
Resolved since last review (1)
Previously missed (1)

In code that hasn't changed since last review

Medium severity Add a gfx942 negative test for unsupported fp8 dot precision

op_tests/​triton_tests/​attention/​test_sparse_mla.py:35

[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.

Copilot AI review requested due to automatic review settings September 21, 2026 09:57

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 gfx942 packed-cache path has an unresolved critical issue, with additional configuration, test, and documentation follow-ups.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 1 High severity · 1 Low severity

Open (2)
Resolved since last review (1)
Previously missed (1)

In code that hasn't changed since last review

Medium severity Add LDS budget boundary regression tests

aiter/​ops/​triton/​attention/​sparse_mla.py:667

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.

Comment thread aiter/ops/triton/attention/sparse_mla.py
Copilot AI review requested due to automatic review settings September 21, 2026 10:15

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

Unresolved findings remain around packed-cache support, launch-config placement, and LDS guard coverage/capacity reuse.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 2 High severity

Open (2)
Resolved since last review (1)

Comment thread aiter/ops/triton/attention/sparse_mla.py Outdated
Copilot AI review requested due to automatic review settings September 21, 2026 10:20

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

🔵 Needs a closer look

Unresolved moderate issues remain in the tuning configuration, LDS boundary coverage, and gfx942 benchmark fidelity.

Review effort: Lite
Findings: None

Resolved since last review (2)
Previously missed (2)

In code that hasn't changed since last review

Medium severity Add boundary test for the 64 KB LDS budget

aiter/​ops/​triton/​attention/​sparse_mla.py:132

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.

Medium severity Benchmark gfx942 bf16 dots with the production fp8 cache

op_tests/​op_benchmarks/​triton/​bench_sparse_mla.py:177

[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.

@zufayu
zufayu requested review from a team and azaidy September 22, 2026 01:51
Copilot AI review requested due to automatic review settings September 22, 2026 12:49

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

Move launch tuning into the config path and add LDS boundary tests before approval.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 1 Medium severity

Open (1)

Comment thread aiter/ops/triton/attention/sparse_mla.py

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 LDS budget calculation can admit overflowing launches, and tuning and hardware-cap configuration should use shared repository mechanisms.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 1 High severity

Open (1)
Resolved since last review (1)

Comment thread aiter/ops/triton/attention/sparse_mla.py

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

🔵 Needs a closer look

Two moderate issues remain in launch configuration management and prefill LDS validation.

Review effort: Lite
Findings: None

jin-amd and others added 11 commits October 2, 2026 11:36
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>
Copilot AI lite review requested due to automatic review settings October 2, 2026 08:36
@jin-amd
jin-amd force-pushed the gfx942-sparse-mla-main branch from 847d014 to cdbecc4 Compare October 2, 2026 08:36

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

Unresolved LDS validation and launch-configuration issues remain.

Review effort: Lite
Findings: 1 High severity

Open (1)

Comment thread op_tests/triton_tests/attention/test_sparse_mla.py Outdated
jin-amd and others added 3 commits October 5, 2026 08:45
…'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>
Copilot AI lite review requested due to automatic review settings October 5, 2026 08:58

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

🔵 Needs a closer look

Architecture-specific kernel behavior and FP8 handling warrant final human review.

Review effort: Lite
Findings: 1 High severity

Open (1)

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