diff --git a/python/sglang/srt/layers/attention/cutedsl_mla_backend.py b/python/sglang/srt/layers/attention/cutedsl_mla_backend.py index bd93ed57e..4c26f5f19 100644 --- a/python/sglang/srt/layers/attention/cutedsl_mla_backend.py +++ b/python/sglang/srt/layers/attention/cutedsl_mla_backend.py @@ -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, diff --git a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py index 91ed1bd50..353cd6ffd 100644 --- a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py +++ b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py @@ -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( diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index d3de3bb0d..8e8fc49ed 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -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, diff --git a/test/registered/dcp/test_tokenspeed_mla_dcp_metadata.py b/test/registered/dcp/test_tokenspeed_mla_dcp_metadata.py deleted file mode 100644 index 40f14d879..000000000 --- a/test/registered/dcp/test_tokenspeed_mla_dcp_metadata.py +++ /dev/null @@ -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() diff --git a/test/registered/dcp/test_trtllm_mla_family_dcp_metadata.py b/test/registered/dcp/test_trtllm_mla_family_dcp_metadata.py new file mode 100644 index 000000000..e278b95e3 --- /dev/null +++ b/test/registered/dcp/test_trtllm_mla_family_dcp_metadata.py @@ -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()