[DeepSeek v3.2] opt Context Parallelism: support fused moe, multi batch and fp8 kvcache (#13959)
This commit is contained in:
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import IntEnum, auto
|
||||
from typing import TYPE_CHECKING, Dict, List, Literal, Optional, TypeAlias
|
||||
from typing import TYPE_CHECKING, Dict, List, Literal, Optional, Tuple, TypeAlias
|
||||
|
||||
import torch
|
||||
|
||||
@@ -25,8 +25,12 @@ from sglang.srt.layers.attention.nsa.utils import (
|
||||
NSA_ENABLE_MTP_PRECOMPUTE_METADATA,
|
||||
NSA_FLASHMLA_BACKEND_DECODE_COMPUTE_FP8,
|
||||
NSA_FUSE_TOPK,
|
||||
can_nsa_prefill_cp_round_robin_split,
|
||||
compute_nsa_seqlens,
|
||||
is_nsa_enable_prefill_cp,
|
||||
nsa_cp_round_robin_split_data,
|
||||
nsa_cp_round_robin_split_q_seqs,
|
||||
pad_nsa_cache_seqlens,
|
||||
)
|
||||
from sglang.srt.layers.attention.trtllm_mla_backend import _concat_mla_absorb_q_general
|
||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||
@@ -125,6 +129,13 @@ class NSAMetadata:
|
||||
# shape: (seq_lens_sum,)
|
||||
topk_indices_offset: Optional[torch.Tensor] = None
|
||||
|
||||
# k_start and k_end in kv cache for each token.
|
||||
indexer_k_start_end: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
|
||||
# seq lens for each batch.
|
||||
indexer_seq_lens_cpu: Optional[torch.Tensor] = None
|
||||
# batch index for each token.
|
||||
token_to_batch_idx: Optional[torch.Tensor] = None
|
||||
|
||||
|
||||
class TopkTransformMethod(IntEnum):
|
||||
# Transform topk indices to indices to the page table (page_size = 1)
|
||||
@@ -172,6 +183,15 @@ class NSAIndexerMetadata(BaseIndexerMetadata):
|
||||
def get_cu_seqlens_k(self) -> torch.Tensor:
|
||||
return self.attn_metadata.cu_seqlens_k
|
||||
|
||||
def get_indexer_kvcache_range(self) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
return self.attn_metadata.indexer_k_start_end
|
||||
|
||||
def get_indexer_seq_len_cpu(self) -> torch.Tensor:
|
||||
return self.attn_metadata.indexer_seq_lens_cpu
|
||||
|
||||
def get_token_to_batch_idx(self) -> torch.Tensor:
|
||||
return self.attn_metadata.token_to_batch_idx
|
||||
|
||||
def topk_transform(
|
||||
self,
|
||||
logits: torch.Tensor,
|
||||
@@ -354,6 +374,13 @@ class NativeSparseAttnBackend(
|
||||
# Centralized dispatch: decide all strategies for this batch
|
||||
self.set_nsa_prefill_impl(forward_batch)
|
||||
topk_transform_method = self.get_topk_transform_method()
|
||||
# Batch indices selected when cp enabled: After splitting multiple sequences,
|
||||
# a certain cp rank may not have some of these sequences.
|
||||
# We use bs_idx_cpu to mark which sequences are finally selected by the current cp rank,
|
||||
# a default value of None indicates that all sequences are selected.
|
||||
bs_idx_cpu = None
|
||||
# seq_len_cpu of selected sequences
|
||||
indexer_seq_lens_cpu = forward_batch.seq_lens_cpu
|
||||
|
||||
if forward_batch.forward_mode.is_decode_or_idle():
|
||||
extend_seq_lens_cpu = [1] * batch_size
|
||||
@@ -441,7 +468,6 @@ class NativeSparseAttnBackend(
|
||||
page_table = torch.repeat_interleave(
|
||||
page_table, repeats=forward_batch.extend_seq_lens, dim=0
|
||||
)
|
||||
|
||||
elif forward_batch.forward_mode.is_extend():
|
||||
assert (
|
||||
forward_batch.extend_seq_lens_cpu is not None
|
||||
@@ -450,18 +476,7 @@ class NativeSparseAttnBackend(
|
||||
), "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)
|
||||
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)
|
||||
)
|
||||
else:
|
||||
max_seqlen_q = max_seqlen_k
|
||||
cu_seqlens_q = cu_seqlens_k
|
||||
extend_seq_lens = forward_batch.extend_seq_lens
|
||||
|
||||
seqlens_expanded = torch.cat(
|
||||
[
|
||||
@@ -479,6 +494,36 @@ class NativeSparseAttnBackend(
|
||||
]
|
||||
)
|
||||
|
||||
if can_nsa_prefill_cp_round_robin_split(forward_batch):
|
||||
seqlens_expanded = nsa_cp_round_robin_split_data(seqlens_expanded)
|
||||
extend_seq_lens_cpu, extend_seq_lens, bs_idx_cpu, bs_idx = (
|
||||
nsa_cp_round_robin_split_q_seqs(
|
||||
extend_seq_lens_cpu, extend_seq_lens
|
||||
)
|
||||
)
|
||||
indexer_seq_lens_cpu = indexer_seq_lens_cpu[bs_idx_cpu]
|
||||
cache_seqlens_int32 = cache_seqlens_int32[bs_idx]
|
||||
cu_seqlens_k = compute_cu_seqlens(cache_seqlens_int32)
|
||||
max_seqlen_k = (
|
||||
int(indexer_seq_lens_cpu.max().item() + draft_token_num)
|
||||
if len(indexer_seq_lens_cpu) != 0
|
||||
else 0
|
||||
)
|
||||
page_table = page_table[bs_idx, :max_seqlen_k]
|
||||
|
||||
if (
|
||||
any(forward_batch.extend_prefix_lens_cpu)
|
||||
or forward_batch.forward_mode == ForwardMode.DRAFT_EXTEND
|
||||
or bs_idx_cpu is not None
|
||||
):
|
||||
max_seqlen_q = (
|
||||
max(extend_seq_lens_cpu) if len(extend_seq_lens_cpu) != 0 else 1
|
||||
)
|
||||
cu_seqlens_q = compute_cu_seqlens(extend_seq_lens.to(torch.int32))
|
||||
else:
|
||||
max_seqlen_q = max_seqlen_k
|
||||
cu_seqlens_q = cu_seqlens_k
|
||||
|
||||
# Check if MHA FP8 dequantization is needed
|
||||
mha_dequantize_needed = (
|
||||
self.use_mha
|
||||
@@ -496,13 +541,13 @@ class NativeSparseAttnBackend(
|
||||
[
|
||||
page_table[i, :kv_len]
|
||||
for i, kv_len in enumerate(
|
||||
forward_batch.seq_lens_cpu.tolist(),
|
||||
indexer_seq_lens_cpu.tolist(),
|
||||
)
|
||||
]
|
||||
)
|
||||
assert (
|
||||
page_table_1_flattened.shape[0] == forward_batch.seq_lens_sum
|
||||
), f"{page_table_1_flattened.shape[0] = } must be the same as {forward_batch.seq_lens_sum = }"
|
||||
assert page_table_1_flattened.shape[0] == sum(
|
||||
indexer_seq_lens_cpu
|
||||
), f"{page_table_1_flattened.shape[0] = } must be the same as {sum(indexer_seq_lens_cpu) = }"
|
||||
|
||||
# Validate indices when logical tokens exceed physical capacity
|
||||
# This is likely to be triggered by PP with high kv reuse & parallelism
|
||||
@@ -520,16 +565,22 @@ class NativeSparseAttnBackend(
|
||||
if topk_transform_method == TopkTransformMethod.RAGGED:
|
||||
topk_indices_offset = torch.repeat_interleave(
|
||||
cu_seqlens_k[:-1],
|
||||
forward_batch.extend_seq_lens,
|
||||
extend_seq_lens,
|
||||
)
|
||||
else:
|
||||
assert False, f"Unsupported {forward_batch.forward_mode = }"
|
||||
|
||||
indexer_k_start_end, token_to_batch_idx = self._cal_indexer_k_start_end(
|
||||
forward_batch, bs_idx_cpu
|
||||
)
|
||||
# 1D, expanded seqlens (1D means cheap to compute, so always compute it)
|
||||
nsa_cache_seqlens_int32 = compute_nsa_seqlens(
|
||||
original_seq_lens=seqlens_expanded,
|
||||
nsa_index_topk=self.nsa_index_topk,
|
||||
)
|
||||
nsa_cache_seqlens_int32 = pad_nsa_cache_seqlens(
|
||||
forward_batch, nsa_cache_seqlens_int32
|
||||
)
|
||||
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))
|
||||
|
||||
@@ -586,10 +637,88 @@ class NativeSparseAttnBackend(
|
||||
real_page_table=self._transform_table_1_to_real(page_table),
|
||||
nsa_max_seqlen_q=1,
|
||||
topk_indices_offset=topk_indices_offset,
|
||||
indexer_k_start_end=indexer_k_start_end,
|
||||
indexer_seq_lens_cpu=indexer_seq_lens_cpu,
|
||||
token_to_batch_idx=token_to_batch_idx,
|
||||
)
|
||||
|
||||
self.forward_metadata = metadata
|
||||
|
||||
def _cal_indexer_k_start_end(
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
bs_idx: Optional[List[int]] = None,
|
||||
):
|
||||
if not forward_batch.forward_mode.is_extend_without_speculative():
|
||||
return None, None
|
||||
if forward_batch.batch_size == 0 or (bs_idx is not None and len(bs_idx) == 0):
|
||||
empty_t = torch.empty(0, dtype=torch.int32, device=self.device)
|
||||
return (empty_t, empty_t), empty_t
|
||||
|
||||
# Suppose there are two requests, with extend_seq_len = [3, 2]
|
||||
# and seq_lens = [10, 4]
|
||||
# The logits matrix looks like this, with * representing the valid logits
|
||||
# and - representing the invalid logits:
|
||||
#
|
||||
# ********--|----
|
||||
# *********-|----
|
||||
# **********|----
|
||||
# ----------|***-
|
||||
# ----------|****
|
||||
#
|
||||
# ks = [0, 0, 0, 10, 10]
|
||||
# ke = [8, 9, 10, 13, 14]
|
||||
ks_list = []
|
||||
ke_list = []
|
||||
token_to_batch_idx = []
|
||||
|
||||
q_offset = 0
|
||||
k_offset = 0
|
||||
|
||||
assert (
|
||||
forward_batch.seq_lens_cpu is not None
|
||||
and forward_batch.extend_seq_lens_cpu is not None
|
||||
)
|
||||
for i in range(forward_batch.batch_size):
|
||||
seq_len = forward_batch.seq_lens_cpu[i].item()
|
||||
assert isinstance(seq_len, int)
|
||||
extend_seq_len = forward_batch.extend_seq_lens_cpu[i]
|
||||
ks = torch.full(
|
||||
(extend_seq_len,), k_offset, dtype=torch.int32, device=self.device
|
||||
)
|
||||
kv_len = seq_len
|
||||
if forward_batch.forward_mode.is_target_verify():
|
||||
kv_len += self.speculative_num_draft_tokens
|
||||
seq_lens_expanded = torch.arange(
|
||||
kv_len - extend_seq_len + 1,
|
||||
kv_len + 1,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
ke = ks + seq_lens_expanded
|
||||
ks_list.append(ks)
|
||||
ke_list.append(ke)
|
||||
|
||||
# bi: The index within the selected batch bs_idx. Entries that were not selected are ignored.
|
||||
bi = bs_idx.index(i) if (bs_idx is not None and i in bs_idx) else i
|
||||
tb = torch.full(
|
||||
(extend_seq_len,), bi, dtype=torch.int32, device=self.device
|
||||
)
|
||||
token_to_batch_idx.append(tb)
|
||||
|
||||
if bs_idx is None or i in bs_idx: # skip batch not included in bs_idx
|
||||
q_offset += extend_seq_len
|
||||
k_offset += seq_len
|
||||
|
||||
ks = torch.cat(ks_list, dim=0)
|
||||
ke = torch.cat(ke_list, dim=0)
|
||||
token_to_batch_idx = torch.cat(token_to_batch_idx, dim=0)
|
||||
if bs_idx is not None:
|
||||
assert can_nsa_prefill_cp_round_robin_split(forward_batch)
|
||||
ks = nsa_cp_round_robin_split_data(ks)
|
||||
ke = nsa_cp_round_robin_split_data(ke)
|
||||
token_to_batch_idx = nsa_cp_round_robin_split_data(token_to_batch_idx)
|
||||
return (ks, ke), token_to_batch_idx
|
||||
|
||||
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
|
||||
"""Initialize CUDA graph state for the attention backend.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user