Skip to content

[ROCm][DSv4.1][Perf] Route gfx942 sparse MLA prefill to AITER pa_prefill_sparse - #60102

Draft
frida-andersson wants to merge 10 commits into
vllm-project:mainfrom
frida-andersson:gfx942-dsv41-pa-prefill-sparse
Draft

frida-andersson wants to merge 10 commits into
vllm-project:mainfrom
frida-andersson:gfx942-dsv41-pa-prefill-sparse

Conversation

@frida-andersson

@frida-andersson frida-andersson commented Oct 5, 2026 •

Copy link
Copy Markdown
Contributor

Summary

On gfx942, fp16/bf16 sparse MLA prefill with at least 512 query rows goes to AITER pa_prefill_sparse. Shorter calls stay on the in-tree kernel. The import and the Using AITER pa_prefill_sparse for sparse MLA prefill on gfx942 log happen only when that kernel launches.

What changed

  • vllm/v1/attention/ops/rocm_aiter_mla_sparse.py: after the gfx950 OPUS attempt, call pa_prefill_sparse when the gfx942 gate passes. One KV pool. The caller output is passed as out when it matches q; a narrower output is copied from the kernel result.
  • tests/kernels/attention/test_rocm_triton_attn_dsv4.py: the dense Triton fallback test stubs the getter, so that case stays in-tree.

Gating

gfx942, AITER enabled, q fp16 or bf16, kv.dtype == q.dtype, and q.shape[0] >= 512 (_GFX942_AITER_PA_PREFILL_SPARSE_MIN_QUERIES).

pa_prefill_sparse does not split the KV. Under 512 rows the call stays on the in-tree ragged prefill, which also does not. DeepSeek-V4.1 decode uses rocm_sparse_attn_decode and does not enter this function. The 2-row and 16-row shapes stay off pa_prefill_sparse. On this branch that is the in-tree kernel, not Gluon split-K.

Depends on

ROCm/aiter#6002. The symbol is imported only after the gate. A missing symbol raises on an eligible call. A short call does not import it.

Performance

DeepSeek-V4.1-Flash

Nightly 0.31.1rc1.dev1+gbb87d227d, 4× MI325X, TP4, spec off, no FlyDSL QuickReduce, no decoder replay. Seed 7. This table is the Oct 7 run, when the gate was 1024. It was not rerun for this commit. Every length below is also above 512, so the prefill kernel is the same. The on arm logged the AITER line.

Input Output Conc Median TTFT (ms) Δ Output tok/s Δ
8192 128 2 1121.58 → 988.77 −11.8% 73.25 → 92.81 +26.7%
8192 128 8 3876.07 → 3453.47 −10.9% 144.75 → 157.37 +8.7%
20480 128 2 2648.53 → 2228.60 −15.9% 50.96 → 58.26 +14.3%
65536 128 2 7966.74 → 7082.31 −11.1% 21.30 → 23.56 +10.6%
65536 1024 2 7950.64 → 7081.32 −10.9% 91.54 → 96.61 +5.5%
262144 1024 2 36728.03 → 33240.25 −9.5% 33.47 → 36.23 +8.2%
262144 1024 4 62228.10 → 56206.90 −9.7% 36.68 → 40.09 +9.3%

This commit only reran 8192 in / 128 out on that nightly, after the gate moved to 512. Median TPOT is 13.88 ms at concurrency 2 (13.88 with the 1024 gate) and 23.83 ms at concurrency 8 (23.89). DeepSeek-V4.1 decode does not use this function, so this does not show a GLM decode change.

The 512 cutoff is from one MI325, H=16, D=512, bf16, against the in-tree ragged prefill. AITER time / in-tree time:

Rows Fanout 128 Fanout 512 Fanout 2048
256 1.04 0.88 0.84
512 0.64 0.49 0.46
1024 0.43 0.32 0.30
2048 0.49 0.39 0.35

GLM-5.3-Flash

zai-org/GLM-5.3-Flash, 4× MI325X, TP4, spec off, 4 prompts, 8192 in / 128 out, concurrency 1, seed 7. Run Oct 7, before any gate. The on arm also sent one-row decode through pa_prefill_sparse, so output tok/s includes those steps. With this head the 8192-token prefill still takes it and the decode stays in-tree.

In-tree pa_prefill_sparse Δ
Output tok/s 65.53 66.66 +1.7%
Mean TTFT (ms) 411.65 384.89 −6.5%
Median TTFT (ms) 506.14 469.44 −7.3%

GSM8K, 1319 questions, 5-shot, spec off. Strict-match and flexible-extract 0.9227 ± 0.0074 → 0.9257 ± 0.0072.

Test plan

python3 -m pytest -q \
  tests/kernels/attention/test_rocm_triton_attn_dsv4.py::test_sparse_attn_prefill_preserves_dense_triton_fallback
1 passed, 14 warnings in 3.80s

…ill_sparse

gfx942 sparse MLA prefill stays on the in-tree Triton kernel. When the
installed aiter has the gfx942 pa_prefill_sparse launch, call that kernel
and keep the in-tree path otherwise.

