[SM120] Use exact query-head widths for DeepSeek-V4 sparse MLA decode (#36655)

This commit is contained in:
Pengyun Lin
2026-09-10 15:40:17 -07:00
committed by GitHub
parent bb15be6d79
commit d076eec427
3 changed files with 160 additions and 24 deletions
@@ -13,7 +13,8 @@ separate region at the end of each page.
import logging import logging
import math import math
from typing import Optional from functools import lru_cache
from typing import FrozenSet, Optional, Tuple
import torch import torch
import triton import triton
@@ -266,6 +267,34 @@ def _flash_mla_sm120_prefill(
return (output.unsqueeze(1), None) return (output.unsqueeze(1), None)
@lru_cache(maxsize=1)
def _flashinfer_dsv4_decode_capabilities() -> Tuple[int, FrozenSet[int]]:
"""Read the installed FlashInfer DSV4 decode capabilities once."""
try:
from flashinfer.mla._sparse_mla_sm120 import (
_DECODE_DSV4_DISPATCH,
_DECODE_MAX_TOKENS,
)
except (AttributeError, ImportError):
return 0, frozenset()
return int(_DECODE_MAX_TOKENS), frozenset(
heads for heads, _ in _DECODE_DSV4_DISPATCH
)
def flashinfer_dsv4_decode_supports_num_heads(num_heads: int, num_tokens: int) -> bool:
"""Return whether FlashInfer supports this DSV4 decode head count.
Keep this capability check fail-closed because SGLang can be used with a
locally installed FlashInfer even though the release dependency is pinned.
The padded 64-head decode path remains the safe fallback for older builds.
Prefill head selection is handled separately by the caller.
"""
decode_max_tokens, supported_heads = _flashinfer_dsv4_decode_capabilities()
return num_tokens <= decode_max_tokens and num_heads in supported_heads
def flash_mla_with_kvcache_sm120(**kwargs): def flash_mla_with_kvcache_sm120(**kwargs):
"""SM120 FlashMLA sparse decode entry point. """SM120 FlashMLA sparse decode entry point.
+48 -23
View File
@@ -164,7 +164,6 @@ from sglang.srt.utils import (
is_gfx95_supported, is_gfx95_supported,
is_gfx942_supported, is_gfx942_supported,
is_gfx1250_supported, is_gfx1250_supported,
is_sm120_supported,
log_info_on_rank0, log_info_on_rank0,
make_layers, make_layers,
) )
@@ -821,17 +820,52 @@ class MqaAttentionBase(nn.Module):
self.register_buffer("freqs_cis", freqs_cis, persistent=False) self.register_buffer("freqs_cis", freqs_cis, persistent=False)
self.freqs_cis: torch.Tensor self.freqs_cis: torch.Tensor
def _local_attn_sink(self) -> torch.Tensor: def _kernel_num_heads(self, num_tokens: int) -> int:
if self.attn_tp_size == 1:
return self.n_local_heads
if get_platform().is_sm120:
# Prefill already accepts the native per-rank query width.
if num_tokens > SM120_DECODE_MAX_TOKENS:
return self.n_local_heads
if envs.SGLANG_SM120_FLASHMLA_BACKEND.get() == "flashinfer":
from sglang.kernels.ops.attention.flash_mla_sm120 import (
flashinfer_dsv4_decode_supports_num_heads,
)
if flashinfer_dsv4_decode_supports_num_heads(
self.n_local_heads, num_tokens
):
return self.n_local_heads
# Other FlashMLA implementations retain their existing padded shape.
return 64 if self.n_local_heads <= 64 else self.n_heads
def _local_attn_sink(self, kernel_num_heads: Optional[int] = None) -> torch.Tensor:
if self.attn_tp_size == 1: if self.attn_tp_size == 1:
return self.attn_sink return self.attn_sink
rank = self.attn_tp_rank
num_heads = self.n_local_heads
padded_num_heads = 64 if num_heads <= 64 else self.n_heads
if kernel_num_heads is None:
# Preserve the legacy contract for subclasses such as DSpark that
# always pad their attention query independently of this helper.
kernel_num_heads = padded_num_heads
assert kernel_num_heads >= num_heads
# Keep one fallback-width allocation and return a view matching Q.
# Prefill and decode can alternate, and CUDA graphs can retain the
# view, so replacing this tensor when the path changes would
# both reallocate every transition and risk invalidating a captured
# pointer.
sink_num_heads = max(kernel_num_heads, padded_num_heads)
if self._attn_sink_local is None: if self._attn_sink_local is None:
rank = self.attn_tp_rank sink = self.attn_sink.new_zeros(sink_num_heads)
num_heads = self.n_local_heads
padded_num_heads = 64 if num_heads <= 64 else self.n_heads
sink = self.attn_sink.new_zeros(padded_num_heads)
sink[:num_heads] = self.attn_sink[rank * num_heads : (rank + 1) * num_heads] sink[:num_heads] = self.attn_sink[rank * num_heads : (rank + 1) * num_heads]
self._attn_sink_local = sink self._attn_sink_local = sink
return self._attn_sink_local return self._attn_sink_local[:kernel_num_heads]
@contextmanager @contextmanager
def maybe_use_decode_attn_tp(self, forward_batch: ForwardBatch): def maybe_use_decode_attn_tp(self, forward_batch: ForwardBatch):
@@ -1624,30 +1658,21 @@ class MQALayer(MqaAttentionBase):
) )
tp_slice, q_padded, q_out = slice(None), None, None tp_slice, q_padded, q_out = slice(None), None, None
# Above this the SM120 route is the prefill kernel, which takes kernel_num_heads = self._kernel_num_heads(x.shape[0])
# arbitrary h_q, so the decode pad below would just be sliced back off. if kernel_num_heads != self.n_local_heads:
skip_decode_pad = is_sm120_supported() and x.shape[0] > SM120_DECODE_MAX_TOKENS # Backends without an exact-head specialization retain the existing
if self.attn_tp_size > 1: # padded shape. attn_sink is sliced to this rank and padded to match.
# FlashMLA's fp8 sparse decode kernel only specializes h_q for {64, 128}.
# Pad the per-rank heads to 64 (not the full n_heads) when they fit, to
# dispatch the cheaper decode::head64 variant; attn_sink is sliced to
# this rank and padded to match.
padded_num_heads = (
self.n_local_heads
if skip_decode_pad
else (64 if self.n_local_heads <= 64 else self.n_heads)
)
# Only [0:n_local_heads] is written below. Uninitialized padded TP # Only [0:n_local_heads] is written below. Uninitialized padded TP
# heads inject NaN into attention on gfx942 (fnuz), so zero-init # heads inject NaN into attention on gfx942 (fnuz), so zero-init
# there; other archs tolerate new_empty and skip the per-forward # there; other archs tolerate new_empty and skip the per-forward
# memset. # memset.
if _is_gfx942_supported: if _is_gfx942_supported:
q_padded = x.new_zeros(x.shape[0], padded_num_heads, self.head_dim) q_padded = x.new_zeros(x.shape[0], kernel_num_heads, self.head_dim)
else: else:
q_padded = x.new_empty(x.shape[0], padded_num_heads, self.head_dim) q_padded = x.new_empty(x.shape[0], kernel_num_heads, self.head_dim)
tp_slice = slice(0, self.n_local_heads) tp_slice = slice(0, self.n_local_heads)
q_out = q_padded[:, tp_slice, :] q_out = q_padded[:, tp_slice, :]
attn_sink = self._local_attn_sink() attn_sink = self._local_attn_sink(kernel_num_heads)
if enable_multi_stream: if enable_multi_stream:
# Multi-stream path always fuses cache write into the K kernel, # Multi-stream path always fuses cache write into the K kernel,
@@ -43,6 +43,7 @@ from sglang.kernels.ops.attention.flash_mla_sm120 import (
_sm120_sparse_decode_fwd, _sm120_sparse_decode_fwd,
_split_kv_pages_to_64, _split_kv_pages_to_64,
flash_mla_with_kvcache_sm120, flash_mla_with_kvcache_sm120,
flashinfer_dsv4_decode_supports_num_heads,
) )
from sglang.kernels.ops.attention.flash_mla_sm120_triton import ( from sglang.kernels.ops.attention.flash_mla_sm120_triton import (
_apply_attn_sink, _apply_attn_sink,
@@ -502,6 +503,87 @@ class TestEntryPointDispatch(CustomTestCase):
rtol=5e-2, rtol=5e-2,
) )
def test_flashinfer_exact_heads_match_padded_64_heads(self):
"""Native TP4/TP8 heads agree with padding across the prefill boundary."""
num_pages, page_size, topk = 2, 64, 128
k_cache, _ = _build_kvcache(num_pages, page_size, device=self.device, seed=17)
extra_cache, _ = _build_kvcache(
num_pages, page_size, device=self.device, seed=23
)
for num_heads in (8, 16):
with self.subTest(heads=num_heads):
if not flashinfer_dsv4_decode_supports_num_heads(num_heads, 1):
self.skipTest(
f"FlashInfer has no {num_heads}-head DSV4 decode specialization"
)
self.assertTrue(
flashinfer_dsv4_decode_supports_num_heads(num_heads, 64)
)
self.assertFalse(
flashinfer_dsv4_decode_supports_num_heads(num_heads, 65)
)
for num_tokens in (1, 64, 65):
for dual_cache in (False, True):
with self.subTest(tokens=num_tokens, dual_cache=dual_cache):
q, indices = _build_q_indices(
num_tokens,
num_heads,
topk,
num_pages,
page_size,
device=self.device,
seed=29,
)
topk_length = torch.full(
(num_tokens,),
topk,
dtype=torch.int32,
device=self.device,
)
sink = torch.linspace(
-1.0,
1.0,
num_heads,
dtype=torch.float32,
device=self.device,
)
q_padded = q.new_zeros(num_tokens, 1, 64, _D)
q_padded[:, :, :num_heads].copy_(q)
sink_padded = sink.new_zeros(64)
sink_padded[:num_heads].copy_(sink)
common = dict(
k_cache=k_cache,
indices=indices,
topk_length=topk_length,
head_dim_v=_D,
softmax_scale=_D**-0.5,
extra_k_cache=extra_cache if dual_cache else None,
extra_indices_in_kvcache=indices
if dual_cache
else None,
extra_topk_length=topk_length if dual_cache else None,
)
with mock.patch.object(
fmod, "_sm120_default_backend", "flashinfer"
):
out_exact, _ = flash_mla_with_kvcache_sm120(
q=q, attn_sink=sink, **common
)
out_padded, _ = flash_mla_with_kvcache_sm120(
q=q_padded, attn_sink=sink_padded, **common
)
self.assertEqual(
out_exact.shape, (num_tokens, 1, num_heads, _D)
)
torch.testing.assert_close(
out_exact.float(),
out_padded[:, :, :num_heads].float(),
atol=5e-2,
rtol=5e-2,
)
@unittest.skipUnless(_IS_SM120, "SM120 (compute capability 12.0) required") @unittest.skipUnless(_IS_SM120, "SM120 (compute capability 12.0) required")
class TestTouchedPageSplit(CustomTestCase): class TestTouchedPageSplit(CustomTestCase):