[Perf] Optimize NSA backend metadata under MTP (#19536)

Co-authored-by: Baidu-AIAK <Baidu_AIAK@163.com>
Co-authored-by: zengpai <zengpai@baidu.com>
This commit is contained in:
Brayden Zhong
2026-03-01 01:59:26 -08:00
committed by GitHub
co-authored by Baidu-AIAK zengpai
parent d098c8dab0
commit 80a6b32703
2 changed files with 85 additions and 64 deletions
@@ -36,6 +36,7 @@ from sglang.srt.layers.attention.nsa.utils import (
from sglang.srt.layers.attention.utils import (
concat_mla_absorb_q_general,
mla_quantize_and_rope_for_fp8,
seqlens_expand_triton,
)
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
@@ -434,24 +435,11 @@ class NativeSparseAttnBackend(
extend_seq_lens_cpu = [self.speculative_num_draft_tokens] * batch_size
forward_batch.extend_seq_lens_cpu = extend_seq_lens_cpu
seqlens_int32_cpu = [
self.speculative_num_draft_tokens + kv_len
for kv_len in forward_batch.seq_lens_cpu.tolist()
]
seqlens_expanded = torch.cat(
[
torch.arange(
kv_len - qo_len + 1,
kv_len + 1,
dtype=torch.int32,
device=device,
)
for qo_len, kv_len in zip(
extend_seq_lens_cpu,
seqlens_int32_cpu,
strict=True,
)
]
seqlens_expanded = seqlens_expand_triton(
torch.tensor(extend_seq_lens_cpu, dtype=torch.int32, device=device),
cache_seqlens_int32,
self.speculative_num_draft_tokens * batch_size,
self.speculative_num_draft_tokens,
)
page_table = torch.repeat_interleave(
page_table, repeats=self.speculative_num_draft_tokens, dim=0
@@ -474,20 +462,12 @@ class NativeSparseAttnBackend(
dtype=torch.int32,
device=device,
)
seqlens_expanded = torch.cat(
[
torch.arange(
kv_len - qo_len + 1,
kv_len + 1,
dtype=torch.int32,
device=device,
)
for qo_len, kv_len in zip(
forward_batch.extend_seq_lens_cpu,
forward_batch.seq_lens_cpu.tolist(),
strict=True,
)
]
seqlens_expanded = seqlens_expand_triton(
forward_batch.extend_seq_lens,
cache_seqlens_int32,
sum(extend_seq_lens_cpu),
self.speculative_num_draft_tokens,
)
if forward_batch.forward_mode.is_draft_extend_v2():
# DRAFT_EXTEND_V2: V2 worker pre-fills draft KV cache with ALL speculated
@@ -1005,24 +985,13 @@ class NativeSparseAttnBackend(
metadata.page_table_1[:, :max_seqlen_k].copy_(page_indices)
extend_seq_lens_cpu = [self.speculative_num_draft_tokens] * bs
seqlens_int32_cpu = [
self.speculative_num_draft_tokens + kv_len
for kv_len in seq_lens_cpu.tolist()
]
seqlens_expanded = torch.cat(
[
torch.arange(
kv_len - qo_len + 1,
kv_len + 1,
dtype=torch.int32,
device=self.device,
)
for qo_len, kv_len in zip(
extend_seq_lens_cpu,
seqlens_int32_cpu,
strict=True,
)
]
seqlens_expanded = seqlens_expand_triton(
torch.tensor(
extend_seq_lens_cpu, dtype=torch.int32, device=self.device
),
cache_seqlens,
self.speculative_num_draft_tokens * bs,
self.speculative_num_draft_tokens,
)
metadata.nsa_seqlens_expanded.copy_(seqlens_expanded)
nsa_cache_seqlens = compute_nsa_seqlens(
@@ -1048,20 +1017,11 @@ class NativeSparseAttnBackend(
page_indices
)
seqlens_expanded = torch.cat(
[
torch.arange(
kv_len - qo_len + 1,
kv_len + 1,
dtype=torch.int32,
device=self.device,
)
for qo_len, kv_len in zip(
extend_seq_lens_cpu,
seq_lens_cpu.tolist(),
strict=True,
)
]
seqlens_expanded = seqlens_expand_triton(
extend_seq_lens,
cache_seqlens,
sum(extend_seq_lens_cpu),
self.speculative_num_draft_tokens,
)
metadata.nsa_seqlens_expanded[: seqlens_expanded.shape[0]].copy_(
seqlens_expanded