[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 04:59:26 -05:00
committed by GitHub
parent d098c8dab0
commit 80a6b32703
2 changed files with 85 additions and 64 deletions

View File

@@ -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

View File

@@ -286,6 +286,67 @@ def pad_sequence_with_mask(
return B, output, attn_mask
@triton.jit
def seqlens_expand_kernel(
extend_seq_lens_ptr, # [N]
seq_lens_ptr, # [N]
offsets_ptr, # [N+1]
output_ptr, # [sum(extend_seq_lens)]
N,
BLOCK: tl.constexpr,
):
pid = tl.program_id(0)
if pid >= N:
return
qo_len = tl.load(extend_seq_lens_ptr + pid)
kv_len = tl.load(seq_lens_ptr + pid)
start = kv_len - qo_len + 1
out_offset = tl.load(offsets_ptr + pid)
offs = tl.arange(0, BLOCK)
mask = offs < qo_len
values = start + offs
tl.store(output_ptr + out_offset + offs, values, mask=mask)
def seqlens_expand_triton(
extend_seq_lens: torch.Tensor,
seq_lens: torch.Tensor,
total_len: int,
max_q_len: int,
):
"""
extend_seq_lens: [N], int32, CUDA
seq_lens: [N], int32, CUDA
"""
assert extend_seq_lens.is_cuda
assert seq_lens.is_cuda
N = extend_seq_lens.numel()
offsets = torch.zeros(N + 1, device=extend_seq_lens.device, dtype=torch.int32)
offsets[1:] = torch.cumsum(extend_seq_lens, dim=0)
output = torch.empty(total_len, device=extend_seq_lens.device, dtype=torch.int32)
BLOCK = triton.next_power_of_2(max_q_len)
grid = (N,)
seqlens_expand_kernel[grid](
extend_seq_lens,
seq_lens,
offsets,
output,
N,
BLOCK=BLOCK,
)
return output
# When num_kv_heads=1, we have tensors with degenerate strides,
# For example, as below, where we have stride[-3] == stride[-2]:
# - shape: [num_pages, 1, 64, 128]