From f46d5f25b4c3b943167c6021aedb60415c4bfcd9 Mon Sep 17 00:00:00 2001 From: Zhangheng Date: Thu, 30 Jul 2026 16:45:32 +0800 Subject: [PATCH] [4/N][CP] Support interleave strategy for cp v2 (#30482) Co-authored-by: Xinyuan Tong Co-authored-by: Baizhou Zhang --- .../srt/layers/attention/dsa/dsa_indexer.py | 31 +- .../sglang/srt/layers/attention/dsa/utils.py | 33 +- .../srt/layers/attention/dsa_backend.py | 54 ++- python/sglang/srt/layers/cp/base.py | 44 ++- python/sglang/srt/layers/cp/interleave.py | 207 +++++++++-- python/sglang/srt/layers/cp/utils.py | 66 +++- python/sglang/srt/layers/cp/zigzag.py | 9 +- .../srt/model_executor/runner/eager_runner.py | 31 +- .../attention_forward_methods/forward_mla.py | 21 +- python/sglang/srt/models/deepseek_nextn.py | 54 +-- python/sglang/srt/models/deepseek_v2.py | 42 +-- python/sglang/srt/server_args.py | 23 +- test/registered/cp/test_cp_strategy_unit.py | 336 ++++++++++++++++++ test/registered/cp/test_dsa_prefill_cp.py | 89 +++++ ..._node.py => test_dsa_prefill_cp_legacy.py} | 1 + 15 files changed, 907 insertions(+), 134 deletions(-) create mode 100644 test/registered/cp/test_dsa_prefill_cp.py rename test/registered/cp/{test_deepseek_v32_cp_single_node.py => test_dsa_prefill_cp_legacy.py} (99%) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index e5268aeab..cee6b5bd0 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -111,6 +111,8 @@ from sglang.srt.distributed import ( from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers.communicator import ScatterMode +from sglang.srt.layers.cp.base import get_cp_strategy +from sglang.srt.layers.cp.utils import is_cp_v2_active from sglang.srt.layers.linear import ReplicatedLinear from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.rotary_embedding import get_rope_wrapper @@ -630,13 +632,20 @@ class Indexer(MultiPlatformOp): self.alt_stream.wait_stream(current_stream) query = self._maybe_rotate(query) + # Gather the full key on alt_stream so the CP all-gather overlaps + # with the query rotate above on the current stream. with torch.cuda.stream(self.alt_stream): - key = cp_all_gather_rerange_output( - key.contiguous(), - self.cp_size, - forward_batch, - torch.cuda.current_stream(), - ) + if is_cp_v2_active(forward_batch): + key = get_cp_strategy().materialize_full_indexer_k_cache( + key, forward_batch + ) + else: + key = cp_all_gather_rerange_output( + key.contiguous(), + self.cp_size, + forward_batch, + torch.cuda.current_stream(), + ) current_stream.wait_stream(self.alt_stream) return query, key, weights_raw else: @@ -644,7 +653,9 @@ class Indexer(MultiPlatformOp): key = self._maybe_rotate(key) # allgather+rerrange - if forward_batch.attn_cp_metadata is not None and self.dsa_enable_prefill_cp: + if is_cp_v2_active(forward_batch): + key = get_cp_strategy().materialize_full_indexer_k_cache(key, forward_batch) + elif forward_batch.attn_cp_metadata is not None and self.dsa_enable_prefill_cp: key = cp_all_gather_rerange_output( key.contiguous(), self.cp_size, @@ -1870,7 +1881,11 @@ class Indexer(MultiPlatformOp): weights = self._apply_q_scale_and_softmax_scale(weights, q_scale) else: query, key, weights_raw = self._get_q_k_bf16( - q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch + q_lora, + x, + positions, + enable_dual_stream, + forward_batch=forward_batch, ) if enable_dual_stream: diff --git a/python/sglang/srt/layers/attention/dsa/utils.py b/python/sglang/srt/layers/attention/dsa/utils.py index 73d10839d..674ad95f2 100644 --- a/python/sglang/srt/layers/attention/dsa/utils.py +++ b/python/sglang/srt/layers/attention/dsa/utils.py @@ -82,7 +82,17 @@ def should_use_dsa_fused_topk( def is_dsa_enable_prefill_cp(): - return get_server_args().enable_dsa_prefill_context_parallel + if not envs.SGLANG_ENABLE_CP_V2.get(): + return get_parallel().enable_dsa_prefill_context_parallel + + # Derive from the runtime CP topology + model arch rather than the legacy + # flag under CP-v2: DSA prefill CP is active when the CP group is on for a + # DeepSeek Sparse Attention model. + if get_parallel().attn_cp_size <= 1: + return False + from sglang.srt.configs.model_config import is_deepseek_dsa + + return is_deepseek_dsa(get_server_args().get_model_config().hf_config) def is_dsa_prefill_cp_in_seq_split(): @@ -206,6 +216,15 @@ def pad_dsa_cache_seqlens(forward_batch: "ForwardBatch", dsa_cache_seqlens): def can_dsa_cp_split(seq_len: int, cp_size: int, use_dsa: bool, forward_batch): + if ( + cp_size <= 1 + or not use_dsa + or not forward_batch.forward_mode.is_context_parallel_extend() + or not is_dsa_enable_prefill_cp() + or sum(forward_batch.extend_seq_lens_cpu) < cp_size + ): + return False + if is_dsa_prefill_cp_round_robin_split(): cur_cp_seq_len = seq_len // cp_size assert ( @@ -216,17 +235,7 @@ def can_dsa_cp_split(seq_len: int, cp_size: int, use_dsa: bool, forward_batch): # Note: (self.cp_size * 2) To achieve load balancing for seq computation, # the seq data needs to be divided and recombined at twice the size of cp_size. cur_cp_seq_len = seq_len // (cp_size * 2) - if ( - cur_cp_seq_len != 0 - and cp_size > 1 - and use_dsa - and forward_batch.forward_mode.is_context_parallel_extend() - and is_dsa_enable_prefill_cp() - and sum(forward_batch.extend_seq_lens_cpu) >= cp_size - ): - return True - else: - return False + return cur_cp_seq_len != 0 from sglang.kernels.ops.attention.dsa.cp_split import ( diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 5931ee9d6..1223de526 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -62,6 +62,8 @@ from sglang.srt.layers.attention.trtllm_mla_backend import ( grow_multi_ctas_kv_counter_buffer_if_needed, make_persistent_multi_ctas_kv_counter_buffer, ) +from sglang.srt.layers.cp.base import get_cp_strategy +from sglang.srt.layers.cp.utils import is_cp_v2_active from sglang.srt.layers.utils.cp_utils import ( cp_all_gather_rerange_output, cp_split_and_rebuild_position, @@ -111,6 +113,23 @@ def _all_gather_dsa_trtllm_fp8_kv( return kv.split((kv_lora_rank, qk_rope_head_dim), dim=-1) +def materialize_full_kv_cp( + attn_mla, + forward_batch: ForwardBatch, + latent_cache: torch.Tensor, + k_nope: torch.Tensor, + k_pe: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + if is_cp_v2_active(forward_batch): + return get_cp_strategy().materialize_full_mla_kv( + forward_batch, + attn_mla.attn_mqa, + k_nope, + k_pe, + ) + return attn_mla.rebuild_cp_kv_cache(latent_cache, forward_batch, k_nope, k_pe) + + _is_hip = is_hip() if _is_hip: @@ -987,12 +1006,19 @@ class DeepseekSparseAttnBackend( ) if can_dsa_prefill_cp_round_robin_split(forward_batch): - seqlens_expanded = dsa_cp_round_robin_split_data(seqlens_expanded) - extend_seq_lens_cpu, extend_seq_lens, bs_idx_cpu, bs_idx = ( - dsa_cp_round_robin_split_q_seqs( - extend_seq_lens_cpu, extend_seq_lens + if is_cp_v2_active(forward_batch): + strategy = get_cp_strategy() + seqlens_expanded = strategy.shard_local_tokens(seqlens_expanded) + extend_seq_lens_cpu, extend_seq_lens, bs_idx_cpu, bs_idx = ( + strategy.shard_per_request(extend_seq_lens_cpu, extend_seq_lens) + ) + else: + seqlens_expanded = dsa_cp_round_robin_split_data(seqlens_expanded) + extend_seq_lens_cpu, extend_seq_lens, bs_idx_cpu, bs_idx = ( + dsa_cp_round_robin_split_q_seqs( + extend_seq_lens_cpu, extend_seq_lens + ) ) - ) indexer_seq_lens_cpu = indexer_seq_lens_cpu[bs_idx_cpu] indexer_seq_lens = indexer_seq_lens[bs_idx] cache_seqlens_int32 = cache_seqlens_int32[bs_idx] @@ -1198,9 +1224,14 @@ class DeepseekSparseAttnBackend( token_to_batch_idx = torch.cat(token_to_batch_idx, dim=0) if bs_idx is not None: assert can_dsa_prefill_cp_round_robin_split(forward_batch) - ks = dsa_cp_round_robin_split_data(ks) - ke = dsa_cp_round_robin_split_data(ke) - token_to_batch_idx = dsa_cp_round_robin_split_data(token_to_batch_idx) + split_per_token = ( + get_cp_strategy().shard_local_tokens + if is_cp_v2_active(forward_batch) + else dsa_cp_round_robin_split_data + ) + ks = split_per_token(ks) + ke = split_per_token(ke) + token_to_batch_idx = split_per_token(token_to_batch_idx) return (ks, ke), token_to_batch_idx def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): @@ -3187,7 +3218,12 @@ class DeepseekSparseAttnBackend( self.qk_rope_head_dim, ) if save_kv_cache and dsa_use_prefill_cp(forward_batch): - k, k_rope = _all_gather_dsa_trtllm_fp8_kv(forward_batch, k, k_rope) + if is_cp_v2_active(forward_batch): + k, k_rope = get_cp_strategy().all_gather_dsa_trtllm_fp8_kv( + forward_batch, k, k_rope + ) + else: + k, k_rope = _all_gather_dsa_trtllm_fp8_kv(forward_batch, k, k_rope) merge_query = False # Save KV cache if requested diff --git a/python/sglang/srt/layers/cp/base.py b/python/sglang/srt/layers/cp/base.py index 6b66a21f4..d0070788d 100644 --- a/python/sglang/srt/layers/cp/base.py +++ b/python/sglang/srt/layers/cp/base.py @@ -67,17 +67,20 @@ class CPAttentionBackendKind(IntEnum): """Attention backend calling convention used by CP strategy dispatch.""" FLASH_ATTENTION = 0 - TRTLLM_MHA = 1 + DSA = 1 + TRTLLM_MHA = 2 @classmethod def from_string(cls, value: str) -> CPAttentionBackendKind: if value in ("fa3", "fa4", "flashinfer"): return cls.FLASH_ATTENTION + if value in ("dsa"): + return cls.DSA if value == "trtllm_mha": return cls.TRTLLM_MHA raise ValueError( f"Unsupported attention_backend={value!r} for CP strategy; expected one " - "of {'fa3', 'fa4', 'flashinfer', 'trtllm_mha'}" + "of {'fa3', 'fa4', 'flashinfer', 'dsa', 'trtllm_mha'}" ) @@ -153,6 +156,25 @@ class ContextParallelStrategy(ABC): f"{self.name} strategy does not support per-request sharding" ) + def shard_local_tokens(self, input_: Any) -> Any: + raise NotImplementedError( + f"{self.name} strategy does not support local-token sharding" + ) + + def materialize_full_indexer_k_cache( + self, key: Any, forward_batch: ForwardBatch + ) -> Any: + raise NotImplementedError( + f"{self.name} strategy does not support DSA indexer key gather" + ) + + def all_gather_dsa_trtllm_fp8_kv( + self, forward_batch: ForwardBatch, k: Any, k_rope: Any + ) -> Any: + raise NotImplementedError( + f"{self.name} strategy does not support DSA trtllm FP8 KV gather" + ) + def split_before_forward( self, forward_batch: ForwardBatch, @@ -185,13 +207,23 @@ class ContextParallelStrategy(ABC): def materialize_full_kv( self, forward_batch: ForwardBatch, - layer: Any, - k: Any, - v: Any, + layer: Any = None, + k: Any = None, + v: Any = None, swa_loc: Optional[Any] = None, - ) -> None: + ) -> Any: """Write full-layout K/V to the backend cache if needed.""" + @abstractmethod + def materialize_full_mla_kv( + self, + forward_batch: ForwardBatch, + layer: Any, + k_nope: Any, + k_rope: Any, + ) -> Any: + """Materialize full-layout MLA K/V for the strategy.""" + def reindex_attn_metadata(self, core_attn_metadata: Any) -> None: """Optional attention metadata rewrite for strategies that need it.""" return None diff --git a/python/sglang/srt/layers/cp/interleave.py b/python/sglang/srt/layers/cp/interleave.py index 91769e223..60c6f2e38 100644 --- a/python/sglang/srt/layers/cp/interleave.py +++ b/python/sglang/srt/layers/cp/interleave.py @@ -31,17 +31,30 @@ from __future__ import annotations from dataclasses import dataclass from typing import Any, List, Optional +import torch + +from sglang.srt.distributed.device_communicators.pynccl_allocator import ( + use_symmetric_memory, +) from sglang.srt.layers.cp.base import ( BaseContextParallelMetadata, ContextParallelStrategy, ContextParallelStrategyKind, CPAttentionBackendKind, ) +from sglang.srt.layers.cp.padding import pad_local_rows +from sglang.srt.layers.dp_attention import ( + attn_cp_all_gather_into_tensor, + is_allocation_symmetric, +) +from sglang.srt.runtime_context import get_parallel @dataclass class InterleaveContextParallelMetadata(BaseContextParallelMetadata): - """Interleave has no per-forward zigzag permutation payload.""" + per_rank_actual_token: Optional[List[int]] = None + max_rank_len: Optional[List[int]] = None + per_rank_logical_token: Optional[List[int]] = None class InterleaveCPStrategy(ContextParallelStrategy): @@ -49,10 +62,11 @@ class InterleaveCPStrategy(ContextParallelStrategy): kind = ContextParallelStrategyKind.INTERLEAVE def can_apply(self, num_tokens: int, forward_batch) -> bool: - if self.cp_size <= 1 or num_tokens < self.cp_size: + if not forward_batch.forward_mode.is_context_parallel_extend(): return False - forward_mode = getattr(forward_batch, "forward_mode", None) - return forward_mode is None or forward_mode.is_context_parallel_extend() + cp_size = self.cp_size + seq_len = sum(forward_batch.extend_seq_lens_cpu) + return seq_len > 0 and seq_len >= cp_size and cp_size > 1 def build_metadata( self, @@ -60,32 +74,150 @@ class InterleaveCPStrategy(ContextParallelStrategy): seqs_len: Optional[List[int]], extend_seqs_len: Optional[List[int]] = None, ) -> InterleaveContextParallelMetadata: + if extend_seqs_len is None: + extend_seqs_len = seqs_len or [num_tokens] + extend_seqs_len = [int(x) for x in extend_seqs_len] + + pad_len = int(num_tokens) - sum(extend_seqs_len) + if pad_len > 0: + extend_seqs_len[-1] += pad_len + + total_seq_lens = sum(extend_seqs_len) + base_len, extra = divmod(total_seq_lens, self.cp_size) + per_rank_actual_token = [ + base_len + (rank < extra) for rank in range(self.cp_size) + ] + return InterleaveContextParallelMetadata( - total_seq_lens=sum(extend_seqs_len or seqs_len or [num_tokens]), - bs=len(extend_seqs_len or seqs_len or [num_tokens]), + per_rank_actual_token=per_rank_actual_token, + max_rank_len=[max(per_rank_actual_token)] * self.cp_size, + total_seq_lens=total_seq_lens, + bs=len(extend_seqs_len), ) def shard_hidden_states(self, x: Any, forward_batch) -> Any: - raise NotImplementedError( - "Interleave hidden-state sharding will land in a follow-up PR" - ) + metadata = forward_batch.attn_cp_metadata + local_x = self._interleave_shard(x[: metadata.total_seq_lens]) + return pad_local_rows(local_x, metadata, dim=0) def shard_position_ids(self, positions: Any, forward_batch) -> Any: - raise NotImplementedError( - "Interleave position-id sharding will land in a follow-up PR" + metadata = forward_batch.attn_cp_metadata + local_positions = self._interleave_shard(positions[: metadata.total_seq_lens]) + return pad_local_rows(local_positions, metadata, dim=0) + + def _interleave_shard(self, input_: Any) -> Any: + cp_size = self.cp_size + cp_rank = self.cp_rank + if isinstance(input_, (tuple, list)): + indices = range(cp_rank, len(input_), cp_size) + return input_[indices] + + tokens = len(input_) + if tokens % cp_size != 0: + cur_len = tokens // cp_size + (tokens % cp_size > cp_rank) + if cur_len == 0: + return input_.new_empty(0, *input_.shape[1:]) + indices = torch.arange(cp_rank, tokens, cp_size, device=input_.device) + return input_[indices] + + return input_.view(-1, cp_size, *input_.shape[1:])[:, cp_rank].contiguous() + + def shard_local_tokens(self, input_: Any) -> Any: + return self._interleave_shard(input_) + + def shard_per_request( + self, + extend_seqs_cpu: List[int], + extend_seqs: Any, + ): + """Build device outputs in the shared kernel to keep the split graph-safe.""" + from sglang.kernels.ops.attention.dsa.cp_split import ( + dsa_cp_round_robin_split_q_seqs_kernel, ) + cp_size = self.cp_size + cp_rank = self.cp_rank + + extra_seq = 0 + q_lens_cpu: List[int] = [] + for cur_len in extend_seqs_cpu: + cur_len += extra_seq + cur_seq = cur_len // cp_size + int(cur_len % cp_size > cp_rank) + q_lens_cpu.append(cur_seq) + extra_seq = cur_len - cur_seq * cp_size + bs_idx_cpu = [i for i, q_len in enumerate(q_lens_cpu) if q_len > 0] + q_lens_cpu = [q_len for q_len in q_lens_cpu if q_len > 0] + + q_lens = torch.empty( + (len(bs_idx_cpu),), device=extend_seqs.device, dtype=extend_seqs.dtype + ) + bs_idx = torch.empty( + (len(bs_idx_cpu),), device=extend_seqs.device, dtype=torch.int32 + ) + dsa_cp_round_robin_split_q_seqs_kernel[(1,)]( + extend_seqs, q_lens, bs_idx, len(extend_seqs), cp_size, cp_rank + ) + return q_lens_cpu, q_lens, bs_idx_cpu, bs_idx + def gather_hidden_states( self, x: Any, forward_batch, stream: Optional[Any] = None ) -> Any: - raise NotImplementedError( - "Interleave hidden-state gather will land in a follow-up PR" - ) + return self._gather_interleaved_tensor(x, forward_batch) def gather_kv_cache( self, x: Any, forward_batch, stream: Optional[Any] = None ) -> Any: - raise NotImplementedError("Interleave KV gather will land in a follow-up PR") + return self._gather_interleaved_tensor(x, forward_batch) + + def _gather_interleaved_tensor(self, x: Any, forward_batch) -> Any: + metadata = getattr(forward_batch, "attn_cp_metadata", None) + if metadata is None: + raise RuntimeError("Interleave CP gather requires attn_cp_metadata.") + + total_tokens = int(metadata.total_seq_lens) + if total_tokens < 0: + raise RuntimeError( + f"Invalid interleave CP total_seq_lens={total_tokens}; expected >= 0." + ) + + logical_rank_lens = ( + metadata.per_rank_logical_token or metadata.per_rank_actual_token + ) + local_logical_len = logical_rank_lens[self.cp_rank] + if x.shape[0] < local_logical_len: + raise RuntimeError( + "Interleave CP gather received an unexpected local token count: " + f"rank={self.cp_rank}, got={x.shape[0]}, " + f"expected_at_least={local_logical_len}, " + f"total={total_tokens}, cp_size={self.cp_size}." + ) + + physical_rank_len = max(metadata.per_rank_actual_token) + if physical_rank_len == 0: + return x.new_empty((0, *x.shape[1:])) + + padded_x = x.new_zeros((physical_rank_len, *x.shape[1:])) + padded_x[:local_logical_len] = x[:local_logical_len] + + with use_symmetric_memory( + get_parallel().attn_cp_group, disabled=not is_allocation_symmetric() + ): + gathered = x.new_empty((self.cp_size * physical_rank_len, *x.shape[1:])) + attn_cp_all_gather_into_tensor(gathered, padded_x.contiguous()) + + flat_indices = torch.arange(total_tokens, device=x.device) + gather_indices = ( + flat_indices % self.cp_size + ) * physical_rank_len + flat_indices // self.cp_size + return gathered.index_select(0, gather_indices) + + def get_supported_attention_backend(self): + return [CPAttentionBackendKind.DSA] + + def materialize_full_indexer_k_cache(self, key: Any, forward_batch) -> Any: + return self.gather_kv_cache( + key.contiguous(), forward_batch, torch.cuda.current_stream() + ) def run_attention( self, @@ -94,14 +226,47 @@ class InterleaveCPStrategy(ContextParallelStrategy): device: Any, attn_fn, attention_backend: CPAttentionBackendKind = CPAttentionBackendKind.FLASH_ATTENTION, + **kwargs, ) -> Any: - raise NotImplementedError( - "Interleave attention dispatch will land in a follow-up PR" - ) + # No-op: run_attention is the FlashAttention/zigzag dispatch hook. + # Interleave serves the DSA backend, which runs attention itself. + return None + + def all_gather_dsa_trtllm_fp8_kv(self, forward_batch, k: Any, k_rope: Any) -> Any: + kv_lora_rank = k.shape[-1] + qk_rope_head_dim = k_rope.shape[-1] + kv_dtype = k.dtype + # Pack → gather in raw bytes to avoid dtype issues with FP8 + kv = torch.cat((k, k_rope), dim=-1).view(torch.uint8) + kv = self.gather_kv_cache( + kv.contiguous(), forward_batch, torch.cuda.current_stream() + ).view(kv_dtype) + return kv.split((kv_lora_rank, qk_rope_head_dim), dim=-1) def materialize_full_kv( - self, forward_batch, layer: Any, k: Any, v: Any, swa_loc: Optional[Any] = None - ) -> None: + self, + forward_batch, + layer: Any = None, + k: Any = None, + v: Any = None, + swa_loc: Optional[Any] = None, + ) -> Any: raise NotImplementedError( - "Interleave KV materialization will land in a follow-up PR" + f"{self.name} strategy does not support dense K/V materialization" ) + + def materialize_full_mla_kv( + self, + forward_batch, + layer: Any, + k_nope: Any, + k_rope: Any, + ) -> Any: + kv_lora_rank = k_nope.shape[-1] + latent_cache = torch.cat([k_nope, k_rope], dim=-1).squeeze(1) + full_latent = self.gather_kv_cache( + latent_cache.contiguous(), forward_batch, torch.cuda.current_stream() + ) + k_nope = full_latent[..., :kv_lora_rank].unsqueeze(1) + k_rope = full_latent[..., kv_lora_rank:].unsqueeze(1) + return k_nope, k_rope diff --git a/python/sglang/srt/layers/cp/utils.py b/python/sglang/srt/layers/cp/utils.py index 1ec8887f8..e5cc4c6fa 100644 --- a/python/sglang/srt/layers/cp/utils.py +++ b/python/sglang/srt/layers/cp/utils.py @@ -14,6 +14,7 @@ """Public import facade and runtime helpers for context parallel strategies.""" +from contextlib import contextmanager from typing import TYPE_CHECKING, Any, Optional, Tuple from sglang.srt.layers.cp.base import ( @@ -40,6 +41,8 @@ if TYPE_CHECKING: CP_V2_DEFAULT_MODEL_CLASSES = frozenset( { + "DeepseekV32ForCausalLM", + "GlmMoeDsaForCausalLM", "GptOssForCausalLM", "MiMoV2FlashForCausalLM", "MiMoV2ForCausalLM", @@ -176,9 +179,7 @@ def prepare_cp_forward(forward_batch) -> None: from sglang.srt.layers.dp_attention import set_local_dp_buffer_len set_local_dp_buffer_len( - forward_batch.attn_cp_metadata.per_rank_actual_token[ - get_parallel().attn_cp_rank - ] + sum(forward_batch.attn_cp_metadata.per_rank_actual_token) ) if getattr(forward_batch, "out_cache_loc", None) is not None: @@ -191,15 +192,31 @@ def cp_split_before_forward( forward_batch, ) -> Tuple[Optional[Any], Optional[Any]]: """Shard embeddings and positions for CP-v2 model-runner forwarding.""" + assert is_cp_v2_active(forward_batch) + assert complete_hidden_states is not None + assert getattr(forward_batch, "attn_cp_metadata", None) is not None + return ( + cp_shard_hidden_states(complete_hidden_states, forward_batch), + cp_shard_position_ids(complete_position_ids, forward_batch), + ) + + +def cp_shard_hidden_states(complete_hidden_states: Any, forward_batch): assert is_cp_v2_active(forward_batch) strategy = get_cp_strategy() assert strategy is not None assert complete_hidden_states is not None assert getattr(forward_batch, "attn_cp_metadata", None) is not None - return ( - strategy.shard_hidden_states(complete_hidden_states, forward_batch), - strategy.shard_position_ids(complete_position_ids, forward_batch), - ) + return strategy.shard_hidden_states(complete_hidden_states, forward_batch) + + +def cp_shard_position_ids(complete_position_ids: Any, forward_batch): + assert is_cp_v2_active(forward_batch) + strategy = get_cp_strategy() + assert strategy is not None + assert complete_position_ids is not None + assert getattr(forward_batch, "attn_cp_metadata", None) is not None + return strategy.shard_position_ids(complete_position_ids, forward_batch) def cp_gather_after_forward(x: Any, forward_batch, stream: Optional[Any] = None): @@ -221,6 +238,38 @@ def cp_gather_after_forward(x: Any, forward_batch, stream: Optional[Any] = None) return strategy.gather_hidden_states(x, forward_batch, stream) +@contextmanager +def cp_shard_model_inputs( + complete_hidden_states: Any, + complete_position_ids: Any, + forward_batch, +): + """Restore the shared batch so logits processing keeps full-batch metadata.""" + assert is_cp_v2_active(forward_batch) + sharded_hidden_states = cp_shard_hidden_states( + complete_hidden_states, forward_batch + ) + sharded_positions = cp_shard_position_ids(complete_position_ids, forward_batch) + + spec_info = getattr(forward_batch, "spec_info", None) + spec_hidden_states = getattr(spec_info, "hidden_states", None) + spec_hidden_states_backup = None + if ( + spec_hidden_states is not None + and spec_hidden_states.shape[0] == complete_hidden_states.shape[0] + ): + spec_hidden_states_backup = spec_hidden_states + spec_info.hidden_states = cp_shard_hidden_states( + spec_hidden_states, forward_batch + ) + + try: + yield sharded_hidden_states, sharded_positions + finally: + if spec_hidden_states_backup is not None: + spec_info.hidden_states = spec_hidden_states_backup + + def _to_int_list(values) -> Optional[list[int]]: if values is None: return None @@ -244,6 +293,9 @@ __all__ = [ "get_cp_strategy", "is_cp_v2_active", "cp_gather_after_forward", + "cp_shard_hidden_states", + "cp_shard_model_inputs", + "cp_shard_position_ids", "cp_split_before_forward", "prepare_cp_forward", "is_glm_dsa_cache_layer_split_enabled", diff --git a/python/sglang/srt/layers/cp/zigzag.py b/python/sglang/srt/layers/cp/zigzag.py index b7fe868d6..7b5154a96 100644 --- a/python/sglang/srt/layers/cp/zigzag.py +++ b/python/sglang/srt/layers/cp/zigzag.py @@ -363,8 +363,13 @@ class ZigzagCPStrategy(ContextParallelStrategy): return result def materialize_full_kv( - self, forward_batch, layer: Any, k: Any, v: Any, swa_loc: Optional[Any] = None - ) -> None: + self, + forward_batch, + layer: Any = None, + k: Any = None, + v: Any = None, + swa_loc: Optional[Any] = None, + ) -> Any: cache_loc = ( forward_batch.out_cache_loc if not layer.is_cross_attention diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index ff1c300e9..3c716c840 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -26,7 +26,7 @@ from sglang.srt.dllm.config import DllmConfig from sglang.srt.environ import envs from sglang.srt.layers.cp.utils import ( cp_gather_after_forward, - cp_split_before_forward, + cp_shard_model_inputs, is_cp_v2_active, prepare_cp_forward, ) @@ -259,7 +259,11 @@ class EagerRunner(BaseRunner): if not self.enable_pdmux: forward_batch = self.load_batch(forward_batch, pp_proxy_tensors) - if forward_batch.needs_forward_metadata_init(): + cp_v2_active = is_cp_v2_active(forward_batch) + if cp_v2_active: + prepare_cp_forward(forward_batch) + + if forward_batch.needs_forward_metadata_init() or cp_v2_active: if model_runner.dcp_size > 1 and hasattr( model_runner.model, "prepare_context_parallel_metadata_for_dcp" ): @@ -285,7 +289,6 @@ class EagerRunner(BaseRunner): model_runner.model.prepare_forward_batch(forward_batch) model_runner.attn_backend.init_forward_metadata(forward_batch) - cp_v2_active = is_cp_v2_active(forward_batch) if not cp_v2_active: forward_batch.attn_cp_metadata = None @@ -341,21 +344,21 @@ class EagerRunner(BaseRunner): """ model = self.model_runner.model - prepare_cp_forward(forward_batch) input_embeds = kwargs.get("input_embeds") if input_embeds is None: input_embeds = model.get_input_embeddings()(forward_batch.input_ids) - input_embeds, positions = cp_split_before_forward( + with cp_shard_model_inputs( input_embeds, forward_batch.positions, forward_batch - ) - - hidden_states = model.model( - forward_batch.input_ids, - positions, - forward_batch, - input_embeds=input_embeds, - pp_proxy_tensors=kwargs.get("pp_proxy_tensors"), - ) + ) as (sharded_input_embeds, sharded_positions): + model_kwargs = {"input_embeds": sharded_input_embeds} + if (pp_proxy_tensors := kwargs.get("pp_proxy_tensors")) is not None: + model_kwargs["pp_proxy_tensors"] = pp_proxy_tensors + hidden_states = model.model( + forward_batch.input_ids, + sharded_positions, + forward_batch, + **model_kwargs, + ) capture_aux_hidden_states = getattr(model, "capture_aux_hidden_states", False) aux_hidden_states = None if capture_aux_hidden_states: 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 80985aa57..604bbcdd2 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 @@ -774,15 +774,24 @@ class DeepseekMLAForwardMixin: dsa_prefill_cp=dsa_prefill_cp, fuse_rope_for_trtllm_mla=fuse_rope_for_trtllm_mla, ) - if ( - (dsa_prefill_cp or mla_prefill_cp) - and not defer_kv_gather_until_after_rope - and not is_cp_v2_active(forward_batch) - ): + if dsa_prefill_cp and not defer_kv_gather_until_after_rope: + from sglang.srt.layers.attention.dsa_backend import materialize_full_kv_cp + + k_nope, k_pe = materialize_full_kv_cp( + self, + forward_batch, + latent_cache, + k_nope, + k_pe, + ) + elif mla_prefill_cp and not is_cp_v2_active(forward_batch): # CP-v1 gathers the latent here; CP-v2 gathers it in the attention # backend via the strategy (materialize_full_mla_kv). k_nope, k_pe = self.rebuild_cp_kv_cache( - latent_cache, forward_batch, k_nope, k_pe + latent_cache, + forward_batch, + k_nope, + k_pe, ) # 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. diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index a2246d22a..741487484 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -35,7 +35,7 @@ from sglang.srt.layers.attention.dsa.utils import ( is_dsa_enable_prefill_cp, is_dsa_prefill_cp_round_robin_split, ) -from sglang.srt.layers.cp.utils import is_cp_v2_active +from sglang.srt.layers.cp.utils import cp_gather_after_forward, is_cp_v2_active from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ReplicatedLinear from sglang.srt.layers.logits_processor import LogitsProcessor @@ -252,11 +252,12 @@ class DeepseekModelNextN(nn.Module): else: hidden_states = self.eh_proj(eh_input) - # CP-v2 shards/gathers at the eager-runner boundary instead. + # CP-v2 shards/gathers hidden states at the eager-runner boundary. + cp_v2_active = is_cp_v2_active(forward_batch) use_cp_v1 = ( dsa_use_prefill_cp(forward_batch, self.dsa_enable_prefill_cp) or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp) - ) and not is_cp_v2_active(forward_batch) + ) and not cp_v2_active if use_cp_v1: hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) positions = cp_split_and_rebuild_position(forward_batch, positions) @@ -304,6 +305,12 @@ class DeepseekModelNextN(nn.Module): forward_batch, torch.cuda.current_stream(), ) + elif ( + cp_v2_active + and should_update_dsa_topk_indices + and topk_indices is not None + ): + topk_indices = cp_gather_after_forward(topk_indices, forward_batch) if should_update_dsa_topk_indices and topk_indices is not None: if forward_batch.reuse_dsa_topk_indices: forward_batch.spec_info.dsa_topk_indices = topk_indices @@ -389,26 +396,27 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM): forward_batch: ForwardBatch, ) -> torch.Tensor: # TODO current just support prefill batch=1 and len(input_ids) > self.cp_size * 2 - if self.dsa_enable_prefill_cp: - if can_dsa_cp_split( - len(input_ids), self.cp_size, self.use_dsa, forward_batch - ): - forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( - len(input_ids), - self.cp_rank, - self.cp_size, - forward_batch.seq_lens_cpu.tolist(), - extend_seqs_len=forward_batch.extend_seq_lens_cpu, - ) - elif self.mla_enable_prefill_cp: - if can_cp_split(len(input_ids), self.cp_size, forward_batch): - forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( - len(input_ids), - self.cp_rank, - self.cp_size, - forward_batch.seq_lens_cpu.tolist(), - extend_seqs_len=forward_batch.extend_seq_lens_cpu, - ) + if not is_cp_v2_active(forward_batch): + if self.dsa_enable_prefill_cp: + if can_dsa_cp_split( + len(input_ids), self.cp_size, self.use_dsa, forward_batch + ): + forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( + len(input_ids), + self.cp_rank, + self.cp_size, + forward_batch.seq_lens_cpu.tolist(), + extend_seqs_len=forward_batch.extend_seq_lens_cpu, + ) + elif self.mla_enable_prefill_cp: + if can_cp_split(len(input_ids), self.cp_size, forward_batch): + forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( + len(input_ids), + self.cp_rank, + self.cp_size, + forward_batch.seq_lens_cpu.tolist(), + extend_seqs_len=forward_batch.extend_seq_lens_cpu, + ) hidden_states = self.model(input_ids, positions, forward_batch) return self.logits_processor( input_ids, hidden_states, self.lm_head, forward_batch diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index dd6da09da..e00c411df 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -3045,6 +3045,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): input_embeds: torch.Tensor = None, pp_proxy_tensors: Optional[PPProxyTensors] = None, ) -> torch.Tensor: + # Multi-modal: input_ids may be None (use input_embeds). # Non-first PP ranks: both are None (activations via pp_proxy_tensors). if input_ids is not None: @@ -3053,26 +3054,27 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): len_input_ids = input_embeds.shape[0] else: len_input_ids = pp_proxy_tensors["hidden_states"].shape[0] - if self.dsa_enable_prefill_cp: - if can_dsa_cp_split( - len_input_ids, self.cp_size, self.use_dsa, forward_batch - ): - forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( - len_input_ids, - self.cp_rank, - self.cp_size, - forward_batch.seq_lens_cpu.tolist(), - extend_seqs_len=forward_batch.extend_seq_lens_cpu, - ) - elif self.mla_enable_prefill_cp: - if can_cp_split(len_input_ids, self.cp_size, forward_batch): - forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( - len_input_ids, - self.cp_rank, - self.cp_size, - forward_batch.seq_lens_cpu.tolist(), - extend_seqs_len=forward_batch.extend_seq_lens_cpu, - ) + if not is_cp_v2_active(forward_batch): + if self.dsa_enable_prefill_cp: + if can_dsa_cp_split( + len_input_ids, self.cp_size, self.use_dsa, forward_batch + ): + forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( + len_input_ids, + self.cp_rank, + self.cp_size, + forward_batch.seq_lens_cpu.tolist(), + extend_seqs_len=forward_batch.extend_seq_lens_cpu, + ) + elif self.mla_enable_prefill_cp: + if can_cp_split(len_input_ids, self.cp_size, forward_batch): + forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( + len_input_ids, + self.cp_rank, + self.cp_size, + forward_batch.seq_lens_cpu.tolist(), + extend_seqs_len=forward_batch.extend_seq_lens_cpu, + ) with get_attn_tp_context().maybe_input_scattered(forward_batch): hidden_states = self.model( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 9275a8526..1182ccbb2 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -6189,15 +6189,26 @@ class ServerArgs: def _handle_context_parallelism(self): if parse_connector_type(self.model_path) != ConnectorType.INSTANCE: + from sglang.srt.configs.model_config import is_deepseek_dsa from sglang.srt.layers.cp.utils import CP_V2_DEFAULT_MODEL_CLASSES model_config = self.get_model_config() - model_arch = model_config.hf_config.architectures[0] - if ( - model_arch in CP_V2_DEFAULT_MODEL_CLASSES - and not envs.SGLANG_ENABLE_CP_V2.is_set() - ): - envs.SGLANG_ENABLE_CP_V2.set(True) + hf_config = model_config.hf_config + model_arch = hf_config.architectures[0] + if model_arch in CP_V2_DEFAULT_MODEL_CLASSES: + if getattr(hf_config, "index_share_for_mtp_iteration", False): + # GLM 5.2 (DSA index-share MTP): CP-v2 is not ready for it + # yet, so default the env to off and keep the legacy CP path. + if not envs.SGLANG_ENABLE_CP_V2.is_set(): + envs.SGLANG_ENABLE_CP_V2.set(False) + else: + is_dsa_default_model = is_deepseek_dsa(hf_config) + # DSA CP-v2 currently supports only the interleave strategy. + enable_default_cp_v2 = not is_dsa_default_model or ( + self.enable_prefill_cp and self.cp_strategy == "interleave" + ) + if enable_default_cp_v2 and not envs.SGLANG_ENABLE_CP_V2.is_set(): + envs.SGLANG_ENABLE_CP_V2.set(True) if ( self.enable_prefill_cp diff --git a/test/registered/cp/test_cp_strategy_unit.py b/test/registered/cp/test_cp_strategy_unit.py index a4c1ba50e..d586416f4 100644 --- a/test/registered/cp/test_cp_strategy_unit.py +++ b/test/registered/cp/test_cp_strategy_unit.py @@ -1,4 +1,5 @@ import unittest +from contextlib import contextmanager from types import SimpleNamespace from unittest.mock import patch @@ -13,6 +14,7 @@ from sglang.srt.layers.cp.base import ( is_interleave, is_zigzag, ) +from sglang.srt.layers.cp.interleave import InterleaveCPStrategy from sglang.srt.layers.cp.padding import ( get_cp_padding_align_size, pad_local_rows, @@ -22,6 +24,7 @@ from sglang.srt.layers.cp.utils import ( cp_split_before_forward, enable_cp_v2, is_cp_v2_active, + prepare_cp_forward, ) from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy from sglang.srt.mem_cache.memory_pool import KVWriteLoc @@ -545,5 +548,338 @@ class TestCPZigzagStrategy(CustomTestCase): self.assertTrue(torch.equal(out, q + 100)) +class TestCPInterleaveStrategy(CustomTestCase): + def setUp(self): + init_cp_strategy( + SimpleNamespace( + enable_prefill_cp=True, + cp_strategy="interleave", + attn_cp_size=4, + attention_backend="fa3", + ) + ) + + def tearDown(self): + init_cp_strategy(SimpleNamespace(enable_prefill_cp=False)) + + def _metadata_for_rank(self, rank, *, cp_size, seq_lens, extend_seq_lens): + strategy = InterleaveCPStrategy(cp_size=cp_size) + with get_parallel().override(attn_cp_rank=rank, attn_cp_size=cp_size): + return strategy.build_metadata( + num_tokens=sum(extend_seq_lens), + seqs_len=seq_lens, + extend_seqs_len=extend_seq_lens, + ) + + def _forward_batch(self, metadata, extend_seq_lens): + return SimpleNamespace( + input_ids=torch.arange(sum(extend_seq_lens)), + forward_mode=_ExtendMode(), + extend_seq_lens_cpu=extend_seq_lens, + attn_cp_metadata=metadata, + ) + + def _rank_tensors(self, x, *, cp_size, seq_lens, extend_seq_lens): + per_rank = [] + metas = [] + with self._patch_legacy_round_robin_mode(): + for rank in range(cp_size): + metadata = self._metadata_for_rank( + rank, + cp_size=cp_size, + seq_lens=seq_lens, + extend_seq_lens=extend_seq_lens, + ) + metas.append(metadata) + fb = self._forward_batch(metadata, extend_seq_lens) + strategy = InterleaveCPStrategy(cp_size=cp_size) + with get_parallel().override(attn_cp_rank=rank, attn_cp_size=cp_size): + per_rank.append(strategy.shard_hidden_states(x, fb)) + return metas, per_rank + + @contextmanager + def _patch_legacy_round_robin_mode(self): + with patch( + "sglang.srt.layers.attention.dsa.utils.is_dsa_prefill_cp_round_robin_split", + return_value=True, + ): + yield + + @contextmanager + def _patch_interleave_all_gather(self, rank_tensors): + def all_gather(output, input_tensor): + del input_tensor + torch.cat(rank_tensors, dim=0, out=output) + + patchers = ( + patch( + "sglang.srt.layers.cp.interleave.attn_cp_all_gather_into_tensor", + side_effect=all_gather, + ), + patch( + "sglang.srt.layers.cp.interleave.is_allocation_symmetric", + return_value=False, + ), + patch( + "sglang.srt.layers.cp.interleave.use_symmetric_memory", + return_value=torch.no_grad(), + ), + patch( + "sglang.srt.layers.attention.dsa.utils.is_dsa_prefill_cp_round_robin_split", + return_value=True, + ), + ) + with ( + patchers[0], + patchers[1], + patchers[2], + patchers[3], + get_parallel().override(attn_cp_group=object()), + ): + yield + + def test_interleave_metadata_supports_shared_padding(self): + metadata = InterleaveCPStrategy(cp_size=4).build_metadata( + num_tokens=10, + seqs_len=[10], + extend_seqs_len=[10], + ) + + self.assertEqual(metadata.per_rank_actual_token, [3, 3, 2, 2]) + with patch( + "sglang.srt.layers.cp.padding.get_cp_padding_align_size", + return_value=4, + ): + pad_logical_token_to_physical(metadata) + + self.assertEqual(metadata.per_rank_logical_token, [3, 3, 2, 2]) + self.assertEqual(metadata.per_rank_actual_token, [4, 4, 4, 4]) + self.assertEqual(metadata.max_rank_len, [4, 4, 4, 4]) + + def test_prepare_cp_forward_sizes_gather_buffer_for_all_cp_ranks(self): + forward_batch = SimpleNamespace( + input_ids=torch.arange(10), + positions=torch.arange(10), + forward_mode=_ExtendMode(), + seq_lens_cpu=[10], + extend_seq_lens_cpu=[10], + attn_cp_metadata=None, + global_num_tokens_cpu=[10], + out_cache_loc=None, + ) + + with ( + get_parallel().override(attn_cp_rank=2, attn_cp_size=4), + patch( + "sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", + return_value=True, + ), + patch( + "sglang.srt.layers.cp.padding.get_cp_padding_align_size", + return_value=4, + ), + patch( + "sglang.srt.layers.dp_attention.set_local_dp_buffer_len" + ) as set_buffer_len, + ): + prepare_cp_forward(forward_batch) + + self.assertEqual( + forward_batch.attn_cp_metadata.per_rank_actual_token, + [4, 4, 4, 4], + ) + set_buffer_len.assert_called_once_with(16) + + def test_interleave_shards_hidden_states_and_position_ids(self): + cp_size = 4 + seq_lens = [8] + extend_seq_lens = [8] + x = torch.arange(sum(extend_seq_lens) * 2).view(sum(extend_seq_lens), 2) + positions = torch.arange(sum(extend_seq_lens)) + + for rank in range(cp_size): + metadata = self._metadata_for_rank( + rank, + cp_size=cp_size, + seq_lens=seq_lens, + extend_seq_lens=extend_seq_lens, + ) + fb = self._forward_batch(metadata, extend_seq_lens) + strategy = InterleaveCPStrategy(cp_size=cp_size) + expected_x = x[rank::cp_size] + expected_positions = positions[rank::cp_size] + + with ( + get_parallel().override( + attn_cp_rank=rank, + attn_cp_size=cp_size, + ), + self._patch_legacy_round_robin_mode(), + ): + local_x = strategy.shard_hidden_states(x, fb) + local_positions = strategy.shard_position_ids(positions, fb) + + with patch( + "sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", + return_value=True, + ): + helper_x, helper_positions = cp_split_before_forward( + x, + positions, + fb, + ) + + self.assertTrue(torch.equal(local_x, expected_x)) + self.assertTrue(torch.equal(local_positions, expected_positions)) + self.assertTrue(torch.equal(helper_x, expected_x)) + self.assertTrue(torch.equal(helper_positions, expected_positions)) + + def test_interleave_padding_preserves_shard_and_gather(self): + cp_size = 4 + total_tokens = 10 + x = torch.arange(total_tokens * 2).view(total_tokens, 2) + rank_tensors = [] + metas = [] + + for rank in range(cp_size): + metadata = self._metadata_for_rank( + rank, + cp_size=cp_size, + seq_lens=[total_tokens], + extend_seq_lens=[total_tokens], + ) + with patch( + "sglang.srt.layers.cp.padding.get_cp_padding_align_size", + return_value=4, + ): + pad_logical_token_to_physical(metadata) + fb = self._forward_batch(metadata, [total_tokens]) + with get_parallel().override(attn_cp_rank=rank, attn_cp_size=cp_size): + local_x = InterleaveCPStrategy(cp_size=cp_size).shard_hidden_states( + x, fb + ) + + logical_len = metadata.per_rank_logical_token[rank] + self.assertEqual(local_x.shape[0], 4) + self.assertTrue(torch.equal(local_x[:logical_len], x[rank::cp_size])) + self.assertTrue(torch.count_nonzero(local_x[logical_len:]) == 0) + metas.append(metadata) + rank_tensors.append(local_x) + + for rank in range(cp_size): + fb = self._forward_batch(metas[rank], [total_tokens]) + with ( + get_parallel().override( + attn_cp_rank=rank, + attn_cp_size=cp_size, + ), + self._patch_interleave_all_gather(rank_tensors), + ): + gathered = InterleaveCPStrategy(cp_size=cp_size).gather_hidden_states( + rank_tensors[rank], fb, stream=None + ) + + self.assertTrue(torch.equal(gathered, x)) + + def test_interleave_gathers_hidden_states_to_original_order(self): + cp_size = 4 + seq_lens = [10] + extend_seq_lens = [10] + x = torch.arange(sum(extend_seq_lens) * 2).view(sum(extend_seq_lens), 2) + metas, rank_tensors = self._rank_tensors( + x, + cp_size=cp_size, + seq_lens=seq_lens, + extend_seq_lens=extend_seq_lens, + ) + max_rank_len = max(t.shape[0] for t in rank_tensors) + padded_rank_tensors = [] + for tensor in rank_tensors: + if tensor.shape[0] < max_rank_len: + padded = tensor.new_zeros((max_rank_len, *tensor.shape[1:])) + padded[: tensor.shape[0]] = tensor + padded_rank_tensors.append(padded) + else: + padded_rank_tensors.append(tensor) + + for rank in range(cp_size): + fb = self._forward_batch(metas[rank], extend_seq_lens) + with ( + get_parallel().override( + attn_cp_rank=rank, + attn_cp_size=cp_size, + ), + self._patch_interleave_all_gather(padded_rank_tensors), + ): + gathered = InterleaveCPStrategy(cp_size=cp_size).gather_hidden_states( + rank_tensors[rank], fb, stream=None + ) + + self.assertTrue(torch.equal(gathered, x)) + + def test_interleave_gathers_kv_cache_to_original_order(self): + cp_size = 4 + seq_lens = [8] + extend_seq_lens = [8] + kv = torch.arange(sum(extend_seq_lens) * 2 * 3).view(sum(extend_seq_lens), 2, 3) + metas, rank_tensors = self._rank_tensors( + kv, + cp_size=cp_size, + seq_lens=seq_lens, + extend_seq_lens=extend_seq_lens, + ) + + for rank in range(cp_size): + fb = self._forward_batch(metas[rank], extend_seq_lens) + with ( + get_parallel().override( + attn_cp_rank=rank, + attn_cp_size=cp_size, + ), + self._patch_interleave_all_gather(rank_tensors), + ): + gathered = InterleaveCPStrategy(cp_size=cp_size).gather_kv_cache( + rank_tensors[rank], fb, stream=None + ) + + self.assertTrue(torch.equal(gathered, kv)) + + def test_interleave_materializes_full_mla_kv(self): + strategy = InterleaveCPStrategy(cp_size=2) + forward_batch = SimpleNamespace() + layer = object() + k_nope = torch.arange(6).view(2, 1, 3) + k_rope = torch.arange(4).view(2, 1, 2) + 10 + full_latent = torch.arange(20).view(4, 5) + + with ( + patch.object( + strategy, + "gather_kv_cache", + return_value=full_latent, + ) as gather, + patch( + "sglang.srt.layers.cp.interleave.torch.cuda.current_stream", + return_value=None, + ), + ): + full_k_nope, full_k_rope = strategy.materialize_full_mla_kv( + forward_batch, + layer, + k_nope, + k_rope, + ) + + gather.assert_called_once() + packed_kv, gathered_forward_batch, stream = gather.call_args.args + self.assertTrue( + torch.equal(packed_kv, torch.cat([k_nope, k_rope], dim=-1).squeeze(1)) + ) + self.assertIs(gathered_forward_batch, forward_batch) + self.assertIsNone(stream) + self.assertTrue(torch.equal(full_k_nope, full_latent[:, :3].unsqueeze(1))) + self.assertTrue(torch.equal(full_k_rope, full_latent[:, 3:].unsqueeze(1))) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/cp/test_dsa_prefill_cp.py b/test/registered/cp/test_dsa_prefill_cp.py new file mode 100644 index 000000000..2bd657549 --- /dev/null +++ b/test/registered/cp/test_dsa_prefill_cp.py @@ -0,0 +1,89 @@ +import unittest +from types import SimpleNamespace + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.run_eval import run_eval +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + is_in_ci, + popen_launch_server, + write_github_step_summary, +) + +register_cuda_ci(est_time=320, stage="extra-b", runner_config="deepep-8-gpu-h200") +GLM52_MODEL_PATH = "zai-org/GLM-5.2-FP8" +SERVER_LAUNCH_TIMEOUT = max(DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, 1800) + + +class TestDSACPV2Interleave(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.model = GLM52_MODEL_PATH + cls.base_url = DEFAULT_URL_FOR_TEST + other_args = [ + "--trust-remote-code", + "--tp", + "8", + "--enable-prefill-cp", + "--cp-strategy", + "interleave", + "--attn-cp-size", + "8", + "--speculative-algorithm", + "EAGLE", + "--speculative-num-steps", + "3", + "--speculative-eagle-topk", + "1", + "--speculative-num-draft-tokens", + "4", + "--mem-frac", + "0.85", + "--cuda-graph-max-bs", + "32", + "--max-running-requests", + "32", + "--model-loader-extra-config", + '{"enable_multithread_load": true, "num_threads": 64}', + ] + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=SERVER_LAUNCH_TIMEOUT, + other_args=other_args, + env={"SGLANG_ENABLE_CP_V2": "1"}, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_a_gsm8k( + self, + ): # Append an "a" to make this test run first (alphabetically) to warm up the server + args = SimpleNamespace( + base_url=self.base_url, + model=self.model, + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=500, + num_threads=32, + num_shots=20, + ) + metrics = run_eval(args) + print(f"{metrics=}") + + if is_in_ci(): + write_github_step_summary( + f"### test_a_gsm8k (dsa-cp-v2-interleave)\n" + f'{metrics["score"]=:.3f}\n' + ) + self.assertGreater(metrics["score"], 0.935) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/cp/test_deepseek_v32_cp_single_node.py b/test/registered/cp/test_dsa_prefill_cp_legacy.py similarity index 99% rename from test/registered/cp/test_deepseek_v32_cp_single_node.py rename to test/registered/cp/test_dsa_prefill_cp_legacy.py index 22ca57172..664e0b85b 100644 --- a/test/registered/cp/test_deepseek_v32_cp_single_node.py +++ b/test/registered/cp/test_dsa_prefill_cp_legacy.py @@ -122,6 +122,7 @@ class TestDeepseekV32CPRoundRobinSplit(CustomTestCase): cls.base_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, other_args=other_args, + env={"SGLANG_ENABLE_CP_V2": "0"}, ) @classmethod