[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:
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user