Skip to content

[Triton/Gluon] [ASM] [HIP] Mha v4: adds bf16 sparse, LSE support, KV varlen, fixes, etc - #5798

Merged
jcaraban merged 66 commits into
mainfrom
mha_v4_fixes_sparse_etc
Sep 28, 2026
Merged

jcaraban merged 66 commits into
mainfrom
mha_v4_fixes_sparse_etc

Conversation

@jcaraban

@jcaraban jcaraban commented Sep 23, 2026 •

Copy link
Copy Markdown
Contributor

⚠️ These MHA kernels are mainly tested & intended for Diffusion Inference workloads.
However, mha_v4 gfx950 bf16 is virtually as accurate as mha_v3 at +100 TFlops 👌

Motivation

Add log-sum-exp output to MHA v4 gfx950 kernels (only dense variants now), so it can run under ring / context parallelism, which merges per-rank partials via LSE. Also extends block-sparse to the BF16 recipes, canonicalizes the MXFP4 rows, and improves K/V quantization accuracy.

Technical Details

Kernels

  • New BF16 and BF16FP8 sorted-sparse kernels.
  • Dense LSE epilogue on all ten hd128 recipes, gated at runtime on s_lse.
  • f8f6, f6f4 and mxfp4 sparse rows moved to FP6-P V, matching their dense siblings.
  • Dense mxfp4 claims the canonical FP6-P V order; the duplicate f4f4 row is disabled.

Host

  • bugfix: per-channel V amax clamped so empty heads cannot quantize to NaN.
  • Optional lse output on mha_v4 / mha_v4_packed, plumbed through asm_mha_v4_fwd.cu.
    • ABI unchanged: ptr_lse / s_lse / s_lse_Hs were already reserved in the kernarg.
  • _LSE_CAPABLE_QV gates the supported format pairs; sorted-sparse still raises.
  • K mean is subtracted before quantizing (K-smoothing), fused into the MX quantizer kernels.
  • Dense MHA v4 accepts per-batch key lengths (ragged seqlen_k).

Minor

  • Retired MXFP4 Q/K + FP8 V and the deprecated mha_v4_mxfp8 alias.
  • bench_sage.py: improve input distributions, diffusion-calibrated default, BF16 sparse modes.
  • Split the block-sparse cases into op_tests/test_mha_v4_sparse.py.

Test Plan

  1. Run the complete MHA v4 test suite on MI355X/gfx950.
  2. Per-recipe dense and sparse validation against Torch; full-LUT sparse bit-exact vs dense.
  3. Compare LSE against float32 torch.logsumexp for every capable format pair.
  4. Compare O bitwise, LSE off vs on and against the previous code objects.
  5. Interleaved A/B throughput, 12 rounds per variant on one idle GPU.

Test Result

  1. 330 passed, 1 skipped in test_mha_v4.py + test_mha_v4_sparse.py on MI355X/gfx950.
  2. Dense validation passes 10/10 recipes; unchanged sparse objects are byte-identical.
  3. LSE vs float32: bf16 max_abs 7.2e-4; quantized rows |mean_off| <= 0.046.
  4. O bitwise unchanged vs the previous objects (40/40 cases); dense throughput at parity
  5. black --check and ruff on PR-touched Python files: passed.

Previous PRs: #5335, #5005, #4967, #4627

image

test_mha_v4_raw_mxfp4_v_supports_unaligned_sequence already ran the shape that
fails, q[1,129,2,128] with k=v[1,257,2,128], but asserted only eager==compiled
and isfinite. A wrong yet stable result satisfies both, and that is exactly what
the MX V rows return: on that shape the cosine against attention is 0.0007 while
both assertions hold. Give it a reference and add a sweep over partial-tile
occupancy, since the fault is invisible at tile-aligned lengths.

The sweep holds the KV length at one full 128-token tile plus a tail so only the
partial tile varies. Recipes with a per-tensor FP8 V are flat across the tail, so
any dependence belongs to the MX V path. Measured today, cosine against
attention at tail=1 and tail=64: mxfp6 0.007 and 0.384, f6f4 -0.012 and 0.382,
f4f4 -0.012 and 0.377, against 0.998 for f6f8 at every tail. f8f6 is milder but
still dips to 0.910, so it is covered too.

23 of the new cases fail today. That is the point: they describe the defect, and
the surrounding 248 tests still pass.
Every existing sparse test uses sequence_k=512, which is four KV tiles. A kernel
whose LUT walk degrades only past that stays green: the mxfp8 sparse object was
correct to six tiles and then fell from cosine 0.998 to 0.925 at seven and 0.765
at thirty-two, and the suite never saw it. The gap was shape coverage, not the
assertion -- the existing cosine>0.99 bound would have caught it at seven tiles.

Sweep 2..32 tiles for every sparse row, both with an all-true LUT and with one
that actually skips, and judge each row against its own accuracy at two tiles
rather than a shared threshold, so a row is held to not degrading with tile count
regardless of its quantization error.

Also assert that an all-true LUT matches dense bitwise. That is strictly stronger
than cosine and worth having: against the previous mxfp8 object this fails at two
tiles with cosine 0.9999995, so the walk was already not reducing to the dense one
long before any cosine threshold would have noticed. Restricted to the rows whose
V stays FP8, since the MX-V rows resolve to a different V packing on their dense
manifest row than on their sparse one and cannot be compared bitwise.

New module rather than an addition to test_mha_v4.py, so this can grow without
touching that file.

Verified to fail on the previous code object and pass on the current one.
The mxfp6, f6f4, mxfp4 and f8f6 dense kernels returned corrupted results
whenever seqlen_k was not a multiple of the 128-element KV tile. Two separate
kernel defects were responsible.

The V buffer descriptor bounded num_records by the unpadded key length. V is
column-major within a KV tile, so that bound cut the final partial tile along
the head-dim channel axis instead of the token axis, silently dropping output
channels. The descriptor now covers whole packed tiles.

The masked tail also rescaled the running numerator by a delta of exactly zero,
discarding every KV tile accumulated before the partial one. Both the running
maximum and the tile scores are now compared in the same domain.

Aligned sequence lengths are unaffected; those paths never ran the masked tail
after real accumulation, which is why the existing tests passed. All ten dense
recipes now score the same at ragged lengths as at aligned ones, and a uniform
attention probe with an all-ones V returns exactly 1.0 at every length. Cosine
is scale-invariant and could not observe the dropped channels on its own.
All six objects (dense and sparse for each format) move to a unified block
schedule: per KV tile each wave runs an M phase holding the deferred PV and the
next QK, then an S phase holding the mask and softmax, with the two wave bands
anti-phased so one band's matrix work covers the other's VALU. This replaces the
older six-rendezvous seam topology and cuts the steady loop to two barriers per
tile.

MXFP8 sparse also fixes a correctness bug: two scalar registers shared one slot,
so the sparse prologue overwrote the probability packing scale with a denormal.
Probabilities flushed to zero, the softmax denominator went to zero, and the
kernel returned NaN on every LUT path. It was not reachable in dense builds.

I8FP8 additionally moves its V operand staging out of the softmax phase, which
is VALU-bound in that format, and rebalances the address update across the QK
matrix gaps. At b=1, hq=5, sq=sk=65536, d=128 on MI355X this leaves it slightly
ahead of the object it replaces, measured with per-GPU rotation across 7 GPUs.

Validation on gfx950: op_tests/test_mha_v4.py 271 passed / 1 skipped; dense
24/24 and sparse 210/210 over the tile, pattern and gather sweeps, every case
deterministic across repeated launches. Full-LUT sparse is bitwise equal to
dense for the FP8-V formats. Each object was rebuilt from source and compared
byte-for-byte against the deployed slot.
Both objects, dense and sparse, move to the same unified block schedule as the
earlier FP8, MXFP8 and I8FP8 refresh: per KV tile each wave runs an M phase
holding the deferred PV and the next QK, then an S phase holding the softmax,
with the two wave bands anti-phased so one band's matrix work covers the other's
VALU. This replaces the older six-rendezvous seam topology and cuts the steady
loop to two barriers per tile.

The V operand reads are rebalanced on top of that. One unit stays resident
before PV, the next is read under PV's two leading denominator MMAs, and the
rest stream one matrix operation behind, which shortens the distance from each
read to the wait that needs it and empties most of the softmax tail. At b=1,
hq=5, sq=sk=65536, d=128 on MI355X that is worth 3.2% dense and 2.6% sparse over
the first ported build, measured with per-GPU rotation across 7 GPUs.

The port also repairs F6F8 below one KV tile, where the compact short schedule
replaces a path that returned cosine 0.62 at sk=64 and NaN at sk=100.

Validation on gfx950: dense 24/24 across the KV length sweep and sparse 70/70
including the full-LUT bitwise comparison against dense, every case
deterministic across repeated launches. Output is bitwise identical to the
object it replaces from sk=512 through sk=8192. Both objects were rebuilt from
source and compared byte-for-byte against the deployed slot.
mha_v4 gains an optional seqlens_k naming how many keys each batch really has.
Callers that pad K/V to a common length and track the real lengths separately -
the diffusion runners all do, for text conditioning - no longer have to choose
between a wrong softmax denominator and falling off the kernel.

The pointer travels in the kernarg slot at 0x1C0, which dense launches leave
null. The refreshed bf16 and bf16fp8 objects test it and skip the read when it
is absent, so the same object serves both and passing no lengths is bitwise
identical to the previous dense result, which a test pins. Cost on the dense
path is seven scalar instructions outside the loop, measured at parity.

Only the two BF16-Q rows read the slot so far; the quantized rows ignore it, and
the sorted-sparse launcher rejects it rather than silently dropping it.
The sorted-sparse F8F6 row was pinned to the canonical V packing because its
code object predated the FP6 P pack. That object now packs P as FP6 like the
dense row, so it needs V staged in the matching layout. Handed canonical V it
reads KV rows with bits 2 and 5 of the row index exchanged, scoring about 0.69
cosine against dense while every structural check still passes, because the
usual V staging probes drive the kernel with uniform attention and uniform
attention cannot see a row permutation.

Move the manifest row to v_pack=1, stop withholding V_FOR_FP6_P from it in
_resolve_raw_recipe, and stage V the same way in the benchmark, which quantizes
its own operands. The MXFP6-Q sparse rows keep canonical V until their objects
are rebuilt.

Sorted-sparse F8F6 is now bit-exact against dense at every tile count, and
measures about 13% higher throughput at sq=sk=65536 than the object it replaces.
Dense and sparse fwd_hd128_f8f6 now use the same two-barrier block schedule as
the other MHA v4 kernels. Both replace seam-schedule objects.

Balanced 7-GPU paired medians at b=1,hq=5,sq=sk=65536,d=128, n=14: dense -0.35%
and sparse +3.84% (winning every lane) against the objects they replace. Sparse
gains come from folding the block-LUT prefetch into the QK MFMA shadows, which
also cuts it from three LDS reads per two tiles to one per tile.

