Nsa trtllm mla sparse fp8 support with Deepseek v3.2 NVFP4 (#18389)

This commit is contained in:
Rain Jiang
2026-02-16 09:29:54 +08:00
committed by GitHub
parent 8290171f52
commit 0ffd0a3995
10 changed files with 352 additions and 183 deletions
+172 -66
View File
@@ -33,7 +33,10 @@ from sglang.srt.layers.attention.nsa.utils import (
nsa_cp_round_robin_split_q_seqs,
pad_nsa_cache_seqlens,
)
from sglang.srt.layers.attention.trtllm_mla_backend import _concat_mla_absorb_q_general
from sglang.srt.layers.attention.utils import (
concat_mla_absorb_q_general,
mla_quantize_and_rope_for_fp8,
)
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.utils import is_cuda, is_hip
@@ -340,6 +343,7 @@ class NativeSparseAttnBackend(
self.device_capability = torch.cuda.get_device_capability()
self.device_sm_major = self.device_capability[0]
self.kv_cache_dtype = model_runner.kv_cache_dtype
# Allocate global workspace buffer for TRT-LLM kernels (ragged attention on SM100/B200, or trtllm decode)
if self.device_sm_major >= 10 or self.nsa_decode_impl == "trtllm":
@@ -1299,8 +1303,41 @@ class NativeSparseAttnBackend(
q_rope: Optional[torch.Tensor] = None,
k_rope: Optional[torch.Tensor] = None,
topk_indices: Optional[torch.Tensor] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
is_neox: Optional[bool] = False,
llama_4_scaling: Optional[torch.Tensor] = None,
) -> torch.Tensor:
causal = not layer.is_cross_attention
metadata = self.forward_metadata
assert causal, "NSA is causal only"
nsa_impl = (
self.nsa_decode_impl
if (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
)
else self.nsa_prefill_impl
)
if nsa_impl == "trtllm" and not self.use_mha:
return self._forward_trtllm(
q,
k,
v,
layer,
forward_batch,
metadata.nsa_cache_seqlens_int32,
save_kv_cache,
q_rope,
k_rope,
topk_indices,
cos_sin_cache,
is_neox,
llama_4_scaling,
)
if k is not None:
assert v is not None
if save_kv_cache:
@@ -1316,10 +1353,6 @@ class NativeSparseAttnBackend(
k_rope,
)
metadata = self.forward_metadata
causal = not layer.is_cross_attention
assert causal, "NSA is causal only"
# Use MHA kernel if in MHA_ONE_SHOT mode
if self.use_mha:
assert k is not None and v is not None
@@ -1381,18 +1414,9 @@ class NativeSparseAttnBackend(
page_size=1,
)
nsa_impl = (
self.nsa_decode_impl
if (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
)
else self.nsa_prefill_impl
)
if nsa_impl == "tilelang":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
q_all = concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_tilelang(
q_all=q_all,
kv_cache=kv_cache,
@@ -1402,7 +1426,7 @@ class NativeSparseAttnBackend(
)
elif nsa_impl == "flashmla_sparse":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
q_all = concat_mla_absorb_q_general(q_nope, q_rope)
if topk_transform_method == TopkTransformMethod.RAGGED:
if any(forward_batch.extend_prefix_lens_cpu):
@@ -1426,7 +1450,7 @@ class NativeSparseAttnBackend(
)
elif nsa_impl == "flashmla_kv":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
q_all = concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_flashmla_kv(
q_all=q_all,
kv_cache=kv_cache,
@@ -1452,21 +1476,6 @@ class NativeSparseAttnBackend(
logit_cap=layer.logit_cap,
page_size=1,
)
elif nsa_impl == "trtllm":
assert forward_batch.forward_mode.is_target_verify() or forward_batch.forward_mode.is_draft_extend(
include_v2=True
), "TRT-LLM NSA only supports target_verify/draft_extend; normal extend untested."
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
# Use expanded seq_lens for per-token decode in target_verify/draft_extend.
return self._forward_trtllm(
q_all=q_all,
kv_cache=kv_cache,
page_table_1=page_table_1,
metadata=metadata,
sm_scale=layer.scaling,
seq_lens=metadata.nsa_cache_seqlens_int32,
)
else:
raise ValueError(f"Unsupported {nsa_impl = }")
@@ -1482,7 +1491,32 @@ class NativeSparseAttnBackend(
q_rope: Optional[torch.Tensor] = None,
k_rope: Optional[torch.Tensor] = None,
topk_indices: Optional[torch.Tensor] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
is_neox: Optional[bool] = False,
llama_4_scaling: Optional[torch.Tensor] = None,
) -> torch.Tensor:
causal = not layer.is_cross_attention
metadata = self.forward_metadata
assert causal, "NSA is causal only"
if self.nsa_decode_impl == "trtllm":
return self._forward_trtllm(
q,
k,
v,
layer,
forward_batch,
metadata.cache_seqlens_int32,
save_kv_cache,
q_rope,
k_rope,
topk_indices,
cos_sin_cache,
is_neox,
llama_4_scaling,
)
if k is not None:
assert v is not None
if save_kv_cache:
@@ -1498,10 +1532,6 @@ class NativeSparseAttnBackend(
k_rope,
)
metadata = self.forward_metadata
causal = not layer.is_cross_attention
assert causal, "NSA is causal only"
# Do absorbed multi-latent attention
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
if q_rope is not None:
@@ -1529,7 +1559,7 @@ class NativeSparseAttnBackend(
if self.nsa_decode_impl == "flashmla_sparse":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
q_all = concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_flashmla_sparse(
q_all=q_all,
kv_cache=kv_cache,
@@ -1539,7 +1569,7 @@ class NativeSparseAttnBackend(
)
elif self.nsa_decode_impl == "flashmla_kv":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
q_all = concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_flashmla_kv(
q_all=q_all,
kv_cache=kv_cache,
@@ -1552,7 +1582,7 @@ class NativeSparseAttnBackend(
)
elif self.nsa_decode_impl == "tilelang":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
q_all = concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_tilelang(
q_all=q_all,
kv_cache=kv_cache,
@@ -1587,18 +1617,6 @@ class NativeSparseAttnBackend(
bs=forward_batch.batch_size,
)
elif self.nsa_decode_impl == "trtllm":
if q_rope is not None:
q_all = _concat_mla_absorb_q_general(q_nope, q_rope)
return self._forward_trtllm(
q_all=q_all,
kv_cache=kv_cache,
page_table_1=page_table_1,
metadata=metadata,
sm_scale=layer.scaling,
seq_lens=metadata.cache_seqlens_int32,
)
else:
assert False, f"Unsupported {self.nsa_decode_impl = }"
@@ -1860,22 +1878,103 @@ class NativeSparseAttnBackend(
def _forward_trtllm(
self,
q_all: torch.Tensor,
kv_cache: torch.Tensor,
page_table_1: torch.Tensor,
metadata: NSAMetadata,
sm_scale: float,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
layer: RadixAttention,
forward_batch: ForwardBatch,
seq_lens: torch.Tensor,
save_kv_cache=True,
# For multi-head latent attention
q_rope: Optional[torch.Tensor] = None,
k_rope: Optional[torch.Tensor] = None,
topk_indices: Optional[torch.Tensor] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
is_neox: Optional[bool] = False,
llama_4_scaling: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Forward using TRT-LLM sparse MLA kernel."""
import flashinfer.decode
metadata = self.forward_metadata
merge_query = q_rope is not None
if self.kv_cache_dtype == torch.float8_e4m3fn:
# For FP8 path, we quantize the query and rope parts and merge them into a single tensor
# Note: rope application in deepseek_v2.py:forward_absorb_prepare is skipped for FP8 decode path of this trtllm_mla backend
assert q_rope is not None, "For FP8 path q_rope should not be None."
assert k_rope is not None, "For FP8 path k_rope should not be None."
assert (
cos_sin_cache is not None
), "For FP8 path cos_sin_cache should not be None."
q, k, k_rope = mla_quantize_and_rope_for_fp8(
q,
q_rope,
k.squeeze(1),
k_rope.squeeze(1),
forward_batch.positions,
cos_sin_cache,
is_neox,
self.kv_lora_rank,
self.qk_rope_head_dim,
)
merge_query = False
# Save KV cache if requested
if save_kv_cache:
assert (
k is not None and k_rope is not None
), "For populating trtllm_mla kv cache, both k_nope and k_rope should be not None."
cache_loc = (
forward_batch.out_cache_loc
if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc
)
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
layer, cache_loc, k, k_rope
)
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
kv_cache = k_cache.view(-1, self.real_page_size, self.kv_cache_dim).unsqueeze(1)
if merge_query:
q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim)
q_rope_reshaped = q_rope.view(
-1, layer.tp_q_head_num, layer.head_dim - layer.v_head_dim
)
q_all = concat_mla_absorb_q_general(q_nope, q_rope_reshaped)
else:
q_all = q.view(-1, layer.tp_q_head_num, layer.head_dim)
# Align topk_indices with q dimensions
if topk_indices is not None:
topk_indices = self._pad_topk_indices(topk_indices, q.shape[0])
if envs.SGLANG_NSA_FUSE_TOPK.get():
page_table_1 = topk_indices
else:
page_table_1 = transform_index_page_table_decode(
page_table=metadata.page_table_1,
topk_indices=topk_indices,
page_size=1,
)
q_scale = 1.0
k_scale = (
layer.k_scale_float
if getattr(layer, "k_scale_float", None) is not None
else 1.0
)
bmm1_scale = q_scale * k_scale * layer.scaling
batch_size = page_table_1.shape[0]
_, num_heads, head_dim = q_all.shape
q = q_all.view(batch_size, 1, num_heads, head_dim)
kv = kv_cache.view(-1, 1, self.real_page_size, self.kv_cache_dim)
block_tables = page_table_1.unsqueeze(1)
seq_lens = metadata.cache_seqlens_int32 if seq_lens is None else seq_lens
out = flashinfer.decode.trtllm_batch_decode_with_kv_cache_mla(
query=q,
@@ -1888,7 +1987,7 @@ class NativeSparseAttnBackend(
seq_lens=seq_lens,
max_seq_len=metadata.max_seq_len_k,
sparse_mla_top_k=self.nsa_index_topk,
bmm1_scale=sm_scale,
bmm1_scale=bmm1_scale,
backend="trtllm-gen",
)
# Output: [batch, q_len=1, heads, v_dim] -> [batch, heads, v_dim]
@@ -1933,12 +2032,19 @@ class NativeSparseAttnBackend(
sum_seq_lens = sum(forward_batch.seq_lens_cpu)
device_sm = get_device_sm()
# when nsa prefill impl is trtllm, use its max chunk capacity as mha max kv len
mha_max_kv_len = (
forward_batch.get_max_chunk_capacity()
if self.nsa_prefill_impl == "trtllm"
else self.nsa_index_topk
)
# 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/SM100 only
and max_kv_len <= self.nsa_index_topk # Short enough for MHA
and max_kv_len <= mha_max_kv_len # Short enough for MHA
and forward_batch.token_to_kv_pool.dtype
in [torch.bfloat16, torch.float8_e4m3fn]
and sum_seq_lens
@@ -2022,7 +2128,7 @@ class NativeSparseAttnMultiStepBackend:
self.topk = topk
self.speculative_num_steps = speculative_num_steps
self.attn_backends = []
for i in range(self.speculative_num_steps):
for i in range(self.speculative_num_steps - 1):
self.attn_backends.append(
NativeSparseAttnBackend(
model_runner,
@@ -2037,11 +2143,11 @@ class NativeSparseAttnMultiStepBackend:
self.attn_backends[i].init_forward_metadata(forward_batch)
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
for i in range(self.speculative_num_steps):
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens)
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
for i in range(self.speculative_num_steps):
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
forward_batch.batch_size,
forward_batch.batch_size * self.topk,
@@ -2068,7 +2174,7 @@ class NativeSparseAttnMultiStepBackend:
# Use multi-backend fused copy when we have 3 or more backends
# This is 3x faster than calling the single-backend copy 3 times
if self.speculative_num_steps >= 3:
if self.speculative_num_steps > 3:
try:
from sglang.jit_kernel.fused_metadata_copy import (
fused_metadata_copy_multi_cuda,
@@ -2187,7 +2293,7 @@ class NativeSparseAttnMultiStepBackend:
)
# Copy remaining backends one by one (if > 3 backends)
for i in range(3, self.speculative_num_steps):
for i in range(3, self.speculative_num_steps - 1):
self.attn_backends[
i
].init_forward_metadata_replay_cuda_graph_from_precomputed(
@@ -2205,7 +2311,7 @@ class NativeSparseAttnMultiStepBackend:
print(
f"Warning: Multi-backend fused metadata copy kernel failed with error: {e}, falling back to loop."
)
for i in range(self.speculative_num_steps):
for i in range(self.speculative_num_steps - 1):
self.attn_backends[
i
].init_forward_metadata_replay_cuda_graph_from_precomputed(
@@ -2215,7 +2321,7 @@ class NativeSparseAttnMultiStepBackend:
)
else:
# Less than 3 backends: copy to each backend individually
for i in range(self.speculative_num_steps):
for i in range(self.speculative_num_steps - 1):
self.attn_backends[
i
].init_forward_metadata_replay_cuda_graph_from_precomputed(
@@ -2225,7 +2331,7 @@ class NativeSparseAttnMultiStepBackend:
)
else:
# Fallback: compute metadata separately for each backend
for i in range(self.speculative_num_steps):
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=forward_batch.req_pool_indices,