[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,