The build also closes a latent cross-band WAR on the V LDS buffer by giving V
three rotating slots, so repeated launches no longer diverge at ragged KV
lengths.

Output is bit-identical at sk 512/1k/2k/4k/8k and sparse remains bit-exact
against dense. op_tests/test_mha_v4.py: 277 passed, 1 skipped.
Resources 249 VGPR, 99 SGPR, 70144 B LDS dense / 100868 B sparse.
Both fwd_hd128_f8f6 objects drop from 249 to 237 VGPRs after pruning dead
register-table entries and regrouping the allocation; occupancy is unchanged at
2 waves/SIMD. The sparse object additionally switches the L-sum MFMA's ones
operand from FP4 to FP6, which the ISA runs at the same 16 cycles.

Balanced 8-GPU paired medians at b=1,hq=5,sq=sk=65536,d=128: dense +0.85% over
the objects they replace, winning every lane; sparse is within run-to-run noise.

Output is bit-identical at sk 512/1k/2k/4k/8k and sparse remains bit-exact
against dense. op_tests/test_mha_v4.py: 277 passed, 1 skipped.
Resources 237 VGPR, 99 SGPR, 70144 B LDS dense / 100868 B sparse.
The F6F8 kernels staged V through two LDS slots, one short of what the
software pipeline needs: the V loads issued for a later tile targeted the same
slot the current tile was still reading, with no barrier separating them. A
repeated dense sweep reproduced nondeterministic output on 4 runs in 10; the
rebuilt objects show 0 in 25. Rotating V through three slots costs no
registers and raises the F6F8 group segment from 59904 to 76544 bytes, which
keeps the same occupancy.

Both formats also rebuild with a reordered prologue, so the accumulator
zeroing and (for F6F8) the Q loads overlap the outstanding memory traffic
instead of running before or after it.

Measured on 8x MI355X, paired per-GPU, against the previously shipped objects:
  fwd_hd128_f6f8         +0.71% at sq=65536 (8/8 lanes), +1.19% at sq=1024
  fwd_hd128_f6f8_sparse  -0.23%, accepted as the cost of the race fix
  fwd_hd128_f8f6         +0.45% at sq=1024 (8/8 lanes), +0.15% at sq=65536

Numerics are unchanged: output digests match the previous objects at
sk 512/1024/2048/4096/8192 for every variant, and op_tests/test_mha_v4.py
reports 277 passed, 1 skipped.
MXFP6 Q/K/V was the only MHA v4 recipe with no sorted-sparse row, so a block
mask on it raised NotImplementedError. This ships fwd_hd128_mxfp6_sparse.co,
adds its manifest row, and drops the guard.

The new object is an FP6-P build, like the dense MXFP6 one, so its V operand
must be repacked to the FP6 P element order. _resolve_raw_recipe previously
forced canonical V for every sparse MXFP6 row, which no longer holds: the
MXFP6-V row is FP6-P in both modes, while the MXFP4-V sparse row is still a
pre-FP6-P build and keeps canonical V. Getting this wrong does not fail
loudly, it just fails find_config on v_pack, so the two are now split
explicitly rather than by mode.

bench_sage could not benchmark the new row either. mha4_mxfp6 was the only
mha4 entry without supports_block_sparse, and the same canonical-V rule was
duplicated in the quantizer and the launcher, where disagreeing copies would
silently mispack V. Both now read one predicate.

The dense object is rebuilt in the same change. Its kernel gives up the
scalar register that the sparse prologue needs, which relocates a single
prologue instruction; measured +0.14% at b1/hq5/sq65536, inside noise.

Validation on 8x MI355X, sparse across 66 cases covering tile counts 1..32,
skipping and gathered patterns: every case at or above its cosine floor with
identical output across repeated launches, and the gathered checks match the
dense kernel exactly (cosine 1.00000) over the same physically gathered tiles.
The dense sweep is unchanged at 24/24. bench_sage --block-sparsity 0.5 against
the torch reference gives cosine 0.99691 for sparse MXFP6 versus 0.99684 for
dense MXFP6, confirming the packing.

op_tests/test_mha_v4.py reports 277 passed, 1 skipped, and
op_tests/test_mha_v4_sparse_tile_scaling.py 20 passed. The two tests that
asserted MXFP6 sparse was unavailable now cover the working path instead.
… to FP6-P V

Measured on 8x MI355X, paired per-GPU, against the shipped objects:
  fwd_hd128_f6f4         -1.24% at sq=65536, +0.69% at sq=1024 (8/8 lanes)
  fwd_hd128_f6f4_sparse  +6.35% (8/8 lanes)
  fwd_hd128_mxfp6        +0.75% (8/8 lanes)
  fwd_hd128_mxfp6_sparse +1.93% (8/8 lanes)

The dense F6F4 regression at long sequence is accepted: the object moves off the
hand-tuned seam schedule onto the block schedule the other rows already use, and
it buys a correctness fix below one KV tile where the old object was wrong
rather than imprecise -- sk=64 cosine 0.61584 -> 0.99178, sk=100 0.10706 ->
0.99158. The gap was -7.2% when the port landed and is being tuned down.

The rebuilt F6F4 sparse object is an FP6-P build, so its manifest row moves to
v_pack=1 and the MXFP4-V exception in _resolve_raw_recipe disappears: every
MXFP6-Q row now uses V_FOR_FP6_P in both modes. bench_sage carried the same rule
in its own predicate and follows. Getting this wrong does not fail loudly, it
either misses find_config on v_pack or silently mispacks V.

Because dense and sparse now agree on V packing, a full LUT must reproduce dense
bitwise for both F6F4 and MXFP6; that check is enabled and passes.

op_tests/test_mha_v4.py reports 297 passed, 1 skipped, and the sparse validation
covers 280 cases across F6F4, MXFP6, F6F8 and F8F6 with identical output over
repeated launches.
The MXFP4-Q raw path ignored recipe.v_pack and always emitted canonical-order
V, so the f4f4 sparse object, which reads V in FP6-P order, saw its KV rows
permuted: bit2 <-> bit5, cosine 0.6999. Honor v_pack there, route all-MXFP4
sparse to V_FOR_FP6_P, and set the f4f4_sparse manifest row to v_pack=1.
Dense all-MXFP4 stays on the canonical-V mxfp4 row, the only all-MXFP4 object
that handles ragged sequences.

Rebuild both f4f4 objects from source: dense is +5.35% (8/8 lanes) at
b1/hq5/sq65536 with identical numerics, sparse is performance-neutral and now
maps KV rows identity.

Validated: f4f4 sparse 66/66, 556/556 across eight sparse recipes, suite
297 passed / 1 skipped.
The MXFP4 kernels now take the canonical V pack in both dense and sparse modes, so
the sparse manifest row moves from FP8 V + per-channel f32 descale
(9,9,3,0,2,5,5,4) to MXFP4 V + E8M0 (9,9,9,0,2,5,5,5), matching the dense row. The
recipe resolver no longer forces the FP6-P pack for an all-MXFP4 signature; that
object stays reachable by passing v_pack=V_FOR_FP6_P explicitly.

Motivation: the canonical-V objects are flat across ragged sequences where the
FP6-P pack falls off, cos 0.98 vs 0.89 at sk=100 and 0.98 vs 0.79 at sq=129 sk=64.
Cost is ~15% throughput at 50% block sparsity, to be recovered separately.

bench_sage drops its FP8-V special case for sparse mxfp4, so both modes quantize V
the same way and the payload-bytes override goes with it.

Refreshed hd128 mxfp4 and f4f4 objects. Validation: 8/8 dense and 66/66
block-sparse cases for mxfp4.
f6f4: tuned hot-loop code placement, worth +0.24..+0.34% at
b=1,hq=5,sq=sk=65536,d=128 on MI355X. The gain is reproduced across three
runs and separated from a same-phase control build, so it is a placement
effect rather than run-to-run noise. Validated 8/8 dense and 70/70 sparse.

mxfp8: rebuilt from current kernel sources so the shipped objects match
them again; the rebuild drops seven dead prologue instructions that no
live code consumed. Throughput is neutral (median 2983.5 -> 2993.4
TFLOP/s at the same shape). Validated 8/8 dense and 70/70 sparse.
Drops a dead bias instruction from the seeded softmax path and retunes the
hot-loop code placement to match the smaller loop. Output is unchanged: all
eight dense cosines are bit-identical to the previous object.

+0.72% at b=1,hq=5,sq=sk=65536,d=128 on MI355X, winning 8/8 GPUs with the
worst lane still +0.36%. Validated 8/8 dense and 70/70 sparse.
The MXFP4 hd128 forward kernel now consumes FP6 E2M3 probabilities against
MXFP4 V, which requires V staged in the FP6-P token order. _resolve_raw_recipe
therefore selects AttentionPack.V_FOR_FP6_P for the all-MXFP4 recipe, and the
manifest rows move to v_pack=1 for both the dense and sparse objects.

The f4f4 rows are commented out: MXFP4 now covers that shape with the same
v_pack=1 V order, so leaving both in place makes the lookup ambiguous.

Host and kernel form one contract here. The refreshed objects expect the
FP6-P V order, so the manifest, the recipe resolver, the benchmark's packing
helper and both .co files have to move together; deploying the objects alone
gives the kernel V in the old row order and silently degrades accuracy.

Objects also carry a reworked V-scale path that loads scales directly into
VGPRs rather than staging them through LDS, which leaves dense and sparse
running identical scale code. The E8M0 scale image layout is unchanged, so no
host-side packing change accompanies it.

Sparse gains 2.17% at b=1,hq=5,sq=sk=65536,d=128 (winning all 8 GPUs, worst
lane +2.00%); dense is unchanged to within noise. Validated across the dense
shape suite including ragged and sub-tile sequences, and the full block-sparse
suite.
_resolve_raw_recipe selected the FP6-P V packing from the recipe kind alone,
unlike the neighbouring clauses, which also test v_format. MXFP4 Q/K with FP8 V
therefore requested a repacked V it cannot consume, and _validate_pack_contract
rejected the launch outright. Test the V format too.

test_mha_v4_resolves_raw_recipe still expected the pre-FP6-P packing for
all-MXFP4, so it now expects V_FOR_FP6_P and covers sparse alongside dense.

The remainder is black and ruff over the files this branch touches.
The sorted-sparse manifest used to carry an FP8-V row for MXFP4 Q/K, because
no MXFP4-V sparse object existed; dense rejected the same combination outright.
Now that the MXFP4-V sparse row ships, that asymmetry has no reason to exist,
and nothing is left using the FP8-V variant.

_resolve_raw_recipe therefore accepts only MXFP4 V for MXFP4 Q/K, in both dense
and sparse, replacing the mode split and its dense-only error with one rule.

The sparse mxfp4 test recipes move to MXFP4 V to match. That makes the separate
all-MXFP4 f4f4 entries in the empty-row launches and the tile-scaling recipe
table exact duplicates of the mxfp4 ones, so they are dropped; the manifest
already routes that shape to the mxfp4 objects.

