[DeepSeekV3.2] Centralize NSA dispatch logic in NativeSparseAttnBackend (#13544)

Co-authored-by: hlu1 <14827759+hlu1@users.noreply.github.com>
This commit is contained in:
YAMY
2025-11-25 11:32:30 -08:00
committed by GitHub
parent 808b6dfdea
commit 5eed5fc0b0
2 changed files with 74 additions and 78 deletions

View File

@@ -20,6 +20,7 @@ from sglang.srt.layers.attention.nsa.utils import (
NSA_FLASHMLA_BACKEND_DECODE_COMPUTE_FP8,
NSA_FUSE_TOPK,
compute_nsa_seqlens,
is_nsa_enable_prefill_cp,
)
from sglang.srt.layers.attention.trtllm_mla_backend import _concat_mla_absorb_q_general
from sglang.srt.layers.dp_attention import get_attention_tp_size
@@ -228,9 +229,6 @@ def compute_cu_seqlens(seqlens: torch.Tensor) -> torch.Tensor:
_NSA_IMPL_T: TypeAlias = Literal["flashmla_sparse", "flashmla_kv", "fa3", "tilelang"]
NSA_PREFILL_IMPL: _NSA_IMPL_T
NSA_DECODE_IMPL: _NSA_IMPL_T
class NativeSparseAttnBackend(AttentionBackend):
def __init__(
@@ -264,10 +262,12 @@ class NativeSparseAttnBackend(AttentionBackend):
assert model_runner.req_to_token_pool is not None
self.req_to_token = model_runner.req_to_token_pool.req_to_token
global NSA_PREFILL_IMPL, NSA_DECODE_IMPL
NSA_PREFILL_IMPL = model_runner.server_args.nsa_prefill_backend
NSA_DECODE_IMPL = model_runner.server_args.nsa_decode_backend
self.enable_auto_select_prefill_impl = NSA_PREFILL_IMPL == "flashmla_auto"
self.use_mha: bool = False
self.nsa_prefill_impl: _NSA_IMPL_T = (
model_runner.server_args.nsa_prefill_backend
)
self.nsa_decode_impl: _NSA_IMPL_T = model_runner.server_args.nsa_decode_backend
self.enable_auto_select_prefill_impl = self.nsa_prefill_impl == "flashmla_auto"
self._arange_buf = torch.arange(16384, device=self.device, dtype=torch.int32)
@@ -339,6 +339,8 @@ class NativeSparseAttnBackend(AttentionBackend):
page_table_1_flattened = None
topk_indices_offset = None
# Centralized dispatch: decide all strategies for this batch
self.set_nsa_prefill_impl(forward_batch)
topk_transform_method = self.get_topk_transform_method()
@@ -455,10 +457,13 @@ class NativeSparseAttnBackend(AttentionBackend):
]
)
# Generate page_table_1_flattened when needed:
# Check if MHA with FP8 needs page_table_1_flattened for dequantization
mha_dequantize_needed = (
self.nsa_kv_cache_store_fp8 and max_seqlen_k <= self.nsa_index_topk
self.use_mha
and forward_batch.token_to_kv_pool.dtype == torch.float8_e4m3fn
)
forward_batch.using_mha_one_shot_fp8_dequant = mha_dequantize_needed
if (
topk_transform_method == TopkTransformMethod.RAGGED
or mha_dequantize_needed
@@ -506,7 +511,7 @@ class NativeSparseAttnBackend(AttentionBackend):
cache_seqlens=nsa_cache_seqlens_int32,
seq_len_q=1,
)
if NSA_DECODE_IMPL == "flashmla_kv"
if self.nsa_decode_impl == "flashmla_kv"
else None
),
nsa_cache_seqlens_int32=nsa_cache_seqlens_int32,
@@ -554,7 +559,7 @@ class NativeSparseAttnBackend(AttentionBackend):
),
seq_len_q=1,
)
if NSA_DECODE_IMPL == "flashmla_kv"
if self.nsa_decode_impl == "flashmla_kv"
else None
),
}
@@ -594,7 +599,7 @@ class NativeSparseAttnBackend(AttentionBackend):
seqlens_expanded = cache_seqlens_int32
nsa_extend_seq_lens_list = [1] * num_tokens
if NSA_DECODE_IMPL == "flashmla_kv":
if self.nsa_decode_impl == "flashmla_kv":
flashmla_metadata = self.decode_cuda_graph_metadata[
"flashmla_metadata"
].slice(slice(0, num_tokens + 1))
@@ -651,7 +656,7 @@ class NativeSparseAttnBackend(AttentionBackend):
)
nsa_extend_seq_lens_list = [1] * bs * self.speculative_num_draft_tokens
if NSA_DECODE_IMPL == "flashmla_kv":
if self.nsa_decode_impl == "flashmla_kv":
flashmla_metadata = self.decode_cuda_graph_metadata[
"flashmla_metadata"
].slice(slice(0, bs * self.speculative_num_draft_tokens + 1))
@@ -834,7 +839,7 @@ class NativeSparseAttnBackend(AttentionBackend):
else:
assert metadata.real_page_table is metadata.page_table_1
if NSA_DECODE_IMPL == "flashmla_kv":
if self.nsa_decode_impl == "flashmla_kv":
flashmla_metadata = metadata.flashmla_metadata.slice(
slice(0, seqlens_expanded_size + 1)
)
@@ -880,16 +885,13 @@ class NativeSparseAttnBackend(AttentionBackend):
causal = not layer.is_cross_attention
assert causal, "NSA is causal only"
# For fa3 interface version compatibility, we put new fields into conditional keyword args
kwargs = {}
# Detect MHA mode: multi KV heads (vs MLA with single KV head)
is_mha_mode = (layer.tp_k_head_num == layer.tp_q_head_num) and (
layer.tp_k_head_num > 1
)
# Use MHA kernel if in MHA_ONE_SHOT mode
if is_mha_mode and k is not None and v is not None and q_rope is None:
if self.use_mha:
assert k is not None and v is not None
assert q_rope is None, "MHA_ONE_SHOT path should not pass q_rope"
assert (
layer.tp_k_head_num == layer.tp_q_head_num > 1
), "MHA_ONE_SHOT requires dense multi-head config"
return self._forward_standard_mha(
q=q,
k=k,
@@ -945,7 +947,7 @@ class NativeSparseAttnBackend(AttentionBackend):
page_size=1,
)
if NSA_PREFILL_IMPL == "tilelang":
if self.nsa_prefill_impl == "tilelang":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_tilelang(
@@ -955,7 +957,7 @@ class NativeSparseAttnBackend(AttentionBackend):
sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim,
)
elif NSA_PREFILL_IMPL == "flashmla_sparse":
elif self.nsa_prefill_impl == "flashmla_sparse":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
@@ -981,7 +983,7 @@ class NativeSparseAttnBackend(AttentionBackend):
sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim,
)
elif NSA_PREFILL_IMPL == "flashmla_kv":
elif self.nsa_prefill_impl == "flashmla_kv":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_flashmla_kv(
@@ -994,7 +996,7 @@ class NativeSparseAttnBackend(AttentionBackend):
metadata=metadata,
page_table_1=page_table_1,
)
elif NSA_PREFILL_IMPL == "fa3":
elif self.nsa_prefill_impl == "fa3":
return self._forward_fa3(
q_rope=q_rope,
kv_cache=kv_cache,
@@ -1010,7 +1012,7 @@ class NativeSparseAttnBackend(AttentionBackend):
page_size=1,
)
else:
raise ValueError(f"Unsupported {NSA_PREFILL_IMPL = }")
raise ValueError(f"Unsupported {self.nsa_prefill_impl = }")
def forward_decode(
self,
@@ -1065,7 +1067,7 @@ class NativeSparseAttnBackend(AttentionBackend):
page_size=1,
)
if NSA_DECODE_IMPL == "flashmla_sparse":
if self.nsa_decode_impl == "flashmla_sparse":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_flashmla_sparse(
@@ -1075,7 +1077,7 @@ class NativeSparseAttnBackend(AttentionBackend):
sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim,
)
elif NSA_DECODE_IMPL == "flashmla_kv":
elif self.nsa_decode_impl == "flashmla_kv":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_flashmla_kv(
@@ -1088,7 +1090,7 @@ class NativeSparseAttnBackend(AttentionBackend):
metadata=metadata,
page_table_1=page_table_1,
)
elif NSA_DECODE_IMPL == "tilelang":
elif self.nsa_decode_impl == "tilelang":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_tilelang(
@@ -1098,7 +1100,7 @@ class NativeSparseAttnBackend(AttentionBackend):
sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim,
)
elif NSA_DECODE_IMPL == "fa3":
elif self.nsa_decode_impl == "fa3":
return self._forward_fa3(
q_rope=q_rope,
kv_cache=kv_cache,
@@ -1113,7 +1115,7 @@ class NativeSparseAttnBackend(AttentionBackend):
logit_cap=layer.logit_cap,
page_size=1,
)
elif NSA_DECODE_IMPL == "aiter":
elif self.nsa_decode_impl == "aiter":
if q_rope is not None:
q_all = torch.cat([q_nope, q_rope], dim=-1)
return self._forward_aiter(
@@ -1126,7 +1128,7 @@ class NativeSparseAttnBackend(AttentionBackend):
)
else:
assert False, f"Unsupported {NSA_DECODE_IMPL = }"
assert False, f"Unsupported {self.nsa_decode_impl = }"
def _forward_fa3(
self,
@@ -1359,11 +1361,36 @@ class NativeSparseAttnBackend(AttentionBackend):
"""Get the fill value for sequence length in CUDA graph."""
return 1
def set_nsa_prefill_impl(self, forward_batch: Optional[ForwardBatch] = None) -> str:
from sglang.srt.utils import is_blackwell
def set_nsa_prefill_impl(self, forward_batch: Optional[ForwardBatch] = None):
"""
Decide all attention prefill dispatch strategies for this batch.
"""
from sglang.srt.utils import get_device_sm, is_blackwell
global NSA_PREFILL_IMPL
if self.enable_auto_select_prefill_impl:
# Decide MHA vs MLA
if forward_batch and forward_batch.forward_mode.is_extend_without_speculative():
# Check if sequence meets criteria for MHA_ONE_SHOT
assert forward_batch.seq_lens_cpu is not None
max_kv_len = forward_batch.seq_lens_cpu.max().item()
sum_seq_lens = sum(forward_batch.seq_lens_cpu)
device_sm = get_device_sm()
# Requirements: H200/B200, short sequences, supported dtype, fits in chunk
self.use_mha = (
device_sm == 90
or (device_sm >= 100 and device_sm < 110) # SM90/SM100f only
and max_kv_len <= self.nsa_index_topk # Short enough for MHA
and forward_batch.token_to_kv_pool.dtype
in [torch.bfloat16, torch.float8_e4m3fn]
and sum_seq_lens
<= forward_batch.get_max_chunk_capacity() # Fits in chunk
and (not is_nsa_enable_prefill_cp()) # CP not enabled
)
else:
self.use_mha = False # Decode/verify always use MLA
# Set MLA implementation only if not using MHA
if not self.use_mha and self.enable_auto_select_prefill_impl:
if self.nsa_kv_cache_store_fp8:
if (
is_blackwell()
@@ -1374,12 +1401,12 @@ class NativeSparseAttnBackend(AttentionBackend):
total_q_tokens = forward_batch.extend_num_tokens
# Heuristic based on benchmarking flashmla_kv vs flashmla_sparse + dequantize_k_cache_paged
if total_kv_tokens < total_q_tokens * 512:
NSA_PREFILL_IMPL = "flashmla_sparse"
self.nsa_prefill_impl = "flashmla_sparse"
return
NSA_PREFILL_IMPL = "flashmla_kv"
self.nsa_prefill_impl = "flashmla_kv"
else:
# bf16 kv cache
NSA_PREFILL_IMPL = "flashmla_sparse"
self.nsa_prefill_impl = "flashmla_sparse"
def get_topk_transform_method(self) -> TopkTransformMethod:
"""
@@ -1389,7 +1416,7 @@ class NativeSparseAttnBackend(AttentionBackend):
if (
# disable for MTP
self.nsa_kv_cache_store_fp8
and NSA_PREFILL_IMPL == "flashmla_sparse"
and self.nsa_prefill_impl == "flashmla_sparse"
):
topk_transform_method = TopkTransformMethod.RAGGED
else:

View File

@@ -414,43 +414,12 @@ def handle_attention_aiter(attn, forward_batch):
def handle_attention_nsa(attn, forward_batch):
"""
Select MHA or MLA based on sequence length for optimal performance.
- Decode: MLA (avoids per-token decompression)
- Prefill <= 2048: MHA (topk ineffective, MHA has lower FLOPs)
- Prefill > 2048: MLA (topk filtering reduces computation significantly)
Dispatch logic is centralized in NativeSparseAttnBackend.set_nsa_prefill_impl and executed
in init_forward_metadata. Read the decision from backend.use_mha.
"""
if forward_batch.forward_mode.is_decode_or_idle():
return AttnForwardMethod.MLA
if forward_batch.forward_mode.is_extend_without_speculative() and (
not is_nsa_enable_prefill_cp()
):
assert forward_batch.seq_lens_cpu is not None
max_kv_len = forward_batch.seq_lens_cpu.max().item()
# MHA path enabled for both H200 (SM90, FA3) and B200 (SM100, TRTLLm ragged)
# B200 uses trtllm_ragged_attention_deepseek kernel instead of FA4
supports_mha = _device_sm in [90, 100]
# MHA supports both BF16 and FP8 KV cache (FP8 will be dequantized on-demand)
kv_dtype_supported = forward_batch.token_to_kv_pool.dtype in [
torch.bfloat16,
torch.float8_e4m3fn,
]
if (
max_kv_len <= attn.indexer.index_topk
and supports_mha
and kv_dtype_supported
):
# NSA backend uses varlen kernel which supports MHA_ONE_SHOT
# Check if total sequence length fits in chunk capacity
sum_seq_lens = sum(forward_batch.seq_lens_cpu)
# Use MHA_ONE_SHOT for best performance
if sum_seq_lens <= forward_batch.get_max_chunk_capacity():
return AttnForwardMethod.MHA_ONE_SHOT
backend = forward_batch.attn_backend
if hasattr(backend, "use_mha") and backend.use_mha:
return AttnForwardMethod.MHA_ONE_SHOT
return AttnForwardMethod.MLA