[Performance] optimize NSA backend metadata computation for multi-step speculative decoding (#14781)
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user