295 passed, 1 skipped.
Every manifest row carrying MXFP4 V now selects the FP6-P pack, so the canonical
DEFAULT-pack path for MXFP4 V can no longer be reached from mha_v4: the MXFP4
recipe always resolves to V_FOR_FP6_P, and MXFP4 Q/K with FP8 V is rejected
during recipe resolution rather than reaching the quantizer.

That leaves the FP8-V arm of the MXFP4 kind, its matching v_view conditional,
and the quantize_v_mxfp4 fallbacks in the MXFP4 and MXFP6 kinds with no way to
execute. Remove them.

quantize_v_mxfp4 itself stays exported: it is still a usable standalone
producer, though no manifest row currently consumes its output.
The recipe table was dense-only and marked MXFP6 V as dense-only, which stopped
being true once the sorted-sparse MXFP6 object landed. Every quantized recipe is
now available in both modes with the same V packing and scale modes, and only
the BF16 rows are dense-only, so the table carries a Modes column and states the
parity directly.

The packing section claimed AttentionPack.DEFAULT was the order used by sparse
kernels. Sparse rows now select the same pack as their dense counterparts, so
the mode no longer influences packing at all.

Also records per-batch key lengths via seqlens_k, which dense accepts and sorted
sparse rejects, and that MXFP4 Q/K requires MXFP4 V.
mha_v4_mxfp8 has been reachable through mha_v4 with FP8 formats and E8M0 Q/K
scale modes for a while, and carried a DeprecationWarning saying so. Drop the
wrapper, its export, and the alias test.

quantize_v_mxfp4 no longer has a consumer: every MXFP4-V manifest row on either
arch selects the FP6-P pack, so its output cannot be launched. Drop the producer
with its fake, its export, and the test that pinned it against the Python
reference. FP6-P V keeps its own coverage, which compares against
pack_v_mxfp4_colmajor_raw rather than the canonical producer, so nothing is
lost. The benchmark's production-quantize helpers lose the fp6_p branch that
only existed to reach it.

quantize_v_mxfp6 is in the same position but stays: the FP6-P MXFP6 layout test
defines the packing as a permutation of the canonical producer's output, so it
is load-bearing as a reference even though no row consumes it. The doc now says
so.

281 passed, 1 skipped.
The shipped distributions are all far more diffuse and far more self-cancelling
than measured traces, so the quantized rows were never exercised where real
workloads live. "clustered" draws tokens from clusters to pin effective softmax
width and keep cancellation low.

"kcommon" additionally gives K a large shared direction. Softmax is shift
invariant in a component common to every key, so it changes no output and is
invisible to effective width, cancellation and logit spread, yet it multiplies
per-tensor quantization error. No other distribution covers that axis.
Softmax is shift invariant in a component shared by every key: it adds a
constant to every logit of a query and cancels exactly. Quantization noise is
not shift invariant, because per-tensor and per-block scales track the magnitude
of q.k rather than its spread across keys, so a large shared component in K
spends range on a term that contributes nothing.

Removing it is free accuracy. On a K common-mode sweep with effective softmax
width, cancellation and logit spread all held constant, FP8 error rose from
13.5% to 49.9% across the sweep; with the mean removed the same axis costs
almost nothing. Every quantized recipe gains, from 1.5x on INT8 to 5.2x on
MXFP4, and the accuracy of the diffuse distributions is unchanged.

The FP8 path fuses the subtraction into the Hadamard rotation kernel, which
already streams every element, so it adds no pass and no measurable time. The
kernel skips it when the mean tensor is empty. Recipes whose K quantizer does
not yet fuse it get a materialised K; recipes that keep K in BF16 skip it
entirely, having nothing to gain.
Nothing caught a silent loss of the K mean subtraction: the diffuse test inputs
carry almost no shared K component, so every existing case passes with or
without it. This adds the missing axis directly, asserting that a direction
common to every key does not inflate the error.

Verified to trip: with the subtraction disabled the relative error grows 6.2x
for FP8, 7.6x for MXFP8 and 12.4x for MXFP4, against the 1.5x bound asserted
here.
The MXFP8, MXFP6 and MXFP4 K quantizers were the last paths still taking a
materialised smoothed K, which cost a full read and write of K before the
rotation pass that reads it again. All three already stream every element, so
the subtraction rides along for free.

K quantization drops from 0.066/0.069/0.070 ms to 0.040/0.041/0.042 ms for
8192x8x128, a 1.66x to 1.72x speedup, with the accuracy gain unchanged. It is
also slightly more accurate than the host path, which had to round K-mean back
to bf16 before quantizing; the kernel keeps the subtraction in fp32.

The two MX K kernels already carried `sequence` and `heads` for their coalesced
output layout, so they only needed the mean pointer. The per-row indexing is now
one shared device helper rather than a copy in each of the four kernels, and the
tensor validation is one shared host helper.

INT8 still takes a materialised K: its quantizer is Triton, not HIP.
test_mha_v4.py had reached 2688 lines and ~70 tests. The sparse cases were one
contiguous 1000-line block of it with their own helpers used nowhere else, so
they move out whole: 1689 lines remain, 1137 move.

test_mha_v4_sparse_tile_scaling.py is absorbed rather than left beside it. It
existed only because it had nowhere better to live, and it reached back into
test_mha_v4.py for a helper; that cross-module import is now local.

Same 284 tests before and after.
Removing K's shared component is only a win once that component dominates. The
subtraction takes energy out of K's RMS but not out of its outliers, so a
per-tensor scale, which is set by amax, gets relatively coarser. Measured on six
real attention traces whose shared component sits at 0.29-0.52 of a typical K
row, unconditional smoothing was neutral to slightly negative across all seven
recipes, and the two per-tensor FP8 rows lost up to 1.6x on individual traces.
The rotated amax/RMS ratio tracks it: 7.4 to 8.5 on one, 8.1 to 10.3 on another.
Block-scaled recipes are immune, since each 32-element block rescales.

So gate on the shared component relative to a typical row. Below 0.7 the
subtraction is skipped and those traces are untouched, exactly 1.00x on every
recipe. Above it the win is kept in full: 1.07x at 0.73 rising to 3.5x on FP8
and 5.1x on MXFP4 at 0.99. Error after smoothing is flat in the shared
component, which is the point -- it bounds the worst case rather than improving
the typical one.

Shift-invariance holds for any constant vector, not just the exact mean, so the
constant is estimated from a strided sample. Cost is flat in sequence length at
~0.051 ms instead of a full reduction, which at S=75600 is 0.104 ms; a full
reduction over K is what made centering too expensive to keep previously. The
sampled estimate is within 0.4-1.5% of the exact mean, and closest exactly where
the gate fires. The gate is computed on device, so it adds no host sync and
fullgraph tracing still holds.
A sequence-parallel shard whose head count does not divide the parallel degree
carries head slices that are entirely zero. The per-channel FP8 V quantizer took
the amax of such a channel, got exactly zero, and divided by it, so the quantized
data came back NaN while the descale itself stayed finite and healthy-looking.

Only f6f8 could reach this: it is the one recipe whose V scale mode is
F32_PER_CHANNEL, so the only caller of quantize_v_fp8. It turned exactly half the
output of an affected shard into NaN. The clamp matches what
_mxfp4_scale_from_amax already does a few lines away in the same file.

The test covers every recipe rather than only the broken one, and needs live
heads beside the empty ones: with all heads zero the per-tensor paths look fine
too, so the populated heads are what make the failure specific. Verified to fail
on f6f8 without the clamp.
The gate was set at 0.7 from a synthetic sweep. Captures from four real diffusion
models put that threshold inside the band where recipes disagree rather than
above it: Wan reaches 0.76 and HunyuanVideo 1.5 0.67, and on the one captured
layer above 0.7 FP8 lost 7% while MXFP4 and F8F6 gained. The win is only
consistent past ~0.85, where it is already 1.16x-1.31x.

At 0.85 every layer Wan, HunyuanVideo and Flux2 produce is left untouched, and
Z-Image, whose shared component reaches 0.97, still gets smoothed: measured
there, FP8 improves 1.12x and nothing regresses. So the gate now fires only
where it pays, with no per-model configuration.
The frozen softmax raises its reference when a tile overflows the conversion
gate, rescaling O and L to match, but the register the epilogue exports as the
LSE's max term stayed at the value seeded from the first K tile. Output is
unaffected, since the rescale keeps O self-consistent, so the object passed
every correctness test while the LSE it exported was low by the accumulated
rollback.

Only rows that actually trip the gate are affected, and they are the peaked
ones: on a captured 75600-token layer that was 15.5% of rows, median effective
support 31 keys, dragging the mean LSE error to -1.14 nats while the median row
sat at +0.018.

Ring attention weights each chunk by exp(lse), so a bias that varies per chunk
does not cancel. Per-chunk spread drops 1.0222 -> 0.001 nats and the merged
output 0.929 -> 0.042 relative L2 against the same kernel run single-shot,
below the recipe's own 0.032 quantization error.

Output is bitwise identical on 4 shapes and single-vs-reference is unchanged at
0.03242; dense validation 8/8; 10 interleaved benchmark rounds move the median
1991.4 -> 1994.5 TFLOPS, inside the round-to-round spread. Register count is
unchanged at 256.
The K and V buffer descriptors size themselves from the same register the
varlen path overwrites with this batch's seqlens_k entry, so an entry longer
than the padded key length widened the descriptors instead of being caught by
them, and the tile loop read past the tensor. A length four times the real one
faulted the GPU; it now returns the dense result.

The clamp sits behind the branch a dense launch takes, so it costs two scalar
ops only when seqlens_k is supplied. The hot loop is byte-identical, register
counts stay at 96 SGPR / 256 VGPR, and 12 interleaved benchmark rounds with the
order rotated leave both objects at a 1412 TFLOPS median.

Host-side validation stays as it was: the entries live on the device, so
checking them there would synchronize every launch.
check_k_mean validated dtype, shape and contiguity but never where the tensor
lived, and the quantizers take it from the caller. A host tensor therefore
reached subtract_k_mean as a device pointer: quantize_mxfp8_k with a CPU mean
faulted the GPU rather than reporting anything. It now fails the same way every
other check in that file does, naming the problem.

All four K quantizers share this validator, so one check covers
quantize_fp8_rotated, quantize_mxfp8_k, quantize_mxfp4_k and quantize_mxfp6_k.
The test drives them out of process, since these checks abort rather than
raise; it fails against the previous module on all four.
Two spots in the K common-mode LSE test were left unformatted.
mha_v4_packed takes an LSE buffer from the caller and never checks where it
lives, so a buffer on another GPU reached a launch on this one as a device
address. The launcher already required it to be a GPU tensor; it now requires
the same GPU, matching the seqlens_k slot beside it.

