[Spec Decoding] Support MTP for dsv3.2 (#11652)
Co-authored-by: Paiiiiiiiiiiiiii <zengpai@baidu.com>
This commit is contained in:
co-authored by
Paiiiiiiiiiiiiii
parent
d658f0497e
commit
efa473348b
@@ -29,6 +29,7 @@ if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.speculative.spec_info import SpecInput
|
||||
|
||||
|
||||
_is_hip = is_hip()
|
||||
|
||||
if _is_hip:
|
||||
@@ -148,7 +149,14 @@ NSA_DECODE_IMPL: _NSA_IMPL_T
|
||||
|
||||
|
||||
class NativeSparseAttnBackend(AttentionBackend):
|
||||
def __init__(self, model_runner: ModelRunner):
|
||||
def __init__(
|
||||
self,
|
||||
model_runner: ModelRunner,
|
||||
skip_prefill: bool = False,
|
||||
speculative_step_id=0,
|
||||
topk=0,
|
||||
speculative_num_steps=0,
|
||||
):
|
||||
super().__init__()
|
||||
self.forward_metadata: NSAMetadata
|
||||
self.device = model_runner.device
|
||||
@@ -185,6 +193,14 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
(max_bs + 1,), dtype=torch.int32, device=model_runner.device
|
||||
)
|
||||
|
||||
# Speculative decoding
|
||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
||||
self.speculative_num_steps = speculative_num_steps
|
||||
self.speculative_num_draft_tokens = (
|
||||
model_runner.server_args.speculative_num_draft_tokens
|
||||
)
|
||||
self.speculative_step_id = speculative_step_id
|
||||
|
||||
def get_device_int32_arange(self, l: int) -> torch.Tensor:
|
||||
if l > len(self._arange_buf):
|
||||
next_pow_of_2 = 1 << (l - 1).bit_length()
|
||||
@@ -208,13 +224,15 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
batch_size = forward_batch.batch_size
|
||||
device = forward_batch.seq_lens.device
|
||||
|
||||
assert (
|
||||
forward_batch.spec_info is None
|
||||
), "Spec decoding is not supported for NSA backend now"
|
||||
cache_seqlens_int32 = forward_batch.seq_lens.to(torch.int32)
|
||||
if forward_batch.forward_mode.is_target_verify():
|
||||
draft_token_num = self.speculative_num_draft_tokens
|
||||
else:
|
||||
draft_token_num = 0
|
||||
|
||||
cache_seqlens_int32 = (forward_batch.seq_lens + draft_token_num).to(torch.int32)
|
||||
cu_seqlens_k = compute_cu_seqlens(cache_seqlens_int32)
|
||||
assert forward_batch.seq_lens_cpu is not None
|
||||
max_seqlen_k = int(forward_batch.seq_lens_cpu.max().item())
|
||||
max_seqlen_k = int(forward_batch.seq_lens_cpu.max().item() + draft_token_num)
|
||||
page_table = forward_batch.req_to_token_pool.req_to_token[
|
||||
forward_batch.req_pool_indices, :max_seqlen_k
|
||||
]
|
||||
@@ -224,6 +242,41 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
max_seqlen_q = 1
|
||||
cu_seqlens_q = self.get_device_int32_arange(batch_size + 1)
|
||||
seqlens_expanded = cache_seqlens_int32
|
||||
elif forward_batch.forward_mode.is_target_verify():
|
||||
max_seqlen_q = self.speculative_num_draft_tokens
|
||||
nsa_max_seqlen_q = self.speculative_num_draft_tokens
|
||||
cu_seqlens_q = torch.arange(
|
||||
0,
|
||||
batch_size * self.speculative_num_draft_tokens + 1,
|
||||
1,
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
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,
|
||||
)
|
||||
]
|
||||
)
|
||||
page_table = torch.repeat_interleave(
|
||||
page_table, repeats=self.speculative_num_draft_tokens, dim=0
|
||||
)
|
||||
elif forward_batch.forward_mode.is_extend():
|
||||
assert (
|
||||
forward_batch.extend_seq_lens_cpu is not None
|
||||
@@ -232,7 +285,11 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
), "All of them must not be None"
|
||||
extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu
|
||||
assert forward_batch.extend_seq_lens is not None
|
||||
if any(forward_batch.extend_prefix_lens_cpu):
|
||||
|
||||
if (
|
||||
any(forward_batch.extend_prefix_lens_cpu)
|
||||
or forward_batch.forward_mode == ForwardMode.DRAFT_EXTEND
|
||||
):
|
||||
max_seqlen_q = max(extend_seq_lens_cpu)
|
||||
cu_seqlens_q = compute_cu_seqlens(
|
||||
forward_batch.extend_seq_lens.to(torch.int32)
|
||||
@@ -277,7 +334,7 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
flashmla_metadata=(
|
||||
self._compute_flashmla_metadata(
|
||||
cache_seqlens=nsa_cache_seqlens_int32,
|
||||
seq_len_q=1, # TODO handle MTP which is not 1
|
||||
seq_len_q=1,
|
||||
)
|
||||
if NSA_DECODE_IMPL == "flashmla_decode"
|
||||
else None
|
||||
@@ -288,6 +345,7 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
nsa_seqlens_expanded=seqlens_expanded,
|
||||
nsa_extend_seq_lens_list=extend_seq_lens_cpu,
|
||||
real_page_table=self._transform_table_1_to_real(page_table),
|
||||
nsa_max_seqlen_q=1,
|
||||
)
|
||||
|
||||
self.forward_metadata = metadata
|
||||
@@ -302,7 +360,9 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
to avoid memory allocations.
|
||||
"""
|
||||
self.decode_cuda_graph_metadata: Dict = {
|
||||
"cache_seqlens": torch.zeros(max_bs, dtype=torch.int32, device=self.device),
|
||||
"cache_seqlens": torch.ones(
|
||||
max_num_tokens, dtype=torch.int32, device=self.device
|
||||
),
|
||||
"cu_seqlens_q": torch.arange(
|
||||
0, max_bs + 1, dtype=torch.int32, device=self.device
|
||||
),
|
||||
@@ -311,7 +371,7 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
),
|
||||
# fake page_table for sparse_prefill
|
||||
"page_table": torch.zeros(
|
||||
max_bs,
|
||||
max_num_tokens,
|
||||
self.max_context_len,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
@@ -319,9 +379,9 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
"flashmla_metadata": (
|
||||
self._compute_flashmla_metadata(
|
||||
cache_seqlens=torch.ones(
|
||||
max_bs, dtype=torch.int32, device=self.device
|
||||
max_num_tokens, dtype=torch.int32, device=self.device
|
||||
),
|
||||
seq_len_q=1, # TODO handle MTP which is not 1
|
||||
seq_len_q=1,
|
||||
)
|
||||
if NSA_DECODE_IMPL == "flashmla_decode"
|
||||
else None
|
||||
@@ -339,50 +399,166 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
spec_info: Optional[SpecInput],
|
||||
):
|
||||
"""Initialize forward metadata for capturing CUDA graph."""
|
||||
assert forward_mode.is_decode_or_idle(), "Only support decode for now"
|
||||
assert (
|
||||
spec_info is None
|
||||
), "Speculative decoding is not supported for NSA backend now"
|
||||
if forward_mode.is_decode_or_idle():
|
||||
# Normal Decode
|
||||
# Get sequence information
|
||||
cache_seqlens_int32 = seq_lens.to(torch.int32)
|
||||
cu_seqlens_k = compute_cu_seqlens(cache_seqlens_int32)
|
||||
|
||||
# Normal Decode
|
||||
# Get sequence information
|
||||
cache_seqlens_int32 = seq_lens.to(torch.int32)
|
||||
cu_seqlens_k = compute_cu_seqlens(cache_seqlens_int32)
|
||||
# Use max context length for seq_len_k
|
||||
page_table_1 = self.decode_cuda_graph_metadata["page_table"][:bs, :]
|
||||
max_seqlen_q = 1
|
||||
max_seqlen_k = page_table_1.shape[1]
|
||||
|
||||
# Use max context length for seq_len_k
|
||||
page_table_1 = self.decode_cuda_graph_metadata["page_table"][:bs, :]
|
||||
max_seq_len_k = page_table_1.shape[1]
|
||||
# Precompute page table
|
||||
# Precompute cumulative sequence lengths
|
||||
|
||||
# Precompute page table
|
||||
# Precompute cumulative sequence lengths
|
||||
# NOTE(dark): this is always arange, since we are decoding
|
||||
cu_seqlens_q = self.decode_cuda_graph_metadata["cu_seqlens_q"][: bs + 1]
|
||||
nsa_cache_seqlens_int32 = compute_nsa_seqlens(
|
||||
cache_seqlens_int32, nsa_index_topk=self.nsa_index_topk
|
||||
)
|
||||
|
||||
seqlens_expanded = cache_seqlens_int32
|
||||
nsa_extend_seq_lens_list = [1] * num_tokens
|
||||
if NSA_DECODE_IMPL == "flashmla_decode":
|
||||
flashmla_metadata = self.decode_cuda_graph_metadata[
|
||||
"flashmla_metadata"
|
||||
].slice(slice(0, num_tokens + 1))
|
||||
flashmla_metadata.copy_(
|
||||
self._compute_flashmla_metadata(
|
||||
cache_seqlens=nsa_cache_seqlens_int32,
|
||||
seq_len_q=1,
|
||||
)
|
||||
)
|
||||
else:
|
||||
flashmla_metadata = None
|
||||
elif forward_mode.is_target_verify():
|
||||
cache_seqlens_int32 = (seq_lens + self.speculative_num_draft_tokens).to(
|
||||
torch.int32
|
||||
)
|
||||
cu_seqlens_k = compute_cu_seqlens(cache_seqlens_int32)
|
||||
max_seqlen_q = 1
|
||||
page_table_1 = self.decode_cuda_graph_metadata["page_table"][
|
||||
: bs * self.speculative_num_draft_tokens, :
|
||||
]
|
||||
max_seqlen_k = page_table_1.shape[1]
|
||||
|
||||
cu_seqlens_q = torch.arange(
|
||||
0,
|
||||
bs * self.speculative_num_draft_tokens + 1,
|
||||
1,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
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.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,
|
||||
)
|
||||
]
|
||||
)
|
||||
nsa_cache_seqlens_int32 = compute_nsa_seqlens(
|
||||
seqlens_expanded, nsa_index_topk=self.nsa_index_topk
|
||||
)
|
||||
nsa_extend_seq_lens_list = [1] * bs * self.speculative_num_draft_tokens
|
||||
|
||||
if NSA_DECODE_IMPL == "flashmla_decode":
|
||||
flashmla_metadata = self.decode_cuda_graph_metadata[
|
||||
"flashmla_metadata"
|
||||
].slice(slice(0, bs * self.speculative_num_draft_tokens + 1))
|
||||
|
||||
flashmla_metadata.copy_(
|
||||
self._compute_flashmla_metadata(
|
||||
cache_seqlens=nsa_cache_seqlens_int32,
|
||||
seq_len_q=1,
|
||||
)
|
||||
)
|
||||
else:
|
||||
flashmla_metadata = None
|
||||
elif forward_mode.is_draft_extend():
|
||||
cache_seqlens_int32 = (seq_lens + self.speculative_num_draft_tokens).to(
|
||||
torch.int32
|
||||
)
|
||||
cu_seqlens_k = compute_cu_seqlens(cache_seqlens_int32)
|
||||
page_table_1 = self.decode_cuda_graph_metadata["page_table"][:bs, :]
|
||||
max_seqlen_k = page_table_1.shape[1]
|
||||
|
||||
extend_seq_lens_cpu = [self.speculative_num_draft_tokens] * bs
|
||||
extend_seq_lens = torch.full(
|
||||
(bs,),
|
||||
self.speculative_num_draft_tokens,
|
||||
device=self.device,
|
||||
dtype=torch.int32,
|
||||
)
|
||||
|
||||
max_seqlen_q = max(extend_seq_lens_cpu)
|
||||
cu_seqlens_q = compute_cu_seqlens(extend_seq_lens.to(torch.int32))
|
||||
|
||||
seqlens_int32_cpu = [
|
||||
self.speculative_num_draft_tokens + kv_len
|
||||
for kv_len in seq_lens.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,
|
||||
)
|
||||
]
|
||||
)
|
||||
nsa_cache_seqlens_int32 = compute_nsa_seqlens(
|
||||
seqlens_expanded, nsa_index_topk=self.nsa_index_topk
|
||||
)
|
||||
nsa_extend_seq_lens_list = [1] * bs
|
||||
|
||||
if NSA_DECODE_IMPL == "flashmla_decode":
|
||||
flashmla_metadata = self.decode_cuda_graph_metadata[
|
||||
"flashmla_metadata"
|
||||
].slice(slice(0, bs * self.speculative_num_draft_tokens + 1))
|
||||
# As the DeepGemm is not support for q_len = 3/4 in Indexer and every token has independent topk_indices,
|
||||
# we made the Q shape [bs * speculative_num_draft_tokens, 1, head_nums, dim].
|
||||
# So seq_len_q is 1 for flashmla_metadata in target_verify and draft_extend mode.
|
||||
flashmla_metadata.copy_(
|
||||
self._compute_flashmla_metadata(
|
||||
cache_seqlens=nsa_cache_seqlens_int32,
|
||||
seq_len_q=1,
|
||||
)
|
||||
)
|
||||
else:
|
||||
flashmla_metadata = None
|
||||
|
||||
# NOTE(dark): this is always arange, since we are decoding
|
||||
cu_seqlens_q = self.decode_cuda_graph_metadata["cu_seqlens_q"][: bs + 1]
|
||||
nsa_cache_seqlens_int32 = compute_nsa_seqlens(
|
||||
cache_seqlens_int32, nsa_index_topk=self.nsa_index_topk
|
||||
)
|
||||
nsa_cu_seqlens_k = compute_cu_seqlens(nsa_cache_seqlens_int32)
|
||||
nsa_cu_seqlens_q = self.get_device_int32_arange(len(nsa_cu_seqlens_k))
|
||||
real_page_table = self._transform_table_1_to_real(page_table_1)
|
||||
|
||||
if NSA_DECODE_IMPL == "flashmla_decode":
|
||||
flashmla_metadata = self.decode_cuda_graph_metadata[
|
||||
"flashmla_metadata"
|
||||
].slice(slice(0, bs + 1))
|
||||
flashmla_metadata.copy_(
|
||||
self._compute_flashmla_metadata(
|
||||
cache_seqlens=nsa_cache_seqlens_int32,
|
||||
seq_len_q=1, # TODO handle MTP which is not 1
|
||||
)
|
||||
)
|
||||
else:
|
||||
flashmla_metadata = None
|
||||
|
||||
metadata = NSAMetadata(
|
||||
page_size=self.real_page_size,
|
||||
cache_seqlens_int32=cache_seqlens_int32,
|
||||
max_seq_len_q=1,
|
||||
max_seq_len_k=max_seq_len_k,
|
||||
max_seq_len_q=max_seqlen_q,
|
||||
max_seq_len_k=max_seqlen_k,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
page_table_1=page_table_1,
|
||||
@@ -390,9 +566,9 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
nsa_cache_seqlens_int32=nsa_cache_seqlens_int32,
|
||||
nsa_cu_seqlens_q=nsa_cu_seqlens_q,
|
||||
nsa_cu_seqlens_k=nsa_cu_seqlens_k,
|
||||
nsa_seqlens_expanded=cache_seqlens_int32,
|
||||
nsa_seqlens_expanded=seqlens_expanded,
|
||||
real_page_table=real_page_table,
|
||||
nsa_extend_seq_lens_list=[1] * bs,
|
||||
nsa_extend_seq_lens_list=nsa_extend_seq_lens_list,
|
||||
)
|
||||
self.decode_cuda_graph_metadata[bs] = metadata
|
||||
self.forward_metadata = metadata
|
||||
@@ -411,33 +587,119 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
):
|
||||
"""Initialize forward metadata for replaying CUDA graph."""
|
||||
assert seq_lens_cpu is not None
|
||||
assert forward_mode.is_decode_or_idle(), "Only support decode for now"
|
||||
assert (
|
||||
spec_info is None
|
||||
), "Speculative decoding is not supported for NSA backend now"
|
||||
|
||||
seq_lens = seq_lens[:bs]
|
||||
seq_lens_cpu = seq_lens_cpu[:bs]
|
||||
req_pool_indices = req_pool_indices[:bs]
|
||||
|
||||
# Normal Decode
|
||||
metadata: NSAMetadata = self.decode_cuda_graph_metadata[bs]
|
||||
max_len = int(seq_lens_cpu.max().item())
|
||||
if forward_mode.is_decode_or_idle():
|
||||
# Normal Decode
|
||||
max_len = int(seq_lens_cpu.max().item())
|
||||
|
||||
cache_seqlens = seq_lens.to(torch.int32)
|
||||
metadata.cache_seqlens_int32.copy_(cache_seqlens)
|
||||
metadata.cu_seqlens_k[1:].copy_(
|
||||
torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32)
|
||||
)
|
||||
page_indices = self.req_to_token[req_pool_indices, :max_len]
|
||||
metadata.page_table_1[:, :max_len].copy_(page_indices)
|
||||
cache_seqlens = seq_lens.to(torch.int32)
|
||||
metadata.cache_seqlens_int32.copy_(cache_seqlens)
|
||||
metadata.cu_seqlens_k[1:].copy_(
|
||||
torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32)
|
||||
)
|
||||
page_indices = self.req_to_token[req_pool_indices, :max_len]
|
||||
metadata.page_table_1[:, :max_len].copy_(page_indices)
|
||||
nsa_cache_seqlens = compute_nsa_seqlens(
|
||||
cache_seqlens, nsa_index_topk=self.nsa_index_topk
|
||||
)
|
||||
metadata.nsa_cache_seqlens_int32.copy_(nsa_cache_seqlens)
|
||||
seqlens_expanded = cache_seqlens
|
||||
elif forward_mode.is_target_verify():
|
||||
max_seqlen_k = int(
|
||||
seq_lens_cpu.max().item() + self.speculative_num_draft_tokens
|
||||
)
|
||||
|
||||
cache_seqlens = (seq_lens + self.speculative_num_draft_tokens).to(
|
||||
torch.int32
|
||||
)
|
||||
metadata.cache_seqlens_int32.copy_(cache_seqlens)
|
||||
metadata.cu_seqlens_k[1:].copy_(
|
||||
torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32)
|
||||
)
|
||||
page_indices = self.req_to_token[req_pool_indices, :max_seqlen_k]
|
||||
page_indices = torch.repeat_interleave(
|
||||
page_indices, repeats=self.speculative_num_draft_tokens, dim=0
|
||||
)
|
||||
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,
|
||||
)
|
||||
]
|
||||
)
|
||||
metadata.nsa_seqlens_expanded.copy_(seqlens_expanded)
|
||||
nsa_cache_seqlens = compute_nsa_seqlens(
|
||||
seqlens_expanded, self.nsa_index_topk
|
||||
)
|
||||
metadata.nsa_cache_seqlens_int32.copy_(nsa_cache_seqlens)
|
||||
elif forward_mode.is_draft_extend():
|
||||
max_seqlen_k = int(seq_lens_cpu.max().item())
|
||||
cache_seqlens = seq_lens.to(torch.int32)
|
||||
metadata.cache_seqlens_int32.copy_(cache_seqlens)
|
||||
metadata.cu_seqlens_k[1:].copy_(
|
||||
torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32)
|
||||
)
|
||||
page_indices = self.req_to_token[req_pool_indices, :max_seqlen_k]
|
||||
metadata.page_table_1[:, :max_seqlen_k].copy_(page_indices)
|
||||
extend_seq_lens_cpu = spec_info.accept_length[:bs].tolist()
|
||||
|
||||
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,
|
||||
)
|
||||
]
|
||||
)
|
||||
metadata.nsa_seqlens_expanded[: seqlens_expanded.size(0)].copy_(
|
||||
seqlens_expanded
|
||||
)
|
||||
nsa_cache_seqlens = compute_nsa_seqlens(
|
||||
seqlens_expanded, self.nsa_index_topk
|
||||
)
|
||||
metadata.nsa_cache_seqlens_int32[: seqlens_expanded.size(0)].copy_(
|
||||
nsa_cache_seqlens
|
||||
)
|
||||
seqlens_expanded_size = seqlens_expanded.size(0)
|
||||
assert (
|
||||
metadata.nsa_cache_seqlens_int32 is not None
|
||||
and metadata.nsa_cu_seqlens_k is not None
|
||||
and self.nsa_index_topk is not None
|
||||
)
|
||||
nsa_cache_seqlens = compute_nsa_seqlens(cache_seqlens, self.nsa_index_topk)
|
||||
metadata.nsa_cache_seqlens_int32.copy_(nsa_cache_seqlens)
|
||||
metadata.nsa_cu_seqlens_k[1:].copy_(
|
||||
|
||||
metadata.nsa_cu_seqlens_k[1 : 1 + seqlens_expanded_size].copy_(
|
||||
torch.cumsum(nsa_cache_seqlens, dim=0, dtype=torch.int32)
|
||||
)
|
||||
# NOTE(dark): (nsa-) cu_seqlens_q is always arange, no need to copy
|
||||
@@ -451,10 +713,13 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
assert metadata.real_page_table is metadata.page_table_1
|
||||
|
||||
if NSA_DECODE_IMPL == "flashmla_decode":
|
||||
metadata.flashmla_metadata.copy_(
|
||||
flashmla_metadata = metadata.flashmla_metadata.slice(
|
||||
slice(0, seqlens_expanded_size + 1)
|
||||
)
|
||||
flashmla_metadata.copy_(
|
||||
self._compute_flashmla_metadata(
|
||||
cache_seqlens=nsa_cache_seqlens,
|
||||
seq_len_q=1, # TODO handle MTP which is not 1
|
||||
seq_len_q=1,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -473,10 +738,7 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
k_rope: Optional[torch.Tensor] = None,
|
||||
topk_indices: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
assert (
|
||||
not forward_batch.forward_mode.is_target_verify()
|
||||
and not forward_batch.forward_mode.is_draft_extend()
|
||||
), "NSA backend doesn't support speculative decoding"
|
||||
|
||||
if k is not None:
|
||||
assert v is not None
|
||||
if save_kv_cache:
|
||||
@@ -884,3 +1146,58 @@ class NativeSparseAttnBackend(AttentionBackend):
|
||||
flashmla_metadata=flashmla_metadata,
|
||||
num_splits=num_splits,
|
||||
)
|
||||
|
||||
|
||||
class NativeSparseAttnMultiStepBackend:
|
||||
|
||||
def __init__(
|
||||
self, model_runner: ModelRunner, topk: int, speculative_num_steps: int
|
||||
):
|
||||
self.model_runner = model_runner
|
||||
self.topk = topk
|
||||
self.speculative_num_steps = speculative_num_steps
|
||||
self.attn_backends = []
|
||||
for i in range(self.speculative_num_steps):
|
||||
self.attn_backends.append(
|
||||
NativeSparseAttnBackend(
|
||||
model_runner,
|
||||
speculative_step_id=i,
|
||||
topk=self.topk,
|
||||
speculative_num_steps=self.speculative_num_steps,
|
||||
)
|
||||
)
|
||||
|
||||
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||
for i in range(self.speculative_num_steps - 1):
|
||||
self.attn_backends[i].init_forward_metadata(forward_batch)
|
||||
|
||||
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
|
||||
for i in range(self.speculative_num_steps):
|
||||
self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens)
|
||||
|
||||
def init_forward_metadata_capture_cuda_graph(self, forward_batch: ForwardBatch):
|
||||
for i in range(self.speculative_num_steps):
|
||||
self.attn_backends[i].init_forward_metadata_capture_cuda_graph(
|
||||
forward_batch.batch_size,
|
||||
forward_batch.batch_size * self.topk,
|
||||
forward_batch.req_pool_indices,
|
||||
forward_batch.seq_lens,
|
||||
encoder_lens=None,
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
spec_info=forward_batch.spec_info,
|
||||
)
|
||||
|
||||
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,
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
spec_info=forward_batch.spec_info,
|
||||
seq_lens_cpu=forward_batch.seq_lens_cpu,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user