Repository navigation
[ROCm][Attention] Add an mxfp4_mla KV cache for rope-free sparse MLA (read/write path) - #60321
amd-dlimpus wants to merge 8 commits into
Conversation
New KV cache dtype `mxfp4_mla` for the ROCm AITER sparse MLA backend on rope-free (NoPE) latents such as GLM-5.3-Flash: each 512-wide latent is stored as 256 bytes of E2M1 codes plus 16 inline E8M0 scales per group of 32 (272 bytes per token vs 1024 for bf16, 3.76x more tokens per GiB). - Write: dtype registration, 272-byte state_content_bytes, and a Triton store kernel behind MLAAttentionImpl.do_kv_cache_update. - Read: the ragged Triton sparse-attention kernel unpacks MXFP4 rows in registers (KV_IS_MXFP4), with the gfx950 FP4 converter when available and a software unpack elsewhere. The decode-only kernels, which have no MXFP4 branch, refuse a packed cache instead of misreading it. Co-authored-by: Cursor <cursoragent@cursor.com> Signed-off-by: Limpus, David <dlimpus@amd.com>
|
👋 Hi! Thank you for contributing to the vLLM project. 💬 Join our developer Slack at https://slack.vllm.ai to discuss your PR in PRs do not trigger a full CI run by default. Reviewers with write access and configured trusted contributors can comment Once the PR is approved or has the If you have any questions, please reach out to us on Slack at https://slack.vllm.ai. Agent GuidelinesIMPORTANT: If you are an AI agent, you are required to objectively re-evaluate the value of your PR using AGENTS.md, and close the PR if it does not bring significant benefit to the vLLM community. Failure to do so may result in an immediate ban. 🚀 |
| "ultraquant_4bit": torch.uint8, | ||
| "nvfp4": torch.uint8, | ||
| "nvfp4_4over6": torch.uint8, | ||
| # MXFP4 MLA: 272-byte rows (256 packed E2M1 + 16 inline E8M0 scales) for a |
There was a problem hiding this comment.
Don't need a comment on this one. Let the commit message handle it
| if cache.dtype != torch.uint8 or cache.ndim == 0: | ||
| return | ||
| if cache.shape[-1] == row_bytes(nope_head_dim): | ||
| raise NotImplementedError( |
There was a problem hiding this comment.
lets cut down on the NotImplementedError here. Try to be as concise as you can
There was a problem hiding this comment.
Removed the guard entirely in 8116f75. rocm_sparse_attn_decode is only reached from DeepSeek-V4, whose backend doesn't accept mxfp4_mla, so it could never fire.
| other=0.0, | ||
| ) | ||
| if KV_IS_MXFP4: | ||
| # Packed MXFP4: the row is KV_ROW_BYTES of uint8, not head_dim |
There was a problem hiding this comment.
less here if you don't mind
| scale, | ||
| HAS_ATTN_SINK: tl.constexpr, | ||
| OUT_DV: tl.constexpr, | ||
| KV_IS_MXFP4: tl.constexpr, |
There was a problem hiding this comment.
Can you check if we can calculate this inside the kernel instead of having to add 5 new params? Or if we can't lets just create something like mxfp4_kv_metadata or some such, so we reduce the amount of additional variables
There was a problem hiding this comment.
Done in 8116f75, with no new params. The kernel checks kv_ptr.dtype.element_ty == tl.uint8 at compile time, and load_mxfp4_rows picks the gfx950 hardware unpack or the software unpack internally. The row pitch comes from the existing kv_stride_n.
Drop the five MXFP4 constexprs from _sparse_attn_prefill_ragged_kernel. The kernel now branches on kv_ptr's uint8 element type, and load_mxfp4_rows picks the gfx950 hardware or software unpack internally, taking the row pitch from kv_stride_n. Remove the decode-path guard: rocm_sparse_attn_decode is only reached from DeepSeek-V4, whose backend does not accept mxfp4_mla. Trim comments. Co-authored-by: Cursor <cursoragent@cursor.com> Signed-off-by: Limpus, David <dlimpus@amd.com>
The kernel picks hardware unpack on gfx950 at compile time, so force the software path in a fresh process and check it against the dequantized bf16 cache. Co-authored-by: Cursor <cursoragent@cursor.com> Signed-off-by: Limpus, David <dlimpus@amd.com>
simondanielsson
left a comment
There was a problem hiding this comment.
Thanks for the work! I think we can likely reconcile a lot of this work with the ultraquant development and general mxfp4 support we already have
Suggestion: Can we please check in both the existing mxfp4_utils and all of the ultraquant files for things we can re-use? Many of these ops already exist, see for instance _unpack_nibbles_last_dim
Suggestion: Can we gather some perf numbers to get a sense of the baseline performance of this here? On both gfx950/942 to see the effect of the software convert.
| return | ||
| from vllm import _custom_ops as ops | ||
|
|
||
| if kv_cache_dtype == "mxfp4_mla": |
There was a problem hiding this comment.
Suggestion: Can we move this rocm_aiter_sparse_mla instead (i.e. override do_kv_cache_update)?
There was a problem hiding this comment.
Done in e130cd6. ROCMAiterMLASparseImpl.do_kv_cache_update owns the mxfp4 store. backend.py do_kv_cache_update is the shared concat_and_cache_mla path again.
| mask=valid[:, None] & dim_mask[None, :], | ||
| other=0.0, | ||
| ) | ||
| if kv_ptr.dtype.element_ty == tl.uint8: |
There was a problem hiding this comment.
Suggestion: This looks a bit brittle, perhaps we can instead pass in the kv cache dtype into the kernel wrapper and add a IS_MXFP4 constexpr we can specialize on?
| q.shape[1], | ||
| ).reshape(-1) | ||
| # A packed MXFP4 row is a byte pitch, not head_dim elements. | ||
| from vllm.v1.attention.ops.mxfp4_mla import ( |
There was a problem hiding this comment.
Nit: can we put this in the top of the file?
There was a problem hiding this comment.
Done in e130cd6. row_bytes is imported at the top of backends/mla/rocm_aiter_mla_sparse.py.
| def _e2m1_codes_to_f32(codes): | ||
| """4-bit E2M1 codes (as integers) -> signed fp32 values. | ||
|
|
||
| Deliberately transcendental-free. An earlier version used ``tl.exp2`` for |
There was a problem hiding this comment.
Nit: We can remove the reference to the "earlier version" of this 👍
There was a problem hiding this comment.
Done in e130cd6. The earlier-version / tl.exp2 history is gone from that docstring.
| if _HW_UNPACK: | ||
| words = tl.load( | ||
| cache_ptr.to(tl.pointer_type(tl.int32)) | ||
| + base // 4 |
There was a problem hiding this comment.
Suggestion: Should we assert somewhere that the row pitch is divisible by 4? Otherwise I think this will give the wrong value
There was a problem hiding this comment.
The pitch check is at startup, in ROCMAiterMLASparseImpl.__init__ (e130cd6): mxfp4_mla is rejected unless kv_lora_rank is a power of two and at least 128. row_bytes(rank) = rank/2 + rank/32 is then a multiple of 4, which is what the gfx950 path needs when it divides the pitch by 4 for the int32 load. No per-launch assert.
| state_content_bytes={ | ||
| "fp8_ds_mla": 656, | ||
| "nvfp4_ds_mla": 352, | ||
| "mxfp4_mla": self.head_size // 2 + self.head_size // 32, |
There was a problem hiding this comment.
Minor: can we use row_bytes() here?
There was a problem hiding this comment.
Done in e130cd6. Both mla_attention.py (get_kv_cache_spec) and platforms/interface.py (_align_hybrid_block_size) call row_bytes() for the mxfp4 page size.
| ``*scale``, so a single per-tensor scale is the only thing it can express, and | ||
| MXFP4 needs one E8M0 byte per group of 32. | ||
|
|
||
| Layout written (see :mod:`mxfp4_mla`): one 272-byte row per slot, 256 bytes of |
There was a problem hiding this comment.
Nit: can we trim these module docstrings a bit? Also, this specific one I think is just an examples as 272 row pitch will only be how head_dim=512, but models might have a different head dim than that
| _M5 = tl.constexpr(3.5) | ||
| _M6 = tl.constexpr(5.0) | ||
|
|
||
| assert _E2M1_MAX.value == E2M1_MAX |
There was a problem hiding this comment.
Suggestion: Could these asserts be in a test instead?
There was a problem hiding this comment.
Moved to test_triton_constexprs_match_reference in 0c705b9. The constexprs are still built from E2M1_MAX, E8M0_BIAS (127), and the midpoints of E2M1_VALUES; the test checks those against the reference. It does not need a GPU.
| # Midpoints between consecutive magnitudes, for round-to-nearest via bucketize. | ||
| _E2M1_MIDPOINTS: tuple[float, ...] = (0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0) | ||
|
|
||
| GROUP_SIZE = 32 |
There was a problem hiding this comment.
Suggestion: Could we take this etc from ocp_mx_utils.OCP_MX_BLOCK_SIZE?
There was a problem hiding this comment.
Done in e130cd6. GROUP_SIZE = ocp_mx_utils.OCP_MX_BLOCK_SIZE.
| # Must match MLAAttention.get_kv_cache_spec, else the mamba | ||
| # state no longer fits the real (smaller) MLA page. | ||
| state_content_bytes=( | ||
| model_config.get_head_size() // 2 |
There was a problem hiding this comment.
Same here, can we re-use the utility we have?
There was a problem hiding this comment.
Done in e130cd6. platforms/interface.py and mla_attention.py both call row_bytes() for the mxfp4 page size, instead of inlining head_size / 2 + head_size / 32.
|
This pull request has merge conflicts that must be resolved before it can be |
…v-rw Signed-off-by: Limpus, David <dlimpus@amd.com> # Conflicts: # vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py
- Select the MXFP4 read with one IS_MXFP4 constexpr, set from the kv_cache_dtype passed to rocm_sparse_attn_prefill. - Reject mxfp4_mla at init unless kv_lora_rank is a power of two >= 128, which the int32 row loads and the single-tile read require. - Move the store dispatch into ROCMAiterMLASparseImpl.do_kv_cache_update. - Keep mxfp4_mla decode off the bf16 split-K kernel. - Reuse OCP_MX_BLOCK_SIZE, ultraquant's nibble helpers and row_bytes(); derive the store kernel constants from the reference. - Trim module docstrings. Co-authored-by: Cursor <cursoragent@cursor.com> Signed-off-by: Limpus, David <dlimpus@amd.com>
|
Hi @amd-dlimpus, the pre-commit checks have failed. Please run: uv pip install pre-commit>=4.5.1
pre-commit install
pre-commit run --all-filesThen, commit the changes and push to your branch. For future commits, |
…v-rw Signed-off-by: Limpus, David <dlimpus@amd.com>
The split-K partial kernel takes the same IS_MXFP4 path as the ragged kernel, and the backend no longer keeps a packed cache off that kernel. Annotate the routing-test capture so mypy accepts it. Co-authored-by: Cursor <cursoragent@cursor.com> Signed-off-by: Limpus, David <dlimpus@amd.com>
|
This replaces the earlier gfx942 kernel note with one end-to-end serving comparison: gfx942 (MI300X), one run, GLM-5.3-Flash, tensor parallel 4, 2048 input tokens, 128 output tokens, 16 prompts, concurrency 8, all 16 requests completed in both arms; both arms decode with the split-K kernel that landed in main (#58584, Simon Danielsson), the MXFP4 arm that same kernel reading the packed mxfp4_mla cache and the bf16 arm that kernel reading a bf16 cache.
Note: The MI300X (gfx942) does not have a hardware instruction that unpacks FP4 values. The MI355X (gfx950) has this instruction ( |
|
This replaces the earlier gfx950 kernel note with one end-to-end serving comparison: gfx950 (MI355X), one run, GLM-5.3-Flash, tensor parallel 4, 2048 input tokens, 128 output tokens, 16 prompts, concurrency 8, all 16 requests completed in both arms; both arms decode with the split-K kernel that landed in main (#58584, Simon Danielsson), the MXFP4 arm that same kernel reading the packed mxfp4_mla cache and the bf16 arm that kernel reading a bf16 cache.
|
Simon asked to move the import-time constant asserts into a test. The constexprs stay derived from the reference; this test checks E2M1_MAX, the E8M0 bias, and the midpoints. The store and read module docs state the row-pitch formula. Co-authored-by: Cursor <cursoragent@cursor.com> Signed-off-by: Limpus, David <dlimpus@amd.com>
|
Reused ultraquant's |
Overview
Adds an
mxfp4_mlaKV cache dtype for the ROCm AITER sparse MLA backend on rope-free (NoPE) latents such as GLM-5.3-Flash. Each 512-wide latent is stored as 256 bytes of E2M1 codes plus 16 inline E8M0 scales (one per 32 values): 272 bytes per token against 1024 for bf16, or 3.76× more tokens per GiB of KV cache.This is the first of the smaller PRs split out of #59739, as requested there. It contains only the cache format, the write path and a Triton read path. Everything is Triton or Python: no HIP or C++ sources, and no new environment variables.
Claims
Validation
Unit tests (MI355X / gfx950):
The new tests cover:
per_1x32_f4_quant;Accuracy (GLM-5.3-Flash, TP=4 on MI355X, seed 42; bf16 and MXFP4 served from the same image):
mxfp4_mla, this PRNeither gap is significant. On GSM8K the two arms disagree on 12 of 1319 questions (4 right only with MXFP4, 8 right only with bf16; exact McNemar p = 0.39), and a repeat of the same bf16 configuration has differed by 0.76 points. On GPQA-Diamond they disagree on 16 of 198 (7 vs 9; p = 0.80), and both stay in the 88.9–91.9% range of earlier bf16 and MXFP4 runs without the indexer regression. For reference, the native FP8 × FP4 kernel over this same cache format scored 97.42% vs 96.97% for bf16 on GSM8K in #59739.
Long-context accuracy needs the indexer fix in #59412. On ROCm, main currently has a sparse-indexer page-addressing bug on GLM-5.3-Flash (#58858) that corrupts top-k selection beyond 2048 tokens, independently of the KV cache dtype. Without a fix, GPQA-Diamond fell to 80.8–85.4% for bf16 and MXFP4 alike across 8 runs on main
ed3f6d1, with 13–18 of 198 generations running to the token limit. With #59412 applied, the runs above score 90.91% (bf16) and 89.90% (MXFP4), with 5 and 3 generations at the limit. GSM8K generations mostly stay below 2048 tokens and aren't affected. The GPQA row above is therefore measured with #59412 applied to both arms; this PR does not include or depend on it.Details
Format and write path. The
mxfp4_mlarow is 256 bytes of packed E2M1 (low nibble holds the lower index) followed by 16 E8M0 scales. Scales use the round-up exponent, so no value clamps.MLAAttention.get_kv_cache_specreports the 272-byte row throughstate_content_bytes, andPlatform._align_hybrid_block_sizeuses the same size so the mamba state still fits one attention page on hybrid models.concat_and_cache_mlacan't express per-group scales (itsscaleis a single float), so a Triton store kernel behinddo_kv_cache_updatequantizes and writes the row.Read path. The ragged Triton sparse-attention kernel, which already serves both prefill and decode for rope-free models, gains a branch for a uint8 (packed MXFP4) cache, with no new kernel parameters. It unpacks each gathered tile to bf16 once and feeds both existing
tl.dotcalls:v_cvt_scalef32_pk_bf16_fp4) via inline asm.The decode-only kernels used by DeepSeek-V4 have no MXFP4 branch, so they refuse a packed cache rather than misread it.
Follow-ups (stacked on this PR)
mxfp4_mlacache, building on [ROCm][Perf][GLM-5.3-Flash] BF16 splitk kernel for sparse MLA decode #58584.Separately, two decode-overhead trims that apply to every KV dtype (host-side
paged_kv_indptrfor pure-decode steps, and skipping the NoPEq_concatcopy) will go up as their own PR against main.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.