Skip to content

W4A16 DSA: NVFP4 KV cache format for SM100 sparse decode - #18

Merged
LucasWilkinson merged 1 commit into
vllm-project:mainfrom
sychen52:W4A8R8_MLA
Aug 27, 2026
Merged

LucasWilkinson merged 1 commit into
vllm-project:mainfrom
sychen52:W4A8R8_MLA

Conversation

@sychen52

@sychen52 sychen52 commented Aug 7, 2026 •

Copy link
Copy Markdown

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_format argument (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).

@sychen52
sychen52 force-pushed the W4A8R8_MLA branch 2 times, most recently from 875e2c2 to 8a2cb3e Compare August 10, 2026 16:39
@sychen52 sychen52 changed the title W4A8R8 DSA: NVFP4 KV cache format for SM100 sparse decode W4A16 DSA: NVFP4 KV cache format for SM100 sparse decode Aug 11, 2026
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 LucasWilkinson left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM; thank you! (CI run on vllm-project/vllm#51724)

@LucasWilkinson
LucasWilkinson merged commit 6bc4941 into vllm-project:main Aug 27, 2026
1 check passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants