[Performance] optimize NSA backend metadata computation for multi-step speculative decoding (#14781)

This commit is contained in:
Johnsonms
2025-12-18 13:48:27 -08:00
committed by GitHub
parent 9749d3e346
commit e0026f7c92
3 changed files with 440 additions and 16 deletions
+111 -16
View File
@@ -10,6 +10,11 @@ from sglang.srt.configs.model_config import get_nsa_index_topk, is_deepseek_nsa
from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.nsa.dequant_k_cache import dequantize_k_cache_paged
from sglang.srt.layers.attention.nsa.nsa_backend_mtp_precompute import (
NativeSparseAttnBackendMTPPrecomputeMixin,
PrecomputedMetadata,
compute_cu_seqlens,
)
from sglang.srt.layers.attention.nsa.nsa_indexer import BaseIndexerMetadata
from sglang.srt.layers.attention.nsa.quant_k_cache import quantize_k_cache
from sglang.srt.layers.attention.nsa.transform_index import (
@@ -17,6 +22,7 @@ from sglang.srt.layers.attention.nsa.transform_index import (
transform_index_page_table_prefill,
)
from sglang.srt.layers.attention.nsa.utils import (
NSA_ENABLE_MTP_PRECOMPUTE_METADATA,
NSA_FLASHMLA_BACKEND_DECODE_COMPUTE_FP8,
NSA_FUSE_TOPK,
compute_nsa_seqlens,
@@ -224,17 +230,12 @@ class NSAIndexerMetadata(BaseIndexerMetadata):
assert False, f"Unsupported {self.topk_transform_method = }"
def compute_cu_seqlens(seqlens: torch.Tensor) -> torch.Tensor:
assert seqlens.dtype == torch.int32
return torch.nn.functional.pad(
torch.cumsum(seqlens, dim=0, dtype=torch.int32), (1, 0)
)
_NSA_IMPL_T: TypeAlias = Literal["flashmla_sparse", "flashmla_kv", "fa3", "tilelang"]
class NativeSparseAttnBackend(AttentionBackend):
class NativeSparseAttnBackend(
NativeSparseAttnBackendMTPPrecomputeMixin, AttentionBackend
):
def __init__(
self,
model_runner: ModelRunner,
@@ -886,6 +887,77 @@ class NativeSparseAttnBackend(AttentionBackend):
self.forward_metadata = metadata
def init_forward_metadata_replay_cuda_graph_from_precomputed(
self,
bs: int,
precomputed: PrecomputedMetadata,
forward_mode: ForwardMode,
):
"""Fast path: copy precomputed metadata to this backend's metadata.
This function only performs copy operations, no computation.
Args:
bs: Batch size
precomputed: Precomputed metadata to copy from
forward_mode: Forward mode
"""
self.set_nsa_prefill_impl(forward_batch=None)
metadata = self.decode_cuda_graph_metadata[bs]
# Copy basic seqlens
metadata.cache_seqlens_int32.copy_(precomputed.cache_seqlens)
metadata.cu_seqlens_k[1:].copy_(precomputed.cu_seqlens_k[1:])
# Mode-specific copy logic
if forward_mode.is_decode_or_idle():
# Decode mode
metadata.page_table_1[:, : precomputed.max_len].copy_(
precomputed.page_indices
)
metadata.nsa_cache_seqlens_int32.copy_(precomputed.nsa_cache_seqlens)
# seqlens_expanded is same as cache_seqlens (already copied)
elif forward_mode.is_target_verify():
# Target verify mode
metadata.page_table_1[:, : precomputed.max_seqlen_k].copy_(
precomputed.page_indices
)
metadata.nsa_seqlens_expanded.copy_(precomputed.seqlens_expanded)
metadata.nsa_cache_seqlens_int32.copy_(precomputed.nsa_cache_seqlens)
elif forward_mode.is_draft_extend():
# Draft extend mode
rows = precomputed.page_indices.shape[0]
cols = precomputed.max_seqlen_k
metadata.page_table_1[:rows, :cols].copy_(precomputed.page_indices)
size = precomputed.seqlens_expanded_size
metadata.nsa_seqlens_expanded[:size].copy_(precomputed.seqlens_expanded)
metadata.nsa_cache_seqlens_int32[:size].copy_(precomputed.nsa_cache_seqlens)
# Copy NSA cu_seqlens
size = precomputed.seqlens_expanded_size
metadata.nsa_cu_seqlens_k[1 : 1 + size].copy_(
precomputed.nsa_cu_seqlens_k[1 : 1 + size]
)
# Copy real page table
if precomputed.real_page_table is not None:
rows, cols = precomputed.real_page_table.shape
metadata.real_page_table[:rows, :cols].copy_(precomputed.real_page_table)
else:
# real_page_table is same as page_table_1 (already copied)
pass
# Copy FlashMLA metadata
if precomputed.flashmla_metadata is not None:
flashmla_metadata = metadata.flashmla_metadata.slice(slice(0, size + 1))
flashmla_metadata.copy_(precomputed.flashmla_metadata)
self.forward_metadata = metadata
def forward_extend(
self,
q: torch.Tensor,
@@ -1587,14 +1659,37 @@ class NativeSparseAttnMultiStepBackend:
def init_forward_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int
):
for i in range(self.speculative_num_steps):
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
seq_lens_sum=-1,
encoder_lens=None,
if NSA_ENABLE_MTP_PRECOMPUTE_METADATA:
# Precompute metadata once (shared across all backends)
precomputed = self.attn_backends[0]._precompute_replay_metadata(
bs=bs,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_cpu=forward_batch.seq_lens_cpu,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
)
# Fast copy to each backend (1-2x faster than computing N times)
for i in range(self.speculative_num_steps):
self.attn_backends[
i
].init_forward_metadata_replay_cuda_graph_from_precomputed(
bs=bs,
precomputed=precomputed,
forward_mode=ForwardMode.DECODE,
)
else:
# Fallback: compute metadata separately for each backend
for i in range(self.speculative_num_steps):
self.attn_backends[i].init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum,
encoder_lens=None,
forward_mode=ForwardMode.DECODE,
spec_info=forward_batch.spec_info,
seq_lens_cpu=forward_batch.seq_lens_cpu,
out_cache_loc=None,
)