diff --git a/python/sglang/kernels/ops/attention/dcp_kernels.py b/python/sglang/kernels/ops/attention/dcp_kernels.py index 395ec28de..44cce302c 100644 --- a/python/sglang/kernels/ops/attention/dcp_kernels.py +++ b/python/sglang/kernels/ops/attention/dcp_kernels.py @@ -76,6 +76,40 @@ def create_triton_kv_indices_for_dcp_triton( # KV-index build (PR #14194, MLA): global prefix+extend layout for the # all-gathered dcp_kv_buffer, plus the per-rank shard/compact kernel. # --------------------------------------------------------------------------- +@triton.jit +def create_mla_kv_page_table_for_dcp( + req_to_token_ptr, + req_pool_indices_ptr, + local_seq_lens_ptr, + block_kv_indices_ptr, + req_to_token_stride: tl.constexpr, + block_table_stride: tl.constexpr, + PHYSICAL_PAGE_SIZE: tl.constexpr, + DCP_SIZE: tl.constexpr, + DCP_RANK: tl.constexpr, + PAGES_PER_BLOCK: tl.constexpr, +): + req = tl.program_id(0) + page_block = tl.program_id(1) + page_offsets = page_block * PAGES_PER_BLOCK + tl.arange(0, PAGES_PER_BLOCK) + local_len = tl.load(local_seq_lens_ptr + req) + local_pages = tl.cdiv(local_len, PHYSICAL_PAGE_SIZE) + mask = page_offsets < local_pages + global_positions = DCP_RANK + page_offsets * PHYSICAL_PAGE_SIZE * DCP_SIZE + req_pool_index = tl.load(req_pool_indices_ptr + req) + virtual_locs = tl.load( + req_to_token_ptr + req_pool_index * req_to_token_stride + global_positions, + mask=mask, + other=0, + ) + physical_pages = virtual_locs // DCP_SIZE // PHYSICAL_PAGE_SIZE + tl.store( + block_kv_indices_ptr + req * block_table_stride + page_offsets, + physical_pages, + mask=mask, + ) + + @triton.jit def create_dcp_kv_indices( kv_indptr, diff --git a/python/sglang/srt/layers/attention/attention_registry.py b/python/sglang/srt/layers/attention/attention_registry.py index 7b511bdd5..776c824f8 100644 --- a/python/sglang/srt/layers/attention/attention_registry.py +++ b/python/sglang/srt/layers/attention/attention_registry.py @@ -89,9 +89,9 @@ def create_tokenspeed_mla_backend(runner): def create_cutedsl_mla_backend(runner): if not runner.use_mla_backend: raise ValueError("cutedsl_mla backend can only be used with MLA models.") - from sglang.srt.layers.attention.trtllm_mla_backend import TRTLLMMLABackend + from sglang.srt.layers.attention.cutedsl_mla_backend import CuteDslMLABackend - return TRTLLMMLABackend(runner, backend="cute-dsl") + return CuteDslMLABackend(runner) @register_attention_backend("aiter") diff --git a/python/sglang/srt/layers/attention/cutedsl_mla_backend.py b/python/sglang/srt/layers/attention/cutedsl_mla_backend.py new file mode 100644 index 000000000..772b69465 --- /dev/null +++ b/python/sglang/srt/layers/attention/cutedsl_mla_backend.py @@ -0,0 +1,451 @@ +""" +Attention backend for the flashinfer cute-dsl MLA decode kernels with decode +context parallelism (DCP). + +Subclasses :class:`TRTLLMMLABackend` (``backend="cute-dsl"``) to reuse its MLA +data preparation, workspace, and prefill plumbing. The flashinfer cute-dsl +monolithic MLA decode kernel natively accepts cyclic DCP metadata +(``enable_dcp`` / ``cp_world`` / ``cp_rank`` / ``causal_seqlens_kv_global``) and +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). +""" + +from __future__ import annotations + +import logging +import math +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.runtime_context import get_parallel +from sglang.srt.utils import is_flashinfer_available + +if is_flashinfer_available(): + import flashinfer + +if TYPE_CHECKING: + from sglang.srt.layers.radix_attention import RadixAttention + from sglang.srt.model_executor.model_runner import ModelRunner + +logger = logging.getLogger(__name__) + +# The flashinfer cute-dsl MLA decode kernel returns a natural-log (base-e) LSE, +# whereas sglang's DCP cross-rank merge (forward_mla: dcp_a2a_lse_reduce / +# cp_lse_ag_out_rs_mla) assumes the FlashInfer-MLA/FlashMLA base-2 convention +# (is_lse_base_on_e=False). Multiplying a natural-log LSE by log2(e) rebases it +# to base-2 (the softmax output is base-invariant; only the LSE value changes). +# CONFIRMED base-e (not base-2), so this rebase is required, not optional: +# the flashinfer-dcp-backport public-API unit test asserts the public +# trtllm_batch_decode_with_kv_cache_mla LSE against a torch.logsumexp +# (natural-log) reference at atol=1e-2 and passes (a base-2 LSE would be +# off by 1/ln2 ~= 44%). GPU job 467640: +# tests/attention/test_cute_dsl_mla_dcp*.py 27/27 + 17/17 pass. +_LSE_BASE2_FROM_NATURAL_LOG = math.log2(math.e) + + +class CuteDslMLABackend(TRTLLMMLABackend): + """flashinfer cute-dsl MLA decode backend with decode context parallelism.""" + + def __init__( + self, + model_runner: ModelRunner, + skip_prefill: bool = False, + kv_indptr_buf: Optional[torch.Tensor] = None, + q_indptr_decode_buf: Optional[torch.Tensor] = None, + ): + super().__init__( + model_runner, + skip_prefill, + kv_indptr_buf, + q_indptr_decode_buf, + 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: + if forward_mode.is_target_verify(): + self.forward_decode_metadata.global_seq_lens_k = torch.zeros_like( + self.forward_decode_metadata.seq_lens_k + ) + 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 + ) + ) + + # ------------------------------------------------------------------ + # Kernel + decode forward. + # ------------------------------------------------------------------ + def _run_decode_kernel( + self, + query: torch.Tensor, + kv_cache: torch.Tensor, + block_tables: torch.Tensor, + seq_lens: torch.Tensor, + max_seq_len: int, + layer: RadixAttention, + *, + causal_seqs: Optional[torch.Tensor] = None, + cp_world: int = 1, + cp_rank: int = 0, + return_lse: bool = False, + ): + """Call the flashinfer cute-dsl MLA decode kernel. + + Without DCP (``cp_world <= 1``) this defers to the base cute-dsl path. + With DCP, ``seq_lens`` are this rank's cyclic-local KV lengths and + ``causal_seqs`` the global per-request KV lengths; the kernel returns a + rank-local ``(out, lse)`` (LSE rebased to base-2 for the sglang merge). + """ + if cp_world <= 1: + return super()._run_decode_kernel( + query, kv_cache, block_tables, seq_lens, max_seq_len, layer + ) + if causal_seqs is None: + raise ValueError( + "causal_seqs (global per-request KV lengths) is required for DCP " + "MLA decode." + ) + bmm1_scale = self._compute_decode_bmm1_scale(layer) + raw_out, lse = flashinfer.decode.trtllm_batch_decode_with_kv_cache_mla( + query=query, + kv_cache=kv_cache, + workspace_buffer=self.workspace_buffer, + qk_nope_head_dim=self.qk_nope_head_dim, + kv_lora_rank=self.kv_lora_rank, + qk_rope_head_dim=self.qk_rope_head_dim, + block_tables=block_tables, + seq_lens=( + seq_lens if seq_lens.dtype == torch.int32 else seq_lens.to(torch.int32) + ), + max_seq_len=max_seq_len, + bmm1_scale=bmm1_scale, + skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), + backend="cute-dsl", + enable_dcp=True, + cp_world=cp_world, + cp_rank=cp_rank, + causal_seqlens_kv_global=( + causal_seqs + if causal_seqs.dtype == torch.int32 + else causal_seqs.to(torch.int32) + ), + return_lse=True, # DCP requires the rank-local LSE for the merge + ) + return raw_out, lse * _LSE_BASE2_FROM_NATURAL_LOG + + def forward_decode( + self, + q: torch.Tensor, # q_nope + k: torch.Tensor, # k_nope + v: torch.Tensor, # not used in this backend + layer: RadixAttention, + forward_batch: ForwardBatch, + save_kv_cache: bool = True, + q_rope: Optional[torch.Tensor] = None, + k_rope: Optional[torch.Tensor] = None, + cos_sin_cache: Optional[torch.Tensor] = None, + is_neox: Optional[bool] = False, + llama_4_scaling: Optional[torch.Tensor] = None, + ): + parallel = get_parallel() + if not parallel.dcp_enabled: + return super().forward_decode( + q, + k, + v, + layer, + forward_batch, + save_kv_cache, + q_rope, + k_rope, + cos_sin_cache, + is_neox, + llama_4_scaling, + ) + + # Query / KV preparation mirrors the base cute-dsl decode (both FP16 and + # FP8 KV), then swaps to the DCP kernel call + rank-local return. + merge_query = q_rope is not None + if self.data_type == torch.float8_e4m3fn: + assert q_rope is not None and k_rope is not None + if cos_sin_cache is None: + q, k, k_rope = mla_quantize_without_rope_for_fp8( + q, q_rope, k.squeeze(1), k_rope.squeeze(1) + ) + else: + q, k, k_rope = mla_quantize_and_rope_for_fp8( + q, + q_rope, + k.squeeze(1), + k_rope.squeeze(1), + forward_batch.positions, + cos_sin_cache, + is_neox, + self.kv_lora_rank, + self.qk_rope_head_dim, + ) + merge_query = False + + if save_kv_cache: + assert k is not None and k_rope is not None + self.token_to_kv_pool.set_mla_kv_buffer( + layer, forward_batch.out_cache_loc, k, k_rope + ) + + if merge_query: + q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim) + q_rope_reshaped = q_rope.view( + -1, layer.tp_q_head_num, layer.head_dim - layer.v_head_dim + ) + query = concat_mla_absorb_q_general(q_nope, q_rope_reshaped) + else: + query = q.view(-1, layer.tp_q_head_num, layer.head_dim) + + if llama_4_scaling is not None: + query = (query.to(self.q_data_type) * llama_4_scaling).to(self.data_type) + if query.dim() == 3: + query = query.unsqueeze(1) + + k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id) + kv_cache = k_cache.view(-1, self.page_size, self.kv_cache_dim).unsqueeze(1) + + metadata = ( + getattr(forward_batch, "decode_trtllm_mla_metadata", None) + or self.forward_decode_metadata + ) + metadata_batch_size = getattr(metadata, "batch_size", None) + if ( + metadata_batch_size is not None + and metadata_batch_size < forward_batch.batch_size + ): + self.init_forward_metadata(forward_batch) + metadata = forward_batch.decode_trtllm_mla_metadata + + global_seq_lens = forward_batch.seq_lens[: forward_batch.batch_size] + local_seq_lens = self._get_dcp_local_seq_lens(global_seq_lens) + 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, + causal_seqs=global_seq_lens, + cp_world=parallel.dcp_size, + cp_rank=parallel.dcp_rank, + 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) + # Zero-KV rows (a request this rank owns no cyclic slice for) get a + # neutral (out=0, lse=-inf) state so the cross-rank merge ignores them. + fixup_zero_kv_rows( + output, + lse, + local_seq_lens, + self.q_indptr_decode[: forward_batch.batch_size + 1], + 1, + ) + return output.flatten(1), lse + + +class CuteDslMLAMultiStepDraftBackend(TRTLLMMLAMultiStepDraftBackend): + """Multi-step draft backend for cutedsl_mla used by EAGLE / DSPARK.""" + + def __init__( + self, model_runner: ModelRunner, topk: int, speculative_num_steps: int + ): + super().__init__(model_runner, topk, speculative_num_steps) + # Parent populates self.attn_backends with TRT-LLM instances; replace + # them with cute-dsl instances sharing the parent's index buffers. + for i in range(self.speculative_num_steps - 1): + self.attn_backends[i] = CuteDslMLABackend( + model_runner, + skip_prefill=True, + kv_indptr_buf=self.kv_indptr[i], + q_indptr_decode_buf=self.q_indptr_decode, + ) diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 87d716fa7..14b06e7ce 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -49,11 +49,11 @@ from sglang.srt.speculative.spec_utils import ( generate_draft_decode_kv_indices, ) from sglang.srt.utils import ( + get_cuda_graph_max_batch_size, get_int_env_var, is_flashinfer_available, is_sm100_supported, next_power_of_2, - require_gathered_buffer, ) if TYPE_CHECKING: @@ -63,18 +63,6 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) -def _cuda_graph_capture_max_bs(server_args, max_bs: int) -> int: - """Pad max_bs to the alignment cuda-graph capture uses (see get_batch_sizes_to_capture).""" - mul_base = 1 - if server_args.enable_two_batch_overlap: - mul_base *= 2 - if require_gathered_buffer(server_args): - mul_base *= get_parallel().attn_tp_size - if mul_base % get_parallel().attn_cp_size != 0: - mul_base *= get_parallel().attn_cp_size - return (max_bs + mul_base - 1) // mul_base * mul_base - - if envs.SGLANG_ENABLE_TORCH_COMPILE.get(): torch._logging.set_logs(dynamo=logging.ERROR) torch._dynamo.config.suppress_errors = True @@ -447,7 +435,7 @@ class FlashInferAttnBackend(AttentionBackend): ) else: self.workspace_buffer = global_workspace_buffer - max_bs = _cuda_graph_capture_max_bs( + max_bs = get_cuda_graph_max_batch_size( model_runner.server_args, model_runner.req_to_token_pool.size ) if kv_indptr_buf is None: @@ -2230,7 +2218,7 @@ class FlashInferMultiStepDraftBackend: self.generate_draft_decode_kv_indices = generate_draft_decode_kv_indices self.page_size = model_runner.page_size - max_bs = _cuda_graph_capture_max_bs( + max_bs = get_cuda_graph_max_batch_size( model_runner.server_args, model_runner.req_to_token_pool.size * self.topk ) self.kv_indptr = torch.zeros( diff --git a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py index 565654342..6bdf80a4f 100644 --- a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py +++ b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py @@ -22,10 +22,10 @@ from __future__ import annotations """Attention backend for the tokenspeed-mla CuTe DSL kernels on Blackwell. -Subclasses :class:`TRTLLMMLABackend` and overrides only ``_run_decode_kernel`` -and ``_run_prefill_kernel``. All metadata, KV-cache layout, CUDA-graph -plumbing, FP8 quantize/rope, draft-extend padding, and chunked-prefix -dispatch are inherited unchanged from the parent. +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. """ import logging @@ -34,14 +34,28 @@ 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, ) +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.runtime_context import get_parallel from sglang.srt.utils import is_flashinfer_available, is_tokenspeed_mla_available if is_flashinfer_available(): @@ -273,6 +287,162 @@ 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, @@ -281,7 +451,12 @@ class TokenspeedMLABackend(TRTLLMMLABackend): seq_lens: torch.Tensor, max_seq_len: int, layer: RadixAttention, - ) -> torch.Tensor: + *, + causal_seqs: Optional[torch.Tensor] = None, + cp_world: int = 1, + cp_rank: int = 0, + return_lse: bool = False, + ): k_scale = getattr(layer, "k_scale_float", None) if k_scale is None: k_scale = 1.0 @@ -303,8 +478,111 @@ class TokenspeedMLABackend(TRTLLMMLABackend): softmax_scale=softmax_scale, output_scale=output_scale, enable_pdl=is_arch_support_pdl(), + return_lse=return_lse, + causal_seqs=causal_seqs, + cp_world=cp_world, + cp_rank=cp_rank, ) + def forward_decode( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + layer: RadixAttention, + forward_batch: ForwardBatch, + save_kv_cache: bool = True, + q_rope: Optional[torch.Tensor] = None, + k_rope: Optional[torch.Tensor] = None, + cos_sin_cache: Optional[torch.Tensor] = None, + is_neox: Optional[bool] = False, + llama_4_scaling: Optional[torch.Tensor] = None, + ): + parallel = get_parallel() + if not parallel.dcp_enabled: + return super().forward_decode( + q, + k, + v, + layer, + forward_batch, + save_kv_cache, + q_rope, + k_rope, + cos_sin_cache, + is_neox, + llama_4_scaling, + ) + + assert q_rope is not None and k_rope is not None + if cos_sin_cache is None: + q, k, k_rope = mla_quantize_without_rope_for_fp8( + q, q_rope, k.squeeze(1), k_rope.squeeze(1) + ) + else: + q, k, k_rope = mla_quantize_and_rope_for_fp8( + q, + q_rope, + k.squeeze(1), + k_rope.squeeze(1), + forward_batch.positions, + cos_sin_cache, + is_neox, + self.kv_lora_rank, + self.qk_rope_head_dim, + ) + + if save_kv_cache: + self.token_to_kv_pool.set_mla_kv_buffer( + layer, forward_batch.out_cache_loc, k, k_rope + ) + + query = q.view(-1, layer.tp_q_head_num, layer.head_dim) + if llama_4_scaling is not None: + query = (query.to(self.q_data_type) * llama_4_scaling).to(self.data_type) + if query.dim() == 3: + query = query.unsqueeze(1) + + k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id) + kv_cache = k_cache.view(-1, self.page_size, self.kv_cache_dim).unsqueeze(1) + metadata = ( + getattr(forward_batch, "decode_trtllm_mla_metadata", None) + or self.forward_decode_metadata + ) + metadata_batch_size = getattr(metadata, "batch_size", None) + if ( + metadata_batch_size is not None + and metadata_batch_size < forward_batch.batch_size + ): + self.init_forward_metadata(forward_batch) + metadata = forward_batch.decode_trtllm_mla_metadata + + global_seq_lens = forward_batch.seq_lens[: forward_batch.batch_size] + local_seq_lens = self._get_dcp_local_seq_lens(global_seq_lens) + 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, + causal_seqs=global_seq_lens, + cp_world=parallel.dcp_size, + cp_rank=parallel.dcp_rank, + 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) + fixup_zero_kv_rows( + output, + lse, + local_seq_lens, + self.q_indptr_decode[: forward_batch.batch_size + 1], + 1, + ) + return output.flatten(1), lse + def _run_prefill_kernel( self, q: torch.Tensor, @@ -323,6 +601,14 @@ class TokenspeedMLABackend(TRTLLMMLABackend): o_sf_scale: float = 1.0, ): # Q/K/V arrive already in FP8 via the model-side fused path # (prepare_prefill_qkv / pack_prefix_chunk_kv); no quantize here. + # Hybrid MLA models resolve the model-side hook through the outer + # HybridLinearAttnBackend, so their fallback MHA path can pass V as a + # last-dimension slice of kv_b_proj (stride(-2) > size(-1)). The + # TokenSpeed prefill kernel requires dense Q/K/V layouts even though + # the public wrapper accepts arbitrary torch tensors. + q = q.contiguous() + k = k.contiguous() + v = v.contiguous() return tokenspeed_mla.tokenspeed_mla_prefill( query=q, key=k, diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index 37016365a..a74889c55 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -150,6 +150,7 @@ class TRTLLMMLADecodeMetadata: cu_seqlens_q: Optional[torch.Tensor] = None seq_lens_q: Optional[torch.Tensor] = None seq_lens_k: Optional[torch.Tensor] = None + global_seq_lens_k: Optional[torch.Tensor] = None class TRTLLMMLABackend(FlashInferMLAAttnBackend): @@ -382,6 +383,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): if forward_mode.is_target_verify(): metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device) + metadata.global_seq_lens_k = torch.zeros( + (bs,), dtype=torch.int32, device=device + ) elif forward_mode.is_draft_extend_v2(): num_tokens_per_req = self.num_draft_tokens metadata.max_seq_len_q = num_tokens_per_req @@ -423,7 +427,12 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): if forward_mode.is_target_verify(): # Intentional int64 -> int32 same-kind out= downcast. - torch.add(seq_lens[:bs], self.num_draft_tokens, out=metadata.seq_lens_k) + torch.add( + seq_lens[:bs], + self.num_draft_tokens, + out=metadata.global_seq_lens_k, + ) + metadata.seq_lens_k.copy_(metadata.global_seq_lens_k) seq_lens = metadata.seq_lens_k elif forward_mode.is_draft_extend_v2(): num_tokens_per_req = self.num_draft_tokens @@ -569,6 +578,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): max_seq = max_seq + self.num_draft_tokens seq_lens = seq_lens + self.num_draft_tokens self.forward_decode_metadata.seq_lens_k = seq_lens.to(torch.int32) + self.forward_decode_metadata.global_seq_lens_k = ( + self.forward_decode_metadata.seq_lens_k + ) elif forward_batch.forward_mode.is_draft_extend_v2(): sum_seq_lens_q = sum(forward_batch.extend_seq_lens_cpu) max_seq_len_q = max(forward_batch.extend_seq_lens_cpu) @@ -951,8 +963,10 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): q = q.to(self.data_type) if forward_batch.forward_mode.is_target_verify(): - max_seq_len = ( - metadata.max_seq_len_k + forward_batch.spec_info.draft_token_num + draft_token_num = forward_batch.spec_info.draft_token_num + dcp_enabled = get_parallel().dcp_enabled + max_seq_len = metadata.max_seq_len_k + ( + 0 if dcp_enabled else draft_token_num ) # For target_verify, all sequences have the same number of draft tokens q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim) @@ -1006,6 +1020,44 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): assert kv_cache.dtype == self.data_type + if ( + forward_batch.forward_mode.is_target_verify() + and get_parallel().dcp_enabled + ): + raw_out, lse = self._run_decode_kernel( + query=q, + kv_cache=kv_cache, + block_tables=metadata.block_kv_indices, + seq_lens=metadata.seq_lens_k, + max_seq_len=max_seq_len, + layer=layer, + causal_seqs=metadata.global_seq_lens_k, + cp_world=get_parallel().dcp_size, + cp_rank=get_parallel().dcp_rank, + return_lse=True, + ) + output = raw_out.view( + bs * draft_token_num, + layer.tp_q_head_num, + layer.v_head_dim, + ) + lse = lse.view(bs * draft_token_num, layer.tp_q_head_num) + dense_q_indptr = torch.arange( + 0, + (bs + 1) * draft_token_num, + draft_token_num, + dtype=torch.int32, + device=q.device, + ) + fixup_zero_kv_rows( + output, + lse, + metadata.seq_lens_k, + dense_q_indptr, + draft_token_num, + ) + return output.flatten(1), lse + raw_out = self._run_decode_kernel( query=q, kv_cache=kv_cache, diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 44e7c5c3b..bfffed96c 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -3634,6 +3634,12 @@ class HybridLinearKVPool(KVCache): def get_kv_size_bytes(self): return self.full_kv_pool.get_kv_size_bytes() + def get_kv_buffer_shape(self) -> Tuple[torch.Size, torch.Size]: + # Hybrid layer ids are global model-layer ids, while the backing pool + # is dense over only full-attention layers. Shape discovery does not + # need a global layer lookup, so delegate it to that backing pool. + return self.full_kv_pool.get_kv_buffer_shape() + def get_contiguous_buf_infos(self): return self.full_kv_pool.get_contiguous_buf_infos() diff --git a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py index 55f6fac78..1ed5b92a1 100644 --- a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py @@ -23,8 +23,11 @@ from contextlib import contextmanager from typing import TYPE_CHECKING, Any, List, Sequence, Tuple from sglang.srt.model_executor.runner.base_runner import BaseRunner -from sglang.srt.runtime_context import get_flags, get_parallel -from sglang.srt.utils import require_gathered_buffer +from sglang.srt.runtime_context import get_flags +from sglang.srt.utils import ( + get_cuda_graph_batch_size_alignment, + get_cuda_graph_max_batch_size, +) if TYPE_CHECKING: from sglang.srt.model_executor.input_buffers import ForwardInputBuffers @@ -68,22 +71,15 @@ def get_batch_sizes_to_capture( capture_bs = list(server_args.cuda_graph_config.decode.bs) num_max_requests = model_runner.req_to_token_pool.size - mul_base = 1 + mul_base = get_cuda_graph_batch_size_alignment(server_args) # TBO splits each request's rows across two micro-batches, so the # alignment constraint applies per request rather than per token row. alignment_width = captured_req_width if server_args.enable_two_batch_overlap: - mul_base *= 2 alignment_width = 1 - if require_gathered_buffer(server_args): - mul_base *= get_parallel().attn_tp_size - - if mul_base % get_parallel().attn_cp_size != 0: - mul_base *= get_parallel().attn_cp_size - # pad `num_max_requests` to avoid being filtered out - num_max_requests = (num_max_requests + mul_base - 1) // mul_base * mul_base + num_max_requests = get_cuda_graph_max_batch_size(server_args, num_max_requests) if max(capture_bs) > num_max_requests: # In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests # is very small. We add more values here to make sure we capture the maximum bs. diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index 21173ba00..55d40dc87 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -50,7 +50,11 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo set_tc_piecewise_forward_context, ) from sglang.srt.utils import is_hip -from sglang.srt.utils.common import ceil_align, require_mlp_sync +from sglang.srt.utils.common import ( + ceil_align, + get_eager_max_batch_size, + require_mlp_sync, +) logger = logging.getLogger(__name__) @@ -101,14 +105,12 @@ class EagerRunner(BaseRunner): # (expand_for_topk_draft) before the eager fallback. max_bs *= sa.speculative_eagle_topk # Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies. - if require_mlp_sync(sa): - from sglang.srt.layers.cp.padding import get_cp_padding_align_size - - max_bs = ceil_align(max_bs, self.attn_tp_size) - max_bs = ceil_align(max_bs, get_cp_padding_align_size()) + max_bs = get_eager_max_batch_size(sa, max_bs) prefill_ceiling = max(mr.max_total_num_tokens, sa.max_prefill_buffer_tokens()) max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req) if require_mlp_sync(sa): + from sglang.srt.layers.cp.padding import get_cp_padding_align_size + max_num_token = ceil_align(max_num_token, self.attn_tp_size) max_num_token = ceil_align(max_num_token, get_cp_padding_align_size()) self._eager_max_bs = max_bs @@ -261,7 +263,9 @@ class EagerRunner(BaseRunner): forward_batch = self.load_batch(forward_batch, pp_proxy_tensors) if forward_batch.needs_forward_metadata_init(): - if hasattr(model_runner.model, "prepare_context_parallel_metadata_for_dcp"): + if model_runner.dcp_size > 1 and hasattr( + model_runner.model, "prepare_context_parallel_metadata_for_dcp" + ): # prepare kv cache buffer for dcp to gather kv cache forward_batch.attn_dcp_metadata = ( model_runner.model.prepare_context_parallel_metadata_for_dcp( diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index 823fbe2eb..63c470b5c 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -88,6 +88,25 @@ class MlaBmmFusionPlan: attn_output_buf: torch.Tensor +def _is_dcp_mla_decode_phase(forward_batch: ForwardBatch) -> bool: + if not get_parallel().dcp_enabled: + return False + if forward_batch.forward_mode.is_decode(): + return True + if not forward_batch.forward_mode.is_target_verify() or not _is_cuda: + return False + + server_args = get_server_args() + decode_backend = ( + server_args.decode_attention_backend or server_args.attention_backend + ) + return ( + server_args.speculative_algorithm == "DSPARK" + and server_args.speculative_attention_mode == "decode" + and decode_backend in ("tokenspeed_mla", "cutedsl_mla") + ) + + if _is_cuda: from sglang.kernels.ops.gemm import bmm_fp8 @@ -254,8 +273,7 @@ class DeepseekMLAForwardMixin: # weights and skip the per-layer Q all-gather (bf16 decode absorb only). q_replicate_active = ( get_server_args().dcp_replicate_q_proj - and get_parallel().dcp_enabled - and forward_batch.forward_mode.is_decode() + and _is_dcp_mla_decode_phase(forward_batch) and not self.use_deep_gemm_bmm and self.w_kc_qrep is not None and self.q_b_proj_qrep_weight is not None @@ -595,12 +613,12 @@ class DeepseekMLAForwardMixin: # all_gather q_pe, q_nope_out,take tp8 as an example, q_pe [B, H, ROPE_DIM], q_nope_out [B, H, NOPE_DIM] gathered to [B, H * dcp_world_size, ROPE_DIM] [B, H * dcp_world_size, NOPE_DIM] for decode batch, and all gather k_pe, k_nope for extend batch. if get_parallel().dcp_enabled: - if forward_batch.forward_mode.is_decode() and not q_replicate_active: - # if forward_batch.forward_mode is decode, gather q - q_nope_out, q_pe = all_gather_q_for_mla_decode( - q_nope_out=q_nope_out, - q_pe=q_pe, - ) + if _is_dcp_mla_decode_phase(forward_batch): + if not q_replicate_active: + q_nope_out, q_pe = all_gather_q_for_mla_decode( + q_nope_out=q_nope_out, + q_pe=q_pe, + ) elif forward_batch.forward_mode.is_extend(): # for extend, gather kv all_gather_kv_cache_for_mla_extend( @@ -748,10 +766,7 @@ class DeepseekMLAForwardMixin: topk_indices=topk_indices, ) attn_output = fusion_plan.attn_output_buf - elif ( - forward_batch.forward_mode.is_decode() - and get_parallel().dcp_enabled - ): + elif _is_dcp_mla_decode_phase(forward_batch): # set return_lse=True to correct attn_output attn_output, lse = self.attn_mqa_for_dcp_decode( q_nope_out, @@ -825,7 +840,7 @@ class DeepseekMLAForwardMixin: ) # correct attn_output with respect to lse from other ranks - if forward_batch.forward_mode.is_decode() and get_parallel().dcp_enabled: + if _is_dcp_mla_decode_phase(forward_batch): attn_output = attn_output.view( -1, self.num_local_heads * get_parallel().attn_dcp_size, diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index 0670f61e3..c2dd7895c 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -16,6 +16,7 @@ from sglang.srt.distributed import ( tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder +from sglang.srt.layers.dcp.planner import prepare_decode_context_parallel_metadata from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( ColumnParallelBatchedLinear, @@ -52,6 +53,20 @@ from sglang.srt.utils import make_layers from sglang.srt.utils.common import BumpAllocator, add_prefix, set_weight_attrs +def _get_kda_local_num_heads(num_heads: int, tp_size: int) -> int: + if num_heads % tp_size != 0: + raise ValueError( + f"KDA num_heads ({num_heads}) must be divisible by global tp_size ({tp_size})" + ) + return num_heads // tp_size + + +def _materialize_residual_stream( + hidden_states: torch.Tensor, residual: Optional[torch.Tensor] +) -> torch.Tensor: + return hidden_states if residual is None else hidden_states + residual + + class KimiMoE(nn.Module): def __init__( self, @@ -188,8 +203,7 @@ class KimiDeltaAttention(nn.Module): self.head_v_dim = config.linear_attn_config["head_dim"] self.layer_idx = layer_idx self.prefix = prefix - assert self.num_heads % self.tp_size == 0 - self.local_num_heads = divide(self.num_heads, self.tp_size) + self.local_num_heads = _get_kda_local_num_heads(self.num_heads, self.tp_size) projection_size = self.head_dim * self.num_heads self.conv_size = config.linear_attn_config["short_conv_kernel_size"] @@ -317,9 +331,9 @@ class KimiDeltaAttention(nn.Module): self.attn = RadixLinearAttention( layer_id=self.layer_idx, - num_q_heads=self.num_k_heads // self.attn_tp_size, - num_k_heads=self.num_k_heads // self.attn_tp_size, - num_v_heads=self.num_v_heads // self.attn_tp_size, + num_q_heads=_get_kda_local_num_heads(self.num_k_heads, self.tp_size), + num_k_heads=_get_kda_local_num_heads(self.num_k_heads, self.tp_size), + num_v_heads=_get_kda_local_num_heads(self.num_v_heads, self.tp_size), head_q_dim=self.head_k_dim, head_k_dim=self.head_k_dim, head_v_dim=self.head_v_dim, @@ -519,6 +533,7 @@ class KimiLinearModel(nn.Module): self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size self.pp_group = get_pp_group() + self.dspark_layers_to_capture: Optional[list[int]] = None if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -581,7 +596,6 @@ class KimiLinearModel(nn.Module): dtype=torch.float32, device=device, ) - # TODO: capture aux hidden states aux_hidden_states = [] for i in range(self.start_layer, self.end_layer): ctx = get_global_expert_distribution_recorder().with_current_layer(i) @@ -594,6 +608,13 @@ class KimiLinearModel(nn.Module): residual=residual, zero_allocator=zero_allocator, ) + if ( + self.dspark_layers_to_capture is not None + and i in self.dspark_layers_to_capture + ): + aux_hidden_states.append( + _materialize_residual_stream(hidden_states, residual) + ) if not self.pp_group.is_last_rank: return PPProxyTensors( @@ -609,10 +630,9 @@ class KimiLinearModel(nn.Module): else: hidden_states, _ = self.norm(hidden_states, residual) - if len(aux_hidden_states) == 0: - return hidden_states - - return hidden_states, aux_hidden_states + if self.dspark_layers_to_capture is not None: + return hidden_states, aux_hidden_states + return hidden_states class KimiLinearForCausalLM(nn.Module): @@ -642,6 +662,22 @@ class KimiLinearForCausalLM(nn.Module): self.lm_head = PPMissingLayer() logit_scale = getattr(self.config, "logit_scale", 1.0) self.logits_processor = LogitsProcessor(config=config, logit_scale=logit_scale) + self.capture_aux_hidden_states = False + + def get_input_embeddings(self): + return self.model.embed_tokens + + def set_dspark_layers_to_capture(self, layer_ids: list[int]) -> None: + if self.pp_group.world_size > 1: + raise NotImplementedError("DSPARK aux hidden capture requires PP=1.") + if not self.pp_group.is_last_rank: + return + if layer_ids is None: + raise ValueError( + "DSPARK requires explicit layer_ids for aux hidden capture." + ) + self.capture_aux_hidden_states = True + self.model.dspark_layers_to_capture = list(layer_ids) @torch.no_grad() def forward( @@ -660,12 +696,47 @@ class KimiLinearForCausalLM(nn.Module): pp_proxy_tensors, ) if self.pp_group.is_last_rank: + aux_hidden_states = None + if self.capture_aux_hidden_states: + hidden_states, aux_hidden_states = hidden_states return self.logits_processor( - input_ids, hidden_states, self.lm_head, forward_batch + input_ids, + hidden_states, + self.lm_head, + forward_batch, + aux_hidden_states, ) else: return hidden_states + def prepare_context_parallel_metadata_for_dcp( + self, + seq_lens: torch.Tensor, + extend_prefix_lens: torch.Tensor, + extend_prefix_lens_cpu: torch.Tensor, + extend_seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + req_to_token: torch.Tensor, + seq_lens_sum: int, + kv_buffer_shape: torch.Size, + kv_cache_dtype, + kv_cache_device, + create_chunked_prefix_cache_kv_indices_fn, + ): + return prepare_decode_context_parallel_metadata( + seq_lens=seq_lens, + extend_prefix_lens=extend_prefix_lens, + extend_prefix_lens_cpu=extend_prefix_lens_cpu, + extend_seq_lens=extend_seq_lens, + req_pool_indices=req_pool_indices, + req_to_token=req_to_token, + seq_lens_sum=seq_lens_sum, + kv_buffer_shape=kv_buffer_shape, + kv_cache_dtype=kv_cache_dtype, + kv_cache_device=kv_cache_device, + create_chunked_prefix_cache_kv_indices_fn=create_chunked_prefix_cache_kv_indices_fn, + ) + def _is_non_local_pp_weight(self, name: str) -> bool: if self.pp_group.world_size == 1: return False diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 9bf05be7e..509bb82d0 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1038,14 +1038,6 @@ class ServerArgs: ), NS("parallel"), ] = 1 - dcp_size: A[ - int, - Arg( - help="The decode context parallelism size.", - aliases=["--decode-context-parallel-size"], - ), - NS("parallel"), - ] = 1 dwdp_size: A[ int, Arg( @@ -1064,20 +1056,24 @@ class ServerArgs: "combine), or 'fi_a2a' (FlashInfer MNNVL All-to-All kernel; requires " "SM90+ and MNNVL fabric memory, e.g. GB200 NVL72).", choices=["ag_rs", "a2a", "fi_a2a"], + resolvable=True, ), NS("parallel"), ] = "ag_rs" dcp_replicate_q_proj: A[ - bool, + Optional[bool], Arg( help="For MLA decode context parallelism with the a2a/fi_a2a " "backend: replicate the Q projection so each DCP rank computes the " "full-head query locally (redundant projection compute), eliminating " "the per-layer head-dim all-gather of Q. Trades a small amount of " - "extra GEMM for one fewer collective per layer.", + "extra GEMM for one fewer collective per layer. Use " + "--no-dcp-replicate-q-proj to disable the model-specific default.", + action=argparse.BooleanOptionalAction, + resolvable=True, ), NS("parallel"), - ] = False + ] = None enable_prefill_cp: A[ bool, "Enable context parallelism for the prefill phase. Select the layout with --cp-strategy.", @@ -3763,13 +3759,35 @@ class ServerArgs: return elif is_cuda(): if self.speculative_algorithm is not None: - raise ValueError( - "Decode context parallel (--dcp-size / " - "--decode-context-parallel-size > 1) on CUDA platform " - "does not support any speculative algorithm, but got " - f"dcp_size={self.dcp_size} on a CUDA platform with " - "speculative decoding enabled." + model_arches = self.get_model_config().hf_config.architectures + decode_backend = self.decode_attention_backend or self.attention_backend + kimi_linear_dspark = ( + self.speculative_algorithm == "DSPARK" + and "KimiLinearForCausalLM" in model_arches + and self.speculative_attention_mode == "decode" + and decode_backend in ("tokenspeed_mla", "cutedsl_mla") ) + if kimi_linear_dspark: + ragged_verify_mode = envs.SGLANG_RAGGED_VERIFY_MODE.get() + if ragged_verify_mode != "static": + raise ValueError( + "Kimi Linear DCP + DSPARK currently requires " + "SGLANG_RAGGED_VERIFY_MODE=static, but got " + f"{ragged_verify_mode!r}." + ) + else: + raise ValueError( + "Decode context parallel (--dcp-size / " + "--decode-context-parallel-size > 1) with speculative " + "decoding on CUDA is supported only for Kimi Linear + " + "DSPARK + --speculative-attention-mode decode + " + "tokenspeed_mla, or experimental cutedsl_mla, but got " + f"architectures={model_arches}, " + f"speculative_algorithm={self.speculative_algorithm!r}, " + "speculative_attention_mode=" + f"{self.speculative_attention_mode!r}, " + f"decode_attention_backend={decode_backend!r}." + ) else: raise ValueError( "Decode context parallel (--dcp-size / " diff --git a/python/sglang/srt/speculative/draft_utils.py b/python/sglang/srt/speculative/draft_utils.py index bb7106eda..43d633e2b 100644 --- a/python/sglang/srt/speculative/draft_utils.py +++ b/python/sglang/srt/speculative/draft_utils.py @@ -239,7 +239,18 @@ class DraftBackendFactory: ) def _create_cutedsl_mla_decode_backend(self): - return self._create_trtllm_mla_decode_backend(backend="cute-dsl") + if not self.draft_model_runner.use_mla_backend: + raise ValueError( + "cutedsl_mla backend requires MLA model (use_mla_backend=True)." + ) + + from sglang.srt.layers.attention.cutedsl_mla_backend import ( + CuteDslMLAMultiStepDraftBackend, + ) + + return CuteDslMLAMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ) def _create_tokenspeed_mla_decode_backend(self): if not self.draft_model_runner.use_mla_backend: diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index 5a284b4f6..63e110d25 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -3608,6 +3608,31 @@ def require_mlp_sync(server_args: ServerArgs): return server_args.enable_dp_attention or require_gathered_buffer(server_args) +def get_cuda_graph_batch_size_alignment(server_args: ServerArgs) -> int: + alignment = 1 + if server_args.enable_two_batch_overlap: + alignment *= 2 + if require_gathered_buffer(server_args): + alignment *= get_parallel().attn_tp_size + if alignment % get_parallel().attn_cp_size != 0: + alignment *= get_parallel().attn_cp_size + return alignment + + +def get_cuda_graph_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int: + return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment(server_args)) + + +def get_eager_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int: + if not require_mlp_sync(server_args): + return max_batch_size + + from sglang.srt.layers.cp.padding import get_cp_padding_align_size + + max_batch_size = ceil_align(max_batch_size, get_parallel().attn_tp_size) + return ceil_align(max_batch_size, get_cp_padding_align_size()) + + def find_local_repo_dir(repo_id: str, revision: Optional[str] = None) -> Optional[str]: import huggingface_hub as hf diff --git a/test/registered/dcp/test_dcp_layout_unit.py b/test/registered/dcp/test_dcp_layout_unit.py index ec053d569..2ee7e44e4 100644 --- a/test/registered/dcp/test_dcp_layout_unit.py +++ b/test/registered/dcp/test_dcp_layout_unit.py @@ -12,12 +12,19 @@ Usage: python test_dcp_layout_unit.py """ +import math import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, patch import torch from sglang.srt.layers.dcp.layout import get_dcp_lens +from sglang.srt.mem_cache.allocator.paged import PagedTokenToKVPoolAllocator +from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator +from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=10, suite="base-a-test-cpu") @@ -36,7 +43,7 @@ def _legacy_inplace_formula(length: int, n: int, rank: int) -> int: return (length - rank - 1) // n + 1 -class TestGetDcpLens(unittest.TestCase): +class TestGetDcpLens(CustomTestCase): def test_start_none_matches_owner_count(self): for n in DCP_SIZES: for rank in range(n): @@ -86,6 +93,156 @@ class TestGetDcpLens(unittest.TestCase): lens = torch.tensor(LENS, dtype=torch.int32) self.assertTrue(torch.equal(get_dcp_lens(lens, 1, 0), lens)) + def test_paged_allocator_exposes_dcp_virtual_capacity(self): + real_kv_size = 1024 + dcp_size = 4 + physical_page_size = 64 + allocator = PagedTokenToKVPoolAllocator( + size=real_kv_size * dcp_size, + page_size=physical_page_size * dcp_size, + dtype=torch.bfloat16, + device="cpu", + kvcache=object(), + need_sort=False, + ) + + allocations = [allocator.alloc(physical_page_size * dcp_size) for _ in range(4)] + self.assertTrue(all(indices is not None for indices in allocations)) + virtual_indices = torch.cat(allocations) + + self.assertEqual(allocator.size, real_kv_size * dcp_size) + self.assertEqual(allocator.page_size, physical_page_size * dcp_size) + self.assertEqual(allocator.num_pages, real_kv_size // physical_page_size) + self.assertEqual( + len(torch.unique(virtual_indices // dcp_size)), + len(virtual_indices) // dcp_size, + ) + self.assertLess( + int((virtual_indices // dcp_size).max()), + real_kv_size + physical_page_size, + ) + + def test_configurator_scales_only_the_virtual_dcp_allocator(self): + physical_kv_size = 1024 + physical_page_size = 64 + physical_kv_cache = SimpleNamespace( + size=physical_kv_size, + page_size=physical_page_size, + ) + sizes = SimpleNamespace( + max_total_num_tokens=physical_kv_size, + full_max_total_num_tokens=None, + swa_max_total_num_tokens=None, + ) + allocators = {} + + for dcp_size in (1, 4): + configurator = SimpleNamespace( + server_args=SimpleNamespace( + disaggregation_mode="null", + enable_hisparse=False, + page_size=physical_page_size, + dcp_size=dcp_size, + ), + hybrid_gdn_config=None, + is_hybrid_swa=False, + kv_cache_dtype=torch.bfloat16, + device="cpu", + is_draft_worker=False, + ) + with patch( + "sglang.srt.mem_cache.kv_cache_configurator.current_platform.is_out_of_tree", + return_value=False, + ): + allocators[dcp_size] = ( + KVCacheConfigurator._build_token_to_kv_pool_allocator( + configurator, + sizes=sizes, + token_to_kv_pool=physical_kv_cache, + is_dsv4_model=False, + req_to_token_pool=object(), + token_to_kv_pool_allocator=None, + ) + ) + + dcp1_allocator = allocators[1] + dcp4_allocator = allocators[4] + self.assertIs(dcp1_allocator.get_kvcache(), physical_kv_cache) + self.assertIs(dcp4_allocator.get_kvcache(), physical_kv_cache) + self.assertEqual(dcp1_allocator.size, 1024) + self.assertEqual(dcp1_allocator.page_size, 64) + self.assertEqual(dcp1_allocator.num_pages, 16) + self.assertEqual(dcp4_allocator.size, 4096) + self.assertEqual(dcp4_allocator.page_size, 256) + self.assertEqual(dcp4_allocator.num_pages, 16) + + def test_live_cell_and_page_ownership_formulas(self): + dcp_size = 4 + physical_page_size = 64 + ragged_lengths = (0, 1, 2, 3, 4, 63, 64, 65, 255, 256, 257, 515) + + per_rank_counts = [] + for rank in range(dcp_size): + expected_counts = [ + length // dcp_size + int(rank < length % dcp_size) + for length in ragged_lengths + ] + actual_counts = [ + _owner_count(length, dcp_size, rank, 0) for length in ragged_lengths + ] + self.assertEqual(actual_counts, expected_counts) + per_rank_counts.append(sum(actual_counts)) + + allocated_pages = [ + math.ceil(length / (physical_page_size * dcp_size)) + for length in ragged_lengths + ] + active_pages = [ + math.ceil(count / physical_page_size) for count in actual_counts + ] + self.assertTrue( + all( + active <= allocated + for active, allocated in zip(active_pages, allocated_pages) + ) + ) + self.assertTrue( + all( + allocated - active <= 1 + for active, allocated in zip(active_pages, allocated_pages) + ) + ) + + self.assertEqual(sum(per_rank_counts), sum(ragged_lengths)) + + aligned_lengths = (256, 512, 768, 1024) + full_replica_cells = sum(aligned_lengths) + full_replica_pages = sum( + length // physical_page_size for length in aligned_lengths + ) + for rank in range(dcp_size): + local_cells = sum( + _owner_count(length, dcp_size, rank, 0) for length in aligned_lengths + ) + local_pages = sum( + math.ceil(_owner_count(length, dcp_size, rank, 0) / physical_page_size) + for length in aligned_lengths + ) + self.assertEqual(local_cells * dcp_size, full_replica_cells) + self.assertEqual(local_pages * dcp_size, full_replica_pages) + + def test_hybrid_pool_reports_the_backing_attention_shape(self): + pool = object.__new__(HybridLinearKVPool) + pool.start_layer = 0 + pool.layer_transfer_counter = None + pool.full_attention_layer_id_mapping = {3: 0, 7: 1} + pool.full_kv_pool = MagicMock() + expected = (torch.Size([1024, 1, 576]), torch.Size([1024, 1, 576])) + pool.full_kv_pool.get_kv_buffer_shape.return_value = expected + + self.assertEqual(pool.get_kv_buffer_shape(), expected) + pool.full_kv_pool.get_kv_buffer_shape.assert_called_once_with() + if __name__ == "__main__": unittest.main() diff --git a/test/registered/dcp/test_kimi_linear_dcp4.py b/test/registered/dcp/test_kimi_linear_dcp4.py new file mode 100644 index 000000000..37ab83ef5 --- /dev/null +++ b/test/registered/dcp/test_kimi_linear_dcp4.py @@ -0,0 +1,129 @@ +"""Four-Blackwell acceptance coverage for Kimi Linear TokenSpeed MLA DCP. + +The captured-shape and eager-shape requests deliberately straddle +``--cuda-graph-max-bs-decode=64``. This guards both the regular CUDA graph +decode path and the full-capacity eager DCP LSE scratch-buffer path. +""" + +import unittest + +import requests +import torch + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.eval_accuracy_kit import GSM8KMixin +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=900, stage="base-c", runner_config="4-gpu-b200") + +KIMI_LINEAR_MODEL = "moonshotai/Kimi-Linear-48B-A3B-Instruct" + + +def _has_four_blackwell_gpus() -> bool: + if not torch.cuda.is_available() or torch.cuda.device_count() < 4: + return False + return all( + torch.cuda.get_device_capability(device_index) >= (10, 0) + for device_index in range(4) + ) + + +@unittest.skipUnless( + _has_four_blackwell_gpus(), + "TokenSpeed MLA DCP acceptance requires four Blackwell GPUs", +) +class TestKimiLinearDCP4(GSM8KMixin, CustomTestCase): + model = KIMI_LINEAR_MODEL + base_url = DEFAULT_URL_FOR_TEST + gsm8k_score_threshold = 0.90 + gsm8k_num_examples = 200 + # Keep accuracy evaluation within the captured decode batch sizes so its + # score is batch-invariant. The separate smoke test still exercises the + # eager path with batch size 65. + gsm8k_num_threads = 4 + gsm8k_num_shots = 5 + + @classmethod + def setUpClass(cls): + cls.process = None + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 5, + other_args=[ + "--tp-size", + "4", + "--dcp-size", + "4", + "--attention-backend", + "tokenspeed_mla", + "--kv-cache-dtype", + "fp8_e4m3", + "--dcp-comm-backend", + "a2a", + "--dcp-replicate-q-proj", + "--trust-remote-code", + "--random-seed", + "0", + "--dtype", + "bfloat16", + "--cuda-graph-max-bs-decode", + "64", + "--cuda-graph-backend-prefill", + "disabled", + "--mem-fraction-static", + "0.80", + ], + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid, wait_timeout=60) + + def _assert_batch_completes(self, batch_size: int): + prompts = [ + f"Reply with one short word for request {index}: the sky is" + for index in range(batch_size) + ] + response = requests.post( + self.base_url + "/generate", + json={ + "text": prompts, + "sampling_params": { + "temperature": 0, + "max_new_tokens": 8, + "ignore_eos": True, + }, + }, + timeout=180, + ) + response.raise_for_status() + outputs = response.json() + self.assertIsInstance(outputs, list) + self.assertEqual(len(outputs), batch_size) + for output in outputs: + self.assertTrue(output["text"].strip()) + self.assertGreater(output["meta_info"]["completion_tokens"], 0) + + def test_decode_cuda_graph_and_eager_batch(self): + # Batch two replays a captured shape; batch 65 is above the configured + # regular CUDA graph maximum and therefore exercises eager decode. + self._assert_batch_completes(2) + self._assert_batch_completes(2) + self._assert_batch_completes(65) + + def test_physical_capacity_sanity(self): + response = requests.get(self.base_url + "/server_info", timeout=30) + response.raise_for_status() + self.assertGreater(response.json()["max_total_num_tokens"], 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index 7fa9275af..ca3520bfe 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -71,6 +71,8 @@ class TestModelOverridableWhitelist(CustomTestCase): "ep_size", "moe_dense_tp_size", "attn_cp_size", + "dcp_comm_backend", + "dcp_replicate_q_proj", "disable_overlap_schedule", "uses_mamba_radix_cache", "mamba_radix_cache_strategy",