Repository navigation
[Triton/Gluon] [ASM] [HIP] Mha v4: adds bf16 sparse, LSE support, KV varlen, fixes, etc - #5798
Conversation
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.
done, thanks @azaidy ! also addressed the copilot findings |
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.
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.
|
E AssertionError: Tensors not close enough! 0.750732% elements exceed tolerance. 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? |
What is the exact test case that failed and what environment you had? |
* [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>
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
s_lse.f8f6,f6f4andmxfp4sparse rows moved to FP6-P V, matching their dense siblings.mxfp4claims the canonical FP6-P V order; the duplicatef4f4row is disabled.Host
lseoutput onmha_v4/mha_v4_packed, plumbed throughasm_mha_v4_fwd.cu.ptr_lse/s_lse/s_lse_Hswere already reserved in the kernarg._LSE_CAPABLE_QVgates the supported format pairs; sorted-sparse still raises.seqlen_k).Minor
MXFP4 Q/K + FP8 Vand the deprecatedmha_v4_mxfp8alias.bench_sage.py: improve input distributions, diffusion-calibrated default, BF16 sparse modes.op_tests/test_mha_v4_sparse.py.Test Plan
torch.logsumexpfor every capable format pair.Test Result
Previous PRs: #5335, #5005, #4967, #4627