You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
[Triton/Gluon] [CI] [gfx942] Tuned sparse MLA prefill for DSv4.1-Flash TP4; fix attn_sink in Triton fallback - #6002
Perf (gfx942, H ∈ {8, 16}, D=512). A pinned launch config in configs/gfx942/triton/attention/sparse_attention_dsv4/DEFAULT.json (BLOCK_H=16, BLOCK_K=16, 1 warp, 2 stages, kpack=2) replaces the autotune default, whose 32-row head tile is half masked at 16 heads.
Two new kernel constexprs, independent of that config. Only the pinned launch sets them for now; both default to False (unchanged behaviour):
EVEN_HD (valid whenever head_dim == BLOCK_D and num_heads % BLOCK_H == 0): head/dim masks fold away, and the KV gather is masked by row only (~10%).
USE_EXP2: base-2 softmax (~2%).
Otherwise the kernel loop is main's code. Other shapes and archs keep the autotuned launch.
attn_sink:attn_sink or torch.empty(...) raised with a real sink, and with None read an uninitialized placeholder as the sink. Covered by test_pa_prefill_sparse_single_source (with_sink True/False); all 48 cases fail with main's code.
int64 indices: were cast to int32, silently wrapping slots ≥ 2^31 onto valid rows. They're now kept as int64.
Invalid KV rows: never loaded, so a NaN/Inf in an unused pool row can't reach the output.
Clear errors for a non-fp16/bf16 q, a KV dtype mismatch (e.g. an fp8 cache, which previously failed to compile) and a non-contiguous out.
Sparse prefill attention per 16k step, vLLM engine, TP4
138.9 ms
67.1 ms
16k-prompt TTFT
977 ms
900 ms
Pinned vs autotune, H=8/16, 32–640 slots
—
1.7–2.1×
Accuracy: error vs fp32 matches main. Engine prompt logprobs are within run-to-run noise, greedy outputs identical. No TTFT/TPOT/acceptance regression serving with DSpark.
Other callers: the non-pinned (autotuned) path is unchanged vs main: six shapes (H=16–128, D=128/512/576) within ±0.2%.
Notes
Re-tuning: the pinned config was swept on Triton 3.7.1 (current vLLM ROCm images). Re-sweeping is a JSON-only change.
Fragile register budget: the autotuned D=512 kernel sits at 255 VGPRs, one below the 2-waves/SIMD cliff. During review, small edits to the loop (e.g. an unsigned validity compare) pushed it to 258 and cost 30–40%, so re-measure after touching the loop.
Benchmarking: allocate kv inside a >2 GiB buffer as vLLM does. Otherwise Triton compiles a different (buffer-op) variant.
…nk in Triton fallback
- _sparse_attn_prefill_kernel: add HAS_INVALID / USE_EXP2 / EVEN_HD
constexprs (defaults keep current behaviour), unmasked KV gather with
invalid lanes redirected to row 0, single unsigned validity compare,
finite running-max start instead of -inf guards, and scale applied after
the row max so the shift fuses into v_fma_f32.
- pa_prefill_sparse: fixed gfx942 config for num_heads <= 16 (BLOCK_H=16,
BLOCK_K=32, 2 warps, 1 stage); honour has_invalid in the Triton branch;
new optional out= buffer; fix attn_sink handling (raised on a real sink,
used an uninitialized placeholder when None).
- Tests for the Triton single-source branch vs a vectorized fp32 reference.
DSv4.1-Flash TP4 shape on MI325X (Triton 3.7.1): 3.64 -> 2.34 ms per call
on captured engine inputs.
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
All standard extended tests (excludes ci:atom_full)
Only add ci:atom_full for FlyDSL or Triton upgrades.
Add labels via the sidebar or gh pr edit 6002 --add-label <label>
One backend per PR:
A PR changes one kernel backend: [Triton/Gluon] (Triton and Gluon count as one), [HIP], [ASM], [CK], [OPUS] or [FlyDSL]. If the title ends up with two backend tags, split the PR -- as stacked pull requests when one part cannot merge without the other.
PR title tags & labels:
Component tags ([Triton/Gluon], [HIP], [CK], [ASM], ...) are added to the PR title and as PR labels automatically from the changed files and re-synced on every push — change-type tags like [fix]/[Perf], op tags like [MLA], and human labels (ci:*) are left untouched. Add the no-auto-title label to stop the title rewrites; labels stay in sync either way.
The fixed config should not apply to every gfx942 call with 16 heads or fewer. It was swept only for H=16, D=512, 640 slots, on Triton 3.7.1. H=8, other head dims, and shorter gathers all use the same BLOCK_K=32 and 2 warps. Correctness was tested. Speed was not. Shapes that were not swept should stay on autotune.
The -1e30 start, scaling after the max, dropping the NaN guards, and the unmasked gather are in the shared kernel. Every machine that uses this Triton fallback gets them, not just the gfx942 fast path. The MI300 Triton CI job did not run.
With has_invalid=False, a -1 slot is treated as valid. The fast path also sets EVEN_HD, so that load has no mask and can read past the pool. vLLM does not pass this flag today, so the check stays on. The autotuned launch never passes has_invalid at all, which does not match the comment.
… has_invalid=False
- gfx942 fixed config now only for num_heads in (8, 16) and head_dim == 512,
the swept shapes; everything else uses the autotuned launch. New config
from a full sweep on Triton 3.7.1: BLOCK_K=16, 1 warp, 2 stages,
waves_per_eu=1, kpack=2 (2.4-2.9x over the autotune default at H=8/16,
32-640 slots per query; 2.34 -> 1.60 ms/call on DSv4.1 engine inputs).
- Restore the -inf running-max start and NaN guards in the shared kernel;
they cost nothing measurable on the fast path.
- has_invalid=False still redirects out-of-pool slots to row 0, so a broken
promise gives wrong results instead of an out-of-pool read. The autotuned
launch now passes has_invalid too.
- Tests: head dims 128/576 (autotuned launch) and a broken has_invalid=False
promise.
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
The fixed config should not apply to every gfx942 call with 16 heads or fewer. It was swept only for H=16, D=512, 640 slots, on Triton 3.7.1. H=8, other head dims, and shorter gathers all use the same BLOCK_K=32 and 2 warps. Correctness was tested. Speed was not. Shapes that were not swept should stay on autotune.
The -1e30 start, scaling after the max, dropping the NaN guards, and the unmasked gather are in the shared kernel. Every machine that uses this Triton fallback gets them, not just the gfx942 fast path. The MI300 Triton CI job did not run.
With has_invalid=False, a -1 slot is treated as valid. The fast path also sets EVEN_HD, so that load has no mask and can read past the pool. vLLM does not pass this flag today, so the check stays on. The autotuned launch never passes has_invalid at all, which does not match the comment.
github-actionsBot
changed the title
[Triton][gfx942] Faster sparse MLA prefill at low head counts (DSv4.1-Flash TP4); fix attn_sink in Triton fallback
[Triton/Gluon] [gfx942] Faster sparse MLA prefill at low head counts (DSv4.1-Flash TP4); fix attn_sink in Triton fallback
Oct 6, 2026
github-actionsBot
changed the title
[Triton/Gluon] [gfx942] Faster sparse MLA prefill at low head counts (DSv4.1-Flash TP4); fix attn_sink in Triton fallback
[Triton/Gluon] [CI] [gfx942] Faster sparse MLA prefill at low head counts (DSv4.1-Flash TP4); fix attn_sink in Triton fallback
Oct 6, 2026
…ecks, CI scope
- Move the gfx942 tile into configs/gfx942/triton/attention/
sparse_attention_dsv4/DEFAULT.json (per shape: H8/H16 x D512), loaded via
get_tuned_kernel_config; shapes without an entry use the autotuned launch.
- EVEN_HD (unmasked gather) only when the KV pool is non-empty.
- Reject softmax_scale <= 0 (the kernel scales after the row max).
- out= must be on the same device as q, with a unit-stride last dim.
- Kernel repr includes HAS_INVALID, USE_EXP2 and EVEN_HD.
- select_triton_tests.py schedules MI300X when this op's gfx942 launch path
changes.
- Tests: empty pool, non-positive scale, strided out=.
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…ench the public entry
- _sparse_attn_prefill_kernel always checks slots (one unsigned compare,
invalid lanes redirected to row 0). The no-check mode was worth 1.4% on the
gfx942 pinned path, and vLLM cannot use it since invalid top-k entries stay
as -1 inside each row. has_invalid on pa_prefill_sparse now only affects
gfx1250.
- Test: out= on whichever branch the device takes (Triton fallback, gfx950,
gfx1250) returns the buffer and matches the allocating call bit for bit.
- select_triton_tests.py: the test file also routes to MI300X.
- bench_sparse_attention_dsv4.py: adds the public pa_prefill_sparse entry
(pinned gfx942 launch) next to the autotuned kernel, DSv4.1 TP4 shapes
(H=8/16, D=512, 640 slots), and allocates KV inside a >2 GiB buffer as
vLLM does.
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
The public out= documentation omits the layout and compatibility requirements enforced below. Callers currently cannot tell from the API contract that the buffer must match q's shape, dtype, and device and have unit stride in the last dimension. Document those constraints here.
The single-source tests now mark slots as -1, exactly num_kv (first row past
the pool) and far past it (up to the int32 max), and the fp32 reference masks
slots outside [0, len(kv)). Accepting slot == num_kv in the kernel now fails
the suite.
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Missing entries are the expected path for every gfx942 shape except H=8/16, D=512, but get_tuned_kernel_config logs a warning whenever it returns its fallback (tuned_config_utils.py:115-122). Thus normal H=32 or other-D calls now emit “No tuned Triton config” warnings even though they correctly use autotuning. Use an optional/quiet lookup for this sparse table (or extend the shared helper with that mode) so expected fallback does not pollute serving logs.
This advertises int64 only for the slot array, but the new index_dtypes argument is applied to both the slots and indptr, and the int64 test relies on preserving both. Document the accepted int64 kv_indptr_prefix as well so callers do not unnecessarily cast it.
The reason will be displayed to describe this comment to others. Learn more.
Review: PR #6002 — faster sparse MLA prefill at <=16 heads; fix attn_sink in Triton fallback
This is a stable kernel, so I looked at regression risk. The PR is well built. The perf flags are gated, and the invalid-row fix is tested well. I have two asks before merge.
What is done well
The perf flags are gated.USE_EXP2=True and EVEN_HD are set only inside the gfx942 pinned-config branch. gfx950 and non-pinned callers keep the defaults (USE_EXP2=False, EVEN_HD=False), so they do not get those two changes.
The invalid-row fix is tested directly.test_pa_prefill_sparse_nonfinite_row0_ignored_by_invalid_slots poisons row 0 with non-finite values, uses invalid slots, and asserts the poison does not reach the output. test_pa_prefill_sparse_huge_pool_rejects_negative_slots checks the unsigned validity compare near 2^31, and int64_indices covers that dtype. The 296-line test is real coverage, not padding.
There is a benchmark, and it reports a speedup.
1. Benchmark the existing path, not only the new fast path
The perf flags are gated, but three changes in the kernel are not gated and run for every Triton caller:
the validity compare rewrite: slot_off.to(tl.uint64, bitcast=True) < num_kv
the invalid-lane redirect to row 0
the finite running-max start instead of the -inf guards
These change the generated code for the existing (non-pinned) path, not just the new gfx942 fast path. Your own comment on the validity block says this path is scheduler-sensitive: a real 64-bit compare runs about 1.5x slower, the cause is not known, and small edits have triggered it, so re-measure after touching the block.
The included benchmark shows the new pinned path is faster. It does not show that the baseline path is not slower. Please run op_tests/op_benchmarks/triton/bench_sparse_attention_dsv4.py on shapes that do not hit the gfx942 pinned config (for example gfx950 falling back to Triton, or gfx942 shapes with no pinned config), on main and on this branch, and post both. The default-path numbers should match within noise.
Please also confirm in the PR that the unsigned validity compare is exactly equivalent to the old (slot >= 0) & (slot < num_kv). It is equivalent when num_kv < 2^63 and valid slots are non-negative, which holds, but it is an always-on change to the guard every caller relies on, so state it.
2. Separate or call out the attn_sink fix
The commit message says the Triton fallback raised on a real attn_sink and used an uninitialized placeholder when the sink was None. That is a real correctness bug in the fallback, separate from the perf work.
It is reasonable to fix it here, because the now-production gfx942 Triton path hits it. But please either split it into its own PR, or call it out clearly in the description as a separate correctness fix, with the test assertion that covers the sink path. Mixing a behavior fix into a perf PR on a stable kernel makes later bisection harder.
Summary
Good change: the perf flags are gated, the invalid-row fix is tested properly, and a benchmark exists. The one thing to confirm before merge is that the ungated hot-path changes (validity compare, invalid-lane redirect, running-max init) do not slow the existing non-pinned path. Post baseline before/after numbers, confirm the validity compare is equivalent, and separate or document the attn_sink fix.
Review measured nothing for the non-pinned path, and it turned out ~30-37%
slower than main at D=512 (autotuned launch, MI325X, Triton 3.7.1): the
validity-compare rewrite, invalid-lane redirect, prefetch padding and
scale-after-max reshaped its schedule. These now apply only when EVEN_HD is
set, i.e. the gfx942 pinned-tile launch. With EVEN_HD=False every line of
the inner loop is main's code; the generic path needs none of the changes
(its signed compare is exact for any pool size and its gather was already
masked, so non-finite rows were never loaded there).
Non-pinned A/B vs main, 6 shapes: within +/-0.2%. Compiled ISA identical for
H=16 (D=128/512/576); the H>=32 D=512 kernel has the same instruction mix,
differing only in register assignment. Pinned path unchanged (1.66 ms).
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Tested on GLM-5.3-Flash shapes on MI325X (16 heads per rank, D=512, rope-free, top-k 2048), through vllm-project/vllm#60102. The KV pool was over 2 GiB, as in a server.
Prefill: the pinned H16/D512 tile wins here too, even though 2048 slots per query is outside the 32-640 sweep. A 16K-row call takes 8.78 ms, against 12.33 ms for aiter's Gluon sparse_mla_fwd on a BF16 cache. Relative L2 error against f32 is 0.17%, the same as Gluon. End to end, the 131K prefill drops from 4.72 to 4.30 s; gsm8k and needle-in-a-haystack are unchanged.
Two things worth handling in this PR:
FP8 caches fail at compile time. The gfx942/Triton branch has no dtype check (the gfx1250 branch does), so an fp8 unified_kv reaches tl.dot(q, tl.trans(kv)) and fails with CompilationError: Unsupported rhs dtype fp8e4b8. That's what an FP8-KV GLM server hits with #60102. The same up-front check the gfx1250 branch has (raise if unified_kv.dtype != q.dtype) would give callers a clear error to fall back on.
Decode-sized calls. With no split-K, 2-16 query rows take about 293 µs per call, against 32-36 µs for a split-K decode kernel at the same shape. A docstring note that this entry is meant for prefill-sized calls would help callers like #60102, which currently sends GLM's decode steps here, costing up to 27% TPOT end to end.
The previous commit gated four loop changes on EVEN_HD, but EVEN_HD only
means "head/dim masks fold away"; tying a different loop to it is a hidden
coupling. Bisected instead (autotuned launch, H=32 D=512, MI325X, Triton
3.7.1), adding each change alone to main's loop:
unsigned 64-bit validity compare +41.8% 255 -> 258 VGPRs
prefetch padding 0 instead of -1 +30.2% 255 -> 260 VGPRs
redirect invalid lanes to row 0 -0.3%
scale after the row max (FMA) -0.2%
The autotuned config runs 4 waves/program at 255 VGPRs; crossing 256 drops
it from 2 to 1 wave per SIMD. On the pinned path all four together were
worth ~2% once the gather became row-masked, so they are dropped, not
gated: the loop is main's code plus the row-only gather mask under EVEN_HD
and exp2 under USE_EXP2.
- Non-pinned path vs main, 6 shapes: +0.1-0.2% (noise); 255 VGPRs again.
- Pinned path: 1.65 -> 1.69 ms (int32 indices); int64 indices 2.23 -> 1.73
ms (the int64 slowdown came from the unsigned compare).
- main's signed check is exact for any pool size and for int64 indices, so
the num_kv < 2^31 concern no longer exists; tests still cover it.
- softmax_scale > 0 check and its test removed: it only existed because the
kernel scaled after the max.
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…only note
An fp8 unified_kv reached tl.dot and failed at compile time with
"Unsupported rhs dtype fp8e4b8" (reported with an FP8-KV GLM server). The
Triton branch now raises the same RuntimeError as the gfx1250 branch, so
callers get a clear error to fall back on. The docstring says the entry is
for prefill-sized calls: one program per query token and no split over KV
slots, so decode-sized calls leave most of the GPU idle.
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@jin-amd Thanks for testing on GLM shapes, and good to see the pinned tile holds up at 2048 slots per query. Both points are in the latest push:
FP8 caches: the Triton branch now does the same up-front checks as gfx1250: q must be fp16/bf16, and unified_kv.dtype must match q's. Otherwise it raises RuntimeError("unified_kv dtype mismatch: ...") instead of failing to compile in tl.dot. test_pa_prefill_sparse_rejects_kv_dtype_mismatch covers an fp8 pool.
Decode-sized calls: the docstring now says the entry is meant for prefill-sized calls. The grid is one program per query token with no split over KV slots, so a few query rows leave most of the GPU idle, and those calls should go to a split-KV decode kernel. @frida-andersson, this probably wants a query-count gate on the dispatch in [ROCm][DSv4.1][Perf] Route gfx942 sparse MLA prefill to AITER pa_prefill_sparse vllm-project/vllm#60102 so GLM's decode steps keep their decode kernel.
The reason will be displayed to describe this comment to others. Learn more.
Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.
juuso-oskari
changed the title
[Triton/Gluon] [CI] [gfx942] Faster sparse MLA prefill at low head counts (DSv4.1-Flash TP4); fix attn_sink in Triton fallback
[Triton][gfx942] Tuned sparse MLA prefill for DSv4.1-Flash TP4; fix attn_sink in Triton fallback
Oct 8, 2026
The reason will be displayed to describe this comment to others. Learn more.
Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.
github-actionsBot
changed the title
[Triton][gfx942] Tuned sparse MLA prefill for DSv4.1-Flash TP4; fix attn_sink in Triton fallback
[Triton/Gluon] [CI] [gfx942] Tuned sparse MLA prefill for DSv4.1-Flash TP4; fix attn_sink in Triton fallback
Oct 8, 2026
@juuso-oskari the query-count gate is on the vLLM dispatch, #60102. rocm_sparse_attn_prefill calls pa_prefill_sparse only when _get_aiter_pa_prefill_sparse() returns it, q.shape[0] >= 1024 (_GFX950_AITER_SPARSE_PREFILL_OPUS_MIN_QUERIES), q is fp16 or bf16, and kv.dtype == q.dtype. A shorter call does not enter the kernel.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
What
Perf (gfx942, H ∈ {8, 16}, D=512). A pinned launch config in
configs/gfx942/triton/attention/sparse_attention_dsv4/DEFAULT.json(BLOCK_H=16, BLOCK_K=16, 1 warp, 2 stages,kpack=2) replaces the autotune default, whose 32-row head tile is half masked at 16 heads.Two new kernel constexprs, independent of that config. Only the pinned launch sets them for now; both default to
False(unchanged behaviour):EVEN_HD(valid wheneverhead_dim == BLOCK_Dandnum_heads % BLOCK_H == 0): head/dim masks fold away, and the KV gather is masked by row only (~10%).USE_EXP2: base-2 softmax (~2%).Otherwise the kernel loop is
main's code. Other shapes and archs keep the autotuned launch.Correctness fixes (
pa_prefill_sparse, Triton branch):attn_sink:attn_sink or torch.empty(...)raised with a real sink, and withNoneread an uninitialized placeholder as the sink. Covered bytest_pa_prefill_sparse_single_source(with_sinkTrue/False); all 48 cases fail withmain's code.int64indices: were cast toint32, silently wrapping slots ≥ 2^31 onto valid rows. They're now kept asint64.q, a KV dtype mismatch (e.g. an fp8 cache, which previously failed to compile) and a non-contiguousout.Also:
out=buffer, used by [ROCm][DSv4.1][Perf] Route gfx942 sparse MLA prefill to AITER pa_prefill_sparse vllm-project/vllm#60102Results (MI325X, Triton 3.7.1)
mainmain. Engine prompt logprobs are within run-to-run noise, greedy outputs identical. No TTFT/TPOT/acceptance regression serving with DSpark.main: six shapes (H=16–128, D=128/512/576) within ±0.2%.Notes
kvinside a >2 GiB buffer as vLLM does. Otherwise Triton compiles a different (buffer-op) variant.🤖 Generated with Claude Code