Skip to content

[Triton/Gluon] [CI] [gfx942] Tuned sparse MLA prefill for DSv4.1-Flash TP4; fix attn_sink in Triton fallback - #6002

Open
juuso-oskari wants to merge 20 commits into
mainfrom
jukorhon/gfx942-sparse-prefill
Open

juuso-oskari wants to merge 20 commits into
mainfrom
jukorhon/gfx942-sparse-prefill

Conversation

@juuso-oskari

@juuso-oskari juuso-oskari commented Sep 30, 2026 •

Copy link
Copy Markdown
Contributor

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

Correctness fixes (pa_prefill_sparse, Triton branch):

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

Also:

Results (MI325X, Triton 3.7.1)

main this PR
DSv4.1 TP4 call (16k tokens, 16 heads, ≤640 slots), via vLLM 3.65 ms 1.79 ms
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.

🤖 Generated with Claude Code

…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>
@github-actions

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every PR:

  • ✅ Pre-checks (submodule verification, code formatting)
  • ✅ Aiter op tests (gfx942 + gfx950)
  • ✅ Triton tests on MI35X (only when aiter/ops/triton/** or related paths are changed)

Extended tests (opt-in via labels):

Label Tests
ci:gfx1250-ffm-triton Run the five-shard gfx1250 FFM Triton test suite
ci:triton-300x Run an additional Triton test job on MI300X in PRs (added automatically when gfx942 configs change); main branch always runs both MI35X and MI300X
ci:triton-355 Run the full Triton test suite on MI35X, not only the tests the change affects
multigpu Aiter multi-GPU tests on the 8-GPU runner
ci:sglang SGLang integration tests: DeepSeek-R1-MXFP4 accuracy, Qwen 3.5 accuracy
ci:atom ATOM benchmark: DeepSeek-R1-0528, GPT-OSS-120B
ci:atom_full ATOM accuracy suite for PR and main models from ATOM models_accuracy.json
ci:vllm vLLM benchmark: GPT-OSS-120B, DeepSeek-R1-0528, Kimi-K2.5
ci:all All standard extended tests (excludes ci:atom_full)

Only add ci:atom_full for FlyDSL or Triton upgrades.
Add labels via the sidebar or gh pr edit 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.

@frida-andersson

Copy link
Copy Markdown
Contributor

Measured on gfx942, DSv4.1 Flash TP4. Rank 0, 2 prompts, 262144 in / 1024 out, concurrency 4:

In-tree This kernel
Longest 608 of 1320 calls 3.588 ms 2.465 ms (−31%)

The published 3.64 → 2.34 ms/call sits inside that set (2.204–2.857 ms).

Serving on 20 prompts of the same shape, concurrency 2, moves mean TTFT −4.1% and throughput +3.7%. GSM8K and RULER niah_single_2 hold.

vLLM dispatch and more tables in vllm-project/vllm#60102

@ChuanLi1101

Copy link
Copy Markdown
Contributor

The sink fix is correct. Three issues:

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

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

  3. 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>
@juuso-oskari

juuso-oskari commented Oct 6, 2026 •

Copy link
Copy Markdown
Contributor Author

The sink fix is correct. Three issues:

  1. 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.
  2. 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.
  3. 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.

Fixed in 99004d7

@juuso-oskari

juuso-oskari commented Oct 6, 2026 •

Copy link
Copy Markdown
Contributor Author

Measured on gfx942, DSv4.1 Flash TP4. Rank 0, 2 prompts, 262144 in / 1024 out, concurrency 4:

In-tree This kernel
Longest 608 of 1320 calls 3.588 ms 2.465 ms (−31%)
The published 3.64 → 2.34 ms/call sits inside that set (2.204–2.857 ms).

Serving on 20 prompts of the same shape, concurrency 2, moves mean TTFT −4.1% and throughput +3.7%. GSM8K and RULER niah_single_2 hold.

vLLM dispatch and more tables in vllm-project/vllm#60102

The latest config (BLOCK_K=16, 1 warp, 2 stages, kpack=2) achieves 2.34 → 1.60 ms/call on those inputs. Could you try aswell?

@juuso-oskari
juuso-oskari marked this pull request as ready for review October 6, 2026 10:00
@juuso-oskari
juuso-oskari requested review from a team and a balanced review from Copilot October 6, 2026 10:00
@github-actions github-actions Bot 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

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot review overview

🟡 Changes recommended

Correctness, output-buffer validation, tuning configuration, and gfx942 CI coverage issues remain unresolved.

Review effort: Balanced
Findings: 1 High severity · 4 Medium severity · 2 Low severity

Open (7)
What changed in this PR

Optimizes gfx942 sparse MLA prefill while correcting Triton sink handling and extending the wrapper with caller-provided output buffers.

Changes:

  • Adds a gfx942 low-head-count launch configuration.
  • Refines Triton masking, scaling, and exponentiation.
  • Adds fallback correctness and edge-case tests.
File Description
aiter/​ops/​triton/​attention/​pa_prefill_sparse.py Updates dispatch, sink handling, and out= support.
aiter/​ops/​triton/​_triton_kernels/​attention/​sparse_attention_dsv4.py Optimizes validity checks, loads, and softmax computation.
op_tests/​triton_tests/​attention/​test_pa_prefill_sparse.py Adds reference-based fallback tests.

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

Comment thread aiter/ops/triton/_triton_kernels/attention/sparse_attention_dsv4.py Outdated
Comment thread aiter/ops/triton/_triton_kernels/attention/sparse_attention_dsv4.py Outdated
Comment thread aiter/ops/triton/_triton_kernels/attention/sparse_attention_dsv4.py Outdated
Comment thread aiter/ops/triton/attention/pa_prefill_sparse.py
Comment thread op_tests/triton_tests/attention/test_pa_prefill_sparse.py
Comment thread aiter/ops/triton/attention/pa_prefill_sparse.py
Comment thread aiter/ops/triton/attention/pa_prefill_sparse.py Outdated
Copilot AI balanced review requested due to automatic review settings October 6, 2026 10:14
@github-actions github-actions Bot 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
@github-actions github-actions Bot added the CI label 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>

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Comment thread aiter/ops/triton/_triton_kernels/attention/sparse_attention_dsv4.py Outdated
Comment thread .github/scripts/select_triton_tests.py
Comment thread aiter/ops/triton/attention/pa_prefill_sparse.py
Comment thread aiter/ops/triton/attention/pa_prefill_sparse.py
…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>
Copilot AI balanced review requested due to automatic review settings October 6, 2026 11:19

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot review overview

🟡 Changes recommended

The unsigned bounds rewrite lacks positive out-of-pool test coverage, and the new output-buffer constraints are incompletely documented.

Review effort: Balanced
Findings: 1 Medium severity

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

In code that hasn't changed since last review

Low severity Document out= buffer layout and compatibility requirements

aiter/​ops/​triton/​attention/​pa_prefill_sparse.py:98

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.

Comment thread aiter/ops/triton/_triton_kernels/attention/sparse_attention_dsv4.py Outdated
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>

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot review overview

🔵 Needs a closer look

The cross-architecture GPU-kernel and numerical changes require final hardware-backed human review.

Review effort: Balanced
Findings: None

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

In code that hasn't changed since last review

Medium severity Avoid warnings for expected sparse tuning fallbacks

aiter/​ops/​triton/​attention/​pa_prefill_sparse.py:59

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.

Low severity Document int64 support for kv_indptr_prefix

aiter/​ops/​triton/​attention/​pa_prefill_sparse.py:89

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.

@juuso-oskari

Copy link
Copy Markdown
Contributor Author

@cagrikymk could you review? We would like to get this merged by tomorrow

Copilot AI balanced review requested due to automatic review settings October 7, 2026 14:37

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@vgokhale vgokhale left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot AI balanced review requested due to automatic review settings October 8, 2026 06:59

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

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>
Copilot AI balanced review requested due to automatic review settings October 8, 2026 07:33

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@jin-amd

jin-amd commented Oct 8, 2026

Copy link
Copy Markdown
Contributor

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:

  1. 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.
  2. 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>
Copilot AI balanced review requested due to automatic review settings October 8, 2026 10:29
@juuso-oskari

juuso-oskari commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor Author

@vgokhale Thanks, point 1 found a real regression. Both points are addressed in the latest push.

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

…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>
Copilot AI balanced review requested due to automatic review settings October 8, 2026 11:06
@juuso-oskari

Copy link
Copy Markdown
Contributor Author

@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:

  1. 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.
  2. 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.

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@juuso-oskari 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
…e-prefill

# Conflicts:
#	.github/scripts/split_tests.sh
Copilot AI balanced review requested due to automatic review settings October 8, 2026 12:31

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@github-actions github-actions Bot 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
@frida-andersson

frida-andersson commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor

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

torch.testing.assert_close(out, ref, atol=1e-2, rtol=1e-2)


def _sparse_prefill_single_source_torch(q, kv, indices, indptr, attn_sink, scale):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This doesnt have to be a blocker, but can we use existing torch reference for this so we have single source of truth?

You can provide a flag there like "has"extra". That would cover this case as well I think?

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

8 participants