[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:
Tri Vo
2026-08-31 19:19:44 -07:00
committed by GitHub
co-authored by kpham-sgl Claude Opus 5
parent 97744189b8
commit f50b4ad7ae
5 changed files with 519 additions and 487 deletions
@@ -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()