Signed-off-by: Frida Andersson <fanderss@amd.com>
ruff E501 rejects the one-line docstring at 94 characters.

Signed-off-by: Frida Andersson <fanderss@amd.com>
The routing test accepted any int32 tensor, so the dense table still passed.
Assert the packed indices and cover the source gate that selects the kernel.

Signed-off-by: Frida Andersson <fanderss@amd.com>
The routing test already checks the dispatch. The source check fails
closed, so it does not need its own import mock.

Signed-off-by: Frida Andersson <fanderss@amd.com>

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

Thanks for the work!

Ask: Can we please add a wider perf sweep, i.e. more concs and ISL/OSLs?

Optional suggestion: We could also consider running microbenchmarks to convince ourselves that this pa kernel is indeed faster in most cases than the ragged kernel as were now changing the default. Perhaps this should be done on the aiter side before merging the aiter PR shipping the kernel

Minor general suggestion: I see in the description

The server also logged FLYDSL_QUICK_REDUCE and decoder SWA bounded replay for layers 21–39.

Those two features are not yet in main. I'd suggest keeping the reported baseline perf values and PR values using the latest nightly image and that same image patched with the PR so we make sure we're getting the expected perf after merging this PR :)

Comment thread vllm/v1/attention/ops/rocm_aiter_mla_sparse.py Outdated
Comment thread vllm/v1/attention/ops/rocm_aiter_mla_sparse.py Outdated
The dispatch looked the function up twice. Pass it in, and drop the
pull-request number from the comments.

Signed-off-by: Frida Andersson <fanderss@amd.com>
@frida-andersson

frida-andersson commented Oct 7, 2026 •

Copy link
Copy Markdown
Contributor Author

Thanks for the review @simondanielsson! I reran on nightly 0.31.1rc1.dev1+gbb87d227d, 4× MI325X, DSv4.1 Flash, dspark off. No FlyDSL QuickReduce, no decoder replay.

The image AITER has no gfx942 launch, so the on arm adds aiter#6002 (fd133e33) plus this PR. It logged Using AITER pa_prefill_sparse for sparse MLA prefill on gfx942. Seed 7, eight prompts (four at 262144, concurrency 4).

Input Output Conc Median TTFT (ms) Δ Output tok/s Δ
8192 128 2 1121.58 → 988.77 −11.8% 73.25 → 92.81 +26.7%
8192 128 8 3876.07 → 3453.47 −10.9% 144.75 → 157.37 +8.7%
20480 128 2 2648.53 → 2228.60 −15.9% 50.96 → 58.26 +14.3%
65536 128 2 7966.74 → 7082.31 −11.1% 21.30 → 23.56 +10.6%
65536 1024 2 7950.64 → 7081.32 −10.9% 91.54 → 96.61 +5.5%
262144 1024 2 36728.03 → 33240.25 −9.5% 33.47 → 36.23 +8.2%
262144 1024 4 62228.10 → 56206.90 −9.7% 36.68 → 40.09 +9.3%

The stand-in only checked that the dispatch called it. Kernel accuracy
stays in aiter.

Signed-off-by: Frida Andersson <fanderss@amd.com>
@frida-andersson

Copy link
Copy Markdown
Contributor Author

Dropped test_sparse_attn_prefill_aiter_gfx942_routing in 1f774e5 - it called a stand-in, not pa_prefill_sparse. aiter #6002 already compares that kernel to _sparse_prefill_single_source_torch

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

IMO when we remove the source inspection and the aiter PR is part of the build then I think this looks good! Thanks

The gfx942 launch ships with the aiter this depends on, so the getter
no longer inspects the kernel source.

Signed-off-by: Frida Andersson <fanderss@amd.com>
@frida-andersson

Copy link
Copy Markdown
Contributor Author

IMO when we remove the source inspection and the aiter PR is part of the build then I think this looks good! Thanks

Thanks! I've removed the source inspection now and verified the PR on GLM-5.3-Flash. Keeping it as draft until the aiter PR is part of the build

@jin-amd

jin-amd commented Oct 8, 2026

Copy link
Copy Markdown
Contributor

Thanks for this. I A/B tested it on GLM-5.3-Flash and found two problems, both fixable with a gate.

Setup: 4× MI325X, TP4, zai-org/GLM-5.3-Flash, image amdsiloai/vllm:vllm-openai-rocm-glm5.3-flash-mi325-07102026. In that image the sparse attend runs on aiter's Gluon sparse_mla_fwd, which has split-K decode. The B arm adds this PR's 66 lines, applied verbatim to the image's rocm_aiter_mla_sparse.py, plus aiter#6002 (ee4390f7b).

1. Decode gets routed here too. For GLM the backend sends the whole batch, decode rows included, through rocm_sparse_attn_prefill, and the new branch takes every call on gfx942. pa_prefill_sparse has no split-K, so decode-sized calls are about 9× slower:

Per call (16 heads, D=512, top-k 2048) Gluon pa_prefill_sparse
2 query rows 32 µs 293 µs
16 query rows 36 µs 294 µs
16K query rows 12.33 ms 8.78 ms

