[DeepSeek v3.2] opt Context Parallelism: support fused moe, multi batch and fp8 kvcache (#13959)

This commit is contained in:
Yongfei Xu
2026-01-02 23:49:14 +08:00
committed by GitHub
parent 0eae831797
commit 0d244116d2
14 changed files with 602 additions and 263 deletions
+149 -20
View File
@@ -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.