The scope list still named LSE as unsupported, which has not been true since
the dense epilogue landed. Say what is actually there instead: dense rows
produce it, sorted sparse rejects it.
ruff 0.16.0, which the workflow pins, enables PLW1510, and the K-mean test
expects its subprocess to fail, so the intent is check=False rather than the
default. Black also wanted the blank line after the mha_v4_quant docstring.

ruff check . and black over the tree, minus the submodule the style jobs do not
check out, are both clean.
The gfx942 objects do carry the LSE epilogue, so the capability is not absent,
but they predate the frozen-max correction and their exported value has never
been compared against torch.logsumexp. That comparison is the only thing that
can establish it: O never reads the LSE, so the bf16fp8 row passed the entire
suite while exporting a value about a nat low, and no output test on gfx942
would say any more than it did there.

The tile is 256x64 rather than 256x128, so how often the conversion gate trips
differs as well, and that frequency is what decided whether the rollback
mattered. Raise on gfx942 rather than return a number nobody has checked.

Lift this by running the existing LSE-vs-logsumexp comparison on MI300; the
gate is one condition and the code objects already have the epilogue.
Three call sites built the same FP32 [batch, heads, Sq] tensor, and the
constraint that makes that shape load-bearing -- the kernel derives the batch
stride as heads * head stride, so the buffer must stay contiguous -- was
documented at only one of them. The two early-returning launcher branches
carried it implicitly.

State it once, where it is enforced. Those two branches each collapse from a
nine-line conditional expression to a single call.
Twelve comment blocks ran to three lines or more, the longest nineteen. Each is
now at most three, keeping what the code cannot show: what a distribution
models, the bug class it exists to catch, and the constants that give a default
its meaning, such as the rollback threshold the maxstair step is chosen against
and the reference size past which cosine stops being trustworthy.

Dropped the mechanism walkthroughs, the regression dates, and the "Tunables
(env)" tables, which restated the os.environ.get calls immediately below them
and were a second place the defaults had to be kept correct.

Comments only, no code line changed; 103 comment lines down to 58.
@jcaraban

jcaraban commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor Author

LGTM from the triton side overall. Please reduce the comments in op_tests/op_benchmarks/triton/bench_sage.py. They should be 2-3 lines maximum.

done, thanks @azaidy ! also addressed the copilot findings

@jcaraban
jcaraban requested a review from azaidy September 24, 2026 19:38
jcaraban and others added 5 commits September 24, 2026 19:40
Two checks existed twice. The GQA head rules ran once in mha_v4_packed and once
in _validate_mha_v4_raw_inputs, with the same three conditions and the same
power-of-two bound written out separately, so the supported ratio lived in two
places. The sorted-sparse LSE rejection sat beside _check_lse_capable at both
call sites rather than inside it, leaving that function able to approve a
combination the caller still had to reject on its own.

Error text is unchanged: the head helper takes the caller's operation name, so
mha_v4_packed still says "MHA v4" and the raw path still says "mha_v4".
MXFP4 and MXFP6 dense reach the kernel through their own custom ops rather than
mha_v4_packed, and each returned from the middle of the recipe dispatch. Every
field added to the call since has had to be remembered separately at those two
points: the LSE output needed threading into both launchers, and seqlens_k was
silently dropped by them because it was only forwarded on the path they skip.

Record which launcher to run instead of returning, and run it once after the
dispatch. The two branches now differ only in the launcher and its arguments,
and the LSE buffer and the K-smoothing correction are applied in one place
rather than three.

Output, LSE and the no-LSE output are bitwise identical to the previous code
for bf16, mxfp4, mxfp6, f6f4 and f6f8 across three shapes.
Five places built the same E8M0 scale tensor inline. Four were fakes, but one
was quantize_mxfp8_k, and that mattered: block_scale_storage exists because the
ASM gathers scales with unguarded loads covering a whole tile, and mxfp8_q,
mxfp4_q, mxfp4_k and mxfp6_q all allocate through it. mxfp8_k not doing so read
as an oversight when it is only visible as a bare new_empty.

Give that allocation a name and a docstring saying the padding is deliberately
absent, so the asymmetry is a stated choice rather than something a reader has
to notice. The fakes use it too, since they only have to agree on shape.

Poisoning the bytes past the logical end and rerunning leaves the output
unchanged at sequence lengths on and off the 128-row KV tile, so the shorter
allocation does not reach the result. That does not prove the kernel never
reads them, only that their value cannot matter; settling that needs a
sanitizer.

Quantizer outputs and attention output are bitwise identical.
The frozen-max rollback shipped because nothing could see it. O never reads the
LSE, so output validation was silent, and the reference check above compares a
single call against logsumexp, where a bias constant across the sequence passes
whatever its size. What ring actually needs is the other property: it weights
each chunk by exp(lse), so a bias cancels only while it is the same in every
chunk.

Random keys cannot exercise this. They leave every row diffuse, the conversion
gate never trips, and a rollback left out of the exported max changes nothing
measurable -- a chunked comparison on random input returns the same numbers
from the defective object as from the fixed one. The keys here escalate along
the sequence so tile 0 seeds the reference low and later tiles overflow it.

Bounds are per recipe because the spread is not all kernel: per-tensor scales
are refit per chunk, so the FP8 Q/K rows legitimately move 0.23 nats where the
BF16 Q/K rows, which is where the defect was, sit at 0.00. Measured 0.00 to
0.24 across the nine rows against 4.55 for the defect, so each bound keeps
roughly four times the margin in both directions.

Fails on the pre-fix code object at 4.55 nats and passes on the shipped one.
Comment thread op_tests/test_mha_v4_sparse.py
Comment thread op_tests/test_mha_v4.py
test_mha_v4_sparse.py had no benchmark at all: no reference, no candidate loop,
no summary table, and no __main__ guard, so it could only ever be run through
pytest and said nothing about what the sparse path costs. It now carries the
standard shape -- a masked Torch reference, a @benchmark function whose call
args are the table's left columns, run_perftest plus checkAllclose per
candidate, us with TFLOPS and TB/s, an itertools.product sweep in main(), and a
markdown table at the end.

Density is a swept axis because it is the only number that decides how much
work the kernel skips, and the FLOP and byte counts are taken from the tiles a
row actually selects rather than the full key length. Dense joins as a second
candidate only where the mask selects every tile: anywhere else it answers a
different question than the reference and would read as fast and correct at
once.

test_mha_v4.py already had the structure but timed a single BF16 candidate,
which is the one thing the recipe work in this branch cannot be judged by. All
nine dense recipes are candidates now, against one reference, so a row shows
what each costs at the same shape. Their quantizers are inside the timing
because mha_v4 runs them per call. The default sweep moves to 1024 queries
against 1024 and 4096 keys; at the old 128-256 sizes quantization dominated and
the table said more about setup than about attention.

Both files: 339 passed, 1 skipped.
@jcaraban
jcaraban requested a review from valarLip September 25, 2026 15:30
@valarLip

Copy link
Copy Markdown
Collaborator

E AssertionError: Tensors not close enough! 0.750732% elements exceed tolerance.
E Greatest absolute difference: 2.5625 at index (68, 0, 0, 124) (up to 0.3 allowed)
E Greatest relative difference: 2572288.0 at index (68, 0, 0, 124) (up to 0.25 allowed)

op_tests/triton_tests/attention/test_fav3_sage.py:163: AssertionError

@jcaraban
jcaraban dismissed azaidy’s stale review September 28, 2026 06:56

comments addressed

@jcaraban

Copy link
Copy Markdown
Contributor Author

E AssertionError: Tensors not close enough! 0.750732% elements exceed tolerance. E Greatest absolute difference: 2.5625 at index (68, 0, 0, 124) (up to 0.3 allowed) E Greatest relative difference: 2572288.0 at index (68, 0, 0, 124) (up to 0.25 allowed)

op_tests/triton_tests/attention/test_fav3_sage.py:163: AssertionError

Thanks @valarLip ! This PR didn't touch the Triton Sage kernels, so I'm surprised it failed. Those kernels were developed under Triton==3.5.0 ; and >=3.6 is known to change their performance profile. It is possible that some race condition has been uncovered now with the environment change.

@juuso-oskari @Chi-Chu319 @hellozhuo-amd can you investigate this?

@jcaraban
jcaraban merged commit 105615e into main Sep 28, 2026
85 of 86 checks passed
@jcaraban
jcaraban deleted the mha_v4_fixes_sparse_etc branch September 28, 2026 08:07
@Chi-Chu319

Copy link
Copy Markdown
Contributor

E AssertionError: Tensors not close enough! 0.750732% elements exceed tolerance. E Greatest absolute difference: 2.5625 at index (68, 0, 0, 124) (up to 0.3 allowed) E Greatest relative difference: 2572288.0 at index (68, 0, 0, 124) (up to 0.25 allowed)

op_tests/triton_tests/attention/test_fav3_sage.py:163: AssertionError

What is the exact test case that failed and what environment you had?