End to end (131K in / 1K out, 20 prompts, concurrency 2-16):

  • TTFT is 8.8-18.3% lower at every concurrency, and the 131K prefill alone drops from 4.72 to 4.30 s.
  • TPOT is 27% higher at concurrency 2 (10.13 → 12.88 ms) and 10% higher on the geomean.
  • Output throughput falls 11% at concurrency 2 and 2% on the geomean (155.9 → 152.8 tok/s).
  • At your GLM point (8K in / 128 out, concurrency 1), TTFT is 7.8% lower, but TPOT rises 7.09 → 9.97 ms and output throughput falls 107.2 → 83.2 tok/s.
  • Accuracy is unchanged: gsm8k strict 0.9735 → 0.9742, needle-in-a-haystack 60/60 on both.

Suggestion: gate the branch the way the gfx950 OPUS path is gated. Take pa_prefill_sparse only when kv.dtype == q.dtype and the call is prefill-sized (for example q.shape[0] >= 1024, like _GFX950_AITER_SPARSE_PREFILL_OPUS_MIN_QUERIES), and fall through otherwise. That should keep the TTFT gain with no decode or FP8 regressions. I'm happy to rerun the A/B on a gated version.

GLM sends decode rows through this function. Reuse the OPUS query cutoff
and dtype checks so those calls stay on the in-tree kernel.

Signed-off-by: Frida Andersson <fanderss@amd.com>
@frida-andersson

frida-andersson commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor Author

Thanks for testing this @jin-amd! rocm_sparse_attn_prefill now calls pa_prefill_sparse only when _get_aiter_pa_prefill_sparse() returns it, q.shape[0] >= _GFX950_AITER_SPARSE_PREFILL_OPUS_MIN_QUERIES (1024), and the dtypes match. Otherwise it stays on the in-tree prefill. Could you rerun the A/B again?

@simondanielsson

Copy link
Copy Markdown
Contributor

Thanks for this. I A/B tested it on GLM-5.3-Flash and found two problems, both fixable with a gate.

Setup: 4× MI325X, TP4, zai-org/GLM-5.3-Flash, image amdsiloai/vllm:vllm-openai-rocm-glm5.3-flash-mi325-07102026. In that image the sparse attend runs on aiter's Gluon sparse_mla_fwd, which has split-K decode. The B arm adds this PR's 66 lines, applied verbatim to the image's rocm_aiter_mla_sparse.py, plus aiter#6002 (ee4390f7b).

1. Decode gets routed here too. For GLM the backend sends the whole batch, decode rows included, through rocm_sparse_attn_prefill, and the new branch takes every call on gfx942. pa_prefill_sparse has no split-K, so decode-sized calls are about 9× slower:

Per call (16 heads, D=512, top-k 2048) Gluon pa_prefill_sparse
2 query rows 32 µs 293 µs
16 query rows 36 µs 294 µs
16K query rows 12.33 ms 8.78 ms
End to end (131K in / 1K out, 20 prompts, concurrency 2-16):

  • TTFT is 8.8-18.3% lower at every concurrency, and the 131K prefill alone drops from 4.72 to 4.30 s.
  • TPOT is 27% higher at concurrency 2 (10.13 → 12.88 ms) and 10% higher on the geomean.
  • Output throughput falls 11% at concurrency 2 and 2% on the geomean (155.9 → 152.8 tok/s).
  • At your GLM point (8K in / 128 out, concurrency 1), TTFT is 7.8% lower, but TPOT rises 7.09 → 9.97 ms and output throughput falls 107.2 → 83.2 tok/s.
  • Accuracy is unchanged: gsm8k strict 0.9735 → 0.9742, needle-in-a-haystack 60/60 on both.

Suggestion: gate the branch the way the gfx950 OPUS path is gated. Take pa_prefill_sparse only when kv.dtype == q.dtype and the call is prefill-sized (for example q.shape[0] >= 1024, like _GFX950_AITER_SPARSE_PREFILL_OPUS_MIN_QUERIES), and fall through otherwise. That should keep the TTFT gain with no decode or FP8 regressions. I'm happy to rerun the A/B on a gated version.

@jin-amd Can you verify the above again with latest nightly? After #58584, pure decodes for glm-5.3-flash should never enter rocm_sparse_attn_prefill anyways so I think this issue should already be fixed wityhout the 1k token heuristic bound.

I think the logic should be that for pure decodes we should never enter rocm_sparse_attn_prefill at all, and we should rather make any changes do decode paths inside rocm_sparse_attn_decode_bf16

The getter imported and logged before the size check, so a short call
could fail on a missing symbol. Use a gfx942 cutoff of 512, measured
against the in-tree ragged prefill, and log only when the kernel launches.

Signed-off-by: Frida Andersson <fanderss@amd.com>
A short call already skips the import. An eligible call should raise if
the symbol is missing, instead of staying on the in-tree kernel with no
log. Match the cutoff comment to the measured table.

Signed-off-by: Frida Andersson <fanderss@amd.com>

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

Labels

DSv4.1 Related to DeepSeek-V4.1 models rocm Related to AMD ROCm

Projects

Status: Todo

Development

Successfully merging this pull request may close these issues.

3 participants