Repository navigation
W4A16 DSA: NVFP4 KV cache format for SM100 sparse decode - #18
Merged
Merged
Conversation
sychen52
force-pushed
the
W4A8R8_MLA
branch
2 times, most recently
from
August 10, 2026 16:39
875e2c2 to
8a2cb3e
Compare
3 of 4 tasks
Adds an NVFP4 KV-cache format to the SM100 sparse MLA decode kernel. Like
every existing format it is dispatched by shape, inferred from the KV
cache's bytes-per-token (its last dim) together with d_qk/d_v, so the op
signature is unchanged and no caller has to be updated:
d_qk=576, d_v=512 -> 656 = V32 (fp8), 352 = NVFP4
d_qk=512, d_v=512 -> 584 = MODEL1
nvfp4.fp8rope (352 B/token)
[0,256) 512 x e2m1 NoPE, packed 2/byte
[256,320) 64 x e4m3 RoPE, unscaled
[320,352) 32 x e4m3 NoPE scale factors (one per 16 elements), stored
permuted: the scale for element block s lives at byte
8*(s & 3) + (s >> 2), an 8x4 -> 4x8 transpose. A dequant thread
needs blocks {4c+q : c=0..7} for fixed q, which in element order
is 8 byte loads at stride 4 and after the transpose is one
contiguous 8-byte load whose scales also convert two at a time
via cvt.rn.f16x2.e4m3x2. Worth 10-12% of decode latency.
vs 656 B/token for the existing V3.2 fp8 format: 1.9x more KV capacity at
4.89 effective bits per value. W4 = 4-bit NoPE, R8 = 8-bit RoPE; the
kernel dequantizes to bf16 in smem, so Q/P and the MMAs stay bf16 (A16).
Design follows the existing self-describing DS-MLA convention: scale
factors live inline in the token record, so no per-tensor scale is
needed.
The RoPE part carries no block scale: e4m3's
4 exponent bits span the RoPE range unaided, whereas e2m1 (2 exponent
bits, max 6) genuinely needs one.
Kernel: the quantized RoPE and all scale factors are TMA-gathered as one
"tail" box into a staging buffer (128 B-aligned per 4-token group), then a
dedicated dequant warpgroup converts e2m1/e4m3 -> bf16 in smem. RoPE is
dequantized first so it unblocks the QK-RoPE UTCMMA. MMAs stay bf16 with
fp32 accumulation, unchanged. SM90 is not supported and rejects this
format via the feature-check mechanism.
Validated on B200: 4908/4908 cases in tests/test_flash_mla_sparse_decoding.py
(4748 pre-existing + 160 new NVFP4 cases covering h_q 64/128, varlen,
invalid indices, attention sink, corner cases).
Signed-off-by: Shiyang Chen <shiychen@nvidia.com>
LucasWilkinson
approved these changes
Aug 27, 2026
Collaborator
There was a problem hiding this comment.
LGTM; thank you! (CI run on vllm-project/vllm#51724)
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.
Corresponding vLLM PR: vllm-project/vllm#51724
Adds an NVFP4 KV-cache format to the SM100 sparse MLA decode kernel, selected via a new
kv_cache_formatargument (0 = existing fp8 behaviour, so callers that omit it are unaffected):nvfp4.fp8rope (352 B/token)
[0,256) 512 x e2m1 NoPE, packed 2/byte
[256,320) 64 x e4m3 RoPE, unscaled
[320,352) 32 x e4m3 NoPE scale factors (one per 16 elements)
vs 656 B/token for the existing V3.2 fp8 format: 1.9x more KV capacity at 4.89 effective bits per value. W4 = 4-bit NoPE, R8 = 8-bit RoPE; the kernel dequantizes to bf16 in smem, so Q/P and the MMAs stay bf16 (A16).
Design follows the existing self-describing DS-MLA convention: scale factors live inline in the token record, so no per-tensor scale is needed.
The RoPE part carries no block scale: e4m3's
4 exponent bits span the RoPE range unaided, whereas e2m1 (2 exponent bits, max 6) genuinely needs one.
Kernel: the quantized RoPE and all scale factors are TMA-gathered as one "tail" box into a staging buffer (128 B-aligned per 4-token group), then a dedicated dequant warpgroup converts e2m1/e4m3 -> bf16 in smem. RoPE is dequantized first so it unblocks the QK-RoPE UTCMMA. MMAs stay bf16 with fp32 accumulation, unchanged. SM90 is not supported and rejects this format via the feature-check mechanism.
Validated on B200: 4908/4908 cases in tests/test_flash_mla_sparse_decoding.py (4748 pre-existing + 160 new NVFP4 cases covering h_q 64/128, varlen, invalid indices, attention sink, corner cases).