Repository navigation
[ROCm][DSv4.1][Perf] Route gfx942 sparse MLA prefill to AITER pa_prefill_sparse - #60102
frida-andersson wants to merge 10 commits into
Conversation
…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
left a comment
There was a problem hiding this comment.
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 :)
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>
|
Thanks for the review @simondanielsson! I reran on nightly The image AITER has no gfx942 launch, so the on arm adds aiter#6002 (
|
The stand-in only checked that the dispatch called it. Kernel accuracy stays in aiter. Signed-off-by: Frida Andersson <fanderss@amd.com>
|
Dropped |
simondanielsson
left a comment
There was a problem hiding this comment.
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>
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 |
|
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, 1. Decode gets routed here too. For GLM the backend sends the whole batch, decode rows included, through
End to end (131K in / 1K out, 20 prompts, concurrency 2-16):
Suggestion: gate the branch the way the gfx950 OPUS path is gated. Take |
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>
|
Thanks for testing this @jin-amd! |
@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>
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 theUsing AITER pa_prefill_sparse for sparse MLA prefill on gfx942log happen only when that kernel launches.What changed
vllm/v1/attention/ops/rocm_aiter_mla_sparse.py: after the gfx950 OPUS attempt, callpa_prefill_sparsewhen the gfx942 gate passes. One KV pool. The caller output is passed asoutwhen it matchesq; 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,
qfp16 or bf16,kv.dtype == q.dtype, andq.shape[0] >= 512(_GFX942_AITER_PA_PREFILL_SPARSE_MIN_QUERIES).pa_prefill_sparsedoes not split the KV. Under 512 rows the call stays on the in-tree ragged prefill, which also does not. DeepSeek-V4.1 decode usesrocm_sparse_attn_decodeand does not enter this function. The 2-row and 16-row shapes stay offpa_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.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:
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 throughpa_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.pa_prefill_sparseGSM8K, 1319 questions, 5-shot, spec off. Strict-match and flexible-extract 0.9227 ± 0.0074 → 0.9257 ± 0.0072.
Test plan