vgokhale added a commit that referenced this pull request Oct 2, 2026
* [Config] Add gfx942 a8w8 blockscale GEMM tunings for Qwen3/Qwen3.5/GLM/DSV4 shapes (#5839)

MI325X (gfx942, 304 CUs): 286 (M, N, K) rows over 28 weight shapes that no
existing model_configs table covers, tuned on main with
gemm_a8w8_blockscale_tune.py --libtype ck --splitK. Keys already present in
model_configs and rows slower than the heuristic default are left out.

The rows go into each model's existing model_configs table; Qwen3-14B and
Qwen3.5-122B-A10B get new tables. A shape two models share is kept in one
table only, since the config merge rejects duplicate keys across files.

Co-authored-by: Cursor <cursoragent@cursor.com>
(cherry picked from commit 8260de6)

* [HIP] [OPUS] [FlyDSL] gfx950 MXFP8 e8m0 GEMM + DeepSeek-V4/V4.1 tuned configs (#5896)

* feat(flydsl): gfx950 MXFP8 e8m0 GEMM on FlyDSL + DeepSeek-V4/V4.1 configs

- New FlyDSL gfx950 MXFP8 batched GEMM (bmm_a8w8_mxscale_gfx950): scaled
  MFMA 16x16x128, async LDS DMA pipeline, 32x32 / 128x128 e8m0 blocks,
  split-K with same-XCD last-arrival reduction, B direct to registers, XCD
  tile order, non-temporal B, per-stage or preloaded scale panels, K % 64
  tail, column-major (blockscale) x_scale. Scale rows that are not whole
  dwords (e.g. 128-wide blocks at K = 384 / 768) are copied a byte per LDS
  slot with exact buffer bounds. check_bmm_config is the single legality
  source.
- Arch-neutral front door flydsl.batched_gemm_a8w8 dispatching by (arch,
  w_scale block, x_scale layout); gfx950 module with kernelName parsing,
  untuned-shape heuristic and a compiled-launcher cache (host 71 -> 21 us).
- gemm_a8w8_blockscale_bpreshuffle on gfx950 routes e8m0 x_scale + w_scale
  to it through the tuned table AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_MXSCALE_
  BPRESHUFFLE; FP32 scales keep CK / asm. gemm_a8w8_blockscale with
  isBpreshuffled forwards native group32 operands there. The V4 wo_a
  batched GEMM uses the same kernel.
- AOT precompiles kernelId=bmm rows (both x_scale layouts for 128x128).
- Tuned gfx950 configs: DeepSeek-V4-Pro 128x128 linears (980 rows, TP1/4/8,
  1.30x vs CK/asm), DeepSeek-V4.1 group32 (1224 rows, 1.22x vs Triton),
  V4 wo_a batched (680 rows).
- Also carries the other local changes in the tree: inverse_rope_group_quant,
  opus policy / bmm tune, gfx1250 bmm wrapper, tensor_shim helpers.

* refactor(flydsl): tidy gfx950 MXFP8 bmm kernel per FlyDSL cleanup guide

- Use fx.copy instead of fx.copy_atom_call for the single-atom copies.
- Build the async-LDS DMA destination from the LDS pointer with
  fx.add_offset + fx.to_llvm_ptr instead of ptrtoint/inttoptr with a
  hand-picked address space. The raw buffer_load_async_lds stays: the
  BufferLoadAsyncLDS atom has no 1-byte size and no cache-policy operand.
- Factor the compile-and-run + leaked ir.Context recovery out of
  tensor_shim._run_compiled into _compile_and_run, and use it for the
  bmm wrapper's per-config compile cache too.

Generated ISA is byte-identical on 9 reference configs.

(cherry picked from commit 40d524b)

* [Triton/Gluon] [ASM] [HIP] Mha v4: adds bf16 sparse, LSE support, KV varlen, fixes, etc (#5798)

Motivation: adds log-sum-exp output to MHA v4 gfx950 kernels (only dense variants now), so it can run under ring / context parallelism, which merges per-rank partials via LSE. Also extends block-sparse to the BF16 recipes, canonicalizes the MXFP4 rows, and improves K/V quantization accuracy.

# Kernels:
- New BF16 and BF16FP8 sorted-sparse kernels.
- Dense LSE epilogue on all ten hd128 recipes, gated at runtime on s_lse.
- f8f6, f6f4 and mxfp4 sparse rows moved to FP6-P V, matching their dense siblings.
- Dense mxfp4 claims the canonical FP6-P V order; the duplicate f4f4 row is disabled.

# Host:
- bugfix: per-channel V amax clamped so empty heads cannot quantize to NaN.
- Optional lse output on mha_v4 / mha_v4_packed, plumbed through asm_mha_v4_fwd.cu.
  + ABI unchanged: ptr_lse / s_lse / s_lse_Hs were already reserved in the kernarg.
- _LSE_CAPABLE_QV gates the supported format pairs; sorted-sparse still raises.
- K mean is subtracted before quantizing (K-smoothing), fused into the MX quantizer kernels.
- Dense MHA v4 accepts per-batch key lengths (ragged seqlen_k).

# Minor:
- Retired MXFP4 Q/K + FP8 V and the deprecated mha_v4_mxfp8 alias.
- bench_sage.py: improve input distributions, diffusion-calibrated default, BF16 sparse modes.
- Split the block-sparse cases into op_tests/test_mha_v4_sparse.py.

(cherry picked from commit 105615e)

* [tuner] Fail a task as soon as its worker process exits (#5841)

A GPU memory fault aborts the worker process, so its task never returns a
result. mp_tuner only noticed at the task timeout (1800s by default) or
never without one. Record which worker started each task and, when that
process is gone, fail the task and restart the pool right away, like the
existing accelerator-error path. Fixes #5840.

Co-authored-by: Cursor <cursoragent@cursor.com>
(cherry picked from commit bcb56d9)

* [CI] Update Aiter artifact downloads to v8.0.1 (#5908)

(cherry picked from commit c03689f)

* [MLA v4 nm] Test fix _run_one_point reading packed BF16 as FP32 partials (#5905)

* [MLA v4 nm] Fix _run_one_point reading packed BF16 as FP32 partials

Since #4311 the v4 nm dispatcher derives out_16_nosplit from num_kv_splits
and ignores the caller's value, so a single-split launch writes packed BF16
into the logits buffer. _run_one_point still read logits_buf[:, 0] as FP32
for num_kv_splits == 1, which compared every other BF16 element (plus the
never-written tail of the buffer) against the reference and printed spurious
`fp8_dequant_ref vs asm` checkAllclose failures in the script-mode sweep.

Read output_buf, which holds the final result for every split count.

Co-authored-by: Cursor <cursoragent@cursor.com>

* [MLA v4 nm] Gate accuracy checks so a failed! fails the test

checkAllclose only raises on a catastrophic delta; otherwise it logs
`failed!` and returns the mismatch fraction. Four of the six accuracy
checks in test_mla_v4_nm.py dropped that return value, which is how the
packed-BF16 readback bug fixed in the previous commit passed both pytest
and the script-mode CI run.

Route them through _gated_allclose, which asserts the mismatch fraction
against the same tol_err_ratio checkAllclose uses for `failed!`. The
script-mode sweep keeps going past a failing shape, prints a summary, and
exits non-zero so aiter_test.sh reports the file as failed.

Co-authored-by: Cursor <cursoragent@cursor.com>

---------

Co-authored-by: Cursor <cursoragent@cursor.com>
(cherry picked from commit 049fae4)

* [AOT] Inline the FlyDSL FP8 FMHA head shapes, drop the config CSVs (#5901)

Follow-up to #5796. The AOT job list for the gfx950 FlyDSL FP8 flash
attention was driven by a header-only aiter/configs/fmha_fp8_aot.csv
merged with aiter/configs/model_configs/*_fmha_fp8_aot.csv. A single row
does not justify the CSV plumbing, so the head shapes now live in a
DEFAULT_SHAPES table in the module, the way mega_moe.py already does it,
and the comments introduced by #5796 are trimmed.

- fmha_fp8.py: DEFAULT_SHAPES replaces parse_csv()/DEFAULT_CSVS;
  default_jobs() replaces the CSV walk; --csv is gone (--shape still
  overrides). The cu_num column is dropped with the CSVs: the kernel is
  gfx950-only and AOT_ARCH is fixed, so every non-gfx950 row was warned
  about and skipped anyway.
- common.py: FMHA_FP8 returns default_jobs() next to MEGA_MOE instead of
  going through collect_aot_jobs().
- jit/core.py: drop AITER_CONFIG_FMHA_FP8_AOT and its config-file
  property, now unused.
- README.md: document the table instead of the CSVs.

_variant_space()/jobs_for_shape() are unchanged, so coverage is
unchanged: --list still emits the same 92 kernel names for Kimi-K3 TP8
(12:12:192:128 varlen_cross) as the CSV-driven version, and setup.py
source builds still compile them via run_aot().

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
(cherry picked from commit 71a31be)

* [FlyDSL] gfx942 fp8_mqa_logits: let _auto_variant choose rows_per_block (#4963)

* [FlyDSL] gfx942 fp8_mqa_logits: let _auto_variant choose rows_per_block

_auto_variant returned f"mfma_r2_w{wpb}", so rows_per_block was pinned at 2 and
seven of the nine registered variants -- including every member of the r4 family
-- could never be selected.

r4 amortizes each KV tile load over twice as many query rows and is faster from
seq_len 8 upward. The gap is widest exactly where it costs most: vLLM chunks
indexer prefill to fit VLLM_SPARSE_INDEXER_MAX_LOGITS_MB (512 MB), which caps
seq_len at 1024 when seq_len_kv is 131072, so the existing seq_len >= 2048 branch
cannot fire at long context and every such call took mfma_r2_w4 -- the median of
the nine by speed, with the best 1.65x faster.

Measured on MI325X (gfx942), seq_len_kv 131072, best variant vs the r2 pick:

  seq_len    1   r2 25.2 us    r4 79.1 us     r4 3.1x worse
  seq_len    4   r2 24.3 us    r4 31.5 us     r4 1.3x worse
  seq_len    8   r2 34.0 us    r4 32.9 us     r4 1.03x better
  seq_len   16   r2 56.5 us    r4 47.5 us     r4 1.19x better
  seq_len 1024   r2 2564.3 us  r4 1520.3 us   r4 1.69x better

Below seq_len 8 the host padding of seq_len up to a multiple of RPB dominates --
at seq_len 1 an r4 kernel computes 4 rows to obtain 1 -- so r2 is kept there and
behaviour for those shapes is unchanged.

Logits are bitwise identical across all variants at every shape tested, so this
is purely a blocking/occupancy change.

End to end on 8x MI325X, TP8, GLM-5.2-FP8, 131072 in / 1024 out, concurrency 8:
median TPOT improves 7.22% and output throughput 6.60%. This kernel is 16.8% of
GPU time at that point.

Signed-off-by: Jin Tao <jin.tao@amd.com>

* [FlyDSL] gfx942 fp8_mqa_logits: pick RPB on element count, and keep it a divisor

Refines the previous commit's rule after a 2-D sweep. That rule keyed RPB off
seq_len alone with a crossover measured only at seq_len_kv=131072; sweeping the
other contexts shows the crossover is not a seq_len threshold at all, and that a
second effect was being read as one.

RPB tracks the logits element count. Over seq_len 1..8192 x seq_len_kv
1024..262144 on MI325X, the boundaries land on the same element count at every
context: RPB=1 wins below 2**19 elements (27/27 shapes), RPB=2 at 2**19 (6/6),
RPB=4 from 2**21 up (38/38), with 2**20 a transition band split 3/3. Keying off
seq_len instead put the previous rule on the wrong side at low context: at
seq_len 16, seq_len_kv 1024 it chose RPB=4 and ran 1.26x slower than RPB=1.

RPB must also divide seq_len. When it does not, the launcher pads with four
torch.cat calls; that is a flat ~44 us of host-side overhead, independent of
seq_len_kv, and it is the whole of the "small seq_len" penalty the previous
commit attributed to wasted rows. At seq_len 1, seq_len_kv 131072: RPB=1 23.1 us,
RPB=2 67.8 us, of which the four cats are 44.1 us and pre-padding by hand
recovers all of it (21.9 us). So the penalty is not proportional to the padding
-- 1 wasted row of 2 costs the same as 3 of 4 -- and it applies to every odd
seq_len, which the old rule sent to RPB=2 unconditionally.

Stepping down to a divisor is only right while the kernel is cheap relative to
that fixed cost, so it is gated to the same 2**21 elements: at seq_len 1025,
seq_len_kv 131072 the dividing RPB=1 takes 3880 us against 2601 us for a padded
RPB=2.

Measured on MI325X, no FLYDSL_FP8_MQA_LOGITS_VARIANT set, old pick vs new:

  seq_len  seq_len_kv        old         new   speedup
     1024      131072   r2_w4 2598.7   r4_w4 1623.5    1.60x
     1025      131072   r2_w4 2609.3   r4_w4 1594.0    1.64x
      512      131072   r2_w4 1190.1   r4_w4  795.3    1.50x
      700       50000   r2_w4  661.6   r4_w4  440.3    1.50x
      333       12000   r2_w4  118.4   r4_w4   97.1    1.22x
       16      131072   r2_w4   58.3   r4_w4   47.4    1.23x
        3      131072   r2_w4   66.4   r1_w4   28.1    2.36x
        1      131072   r2_w4   69.4   r1_w4   23.2    2.99x
        1        1024   r2_w4   66.2   r1_w4   22.2    2.98x

No shape measured regresses; the smallest gain is 1.04x. Against the best of the
nine variants at each shape, pooled over held-out data (non-power-of-two shapes,
a fine seq_len sweep, and head counts 16 and 64), the geometric mean cost falls
from 1.45x to 1.03x and the worst case from 3.17x to 1.41x.

Logits are bitwise identical across all nine variants at all 180 shapes swept
(1620 timings), so this remains purely a blocking/occupancy change.

WPB is deliberately left alone. It is worth a few percent at most here, and
unlike RPB its optimum moves with the head count -- at 64 heads the current
WPB rule costs 1.61x worst case where a fixed WPB=4 costs 1.07x -- so it needs
its own sweep rather than a change fitted to one head count.

Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>

* [FlyDSL] gfx942 fp8_mqa_logits: trim _auto_variant comments per review

Move RPB element-count thresholds into _auto_variant and shorten the
docstring; logic unchanged.

Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>

* [FlyDSL] gfx942 fp8_mqa_logits: scope the _auto_variant step-down note

The docstring stated the divisor step-down as an unconditional rule, but it
only applies in the middle band: above the top threshold RPB stays 4 and the
padding is accepted. Say which band it applies to, and why the top band is
exempt. No logic change.

Co-authored-by: Cursor <cursoragent@cursor.com>

* [FlyDSL] gfx942 fp8_mqa_logits: unit-test the variant selector

The shape sweep in test_flydsl_fp8_mqa_logits.py never reaches r4 -- its
largest default shape is 1024 x 1560, below the 2**21 threshold -- and nothing
asserted _auto_variant directly, so a regression could re-pin RPB to 2 with
every correctness test still green.

Cover the RPB bands at both thresholds, that the edges track
seq_len * seq_len_kv rather than seq_len alone, the middle-band step-down on
odd seq_len (and that it stops above the top threshold), the production
1024 x 131072 shape, the unchanged WPB rule, and _resolve_variant precedence
so the auto path is confirmed to be the default. Pure shape arithmetic, so no
kernel launch.

Co-authored-by: Cursor <cursoragent@cursor.com>

* [FlyDSL] gfx942 fp8_mqa_logits: put selector tests on the CI path

CI only discovers op_tests/test_*.py, so the r1/r2/r4 selector coverage
never ran. Move it there, run pytest from __main__, and add two r4
shapes to the GPU sweep so auto-selected padding is launched.

Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>

* [FlyDSL] gfx942 fp8_mqa_logits: fold the selector pins into the op test

aiter op tests are plain scripts, not pytest, so drop
test_flydsl_fp8_mqa_logits_variant.py and pin the gfx942 auto-selected
variant in test_flydsl_fp8_mqa_logits.py instead: a host-only
verify_auto_variant table at both RPB band edges, either side of the
odd-seq_len step-down, the long-context prefill shape and the WPB switch.
It runs in the default verify scenario on gfx942, and its failures count
toward the same exit code as the kernel sweep. The selector import sits in
the file's existing ImportError guard.

Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>

---------

Signed-off-by: Jin Tao <jin.tao@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: jin.tao@amd.com <jin.tao@amd.com@tus1-p15-g43.tus.tensorwave.lan>
Co-authored-by: Felix Li <felix.li@amd.com>
(cherry picked from commit bbaccdf)

* [Triton] [MHA] Add a tuned gfx1101 config, split small_head/default (#4493)

gfx1101 (RDNA3, e.g. RX 7800 XT) ships no MHA config, so `_get_config` in
`_triton_kernels/attention/mha.py` finds no
`configs/gfx1101/triton/attention/mha/DEFAULT.json` and every call to
`aiter.ops.triton.attention.mha.flash_attn_func` fails before a kernel runs.
Of the architectures in RDNA_ARCHS, only gfx1151 ships one today.

Nine of the eleven entries are taken verbatim from the gfx1151 donor
(RDNA3.5, added in #3423, tuned in #3560), which is the nearest tuned
architecture. Two forward entries are tuned on gfx1101 instead of inherited,
and they differ from each other only in `num_warps` and `num_stages`:

  fwd/default     BLOCK_M 128, num_warps 8, num_stages 3   (donor: 64 / 4 / 2)
  fwd/small_head  BLOCK_M 128, num_warps 4, num_stages 1

The split uses the `small_head` bucket added in #4414, which is opt-in per
architecture by the mere presence of the key, so this stays a data-only
change. It is needed because one `fwd/default` cannot serve both halves on
this card: measured against the donor, `M128 w4 s1` is 0.935x on head_dim 64
but 1.269x on head_dim 128, while `M128 w8 s3` is 0.914x on head_dim 128 but
1.062x on head_dim 64.

Measured on Windows native ROCm, triton 3.8.0, fp16, 10 independent repeats
of 20 iterations after 5 warmups, `torch.cuda.synchronize()` per iteration;
a result counts only when the [min, max] intervals across repeats are
disjoint. Ratios are against the gfx1151 donor entry, i.e. against what a
donor-inherited config would do.

  head_dim 128, `default`     Flux joint 0.914x / 0.891x / 0.892x on three
                              independent torch+ROCm stacks (2.11/7.15,
                              2.11/10.1, 2.15/10.1), all pinned to the same
                              triton; llama3-8B 0.882x, mixtral-7B 0.880x,
                              kimik25-tp4 0.849x at seqlen 16384
  head_dim <= 64, `small_head` SDXL self-attn 0.935x; deepseek-V3 0.658x and
                              glm47fp8-tp4 0.857x at seqlen 16384

The LLM shapes come from `op_tests/op_benchmarks/triton/utils/model_configs.json`
(prefill, causal, GQA, batch 1); the sequence lengths are not in that file and
are chosen here. Not covered: batch > 1, varlen/thd, sliding window, decode.

Note that the `small_head` comment in `_get_config` does not describe gfx1101.
It states that 16 < d <= 64 suffers a num_stages=1 pipelining pathology which
num_stages=3 cures. On this card the ordering is the opposite -- on the tuned
M128/N32/w4 tile, num_stages 1/2/3 measured 2.300 / 2.399 / 2.492 ms. The
bucket is still the right mechanism here, for a different reason: the two
head_dim ranges want a different `num_warps`, not a different `num_stages`.

Signed-off-by: Martin Domanský <ragua@email.cz>
Co-authored-by: Mehmet Cagri <mehmet.kaymak@amd.com>
(cherry picked from commit f1f95ce)

* [Triton/Gluon] Consolidate tuning harnesses (#5874)

* [Triton/Gluon] Consolidate tuning harnesses

* [Triton/Gluon] Simplify tuning harness

(cherry picked from commit af3514a)

* [Bugfix][Gluon][MLA] Follow-up: fix stale comments + test hardening (#5860)

Addresses review feedback on #5648 (kept separate to not disturb the approved PR):
- mla_gluon.py: the >2GB global_load calls now carry a bounds mask + other=0.0,
  so the old "No mask needed / in-bounds" comments above them were stale and
  gave the opposite (unsafe) guidance. Replace with an accurate one-liner.
- test_mla.py: allocate the output with device=q_nope.device instead of relying
  on the module default device; call torch.cuda.empty_cache() before the
  mem_get_info() free-memory gate to avoid allocator-state flakiness; and assert
  exact parity (atol=rtol=0) between the >2GB and <2GB paths, which read
  identical KV and must be bit-identical (a loose tolerance could hide a
  regression).

Signed-off-by: Rohan138 <rohanpotdar138@gmail.com>
(cherry picked from commit a966245)

* [Triton/Gluon] [Config] gemm_a16w16: gfx950 tuned per-shape configs (#5830)

* [Triton] gfx950: tuned defaults for gemm_a16w16, gmm, MHA fwd and PA decode

Tuned on MI355X (gfx950) and validated for correctness and no regressions on
broader shape sets than tuned. All changes are gated to gfx950 configs/arch.

- gemm_a16w16: per-shape configs for N,K = 2048/2048, 4096/4096, 8192/8192,
  10240/8192, 57344/8192, 8192/28672 (non-persistent and persistent). Only the
  tuned M bucket differs from DEFAULT.json. 1.07-2.5x.
- gmm: new "large_kn" config (8 warps), selected for K, N >= 4096 with
  >= 256 rows per group. 1.14-1.26x there; smaller problems keep "default".
- MHA fwd (Triton): new "mid_head" config (128x128, 2 stages) for bf16/fp16
  with 64 < head_dim <= 128. 1.03-1.16x; d<=64 and d>128 unchanged.
- PA decode (Triton): use v2 whenever there is more than one partition for
  bf16/fp16 KV. 1.9-5.3x for the batch/context sizes that previously hit v1.

Unit tests: test_gemm_a16w16 360 passed, test_gmm 48 passed,
test_pa_decode 816 passed, test_mha 2060 passed.

* [Triton] pa_decode: sort imports (ruff I001)

* Move gmm, MHA fwd and PA decode changes to their own PRs

Per review, each kernel gets its own PR; this PR keeps only the gemm_a16w16
tuned configs. The removed changes are on branches gfx950-tuned-gmm,
gfx950-tuned-mha-fwd and gfx950-pa-decode-v2.

* [Triton] gemm_a16w16_persistent: gfx950 tuned config for N=K=2048

(cherry picked from commit a651db0)

* replacing the fp32 mfma with two bf16 mfma. fn is split inside the kernel into hi = bf16(fn) and lo = bf16(fn -hi). a simple bf16 downcast for fn inside the kernel drops accuracy more and performs worse for large M. also change the concatenation of the four streams from column major to stream major (k = stream * TILE_K + column). retuned the gfx950 configs. (#5885)

(cherry picked from commit f10cd2a)

* perf(fused-moe): add tuned DSV4.1 TP4 configs (#5919)

(cherry picked from commit 977ae79)

* [Config] Add Qwen3.8-27B TP1 a8w8 blockscale GEMM tunings for gfx942 (#5585)

* [Config] Add Qwen3.8-27B TP1 a8w8 blockscale GEMM tunings for gfx942

#3324 covered this model family at TP=2/4/8 only, so the five widths
Qwen3.8-27B drives at TP1 have no tuned entries for gfx942/cu_num=304
and every call falls back to the default -- 275 "not found tuned
config" messages per profile run, now 4. This adds 641 rows in one
file, 249 decode (M <= 512) and 392 prefill (M 907-65536), from a shape
list extracted out of real vLLM server logs rather than generated as an
M ladder.

Decode, against AITER's own default, measured in a real vLLM serving
run -- Qwen3.8-27B-FP8 at 128K context, max-concurrency 1, on one
MI325X, byte-identical images, the config bind-mount the only variable,
three interleaved repeats per arm on an idle node:

  decode step   15.444 ms -> 12.449 ms   1.24x

Per-shape decode GEMM vs the default: down 2.33x, out_proj 1.51x,
in_proj 1.12x, qkv 1.00x, gate_up 1.00x. No decode shape regresses.
Decode-step spread across repeats was 0.037 ms and 0.012 ms.

The 392 prefill rows ship and are tuned, but their end-to-end effect is
currently unmeasurable on this workload: gemm_a8w8_blockscale silently
stores only M*N mod 2^31 output elements once M*N reaches 2^31, and at
a 65,536-token prefill chunk the 34,816-wide gate_up projection crosses
that limit. The truncation depends only on M*N and not on which kernel
the config selects -- AITER's default truncates identically -- so it is
neither introduced nor worsened by this change.

Tuned with gemm_a8w8_blockscale_tune.py --libtype ck --splitK, decode
rows then re-timed interleaved because the tuner's min() over short
samples is noisy on near-ties; only rows beating the default by >=3%
were kept, following #5421. test_csv_validation (15 tests) and
test_config_shape_collision (17) both pass, there are zero duplicate
(gfx, cu_num, M, N, K) keys against the 12 other sources in the runtime
merge set on main, and max relative Frobenius error is 4.11e-3 against
bf16's 2^-8 = 3.9e-3 rounding floor.

Scope: only the decode rows have in-situ evidence, and MI300X shares
the gfx942/304 key but is untested here. Two rows pin AITER's default
kernel (qkv and gate_up at M=1) after the trace measured the tuned
picks slower -- pinned, not deleted, because getPaddedM(1,N,K,0) == 16.
Six decode near-ties ship no row. gate_up ships 72 of 80 prefill rows;
the 8 missing have M*N > INT32_MAX and fault during tuning, but inherit
the tuned M=8192 row at runtime via the same padding collapse.

Signed-off-by: Pham Binh <phamhuuthanh.binh@amd.com>
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
Co-authored-by: Cursor <cursoragent@cursor.com>

* [Config] Thin Qwen3.8-27B TP1 GEMM rows to the padded M ladder

Lookup already rounds M through getPaddedM, so a row per traced M is redundant.
Keep the power-of-two M list used by the other gfx942 blockscale tables, and
omit shapes that main already ships.

---------

Signed-off-by: Pham Binh <phamhuuthanh.binh@amd.com>
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: akii96 <aakif.nawaz@amd.com>
(cherry picked from commit 330f127)

* [HIP] [OPUS] [JIT] Unify gfx950 + gfx1250 MXFP4 paged MQA-logits into one module, schedule and API (#5761)

Fold the gfx950 (#5332) and gfx1250 (#5656) OPUS MXFP4 paged MQA-logits ops into one
implementation. The device math of both bodies is unchanged.

- One JIT module, module_pa_mqa_logits_mxfp4_opus. Sources move to
  csrc/kernels/opus_mqa_logits/pa_mqa_logits_mxfp4/:
  - _opus.h: ABI, kargs, traits;
  - _sched.cuh: shared builder;
  - _gfx950.cuh / _gfx1250.cuh: arch bodies, each an empty stub on the other arch;
  - _kernels.cu: launcher.
  Namespaces are opus_logits::gfx950 / ::gfx1250, with generic names (pa_mqa_logits_mxfp4_traits
  and _kernel). The per-arch sources and module_pa_mqa_logits_mxfp4_gfx1250_opus are deleted.
- One schedule: gfx1250's build_tiles + build_sched serve both arches. build_tiles is skipped at
  q_per_block == 1, where the cut is the identity, which saves one launch per forward on every
  gfx950 call and on gfx1250's MTP=1 decode.
- One runtime-arch dispatch: fwd_sched probes the arch once per process, then dispatches on
  (q_per_block, block_k). An unmatched config raises.
- Host-visible bounds on both arches:
  - num_rows is checked against q, weights and out, and against the lengths of local_ends,
    local_starts and row_to_batch;
  - block_tables width is checked in KV tiles;
  - both kernels drop records with row_id >= num_rows.
  The caller contract (four conditions) is documented in the module docstring and the header.
- One Python API (aiter.ops.opus.pa_mqa_logits_mxfp4: plan_buffers / plan / pa_mqa_logits_mxfp4),
  with per-arch MqaLogitsVariant instances. gfx1250: qlen4_kv64 / qlen1_kv64. gfx950: qlen1_kv64 /
  qlen1_kv256, cta_resident 1024.
  - Retired: #5332's gfx950 entry points (pa_mqa_logits_mxfp4_sched / _build_sched /
    _sched_slots / _sched_buffer_ints).
  - Kept: #5656's public gfx1250 API, except the gfx950-only plan(block_k=).
- -mllvm -enable-post-misched=1 is applied on gfx950 only. It is load-bearing there; on gfx1250
  it slowed chunked prefill.
- kargs drops six unread fields (144 -> 112 B).
- One op test for both arches, op_tests/test_pa_mqa_logits_mxfp4_opus.py. It adds ATOM-shaped
  padded decode, host-raise and raw row-guard cases, a real NaN control, and on gfx950 a FlyDSL
  cross-check.

(cherry picked from commit 9b3885f)

* [FlyDSL] feat(mega_moe/gfx1250): bind mori tokoff-ext allocator on the mori dispatch path (#5810)

* feat(mega_moe/gfx1250): bind mori tokoff-ext allocator on the mori dispatch path

mori's op layer builds the tokoff-ext slot allocator in
EpDispatchCombineOpHip.__init__; driving EpDispatchPlan directly bypasses
that and leaves EpArgs.tokOffPeers null, keeping dispatch on the
serializing cco-window atomic. Build it here from mori's TokOffExt, gated
exactly like mori (default on; MORI_EP_TOKOFF_EXT=0 opts out), pass its
peer-pointer array to plan.launch, and free it in close().

Needs mori with a public TokOffExt (dispatch_combine_v2.hip_backend).

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>

* style(mega_moe/gfx1250): black-format the tok_off_peers ternary

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Co-authored-by: HaonanWang98 <hwang@amd.com>
(cherry picked from commit 475cf0f)

* [Triton] Do not write segment softmax state at NUM_SEGMENTS_PER_SEQ == 1 (#5935)

mla_decode_fwd splits each sequence's KV into NUM_SEGMENTS_PER_SEQ
segments and lets a reduce kernel merge the per-segment softmax max and
expsum. At one segment there is nothing to merge, so the host skips the
reduce kernel and hands both scratch pointers the output buffer itself:

    else:
        segm_output = out
        segm_max = out  # dummy ptr
        segm_expsum = out  # dummy ptr

The decode kernels store M and L unconditionally, so at one segment they
write the softmax state straight over the attention output. Nothing
faults and nothing warns; the result is simply wrong.

Guard both stores. NUM_SEGMENTS_PER_SEQ is a constexpr, so the branch
costs nothing when it is greater than one.

Only the gfx1250 Gluon kernel reaches this today: select_3d_config ends
its gfx12 branch at max(1, ...), while the other branch floors at
MIN_SEGMENTS >= 8. The segment count falls as batch x heads grows, so on
gfx1250 it reaches one at large batch -- a DeepSeek-R1 TP2 serving run at
--max-running-requests 256 sits in that range throughout, and its decode
is silently corrupted. The plain Triton kernel shares both the unguarded
store and the host-side aliasing, so guard it as well, before a future
tuning change makes it reachable.

test_mla_decode_fwd stays green either way, which is why this went
unnoticed. Its grid does reach one segment for the larger head count,
but the corruption lands only at out.flat[token * num_query_heads +
head] -- one element in kv_lora_rank, measured at 0.04% to 0.12% of the
output -- while the assertion allows a tol_err_ratio of 0.01. Catching
it needs a case with no error budget.

Checked on gfx1250. DeepSeek-R1 TP2 at page size 64, gsm8k over 2000
questions, scores 0.945 with the fix. A standalone sweep over batch 1 to
256, bf16 and fp8 e4m3 caches, sequence lengths that are not page
multiples and pages scattered through the pool gives a worst relative
error of 0.0033 for decode and prefill alike. Compared against a torch
reference with no error budget, a one-segment batch differs at a 0.08%
ratio before this change -- softmax maxima around 100 sitting where
attention outputs near 0.02 belong -- and matches to the last element
after it. The plain Triton kernel, forced on by disabling the gfx12
branch, matches the same reference at batch 1 to 256, confirming the
guard leaves the multi-segment reduce path intact. test_mla.py is
unchanged: its 128 pre-existing failures are all shuffled_kv_cache=True
on the pipelined kernel, identical before and after.

Signed-off-by: Lin, Soga <soga.lin@amd.com>
(cherry picked from commit 2b6ff3d)

* [Triton/Gluon] Enable unified-attention skip-mask for gfx950 hd256 FP8 prefill (#5598)

Long FP8 prefill at head size 256 matches `D_GEQ_256` (`BLOCK_M=16`, no `SPLIT_UNMASKED_LOOP`) because D outranks Q in the lookup, so it never reaches the hd128 skip key. Add the same `Q_GEQ_256` family under `D_GEQ_256`, with `SW` and `SHUF` siblings left skip-off, so the existing kernel fast path fires only for non-windowed, non-shuffled prefill.

(cherry picked from commit e289613)

* [Triton/Gluon] [Kimi-K3][ROCm] Add merged MoE front (#5321)

* feat(kimi-k3): add minimal large-M MoE front

Signed-off-by: jiacao-amd <jiahui.cao@amd.com>

* feat(kimi-k3): enable merged front for M=7 decode

Signed-off-by: jiacao-amd <jiahui.cao@amd.com>

* feat(kimi-k3): tune merged front for M=14 decode

Signed-off-by: jiacao-amd <jiahui.cao@amd.com>

* perf(kimi-k3): tune merged front decode bucket

Signed-off-by: jiacao-amd <jiahui.cao@amd.com>

* test(kimi-k3): validate the merged front across the MTP decode bucket

test_decode_full_front_matches_reference stopped at m=16, so the shapes an MTP
server actually decodes at were never checked against the torch.mm reference:
with num_speculative_tokens=3 a pure-decode step submits num_seqs * 4 tokens,
which is 32 and up for any server past four concurrent sequences.

Extend the parametrization to 32, 48, 64, 80, 96, 112, 128 and 192 -- the shapes
the companion vLLM change enables, and for which kimik3_bf16_tuned_gemm.csv
already ships tuned solutions at N=6016, K=7168, gfx950, cu_num=256. No kernel
or config change is needed; this only closes the correctness gap left by the
old parametrization.

gfx950: 9 passed -> 17 passed.

Signed-off-by: jiacao-amd <jiahui.cao@amd.com>

* style(kimi-k3): satisfy the pinned ruff on this PR's own files

`ruff==0.16.0`, the version .github/workflows/pre-checks.yaml pins, reports two
errors in files this PR introduces: RUF022 on the `__all__` list in
kimi_k3_moe_front.py and I001 on the import block in its test. Both are
autofixes and neither changes behaviour -- `__all__` ordering only affects
`import *`, and every symbol here is imported by name.

The two EXE001 findings that remain in the repo are in unrelated flydsl files
and predate this PR.

Signed-off-by: jiacao-amd <jiahui.cao@amd.com>

* perf(kimi-k3): retune the M=16 merged-front FP32 row

The decode shape is 4 * concurrency, and getPaddedM(gl=0) rounds every M <= 256
up to 16, so this row serves both the C1 (M=4) and C4 (M=16) decode buckets --
about 40% of iterations in the 8k/1k replay.

It was tuned against a single weight buffer. The merged front's weight is
6016x7168 BF16 = 86.2 MB and MI355X has 256 MB of LLC, so that regime serves
most of the GEMM from cache; the shipped 16.005 us is faster than anything the
shape can reach when the weight is actually streamed. In the real model 92
distinct MoE layers stream 7.9 GB per decode step and nothing is resident.

Re-searched all 681 hipblaslt solutions with the weight chained across 604 MB
of distinct buffers (past LLC) and timed inside a CUDA graph, which is how vLLM
runs decode:

  solidx 443935 (shipped)  22.364 us  3857 GB/s
  solidx 443486 (this)     19.279 us  4474 GB/s   -13.8%

443486 is not a new solution -- the table already carries it at M=7. The right
kernel was present and keyed to the wrong M.

Max relative error vs a torch.mm FP32 reference is 4.9e-07, so this is a
dispatch change only. The us/tflops/bw columns are the streamed measurement and
are therefore not comparable to the cache-hot numbers in neighbouring rows.

* perf(kimi-k3): size the front-GEMM config cache to the whitelist

_KIMI_K3_MERGED_FRONT_TOKEN_COUNTS admits 28 distinct token counts but the
config cache held 16, so the M values a mixed decode/prefill server cycles
through evicted each other.

Minor: the tuned table is cached upstream, so a miss costs a few get_padded_m()
extension calls rather than a CSV parse, and in the graphed decode path it is
paid at capture instead of replay. This just makes the cache cover the set it
was sized for.

* refactor(kimi-k3): separate Triton epilogue layers

Move the launchable epilogue kernel under _triton_kernels/moe, keep only the public wrapper in ops/triton/moe, and add a config-aware kernel repr and CUDA Graph benchmark. Remove Kimi-specific weight packing and hipBLASLt orchestration from the Triton module.

Signed-off-by: jiacao-amd <jiahui.cao@amd.com>

* refactor(moe): generalize the SiTU epilogue

Rename the op to describe the SiTU epilogue it actually performs rather than implying that it projects the incoming activation. Make all branch widths caller-provided, mask partial tiles for arbitrary shapes, and reuse the shared Triton tanh helper.

Update the numerical and CUDA Graph coverage plus the benchmark for the generic API.

Signed-off-by: jiacao-amd <jiahui.cao@amd.com>

* test(triton): use Triton arch helper for SiTU epilogue

Signed-off-by: jiacao-amd <jiahui.cao@amd.com>

* fix(triton): avoid branch-local mask type collisions

Signed-off-by: jiacao-amd <jiahui.cao@amd.com>

---------

Signed-off-by: jiacao-amd <jiahui.cao@amd.com>
Co-authored-by: Jiahui Cao <jiacao@crs-m2m-cpu-spur-014.us-east2-a.compute.internal>
(cherry picked from commit aa84815)

* [Triton/Gluon] Add SonicMoE pure-Triton grouped GEMM MoE (#5725)

* [Triton] Add SonicMoE pure-Triton grouped GEMM MoE

Add a pure-Triton, expert-major grouped GEMM MoE with full autograd
support (SonicMoE), covering top-k and general routing, GLU/elementwise
activations, and blockwise FP8 scales for the grouped GEMM. Includes
gfx942/gfx950 tuned configs, unit tests, and a benchmark script.

Public entry points: moe_TC_softmax_topk_layer, moe_general_routing_inputs,
moe_pre_routed_inputs (aiter.ops.triton.sonicmoe).

This is the first of two PRs splitting the SonicMoE contribution; a
follow-up PR stacks the hipBLASLt/multistream grouped GEMM backend on
top of this pure-Triton implementation.

* [Triton] Drop unused SonicMoE stream_id dels and restating comments

Keep the positional stream argument for caller compatibility as
_stream_id, and remove comments that only restated the next GEMM call.

* [Triton] Move SonicMoE autograd API out of _triton_kernels

Keep @triton.jit kernels under _triton_kernels/moe/sonicmoe and put
the autograd wrappers plus public entry points in aiter.ops.triton.sonicmoe.

* [Triton] Use AMD copyright headers on SonicMoE kernels

Replace third-party author banners and drop external source-link
comments so new files match aiter's SPDX header.

* Address SonicMoE review feedback

* [Triton] Move SonicMoE routing kernels into moe_routing

* [Triton] Collapse SonicMoE host wrappers into one module

Keep the public API and tests in a single file instead of a set of sibling wrappers.

(cherry picked from commit 2e62094)

* [Triton/Gluon] Move the KDA_DECODE configs into the nested config layout (#5941)

(cherry picked from commit 92192fd)

* [Triton/Gluon] Move _triton_kernels/gated_delta_rule/ to _triton_kernels/gated_delta_net/ (#5943)

(cherry picked from commit cf89ffd)

* [Triton/Gluon] Move _gluon_kernels/gfx1250/norm/ to _gluon_kernels/gfx1250/normalization/ (#5944)

(cherry picked from commit 9ef0d08)

* [CI] Select impacted Triton and Gluon unit tests (#5878)

Select impacted Triton and Gluon unit tests

(cherry picked from commit 32b1cb0)

* [Triton/Gluon] [Config] Drop no-op kpack from RDNA GEMM configs (#5917)

* [Triton] [Config] Drop no-op kpack from RDNA GEMM configs

* [Triton] [Config] Drop redundant matrix_instr_nonkdim from selected RDNA GEMM configs

* [Triton] [Config] Extend RDNA matrix_instr_nonkdim cleanup

* [Triton] [Config] Retain matrix_instr_nonkdim for shared GEMM consumers

(cherry picked from commit 9ef3e41)

* mxfp8 gemm cga update, update (#5055)

Co-authored-by: Satya Nikhil Kodukula <nikhil.kodukula@gmail.com>
(cherry picked from commit 693701c)

* stop fused_bmm_rope_kv_cache from using batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant configs (#5877)

(cherry picked from commit a2650cd)

* [Triton/Gluon] Add Triton-based Conv3D kernels (#5952)

* conv3d implementation on Triton

* Remove unnecessary convolution compatibility aliases

* Update the README file.

* Add Conv2D weight-pack cache cleanup and updated the tests.

* Simplify the Conv2D Winograd test guard

* Move Conv3D device validation out of shape helper

* Add diagnostics to convolution test assertions

* Removing benchmark related tests from unit test

* Add context to Conv3D numerical assertions

* Deduplicate Conv2D and Conv3D blocked-layout kernels

* Simplified benchmark and fixed formatting issue on test file

* Split Conv2D and Conv3D cache-clear tests into their owning suites

* Deduplicate Conv2D and Conv3D Winograd transforms

* Deduplicate the Conv2D and Conv3D Winograd filter transform

* Consolidate Conv2D Winograd launch paths

* Deduplicate Conv2D and Conv3D prepack helpers and cache wrappers

* Black formatting

* Replace convolution test helper prints with logging

* Use safe configs for Conv3D on CDNA

(cherry picked from commit c8325e0)

---------

Signed-off-by: Jin Tao <jin.tao@amd.com>
Signed-off-by: Martin Domanský <ragua@email.cz>
Signed-off-by: Rohan138 <rohanpotdar138@gmail.com>
Signed-off-by: Pham Binh <phamhuuthanh.binh@amd.com>
Signed-off-by: Lin, Soga <soga.lin@amd.com>
Signed-off-by: jiacao-amd <jiahui.cao@amd.com>
Co-authored-by: siliangchen-amd <SiLiang.Chen@amd.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: Lingpeng Jin <103567126+valarLip@users.noreply.github.com>
Co-authored-by: Jesús Carabaño <jcaraban@users.noreply.github.com>
Co-authored-by: Leo <drleonid@amd.com>
Co-authored-by: liyjiang <liying.jiang@amd.com>
Co-authored-by: gbyu-amd <Guanbao.Yu@amd.com>
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
Co-authored-by: Jin Tao <jintao12@amd.com>
Co-authored-by: jin.tao@amd.com <jin.tao@amd.com@tus1-p15-g43.tus.tensorwave.lan>
Co-authored-by: Felix Li <felix.li@amd.com>
Co-authored-by: Martin Domanský <8312516+Ragua1@users.noreply.github.com>
Co-authored-by: Mehmet Cagri <mehmet.kaymak@amd.com>
Co-authored-by: Satya Nikhil Kodukula <nikhil.kodukula@gmail.com>
Co-authored-by: Rohan Potdar <rohanpotdar138@gmail.com>
Co-authored-by: Nimit Patel <61071220+NimitPtl@users.noreply.github.com>
Co-authored-by: Muhammad Ahmed <mm.ahmed2202@gmail.com>
Co-authored-by: yifehuan-amd <yifehuan@amd.com>
Co-authored-by: Pham Binh <phamhuuthanh.binh@amd.com>
Co-authored-by: akii96 <aakif.nawaz@amd.com>
Co-authored-by: li,xiangxiang <xiangxli@amd.com>
Co-authored-by: jhchouuu <jiahzhou@amd.com>
Co-authored-by: HaonanWang98 <hwang@amd.com>
Co-authored-by: sogalin_codegen <39478626+sogalin@users.noreply.github.com>
Co-authored-by: vorapolsiloai <115975949+vorapolsiloai@users.noreply.github.com>
Co-authored-by: jiacao-amd <jiahui.cao@amd.com>
Co-authored-by: Jiahui Cao <jiacao@crs-m2m-cpu-spur-014.us-east2-a.compute.internal>
Co-authored-by: WuLei-AMD <leiwu@amd.com>
Co-authored-by: Hyunjune Kim <132782704+hyjuunn@users.noreply.github.com>
Co-authored-by: Shao-Chun Lee <Shao-Chun.Lee@amd.com>
Co-authored-by: Saeid Rostami <123997133+saeid-rostami@users.noreply.github.com>
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.

5 participants