[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,
|
||||
|
||||
Reference in New Issue
Block a user