[DCP] Support decode context parallelism on the trtllm_mla decode path (#33926)
Co-authored-by: kpham-sgl <khoa.pham@radixark.ai> Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
kpham-sgl
Claude Opus 5
parent
97744189b8
commit
f50b4ad7ae
@@ -10,10 +10,8 @@ returns the rank-local ``(out, lse)`` needed by the cross-rank merge in
|
||||
``deepseek_common/attention_forward_methods/forward_mla.py``.
|
||||
|
||||
Non-DCP (``dcp_size == 1``) decode falls through to the base cute-dsl path
|
||||
unchanged. The DCP metadata helpers below are intentionally duplicated from
|
||||
:mod:`tokenspeed_mla_backend` (they are kernel-agnostic) so that TokenSpeed
|
||||
stays untouched; both should collapse into the base once the cute-dsl decode
|
||||
path is stable (see the TODO in tokenspeed_mla_backend.py).
|
||||
unchanged. The DCP metadata helpers live on :class:`TRTLLMMLABackend`; this
|
||||
module only supplies the cute-dsl kernel call and its decode forward.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -23,26 +21,18 @@ from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention.dcp_kernels import (
|
||||
create_mla_kv_page_table_for_dcp,
|
||||
)
|
||||
from sglang.kernels.ops.attention.fixup_zero_kv import fixup_zero_kv_rows
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
concat_mla_absorb_q_general,
|
||||
mla_quantize_and_rope_for_fp8,
|
||||
mla_quantize_without_rope_for_fp8,
|
||||
)
|
||||
from sglang.kernels.ops.kvcache.kv_indices import (
|
||||
get_num_kv_index_blocks_flashmla,
|
||||
get_num_page_per_block_flashmla,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.trtllm_mla_backend import (
|
||||
TRTLLMMLABackend,
|
||||
TRTLLMMLAMultiStepDraftBackend,
|
||||
)
|
||||
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.layers.logits_processor import get_in_autotune_dummy_run
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import is_flashinfer_available
|
||||
|
||||
@@ -51,6 +41,7 @@ if is_flashinfer_available():
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.radix_attention import RadixAttention
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -74,195 +65,6 @@ class CuteDslMLABackend(TRTLLMMLABackend):
|
||||
backend="cute-dsl",
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# DCP metadata (rank-local KV lengths + page table).
|
||||
# Duplicated from TokenspeedMLABackend — kernel-agnostic, keyed only on
|
||||
# dcp_size / dcp_rank / page_size / req_to_token.
|
||||
# ------------------------------------------------------------------
|
||||
def _get_dcp_local_seq_lens(self, seq_lens: torch.Tensor) -> torch.Tensor:
|
||||
parallel = get_parallel()
|
||||
if not parallel.dcp_enabled:
|
||||
return seq_lens
|
||||
return get_dcp_lens(seq_lens, parallel.dcp_size, parallel.dcp_rank).to(
|
||||
torch.int32
|
||||
)
|
||||
|
||||
def _get_dcp_local_max_seq_len(self, max_seq_len: int) -> int:
|
||||
parallel = get_parallel()
|
||||
if not parallel.dcp_enabled:
|
||||
return max_seq_len
|
||||
local_max = max_seq_len // parallel.dcp_size + int(
|
||||
parallel.dcp_rank < max_seq_len % parallel.dcp_size
|
||||
)
|
||||
# A positive scheduling bound is required even when every sequence in a
|
||||
# padded graph row is empty on this rank.
|
||||
return max(local_max, 1)
|
||||
|
||||
def _fill_dcp_block_kv_indices(
|
||||
self,
|
||||
block_kv_indices: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
local_seq_lens: torch.Tensor,
|
||||
) -> None:
|
||||
parallel = get_parallel()
|
||||
pages_per_block = get_num_page_per_block_flashmla(self.page_size)
|
||||
create_mla_kv_page_table_for_dcp[
|
||||
(
|
||||
block_kv_indices.shape[0],
|
||||
get_num_kv_index_blocks_flashmla(
|
||||
block_kv_indices.shape[1], self.page_size
|
||||
),
|
||||
)
|
||||
](
|
||||
self.req_to_token,
|
||||
req_pool_indices,
|
||||
local_seq_lens,
|
||||
block_kv_indices,
|
||||
self.req_to_token.stride(0),
|
||||
block_kv_indices.stride(0),
|
||||
PHYSICAL_PAGE_SIZE=self.page_size,
|
||||
DCP_SIZE=parallel.dcp_size,
|
||||
DCP_RANK=parallel.dcp_rank,
|
||||
PAGES_PER_BLOCK=pages_per_block,
|
||||
)
|
||||
|
||||
def _create_block_kv_indices(
|
||||
self,
|
||||
batch_size: int,
|
||||
max_blocks: int,
|
||||
req_pool_indices: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
device: torch.device,
|
||||
) -> torch.Tensor:
|
||||
if not get_parallel().dcp_enabled:
|
||||
return super()._create_block_kv_indices(
|
||||
batch_size,
|
||||
max_blocks,
|
||||
req_pool_indices,
|
||||
seq_lens,
|
||||
device,
|
||||
)
|
||||
block_kv_indices = torch.full(
|
||||
(batch_size, max_blocks), -1, dtype=torch.int32, device=device
|
||||
)
|
||||
self._fill_dcp_block_kv_indices(
|
||||
block_kv_indices,
|
||||
req_pool_indices,
|
||||
self._get_dcp_local_seq_lens(seq_lens),
|
||||
)
|
||||
return block_kv_indices
|
||||
|
||||
def _init_cuda_graph_metadata(
|
||||
self,
|
||||
bs: int,
|
||||
num_tokens: int,
|
||||
forward_mode,
|
||||
seq_lens: torch.Tensor,
|
||||
device: torch.device,
|
||||
):
|
||||
super()._init_cuda_graph_metadata(
|
||||
bs, num_tokens, forward_mode, seq_lens, device
|
||||
)
|
||||
if get_parallel().dcp_enabled:
|
||||
metadata = self.forward_decode_metadata
|
||||
if metadata.global_seq_lens_k is None:
|
||||
# Plain decode under DCP also keeps the int32 GLOBAL lens in a
|
||||
# capture-stable buffer (super allocates it only for verify):
|
||||
# the DCP kernel consumes both the rank-local and the global
|
||||
# lens every MLA layer, so both are maintained once per step.
|
||||
metadata.global_seq_lens_k = torch.zeros(
|
||||
(bs,), dtype=torch.int32, device=device
|
||||
)
|
||||
metadata.max_seq_len_k = self._get_dcp_local_max_seq_len(
|
||||
self.max_context_len
|
||||
+ (self.num_draft_tokens if forward_mode.is_target_verify() else 0)
|
||||
)
|
||||
|
||||
def _apply_cuda_graph_metadata(
|
||||
self,
|
||||
bs: int,
|
||||
req_pool_indices: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
forward_mode,
|
||||
):
|
||||
if not get_parallel().dcp_enabled:
|
||||
return super()._apply_cuda_graph_metadata(
|
||||
bs,
|
||||
req_pool_indices,
|
||||
seq_lens,
|
||||
forward_mode,
|
||||
)
|
||||
|
||||
metadata = self.decode_cuda_graph_metadata[bs]
|
||||
if forward_mode.is_target_verify():
|
||||
torch.add(
|
||||
seq_lens[:bs],
|
||||
self.num_draft_tokens,
|
||||
out=metadata.global_seq_lens_k,
|
||||
)
|
||||
metadata.seq_lens_k.copy_(
|
||||
self._get_dcp_local_seq_lens(metadata.global_seq_lens_k)
|
||||
)
|
||||
local_seq_lens = metadata.seq_lens_k
|
||||
elif forward_mode.is_draft_extend_v2():
|
||||
num_tokens_per_req = self.num_draft_tokens
|
||||
metadata.max_seq_len_q = num_tokens_per_req
|
||||
metadata.sum_seq_lens_q = num_tokens_per_req * bs
|
||||
seq_lens = seq_lens[:bs]
|
||||
metadata.seq_lens_k.copy_(seq_lens)
|
||||
local_seq_lens = self._get_dcp_local_seq_lens(seq_lens)
|
||||
else:
|
||||
seq_lens = seq_lens[:bs]
|
||||
# Hoist: refresh the int32 global + rank-local lens once per step
|
||||
# into the capture-stable buffers; forward_decode reads them
|
||||
# instead of recomputing get_dcp_lens + two int32 casts per MLA
|
||||
# layer.
|
||||
metadata.global_seq_lens_k.copy_(seq_lens)
|
||||
metadata.seq_lens_k.copy_(self._get_dcp_local_seq_lens(seq_lens))
|
||||
local_seq_lens = metadata.seq_lens_k
|
||||
|
||||
self._fill_dcp_block_kv_indices(
|
||||
metadata.block_kv_indices,
|
||||
req_pool_indices[:bs],
|
||||
local_seq_lens,
|
||||
)
|
||||
|
||||
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||
super().init_forward_metadata(forward_batch)
|
||||
if (
|
||||
get_parallel().dcp_enabled
|
||||
and self.forward_decode_metadata is not None
|
||||
and (
|
||||
forward_batch.forward_mode.is_decode_or_idle()
|
||||
or forward_batch.forward_mode.is_target_verify()
|
||||
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||
)
|
||||
):
|
||||
if forward_batch.forward_mode.is_target_verify():
|
||||
metadata = self.forward_decode_metadata
|
||||
metadata.global_seq_lens_k = metadata.seq_lens_k
|
||||
metadata.seq_lens_k = self._get_dcp_local_seq_lens(
|
||||
metadata.global_seq_lens_k
|
||||
)
|
||||
elif (
|
||||
forward_batch.forward_mode.is_decode_or_idle()
|
||||
and self.forward_decode_metadata.seq_lens_k is not None
|
||||
):
|
||||
# Same hoist as verify: the parent stored the int32 GLOBAL
|
||||
# lens in seq_lens_k; keep it as global_seq_lens_k and derive
|
||||
# the rank-local view once per step (forward_decode consumes
|
||||
# both every MLA layer).
|
||||
metadata = self.forward_decode_metadata
|
||||
metadata.global_seq_lens_k = metadata.seq_lens_k
|
||||
metadata.seq_lens_k = self._get_dcp_local_seq_lens(
|
||||
metadata.global_seq_lens_k
|
||||
)
|
||||
self.forward_decode_metadata.max_seq_len_k = (
|
||||
self._get_dcp_local_max_seq_len(
|
||||
self.forward_decode_metadata.max_seq_len_k
|
||||
)
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Kernel + decode forward.
|
||||
# ------------------------------------------------------------------
|
||||
@@ -289,7 +91,16 @@ class CuteDslMLABackend(TRTLLMMLABackend):
|
||||
"""
|
||||
if cp_world <= 1:
|
||||
return super()._run_decode_kernel(
|
||||
query, kv_cache, block_tables, seq_lens, max_seq_len, layer
|
||||
query,
|
||||
kv_cache,
|
||||
block_tables,
|
||||
seq_lens,
|
||||
max_seq_len,
|
||||
layer,
|
||||
causal_seqs=causal_seqs,
|
||||
cp_world=cp_world,
|
||||
cp_rank=cp_rank,
|
||||
return_lse=return_lse,
|
||||
)
|
||||
if causal_seqs is None:
|
||||
raise ValueError(
|
||||
@@ -339,6 +150,8 @@ class CuteDslMLABackend(TRTLLMMLABackend):
|
||||
llama_4_scaling: Optional[torch.Tensor] = None,
|
||||
):
|
||||
parallel = get_parallel()
|
||||
if parallel.dcp_enabled and get_in_autotune_dummy_run():
|
||||
return self._dummy_dcp_decode_for_autotune(q, layer)
|
||||
if not parallel.dcp_enabled:
|
||||
return super().forward_decode(
|
||||
q,
|
||||
|
||||
@@ -22,10 +22,11 @@ from __future__ import annotations
|
||||
|
||||
"""Attention backend for the tokenspeed-mla CuTe DSL kernels on Blackwell.
|
||||
|
||||
Subclasses :class:`TRTLLMMLABackend` to share its MLA data preparation and
|
||||
prefill plumbing. Decode-context parallelism is implemented here because the
|
||||
TokenSpeed decode kernel natively accepts CP rank/world metadata and returns
|
||||
the partial log-sum-exp needed by the cross-rank merge.
|
||||
Subclasses :class:`TRTLLMMLABackend` to share its MLA data preparation, prefill
|
||||
plumbing, and DCP metadata (rank-local KV lengths and page table). The decode
|
||||
forward lives here because the TokenSpeed decode kernel natively accepts CP
|
||||
rank/world metadata and returns the partial log-sum-exp needed by the cross-rank
|
||||
merge.
|
||||
"""
|
||||
|
||||
import logging
|
||||
@@ -34,9 +35,6 @@ from typing import TYPE_CHECKING, Optional
|
||||
import torch
|
||||
|
||||
from sglang.kernels.jit.utils import is_arch_support_pdl
|
||||
from sglang.kernels.ops.attention.dcp_kernels import (
|
||||
create_mla_kv_page_table_for_dcp,
|
||||
)
|
||||
from sglang.kernels.ops.attention.fixup_zero_kv import fixup_zero_kv_rows
|
||||
from sglang.kernels.ops.attention.mla_kv_pack_quantize_fp8 import (
|
||||
mla_kv_pack_quantize_fp8,
|
||||
@@ -45,16 +43,11 @@ from sglang.kernels.ops.attention.utils import (
|
||||
mla_quantize_and_rope_for_fp8,
|
||||
mla_quantize_without_rope_for_fp8,
|
||||
)
|
||||
from sglang.kernels.ops.kvcache.kv_indices import (
|
||||
get_num_kv_index_blocks_flashmla,
|
||||
get_num_page_per_block_flashmla,
|
||||
)
|
||||
from sglang.kernels.ops.quantization.fp8_quantize import fp8_quantize
|
||||
from sglang.srt.layers.attention.trtllm_mla_backend import (
|
||||
TRTLLMMLABackend,
|
||||
TRTLLMMLAMultiStepDraftBackend,
|
||||
)
|
||||
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
||||
from sglang.srt.layers.logits_processor import get_in_autotune_dummy_run
|
||||
from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
@@ -302,165 +295,6 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
||||
k_nope, k_pe, v, enable_pdl=is_arch_support_pdl()
|
||||
)
|
||||
|
||||
def _get_dcp_local_seq_lens(self, seq_lens: torch.Tensor) -> torch.Tensor:
|
||||
parallel = get_parallel()
|
||||
if not parallel.dcp_enabled:
|
||||
return seq_lens
|
||||
return get_dcp_lens(seq_lens, parallel.dcp_size, parallel.dcp_rank).to(
|
||||
torch.int32
|
||||
)
|
||||
|
||||
def _get_dcp_local_max_seq_len(self, max_seq_len: int) -> int:
|
||||
parallel = get_parallel()
|
||||
if not parallel.dcp_enabled:
|
||||
return max_seq_len
|
||||
local_max = max_seq_len // parallel.dcp_size + int(
|
||||
parallel.dcp_rank < max_seq_len % parallel.dcp_size
|
||||
)
|
||||
# TokenSpeed requires a positive scheduling bound even when every
|
||||
# sequence in a padded graph row is empty on this rank.
|
||||
return max(local_max, 1)
|
||||
|
||||
def _fill_dcp_block_kv_indices(
|
||||
self,
|
||||
block_kv_indices: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
local_seq_lens: torch.Tensor,
|
||||
) -> None:
|
||||
parallel = get_parallel()
|
||||
pages_per_block = get_num_page_per_block_flashmla(self.page_size)
|
||||
create_mla_kv_page_table_for_dcp[
|
||||
(
|
||||
block_kv_indices.shape[0],
|
||||
get_num_kv_index_blocks_flashmla(
|
||||
block_kv_indices.shape[1], self.page_size
|
||||
),
|
||||
)
|
||||
](
|
||||
self.req_to_token,
|
||||
req_pool_indices,
|
||||
local_seq_lens,
|
||||
block_kv_indices,
|
||||
self.req_to_token.stride(0),
|
||||
block_kv_indices.stride(0),
|
||||
PHYSICAL_PAGE_SIZE=self.page_size,
|
||||
DCP_SIZE=parallel.dcp_size,
|
||||
DCP_RANK=parallel.dcp_rank,
|
||||
PAGES_PER_BLOCK=pages_per_block,
|
||||
)
|
||||
|
||||
def _create_block_kv_indices(
|
||||
self,
|
||||
batch_size: int,
|
||||
max_blocks: int,
|
||||
req_pool_indices: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
device: torch.device,
|
||||
) -> torch.Tensor:
|
||||
if not get_parallel().dcp_enabled:
|
||||
return super()._create_block_kv_indices(
|
||||
batch_size,
|
||||
max_blocks,
|
||||
req_pool_indices,
|
||||
seq_lens,
|
||||
device,
|
||||
)
|
||||
|
||||
block_kv_indices = torch.full(
|
||||
(batch_size, max_blocks), -1, dtype=torch.int32, device=device
|
||||
)
|
||||
self._fill_dcp_block_kv_indices(
|
||||
block_kv_indices,
|
||||
req_pool_indices,
|
||||
self._get_dcp_local_seq_lens(seq_lens),
|
||||
)
|
||||
return block_kv_indices
|
||||
|
||||
def _init_cuda_graph_metadata(
|
||||
self,
|
||||
bs: int,
|
||||
num_tokens: int,
|
||||
forward_mode,
|
||||
seq_lens: torch.Tensor,
|
||||
device: torch.device,
|
||||
):
|
||||
super()._init_cuda_graph_metadata(
|
||||
bs, num_tokens, forward_mode, seq_lens, device
|
||||
)
|
||||
if get_parallel().dcp_enabled:
|
||||
self.forward_decode_metadata.max_seq_len_k = (
|
||||
self._get_dcp_local_max_seq_len(
|
||||
self.max_context_len
|
||||
+ (self.num_draft_tokens if forward_mode.is_target_verify() else 0)
|
||||
)
|
||||
)
|
||||
|
||||
def _apply_cuda_graph_metadata(
|
||||
self,
|
||||
bs: int,
|
||||
req_pool_indices: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
forward_mode,
|
||||
):
|
||||
if not get_parallel().dcp_enabled:
|
||||
return super()._apply_cuda_graph_metadata(
|
||||
bs,
|
||||
req_pool_indices,
|
||||
seq_lens,
|
||||
forward_mode,
|
||||
)
|
||||
|
||||
metadata = self.decode_cuda_graph_metadata[bs]
|
||||
if forward_mode.is_target_verify():
|
||||
torch.add(
|
||||
seq_lens[:bs],
|
||||
self.num_draft_tokens,
|
||||
out=metadata.global_seq_lens_k,
|
||||
)
|
||||
metadata.seq_lens_k.copy_(
|
||||
self._get_dcp_local_seq_lens(metadata.global_seq_lens_k)
|
||||
)
|
||||
local_seq_lens = metadata.seq_lens_k
|
||||
elif forward_mode.is_draft_extend_v2():
|
||||
num_tokens_per_req = self.num_draft_tokens
|
||||
metadata.max_seq_len_q = num_tokens_per_req
|
||||
metadata.sum_seq_lens_q = num_tokens_per_req * bs
|
||||
seq_lens = seq_lens[:bs]
|
||||
metadata.seq_lens_k.copy_(seq_lens)
|
||||
local_seq_lens = self._get_dcp_local_seq_lens(seq_lens)
|
||||
else:
|
||||
seq_lens = seq_lens[:bs]
|
||||
local_seq_lens = self._get_dcp_local_seq_lens(seq_lens)
|
||||
|
||||
self._fill_dcp_block_kv_indices(
|
||||
metadata.block_kv_indices,
|
||||
req_pool_indices[:bs],
|
||||
local_seq_lens,
|
||||
)
|
||||
|
||||
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||
super().init_forward_metadata(forward_batch)
|
||||
if (
|
||||
get_parallel().dcp_enabled
|
||||
and self.forward_decode_metadata is not None
|
||||
and (
|
||||
forward_batch.forward_mode.is_decode_or_idle()
|
||||
or forward_batch.forward_mode.is_target_verify()
|
||||
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||
)
|
||||
):
|
||||
if forward_batch.forward_mode.is_target_verify():
|
||||
metadata = self.forward_decode_metadata
|
||||
metadata.global_seq_lens_k = metadata.seq_lens_k
|
||||
metadata.seq_lens_k = self._get_dcp_local_seq_lens(
|
||||
metadata.global_seq_lens_k
|
||||
)
|
||||
self.forward_decode_metadata.max_seq_len_k = (
|
||||
self._get_dcp_local_max_seq_len(
|
||||
self.forward_decode_metadata.max_seq_len_k
|
||||
)
|
||||
)
|
||||
|
||||
def _run_decode_kernel(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
@@ -517,24 +351,8 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
||||
llama_4_scaling: Optional[torch.Tensor] = None,
|
||||
):
|
||||
parallel = get_parallel()
|
||||
# FlashInfer autotunes MoE kernels with a synthetic full-model decode
|
||||
# and discards the attention/logits result. On multi-node GB300, the
|
||||
# synthetic full-head DCP metadata can make both the TokenSpeed and
|
||||
# TRTLLM decode kernels surface cudaErrorNvlinkUncorrectable. Skip
|
||||
# attention only inside that explicitly scoped dummy pass. Real
|
||||
# requests and CUDA graph capture continue through TokenSpeed below.
|
||||
if parallel.dcp_enabled and get_in_autotune_dummy_run():
|
||||
output = torch.zeros(
|
||||
(q.shape[0], layer.tp_q_head_num * layer.v_head_dim),
|
||||
dtype=self.q_data_type,
|
||||
device=q.device,
|
||||
)
|
||||
lse = torch.zeros(
|
||||
(q.shape[0], layer.tp_q_head_num),
|
||||
dtype=torch.float32,
|
||||
device=q.device,
|
||||
)
|
||||
return output, lse
|
||||
return self._dummy_dcp_decode_for_autotune(q, layer)
|
||||
|
||||
if not parallel.dcp_enabled:
|
||||
return super().forward_decode(
|
||||
|
||||
@@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Optional, Union
|
||||
import torch
|
||||
import triton
|
||||
|
||||
from sglang.kernels.ops.attention.dcp_kernels import create_mla_kv_page_table_for_dcp
|
||||
from sglang.kernels.ops.attention.fixup_zero_kv import fixup_zero_kv_rows
|
||||
from sglang.kernels.ops.attention.pad import (
|
||||
pad_draft_extend_query as pad_draft_extend_query_triton,
|
||||
@@ -50,6 +51,8 @@ from sglang.srt.layers.attention.flashinfer_mla_backend import (
|
||||
FlashInferMLAMultiStepDraftBackend,
|
||||
)
|
||||
from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask
|
||||
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
||||
from sglang.srt.layers.logits_processor import get_in_autotune_dummy_run
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
|
||||
is_in_breakable_cuda_graph,
|
||||
@@ -218,6 +221,10 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
self.num_q_heads = config.num_attention_heads // get_parallel().attn_tp_size
|
||||
self.num_kv_heads = config.get_num_kv_heads(get_parallel().attn_tp_size)
|
||||
self.num_local_heads = config.num_attention_heads // get_parallel().attn_tp_size
|
||||
# A DCP decode attends with the query all-gathered across the DCP
|
||||
# group, so the kernel sees attn_dcp_size x this rank's heads. Anything
|
||||
# sized per decode head must use this, not num_q_heads.
|
||||
self.num_decode_q_heads = self.num_q_heads * get_parallel().attn_dcp_size
|
||||
|
||||
# MLA-specific dimensions
|
||||
self.kv_lora_rank = config.kv_lora_rank
|
||||
@@ -259,7 +266,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
self._multi_ctas_kv_counter_buffer = (
|
||||
make_persistent_multi_ctas_kv_counter_buffer(
|
||||
torch.device(self.device),
|
||||
self.num_q_heads,
|
||||
self.num_decode_q_heads,
|
||||
max_batch_size=model_runner.max_running_requests,
|
||||
)
|
||||
)
|
||||
@@ -292,9 +299,15 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
# instead of set_mla_kv_buffer + concat_mla_absorb_q). Disabled under
|
||||
# async asserts: the fused path writes the pool directly and would
|
||||
# skip the pool's OOB probe.
|
||||
# Also disabled under DCP: unlike its fp8 sibling, set_mla_kv_concat_q
|
||||
# takes no dcp_world_size/dcp_rank, so it writes at the raw virtual
|
||||
# out_cache_loc from every rank, while the reader expects the compacted
|
||||
# row loc // dcp_size written only by the owner. Fall back to the pool's
|
||||
# DCP-aware set_mla_kv_buffer.
|
||||
self._fused_set_kv_concat_q = (
|
||||
self.data_type == torch.bfloat16
|
||||
and not envs.SGLANG_ENABLE_ASYNC_ASSERT.get()
|
||||
and not get_parallel().dcp_enabled
|
||||
and can_use_set_mla_kv_concat_q(
|
||||
self.kv_lora_rank * 2, self.qk_rope_head_dim * 2
|
||||
)
|
||||
@@ -333,6 +346,57 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
blocks = triton.cdiv(blocks, constraint_lcm) * constraint_lcm
|
||||
return blocks
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# DCP metadata (rank-local KV lengths + page table). Kernel-agnostic, so
|
||||
# the whole trtllm_mla family shares it. A no-op when DCP is off.
|
||||
# ------------------------------------------------------------------
|
||||
def _get_dcp_local_seq_lens(self, seq_lens: torch.Tensor) -> torch.Tensor:
|
||||
parallel = get_parallel()
|
||||
if not parallel.dcp_enabled:
|
||||
return seq_lens
|
||||
return get_dcp_lens(seq_lens, parallel.dcp_size, parallel.dcp_rank).to(
|
||||
torch.int32
|
||||
)
|
||||
|
||||
def _get_dcp_local_max_seq_len(self, max_seq_len: int) -> int:
|
||||
parallel = get_parallel()
|
||||
if not parallel.dcp_enabled:
|
||||
return max_seq_len
|
||||
local_max = max_seq_len // parallel.dcp_size + int(
|
||||
parallel.dcp_rank < max_seq_len % parallel.dcp_size
|
||||
)
|
||||
# A positive scheduling bound is required even when every sequence in a
|
||||
# padded graph row is empty on this rank.
|
||||
return max(local_max, 1)
|
||||
|
||||
def _fill_dcp_block_kv_indices(
|
||||
self,
|
||||
block_kv_indices: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
local_seq_lens: torch.Tensor,
|
||||
) -> None:
|
||||
parallel = get_parallel()
|
||||
pages_per_block = get_num_page_per_block_flashmla(self.page_size)
|
||||
create_mla_kv_page_table_for_dcp[
|
||||
(
|
||||
block_kv_indices.shape[0],
|
||||
get_num_kv_index_blocks_flashmla(
|
||||
block_kv_indices.shape[1], self.page_size
|
||||
),
|
||||
)
|
||||
](
|
||||
self.req_to_token,
|
||||
req_pool_indices,
|
||||
local_seq_lens,
|
||||
block_kv_indices,
|
||||
self.req_to_token.stride(0),
|
||||
block_kv_indices.stride(0),
|
||||
PHYSICAL_PAGE_SIZE=self.page_size,
|
||||
DCP_SIZE=parallel.dcp_size,
|
||||
DCP_RANK=parallel.dcp_rank,
|
||||
PAGES_PER_BLOCK=pages_per_block,
|
||||
)
|
||||
|
||||
def _create_block_kv_indices(
|
||||
self,
|
||||
batch_size: int,
|
||||
@@ -358,6 +422,14 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
(batch_size, max_blocks), -1, dtype=torch.int32, device=device
|
||||
)
|
||||
|
||||
if get_parallel().dcp_enabled:
|
||||
self._fill_dcp_block_kv_indices(
|
||||
block_kv_indices,
|
||||
req_pool_indices,
|
||||
self._get_dcp_local_seq_lens(seq_lens),
|
||||
)
|
||||
return block_kv_indices
|
||||
|
||||
if self.kv_index_translator.is_translating:
|
||||
self.kv_index_translator.fill_read_table(
|
||||
out=block_kv_indices,
|
||||
@@ -502,6 +574,18 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
metadata.block_kv_indices = block_kv_indices
|
||||
metadata.max_seq_len_k = self.max_context_len
|
||||
|
||||
if get_parallel().dcp_enabled:
|
||||
if metadata.global_seq_lens_k is None:
|
||||
# A DCP decode consumes both the rank-local and the global
|
||||
# lens, and the branches above allocate this only for verify.
|
||||
metadata.global_seq_lens_k = torch.zeros(
|
||||
(bs,), dtype=torch.int32, device=device
|
||||
)
|
||||
metadata.max_seq_len_k = self._get_dcp_local_max_seq_len(
|
||||
self.max_context_len
|
||||
+ (self.num_draft_tokens if forward_mode.is_target_verify() else 0)
|
||||
)
|
||||
|
||||
self.decode_cuda_graph_metadata[bs] = metadata
|
||||
self.forward_decode_metadata = metadata
|
||||
|
||||
@@ -519,6 +603,11 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
"""
|
||||
metadata = self.decode_cuda_graph_metadata[bs]
|
||||
|
||||
if get_parallel().dcp_enabled:
|
||||
return self._apply_dcp_cuda_graph_metadata(
|
||||
bs, req_pool_indices, seq_lens, forward_mode, metadata
|
||||
)
|
||||
|
||||
if forward_mode.is_target_verify():
|
||||
# Intentional int64 -> int32 same-kind out= downcast.
|
||||
torch.add(
|
||||
@@ -565,6 +654,50 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
PAGED_SIZE=self.page_size,
|
||||
)
|
||||
|
||||
def _apply_dcp_cuda_graph_metadata(
|
||||
self,
|
||||
bs: int,
|
||||
req_pool_indices: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
forward_mode: ForwardMode,
|
||||
metadata: TRTLLMMLADecodeMetadata,
|
||||
):
|
||||
"""DCP variant of the capture+replay body.
|
||||
|
||||
Refreshes the global and rank-local lengths into the capture-stable
|
||||
buffers once per step, and rebuilds the page table over this rank's
|
||||
cyclic slice.
|
||||
"""
|
||||
if forward_mode.is_target_verify():
|
||||
torch.add(
|
||||
seq_lens[:bs],
|
||||
self.num_draft_tokens,
|
||||
out=metadata.global_seq_lens_k,
|
||||
)
|
||||
metadata.seq_lens_k.copy_(
|
||||
self._get_dcp_local_seq_lens(metadata.global_seq_lens_k)
|
||||
)
|
||||
local_seq_lens = metadata.seq_lens_k
|
||||
elif forward_mode.is_draft_extend_v2():
|
||||
num_tokens_per_req = self.num_draft_tokens
|
||||
metadata.max_seq_len_q = num_tokens_per_req
|
||||
metadata.sum_seq_lens_q = num_tokens_per_req * bs
|
||||
seq_lens = seq_lens[:bs]
|
||||
metadata.global_seq_lens_k.copy_(seq_lens)
|
||||
metadata.seq_lens_k.copy_(self._get_dcp_local_seq_lens(seq_lens))
|
||||
local_seq_lens = metadata.seq_lens_k
|
||||
else:
|
||||
seq_lens = seq_lens[:bs]
|
||||
metadata.global_seq_lens_k.copy_(seq_lens)
|
||||
metadata.seq_lens_k.copy_(self._get_dcp_local_seq_lens(seq_lens))
|
||||
local_seq_lens = metadata.seq_lens_k
|
||||
|
||||
self._fill_dcp_block_kv_indices(
|
||||
metadata.block_kv_indices,
|
||||
req_pool_indices[:bs],
|
||||
local_seq_lens,
|
||||
)
|
||||
|
||||
def get_cuda_graph_seq_len_fill_value(self) -> int:
|
||||
"""Get the fill value for sequence lengths in CUDA graph."""
|
||||
return 1
|
||||
@@ -766,6 +899,24 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
self.forward_decode_metadata.max_seq_len_k = int(max_seq)
|
||||
self.forward_decode_metadata.batch_size = bs
|
||||
|
||||
if get_parallel().dcp_enabled:
|
||||
metadata = self.forward_decode_metadata
|
||||
if (
|
||||
forward_batch.forward_mode.is_target_verify()
|
||||
or forward_batch.forward_mode.is_decode_or_idle()
|
||||
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||
) and metadata.seq_lens_k is not None:
|
||||
# The branches above stored the global lengths in
|
||||
# seq_lens_k; keep them as global_seq_lens_k and derive the
|
||||
# rank-local view once per step rather than per MLA layer.
|
||||
metadata.global_seq_lens_k = metadata.seq_lens_k
|
||||
metadata.seq_lens_k = self._get_dcp_local_seq_lens(
|
||||
metadata.global_seq_lens_k
|
||||
)
|
||||
metadata.max_seq_len_k = self._get_dcp_local_max_seq_len(
|
||||
metadata.max_seq_len_k
|
||||
)
|
||||
|
||||
forward_batch.decode_trtllm_mla_metadata = self.forward_decode_metadata
|
||||
else:
|
||||
return super().init_forward_metadata(forward_batch)
|
||||
@@ -863,15 +1014,29 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
"""Hook for subclasses to swap the decode/spec-verify kernel.
|
||||
|
||||
The DCP arguments belong to the hook contract because forward_extend
|
||||
passes them on the DCP target-verify path. This implementation does not
|
||||
forward them to the kernel and returns no LSE, so only the DCP-capable
|
||||
subclasses serve them."""
|
||||
if cp_world > 1 or return_lse:
|
||||
passes them on the DCP target-verify path.
|
||||
|
||||
The trtllm-gen kernel has no in-kernel DCP support (flashinfer rejects
|
||||
``enable_dcp=True`` for every backend except ``cute-dsl``), but a
|
||||
``q_len == 1`` decode does not need it: the per-query global causal
|
||||
bound only varies across query rows when ``q_len > 1``, so for a single
|
||||
query token the rank-local page table and ``seq_lens`` already describe
|
||||
the shard completely.
|
||||
|
||||
Every other DCP path is refused. ``q_len > 1`` catches multi-token
|
||||
verify / draft-extend; ``causal_seqs`` and ``return_lse`` catch the
|
||||
single-token verify and draft-extend that ``q_len`` alone lets through
|
||||
(only plain decode requests the LSE the cross-rank merge needs).
|
||||
"""
|
||||
q_len = query.shape[1] if query.dim() == 4 else 1
|
||||
if get_parallel().dcp_enabled and (
|
||||
q_len > 1 or causal_seqs is not None or not return_lse
|
||||
):
|
||||
raise NotImplementedError(
|
||||
"trtllm_mla does not forward the cyclic DCP metadata to its "
|
||||
"decode kernel and returns no rank-local LSE for the cross-rank "
|
||||
"merge; select cutedsl_mla or tokenspeed_mla for a DCP "
|
||||
"target-verify run"
|
||||
"trtllm_mla cannot forward a global causal bound to its decode "
|
||||
"kernel, which is required for DCP with q_len > 1 (speculative "
|
||||
"target-verify / draft-extend); select cutedsl_mla or "
|
||||
"tokenspeed_mla for a DCP speculative run"
|
||||
)
|
||||
|
||||
# Scale computation for TRTLLM MLA kernel BMM1 operation:
|
||||
@@ -903,6 +1068,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
seq_lens=seq_lens_i32,
|
||||
max_seq_len=max_seq_len,
|
||||
bmm1_scale=bmm1_scale,
|
||||
return_lse=return_lse,
|
||||
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(),
|
||||
**extra_kwargs,
|
||||
)
|
||||
@@ -1045,6 +1211,28 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
dcp_rank=parallel.attn_dcp_rank,
|
||||
)
|
||||
|
||||
def _dummy_dcp_decode_for_autotune(
|
||||
self, q: torch.Tensor, layer: RadixAttention
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Skip decode during FlashInfer MoE autotune dummy forwards.
|
||||
|
||||
That pass discards attention/logits. Under DCP the synthetic
|
||||
full-head metadata can overflow the trtllm-gen workspace (and on
|
||||
multi-node GB300 has also produced NVLink errors). Real requests
|
||||
and CUDA-graph capture must not take this path.
|
||||
"""
|
||||
output = torch.zeros(
|
||||
(q.shape[0], layer.tp_q_head_num * layer.v_head_dim),
|
||||
dtype=self.q_data_type,
|
||||
device=q.device,
|
||||
)
|
||||
lse = torch.zeros(
|
||||
(q.shape[0], layer.tp_q_head_num),
|
||||
dtype=torch.float32,
|
||||
device=q.device,
|
||||
)
|
||||
return output, lse
|
||||
|
||||
def forward_decode(
|
||||
self,
|
||||
q: torch.Tensor, # q_nope
|
||||
@@ -1060,6 +1248,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
llama_4_scaling: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Run forward for decode using TRTLLM MLA kernel."""
|
||||
if get_parallel().dcp_enabled and get_in_autotune_dummy_run():
|
||||
return self._dummy_dcp_decode_for_autotune(q, layer)
|
||||
|
||||
merge_query = q_rope is not None
|
||||
fused_fp8_query = None
|
||||
if self.data_type == torch.float8_e4m3fn:
|
||||
@@ -1181,6 +1372,11 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
self.init_forward_metadata(forward_batch)
|
||||
metadata = forward_batch.decode_trtllm_mla_metadata
|
||||
|
||||
if get_parallel().dcp_enabled:
|
||||
return self._forward_decode_dcp(
|
||||
query, kv_cache, metadata, layer, forward_batch
|
||||
)
|
||||
|
||||
raw_out = self._run_decode_kernel(
|
||||
query=query,
|
||||
kv_cache=kv_cache,
|
||||
@@ -1198,6 +1394,48 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
output = raw_out.view(-1, layer.tp_q_head_num * layer.v_head_dim)
|
||||
return output
|
||||
|
||||
def _forward_decode_dcp(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
metadata: TRTLLMMLADecodeMetadata,
|
||||
layer: RadixAttention,
|
||||
forward_batch: ForwardBatch,
|
||||
):
|
||||
"""Rank-local MLA decode under DCP, returning ``(out, lse)``.
|
||||
|
||||
The cross-rank merge lives in the model
|
||||
(``deepseek_common/attention_forward_methods/forward_mla.py``), so this
|
||||
returns the rank-local attention state rather than a final output.
|
||||
"""
|
||||
bs = forward_batch.batch_size
|
||||
if metadata.seq_lens_k is not None:
|
||||
local_seq_lens = metadata.seq_lens_k[:bs]
|
||||
else:
|
||||
local_seq_lens = self._get_dcp_local_seq_lens(forward_batch.seq_lens[:bs])
|
||||
raw_out, lse = self._run_decode_kernel(
|
||||
query=query,
|
||||
kv_cache=kv_cache,
|
||||
block_tables=metadata.block_kv_indices,
|
||||
seq_lens=local_seq_lens,
|
||||
max_seq_len=metadata.max_seq_len_k,
|
||||
layer=layer,
|
||||
return_lse=True,
|
||||
)
|
||||
|
||||
output = raw_out.view(-1, layer.tp_q_head_num, layer.v_head_dim)
|
||||
lse = lse.view(-1, layer.tp_q_head_num)
|
||||
# A rank that owns no slice of a request must contribute a neutral
|
||||
# (out=0, lse=-inf) state, or its garbage rows poison the merge.
|
||||
fixup_zero_kv_rows(
|
||||
output,
|
||||
lse,
|
||||
local_seq_lens,
|
||||
self.q_indptr_decode[: bs + 1],
|
||||
1,
|
||||
)
|
||||
return output.flatten(1), lse
|
||||
|
||||
def forward_extend(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
|
||||
@@ -1,87 +0,0 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention import tokenspeed_mla_backend as backend_module
|
||||
from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend
|
||||
from sglang.srt.layers.attention.trtllm_mla_backend import TRTLLMMLADecodeMetadata
|
||||
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cuda_ci(est_time=60, stage="base-b", runner_config="4-gpu-b200")
|
||||
|
||||
NUM_DRAFT_TOKENS = 8
|
||||
DCP_SIZE = 4
|
||||
DCP_RANK = 2
|
||||
|
||||
|
||||
def _make_backend(bs: int):
|
||||
backend = object.__new__(TokenspeedMLABackend)
|
||||
backend.num_draft_tokens = NUM_DRAFT_TOKENS
|
||||
metadata = TRTLLMMLADecodeMetadata(
|
||||
block_kv_indices=torch.full((bs, 4), -1, dtype=torch.int32, device="cuda"),
|
||||
seq_lens_k=torch.zeros(bs, dtype=torch.int32, device="cuda"),
|
||||
global_seq_lens_k=torch.zeros(bs, dtype=torch.int32, device="cuda"),
|
||||
)
|
||||
backend.decode_cuda_graph_metadata = {bs: metadata}
|
||||
return backend, metadata
|
||||
|
||||
|
||||
def _apply(backend, *, bs: int, seq_lens: torch.Tensor, forward_mode):
|
||||
parallel = SimpleNamespace(dcp_enabled=True, dcp_size=DCP_SIZE, dcp_rank=DCP_RANK)
|
||||
with (
|
||||
patch.object(backend_module, "get_parallel", return_value=parallel),
|
||||
patch.object(backend, "_fill_dcp_block_kv_indices") as fill,
|
||||
):
|
||||
backend._apply_cuda_graph_metadata(
|
||||
bs=bs,
|
||||
req_pool_indices=torch.arange(bs, dtype=torch.int32, device="cuda"),
|
||||
seq_lens=seq_lens,
|
||||
forward_mode=forward_mode,
|
||||
)
|
||||
return fill
|
||||
|
||||
|
||||
@unittest.skipUnless(torch.cuda.is_available(), "DCP metadata buffers live on CUDA")
|
||||
class TestTokenspeedMLADCPMetadata(CustomTestCase):
|
||||
def test_target_verify_splits_global_and_local_lengths(self):
|
||||
bs = 3
|
||||
backend, metadata = _make_backend(bs)
|
||||
prefix_lens = torch.tensor([10, 20, 30], dtype=torch.int32, device="cuda")
|
||||
|
||||
fill = _apply(
|
||||
backend,
|
||||
bs=bs,
|
||||
seq_lens=prefix_lens,
|
||||
forward_mode=ForwardMode.TARGET_VERIFY,
|
||||
)
|
||||
|
||||
expected_global = prefix_lens + NUM_DRAFT_TOKENS
|
||||
expected_local = get_dcp_lens(expected_global, DCP_SIZE, DCP_RANK).to(
|
||||
torch.int32
|
||||
)
|
||||
torch.testing.assert_close(metadata.global_seq_lens_k, expected_global)
|
||||
torch.testing.assert_close(metadata.seq_lens_k, expected_local)
|
||||
fill.assert_called_once()
|
||||
torch.testing.assert_close(fill.call_args.args[2], expected_local)
|
||||
|
||||
def test_decode_does_not_add_draft_tokens(self):
|
||||
bs = 3
|
||||
backend, _ = _make_backend(bs)
|
||||
seq_lens = torch.tensor([10, 20, 30], dtype=torch.int32, device="cuda")
|
||||
|
||||
fill = _apply(
|
||||
backend, bs=bs, seq_lens=seq_lens, forward_mode=ForwardMode.DECODE
|
||||
)
|
||||
|
||||
expected_local = get_dcp_lens(seq_lens, DCP_SIZE, DCP_RANK).to(torch.int32)
|
||||
fill.assert_called_once()
|
||||
torch.testing.assert_close(fill.call_args.args[2], expected_local)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,250 @@
|
||||
"""DCP cuda-graph metadata for the trtllm_mla backend family.
|
||||
|
||||
The rank-local KV-length and page-table plumbing lives on
|
||||
:class:`TRTLLMMLABackend`, so the same expectations are asserted for the base
|
||||
backend and for both subclasses that inherit it.
|
||||
"""
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention import trtllm_mla_backend as backend_module
|
||||
from sglang.srt.layers.attention.cutedsl_mla_backend import CuteDslMLABackend
|
||||
from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend
|
||||
from sglang.srt.layers.attention.trtllm_mla_backend import (
|
||||
TRTLLMMLABackend,
|
||||
TRTLLMMLADecodeMetadata,
|
||||
)
|
||||
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cuda_ci(est_time=60, stage="base-b", runner_config="4-gpu-b200")
|
||||
|
||||
NUM_DRAFT_TOKENS = 8
|
||||
DCP_SIZE = 4
|
||||
DCP_RANK = 2
|
||||
|
||||
|
||||
def _make_backend(backend_cls, bs: int):
|
||||
backend = object.__new__(backend_cls)
|
||||
backend.num_draft_tokens = NUM_DRAFT_TOKENS
|
||||
metadata = TRTLLMMLADecodeMetadata(
|
||||
block_kv_indices=torch.full((bs, 4), -1, dtype=torch.int32, device="cuda"),
|
||||
seq_lens_k=torch.zeros(bs, dtype=torch.int32, device="cuda"),
|
||||
global_seq_lens_k=torch.zeros(bs, dtype=torch.int32, device="cuda"),
|
||||
)
|
||||
backend.decode_cuda_graph_metadata = {bs: metadata}
|
||||
return backend, metadata
|
||||
|
||||
|
||||
def _apply(backend, *, bs: int, seq_lens: torch.Tensor, forward_mode):
|
||||
parallel = SimpleNamespace(dcp_enabled=True, dcp_size=DCP_SIZE, dcp_rank=DCP_RANK)
|
||||
with (
|
||||
patch.object(backend_module, "get_parallel", return_value=parallel),
|
||||
patch.object(backend, "_fill_dcp_block_kv_indices") as fill,
|
||||
):
|
||||
backend._apply_cuda_graph_metadata(
|
||||
bs=bs,
|
||||
req_pool_indices=torch.arange(bs, dtype=torch.int32, device="cuda"),
|
||||
seq_lens=seq_lens,
|
||||
forward_mode=forward_mode,
|
||||
)
|
||||
return fill
|
||||
|
||||
|
||||
@unittest.skipUnless(torch.cuda.is_available(), "DCP metadata buffers live on CUDA")
|
||||
class _DCPMetadataTests:
|
||||
backend_cls = None
|
||||
|
||||
def test_target_verify_splits_global_and_local_lengths(self):
|
||||
bs = 3
|
||||
backend, metadata = _make_backend(self.backend_cls, bs)
|
||||
prefix_lens = torch.tensor([10, 20, 30], dtype=torch.int32, device="cuda")
|
||||
|
||||
fill = _apply(
|
||||
backend,
|
||||
bs=bs,
|
||||
seq_lens=prefix_lens,
|
||||
forward_mode=ForwardMode.TARGET_VERIFY,
|
||||
)
|
||||
|
||||
expected_global = prefix_lens + NUM_DRAFT_TOKENS
|
||||
expected_local = get_dcp_lens(expected_global, DCP_SIZE, DCP_RANK).to(
|
||||
torch.int32
|
||||
)
|
||||
torch.testing.assert_close(metadata.global_seq_lens_k, expected_global)
|
||||
torch.testing.assert_close(metadata.seq_lens_k, expected_local)
|
||||
fill.assert_called_once()
|
||||
torch.testing.assert_close(fill.call_args.args[2], expected_local)
|
||||
|
||||
def test_decode_does_not_add_draft_tokens(self):
|
||||
bs = 3
|
||||
backend, metadata = _make_backend(self.backend_cls, bs)
|
||||
seq_lens = torch.tensor([10, 20, 30], dtype=torch.int32, device="cuda")
|
||||
|
||||
fill = _apply(
|
||||
backend, bs=bs, seq_lens=seq_lens, forward_mode=ForwardMode.DECODE
|
||||
)
|
||||
|
||||
expected_local = get_dcp_lens(seq_lens, DCP_SIZE, DCP_RANK).to(torch.int32)
|
||||
fill.assert_called_once()
|
||||
torch.testing.assert_close(fill.call_args.args[2], expected_local)
|
||||
# Plain decode keeps both views in the capture-stable buffers.
|
||||
torch.testing.assert_close(metadata.global_seq_lens_k, seq_lens)
|
||||
torch.testing.assert_close(metadata.seq_lens_k, expected_local)
|
||||
|
||||
|
||||
class TestTRTLLMMLADCPMetadata(_DCPMetadataTests, CustomTestCase):
|
||||
backend_cls = TRTLLMMLABackend
|
||||
|
||||
|
||||
class TestTokenspeedMLADCPMetadata(_DCPMetadataTests, CustomTestCase):
|
||||
backend_cls = TokenspeedMLABackend
|
||||
|
||||
|
||||
class TestCuteDslMLADCPMetadata(_DCPMetadataTests, CustomTestCase):
|
||||
backend_cls = CuteDslMLABackend
|
||||
|
||||
|
||||
@unittest.skipUnless(torch.cuda.is_available(), "needs the flashinfer decode hook")
|
||||
class TestTRTLLMMLARejectsDcpMultiTokenQuery(CustomTestCase):
|
||||
"""``_run_decode_kernel`` must refuse every spec path under DCP.
|
||||
|
||||
trtllm-gen takes no global causal bound, which a ``q_len > 1`` batch needs.
|
||||
Single-token spec batches still read as ``q_len == 1``, so the refusal also
|
||||
keys on ``causal_seqs`` and ``return_lse``.
|
||||
"""
|
||||
|
||||
def _make_backend(self):
|
||||
backend = object.__new__(TRTLLMMLABackend)
|
||||
backend.backend = "trtllm-gen"
|
||||
backend.qk_nope_head_dim = 128
|
||||
backend.kv_lora_rank = 512
|
||||
backend.qk_rope_head_dim = 64
|
||||
backend.workspace_buffer = None
|
||||
backend._multi_ctas_kv_counter_buffer = None
|
||||
return backend
|
||||
|
||||
def _call(
|
||||
self,
|
||||
backend,
|
||||
q_len,
|
||||
*,
|
||||
dcp_enabled,
|
||||
kernel=None,
|
||||
causal_seqs=None,
|
||||
return_lse=False,
|
||||
):
|
||||
parallel = SimpleNamespace(
|
||||
dcp_enabled=dcp_enabled,
|
||||
dcp_size=DCP_SIZE if dcp_enabled else 1,
|
||||
dcp_rank=DCP_RANK if dcp_enabled else 0,
|
||||
)
|
||||
flashinfer_stub = SimpleNamespace(
|
||||
decode=SimpleNamespace(
|
||||
trtllm_batch_decode_with_kv_cache_mla=kernel or (lambda **kw: None)
|
||||
)
|
||||
)
|
||||
with (
|
||||
patch.object(backend_module, "get_parallel", return_value=parallel),
|
||||
patch.object(backend_module, "flashinfer", flashinfer_stub),
|
||||
patch.object(backend, "_compute_decode_bmm1_scale", return_value=1.0),
|
||||
):
|
||||
return backend._run_decode_kernel(
|
||||
query=torch.zeros(
|
||||
(2, q_len, 16, 576), dtype=torch.bfloat16, device="cuda"
|
||||
),
|
||||
kv_cache=torch.zeros((4, 1, 64, 576), dtype=torch.bfloat16),
|
||||
block_tables=torch.zeros((2, 4), dtype=torch.int32, device="cuda"),
|
||||
seq_lens=torch.ones(2, dtype=torch.int32, device="cuda"),
|
||||
max_seq_len=64,
|
||||
layer=SimpleNamespace(scaling=1.0, k_scale_float=None),
|
||||
causal_seqs=causal_seqs,
|
||||
return_lse=return_lse,
|
||||
)
|
||||
|
||||
def test_multi_token_query_under_dcp_raises(self):
|
||||
with self.assertRaises(NotImplementedError):
|
||||
self._call(self._make_backend(), q_len=NUM_DRAFT_TOKENS, dcp_enabled=True)
|
||||
|
||||
def test_explicit_causal_bound_under_dcp_raises_at_q_len_one(self):
|
||||
# Nothing forces speculative_num_draft_tokens > 1, so a q_len == 1
|
||||
# target-verify is expressible and must not slip past the q_len proxy.
|
||||
with self.assertRaises(NotImplementedError):
|
||||
self._call(
|
||||
self._make_backend(),
|
||||
q_len=1,
|
||||
dcp_enabled=True,
|
||||
causal_seqs=torch.ones(2, dtype=torch.int32, device="cuda"),
|
||||
)
|
||||
|
||||
def test_single_token_query_under_dcp_is_allowed(self):
|
||||
# The premise of trtllm_mla DCP decode: q_len == 1 needs no global
|
||||
# bound, so the guard must not swallow the path it exists to protect.
|
||||
calls = []
|
||||
self._call(
|
||||
self._make_backend(),
|
||||
q_len=1,
|
||||
dcp_enabled=True,
|
||||
kernel=lambda **kw: calls.append(kw),
|
||||
return_lse=True,
|
||||
)
|
||||
self.assertEqual(len(calls), 1)
|
||||
|
||||
def test_single_token_draft_extend_under_dcp_raises(self):
|
||||
# A single-token draft-extend reads as q_len == 1 with no causal_seqs;
|
||||
# skipping the cross-rank merge (no LSE requested) is the only signal.
|
||||
with self.assertRaises(NotImplementedError):
|
||||
self._call(self._make_backend(), q_len=1, dcp_enabled=True)
|
||||
|
||||
def test_multi_token_query_without_dcp_is_allowed(self):
|
||||
calls = []
|
||||
self._call(
|
||||
self._make_backend(),
|
||||
q_len=NUM_DRAFT_TOKENS,
|
||||
dcp_enabled=False,
|
||||
kernel=lambda **kw: calls.append(kw),
|
||||
)
|
||||
self.assertEqual(len(calls), 1)
|
||||
|
||||
|
||||
class TestDcpDecodeLayout(CustomTestCase):
|
||||
"""Rank-local length math the decode page table above is built from."""
|
||||
|
||||
SIZES = [1, 2, 3, 4, 8]
|
||||
LENS = list(range(0, 41))
|
||||
|
||||
def test_ranks_partition_the_global_length(self):
|
||||
lens = torch.tensor(self.LENS, dtype=torch.int32)
|
||||
for n in self.SIZES:
|
||||
total = sum(
|
||||
get_dcp_lens(lens, n, rank).to(torch.int64) for rank in range(n)
|
||||
)
|
||||
self.assertTrue(
|
||||
torch.equal(total, lens.to(torch.int64)),
|
||||
f"per-rank lengths do not sum to the global length at n={n}",
|
||||
)
|
||||
|
||||
def test_newest_token_is_owned_by_exactly_one_rank(self):
|
||||
# A decode step appends one token; the cross-rank merge double counts
|
||||
# or drops it unless exactly one rank sees its length grow.
|
||||
for n in self.SIZES:
|
||||
for global_len in self.LENS[1:]:
|
||||
prev = torch.tensor([global_len - 1], dtype=torch.int32)
|
||||
cur = torch.tensor([global_len], dtype=torch.int32)
|
||||
grew = [
|
||||
int(get_dcp_lens(cur, n, rank).item())
|
||||
- int(get_dcp_lens(prev, n, rank).item())
|
||||
for rank in range(n)
|
||||
]
|
||||
self.assertEqual(sum(grew), 1, f"n={n}, global_len={global_len}")
|
||||
self.assertEqual(grew[(global_len - 1) % n], 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user