Repository navigation
Conversation
…ing FP8 KV VLLM_ROCM_USE_AITER_TRITON_SPARSE_MLA now also takes effect on gfx942 for the ROCM_AITER_MLA_SPARSE backend's rope-free layout (GLM-5.3-Flash), when the installed aiter's sparse_mla lists gfx942 (ROCm/aiter#5721). An fp8 KV cache also needs an aiter that reads it on gfx942 (ROCm/aiter#6199). That kernel takes no fp8 q, so q stays in the model dtype and the kernel quantizes it under fp8 dots. gfx950 and DeepSeek V4/V4.1 are unchanged. Signed-off-by: Jin Tao <jin.tao@amd.com>
This branch has not been deployed
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
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
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.
Overview
Lets
VLLM_ROCM_USE_AITER_TRITON_SPARSE_MLA(#53492) take effect on gfx942 (MI300X/MI325X) for rope-free sparse MLA, GLM-5.3-Flash's layout, with a BF16 or FP8 KV cache. It needs ROCm/aiter#5721 (merged, not yet in an AITER release) and, for FP8, ROCm/aiter#6199 (in review), so it stays a draft until those ship.Claims
sparse_mla_fwdinstead of vLLM's Triton kernel. On our MI325X image, calling the same kernel with the same arguments: output throughput +7.1% (7/7 concurrencies), TPOT 5-17% lower, decode attention 7.7x faster per call, gsm8k unchanged.--kv-cache-dtype fp8runs on that kernel too, with FP8 dots. Today this combination reaches AITER's asm MLA decode on gfx942, which is built for 576-wide rows and faults. On our image: +6.9% throughput against BF16 KV, KV capacity 12.2M to 23.1M tokens, gsm8k and needle-in-a-haystack unchanged.Validation
Unit tests, on MI325X (gfx942):
pytest tests/v1/attention/test_rocm_aiter_triton_sparse_mla.py, with this change applied tovllm/vllm-openai-rocm:nightly-43b4aaea3e40e20ef53ce2bfecfa1b72060bf2b5. 17 passed. The new cases cover:forward_mqawith an FP8 cache on gfx942: q is passed unquantized, withdot_precision="fp8".pre-commit (ruff, mypy 3.12 and the local hooks) is clean on the changed files.
End to end: not yet run on this branch. No AITER release has #5721 yet, and #6199 is in review. The numbers below come from our GLM-5.3-Flash MI325X images. Their vLLM calls the same kernel with the same arguments through its own dispatch, on AITER 0.1.22.post1 with #5721 and #6199 applied. I'll rerun them on this branch once an AITER build has both PRs.
Setup: GLM-5.3-Flash FP8, TP4 on 4x MI325X,
vllm bench serverandom 131072 in / 1024 out, 20 prompts per concurrency, medians, no MTP. gsm8k is 5-shot over chat completions, all 1319 questions.AITER sparse MLA vs vLLM's Triton kernel, BF16 KV (gsm8k strict 0.9719 vs 0.9712, TTFT flat):
FP8 KV vs BF16 KV, both on AITER sparse MLA, geomean over concurrency 2-16:
To reproduce on this branch, once AITER has #5721 and #6199 (drop
--kv-cache-dtype fp8for BF16 KV):Details
is_triton_sparse_mla_enabled()also gates DeepSeek V4/V4.1, whose packed caches are gfx950-only in AITER. Agfx942_okkeyword lets onlyROCM_AITER_MLA_SPARSEaccept gfx942. It still needs the flag, and it checks that the installed AITER'ssparse_mlalists gfx942 inSUPPORTED_ARCHS._aiter_sparse_mla_unsupported_reasonreports this, and the backend falls back.forward_mqaskips quantizing q,_forward_mla_aiterpassesdot_precision="fp8", and the kernel quantizes q per tile. gfx950 keeps the pre-quantized q.triton_sparse_mla_fwdgains an optionaldot_precision. When it is omitted, the precision is chosen from q's dtype as before.FP8_SCALAR_ARCHS, which [Triton/Gluon] [gfx942] Support fp8 KV caches and fp8 dots in sparse_mla_fwd on gfx942 ROCm/aiter#6199 adds; otherwise it warns and falls back.forward_mqa, so whichever lands second rebases.Pull Request Checklist
I used vLLM's
/pr-checklistskill. (Mandatory for agents, optional for humans).AI assistance was used during the creation of this PR.
Design Fit: Minimizes impact on core components, reuses existing functionality, and justifies added complexity.
Testing and Validation: Validates the change and ensures any added tests are meaningful and reliable, with CI coverage or documented CI resource constraints and validation performed outside CI.
Code Quality and Style: Keeps code and comments clear and concise, and updates relevant documentation and examples.
Pull Request Contents: Includes a brief summary and relevant links, supports claims with evidence, explains root causes and implementation trade-offs, and follows the contributing guide.