diff --git a/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py b/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py index ce6b038e7..efe579f1a 100644 --- a/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py +++ b/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py @@ -5,7 +5,7 @@ import triton import triton.language as tl -@triton.jit(do_not_specialize=["bs", "c128_cur_max_seq_len"]) +@triton.jit(do_not_specialize=["bs", "num_write_tokens", "c128_cur_max_seq_len"]) def _init_compressed_attn_metadata_kernel( seq_lens_ptr, positions_ptr, @@ -21,6 +21,7 @@ def _init_compressed_attn_metadata_kernel( c128_seq_lens_clamp1_ptr, c128_page_indices_ptr, bs, + num_write_tokens, max_pages, c128_cur_max_seq_len, c128_page_size: tl.constexpr, @@ -33,7 +34,8 @@ def _init_compressed_attn_metadata_kernel( seq_len = tl.load(seq_lens_ptr + batch_id) position = tl.load(positions_ptr + batch_id) - raw_out_loc = tl.load(raw_out_loc_ptr + batch_id) + is_write_token = batch_id < num_write_tokens + raw_out_loc = tl.load(raw_out_loc_ptr + batch_id, mask=is_write_token, other=0) c4_should_compress = (seq_len % 4) == 0 c4_out_loc = tl.where(c4_should_compress, raw_out_loc // 4, 0) @@ -41,7 +43,7 @@ def _init_compressed_attn_metadata_kernel( c4_seq_lens_raw = seq_len // 4 c4_seq_lens_clamp1 = tl.maximum(c4_seq_lens_raw, 1) - tl.store(c4_out_loc_ptr + batch_id, c4_out_loc) + tl.store(c4_out_loc_ptr + batch_id, c4_out_loc, mask=is_write_token) tl.store(c4_positions_ptr + batch_id, c4_positions) tl.store(c4_seq_lens_raw_ptr + batch_id, c4_seq_lens_raw) tl.store(c4_seq_lens_clamp1_ptr + batch_id, c4_seq_lens_clamp1) @@ -52,7 +54,7 @@ def _init_compressed_attn_metadata_kernel( c128_seq_lens_raw = seq_len // 128 c128_seq_lens_clamp1 = tl.maximum(c128_seq_lens_raw, 1) - tl.store(c128_out_loc_ptr + batch_id, c128_out_loc) + tl.store(c128_out_loc_ptr + batch_id, c128_out_loc, mask=is_write_token) tl.store(c128_positions_ptr + batch_id, c128_positions) tl.store(c128_seq_lens_raw_ptr + batch_id, c128_seq_lens_raw) tl.store(c128_seq_lens_clamp1_ptr + batch_id, c128_seq_lens_clamp1) @@ -104,14 +106,21 @@ def _init_compressed_attn_metadata_triton( Optional[torch.Tensor], ]: bs = seq_lens.shape[0] + # CP-v2 may add padding rows to the attention metadata, but those rows have + # no cache-write locations. Keep the write buffers unpadded and mask those + # rows in the kernel. + num_write_tokens = raw_out_loc.shape[0] + assert ( + num_write_tokens <= bs + ), f"raw_out_loc has {num_write_tokens} rows, expected at most {bs} metadata rows" device = seq_lens.device - c4_out_loc = torch.empty(bs, dtype=torch.int64, device=device) + c4_out_loc = torch.empty(num_write_tokens, dtype=torch.int64, device=device) c4_positions = torch.empty(bs, dtype=torch.int32, device=device) c4_seq_lens_raw = torch.empty(bs, dtype=torch.int32, device=device) c4_seq_lens_clamp1 = torch.empty(bs, dtype=torch.int32, device=device) - c128_out_loc = torch.empty(bs, dtype=torch.int64, device=device) + c128_out_loc = torch.empty(num_write_tokens, dtype=torch.int64, device=device) c128_positions = torch.empty(bs, dtype=torch.int32, device=device) c128_seq_lens_raw = torch.empty(bs, dtype=torch.int32, device=device) c128_seq_lens_clamp1 = torch.empty(bs, dtype=torch.int32, device=device) @@ -159,6 +168,7 @@ def _init_compressed_attn_metadata_triton( else torch.empty(0, dtype=torch.int32, device=device) ), bs, + num_write_tokens, max_pages, c128_cur_max_seq_len, c128_page_size, diff --git a/python/sglang/srt/arg_groups/deepseek_v4_hook.py b/python/sglang/srt/arg_groups/deepseek_v4_hook.py index 8908b9bb2..da21b91d9 100644 --- a/python/sglang/srt/arg_groups/deepseek_v4_hook.py +++ b/python/sglang/srt/arg_groups/deepseek_v4_hook.py @@ -167,6 +167,7 @@ def validate_deepseek_v4_cp(server_args: ServerArgs) -> None: ) server_args.enable_dsa_prefill_context_parallel = True + server_args.enable_prefill_context_parallel = False server_args.dsa_prefill_cp_mode = "round-robin-split" server_args.enable_dp_attention = True server_args.moe_dense_tp_size = 1 diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index 402d3baa0..0b7cf9195 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -58,6 +58,7 @@ from sglang.srt.layers.attention.dsv4.sparse_prefill_utils import ( SparsePrefillWorkspace, ) from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask +from sglang.srt.layers.cp.utils import is_cp_v2_active from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.runtime_context import get_parallel, get_spec @@ -277,11 +278,16 @@ class DSV4AttnMetadata: for field_name in reference_assign_fields: setattr(self, field_name, getattr(other, field_name)) - def init_compression_metadata(self): + def init_compression_metadata(self, num_tokens: Optional[int] = None) -> None: assert self.page_table.dim() == 2 + # CP-v2 pads causal metadata for per-rank partitioning, while cache-write + # locations remain one-per-logical-token. num_tokens tracks that unpadded + # length; legacy paths use the metadata length. + if num_tokens is None: + num_tokens = self.seq_lens_casual.shape[0] assert ( - self.raw_out_loc.shape == self.seq_lens_casual.shape - ), f"{self.raw_out_loc.shape=}, {self.seq_lens_casual.shape=}" + self.raw_out_loc.shape[0] == num_tokens + ), f"{self.raw_out_loc.shape=}, {num_tokens=}" ( self.c4_out_loc, @@ -305,6 +311,8 @@ class DSV4AttnMetadata: self.c128_page_indices = _pad_last_dim(self.c128_page_indices) self.swa_page_indices = _pad_last_dim(self.swa_page_indices) + # Cache-write locations stay in global logical order and are intentionally + # excluded from CP reindexing. _CP_REINDEX_FIELDS = [ "seq_lens_casual", "positions_casual", @@ -323,7 +331,7 @@ class DSV4AttnMetadata: "c128_out_loc", ] - def apply_cp_reindex(self) -> None: + def apply_cp_reindex(self, num_tokens: Optional[int] = None) -> None: cp_rank = get_parallel().attn_cp_rank cp_size = get_parallel().attn_cp_size idx = slice(cp_rank, None, cp_size) @@ -333,6 +341,8 @@ class DSV4AttnMetadata: "CP round-robin requires padding to ensure divisibility." ) expected_local_len = pre_global_len // cp_size + if num_tokens is None: + num_tokens = pre_global_len for field_name in self._CP_REINDEX_FIELDS: val = getattr(self, field_name, None) assert isinstance( @@ -350,9 +360,9 @@ class DSV4AttnMetadata: val = getattr(self, field_name, None) if val is None: continue - assert val.shape[0] == pre_global_len, ( + assert val.shape[0] == num_tokens, ( f"apply_cp_reindex post-condition: global field {field_name}.shape[0]={val.shape[0]} " - f"!= pre_global_len={pre_global_len} (must remain global for compressor write path)" + f"!= num_tokens={num_tokens} (must remain global for compressor write path)" ) def init_flashmla_related(self, is_prefill: bool = False): @@ -721,13 +731,21 @@ class DeepseekV4AttnBackend( use_prefill_cuda_graph: bool = False, online_c128_state_slot_offset: int = 0, dspark_block_size: Optional[int] = None, + forward_batch: Optional[ForwardBatch] = None, ) -> DSV4Metadata: + padded_num_tokens = out_cache_loc.shape[0] + cp_v2_active = forward_batch is not None and is_cp_v2_active(forward_batch) + if cp_v2_active: + cp_metadata = forward_batch.attn_cp_metadata + assert cp_metadata is not None + padded_num_tokens = sum(cp_metadata.per_rank_actual_token) + seq_lens_casual, req_pool_indices_repeated = self.expand_prefill_casually( num_tokens=num_tokens, seq_lens=seq_lens_cpu, extend_seq_lens=extend_seq_lens_cpu, req_pool_indices=req_pool_indices, - padded_num_tokens=out_cache_loc.shape[0], + padded_num_tokens=padded_num_tokens, seq_lens_tensor=seq_lens, extend_seq_lens_tensor=extend_seq_lens, extend_start_loc=extend_start_loc, @@ -741,7 +759,11 @@ class DeepseekV4AttnBackend( need_compress=need_compress, is_prefill=True, dspark_block_size=dspark_block_size, + num_tokens=num_tokens if cp_v2_active else None, ) + if cp_v2_active: + core_attn_metadata.apply_cp_reindex(num_tokens=num_tokens) + core_attn_metadata.init_flashmla_related(is_prefill=True) indexer_metadata = ( self.init_forward_metadata_indexer( core_attn_metadata, @@ -1458,6 +1480,7 @@ class DeepseekV4AttnBackend( extend_start_loc=forward_batch.extend_start_loc, need_compress=True, use_prefill_cuda_graph=use_prefill_cuda_graph, + forward_batch=forward_batch, ) else: raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}") @@ -1945,6 +1968,7 @@ class DeepseekV4AttnBackend( need_compress: bool = True, is_prefill: bool = False, dspark_block_size: Optional[int] = None, + num_tokens: Optional[int] = None, ) -> DSV4AttnMetadata: assert self.swa_page_size == SWA_WINDOW @@ -2001,7 +2025,7 @@ class DeepseekV4AttnBackend( ) if need_compress: - core_attn_metadata.init_compression_metadata() + core_attn_metadata.init_compression_metadata(num_tokens) core_attn_metadata.init_flashmla_related(is_prefill=is_prefill) else: core_attn_metadata.c4_sparse_topk_lengths = None diff --git a/python/sglang/srt/layers/attention/dsa/utils.py b/python/sglang/srt/layers/attention/dsa/utils.py index 2a7f4a2d8..d202f333a 100644 --- a/python/sglang/srt/layers/attention/dsa/utils.py +++ b/python/sglang/srt/layers/attention/dsa/utils.py @@ -111,9 +111,10 @@ def is_dsa_enable_prefill_cp(): # DeepSeek Sparse Attention model. if get_parallel().attn_cp_size <= 1: return False - from sglang.srt.configs.model_config import is_deepseek_dsa + from sglang.srt.configs.model_config import is_deepseek_dsa, is_deepseek_v4 - return is_deepseek_dsa(get_server_args().get_model_config().hf_config) + hf_config = get_server_args().get_model_config().hf_config + return is_deepseek_dsa(hf_config) or is_deepseek_v4(hf_config) def is_dsa_prefill_cp_in_seq_split(): diff --git a/python/sglang/srt/layers/attention/dsv4/compressor.py b/python/sglang/srt/layers/attention/dsv4/compressor.py index f27eb00d1..ac041d599 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor.py @@ -23,16 +23,15 @@ from sglang.kernels.ops.attention.dsv4.quant_k_cache import ( from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config from sglang.srt.environ import envs from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp +from sglang.srt.layers.cp.utils import cp_materialize_global_token_order from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ReplicatedLinear -from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output from sglang.srt.mem_cache.deepseek_v4_compress_state import ( CompressStatePool, ) from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.models.deepseek_v2 import _is_hip -from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import add_prefix, is_npu, set_weight_attrs _is_npu = is_npu() @@ -426,9 +425,8 @@ class Compressor(BaseFusedOp): # CUDA path: delegate to backend if dsa_use_prefill_cp(forward_batch): - kv_score = cp_all_gather_rerange_output( + kv_score = cp_materialize_global_token_order( kv_score, - get_parallel().attn_cp_size, forward_batch, torch.cuda.current_stream(), ) @@ -473,9 +471,8 @@ class Compressor(BaseFusedOp): return x.new_empty(0, self.head_dim) if dsa_use_prefill_cp(forward_batch): - x = cp_all_gather_rerange_output( + x = cp_materialize_global_token_order( x, - get_parallel().attn_cp_size, forward_batch, torch.cuda.current_stream(), ) diff --git a/python/sglang/srt/layers/cp/utils.py b/python/sglang/srt/layers/cp/utils.py index e5cc4c6fa..b397c4fed 100644 --- a/python/sglang/srt/layers/cp/utils.py +++ b/python/sglang/srt/layers/cp/utils.py @@ -34,6 +34,7 @@ from sglang.srt.layers.cp.zigzag import ( ZigzagContextParallelMetadata, ZigzagCPStrategy, ) +from sglang.srt.layers.moe.utils import get_moe_a2a_backend from sglang.srt.runtime_context import get_parallel if TYPE_CHECKING: @@ -219,6 +220,17 @@ def cp_shard_position_ids(complete_position_ids: Any, forward_batch): return strategy.shard_position_ids(complete_position_ids, forward_batch) +def cp_round_robin_input_ids_v2(input_ids: Any, forward_batch): + assert is_cp_v2_active(forward_batch) + if not get_moe_a2a_backend().is_none(): + return cp_shard_hidden_states(input_ids, forward_batch) + + physical_tokens = sum(forward_batch.attn_cp_metadata.per_rank_actual_token) + padded_input_ids = input_ids.new_zeros(physical_tokens) + padded_input_ids[: input_ids.shape[0]] = input_ids + return padded_input_ids.view(-1, get_parallel().attn_cp_size).T.flatten() + + def cp_gather_after_forward(x: Any, forward_batch, stream: Optional[Any] = None): """Gather CP-v2 hidden states at the model boundary when this batch is active.""" assert is_cp_v2_active(forward_batch) @@ -226,18 +238,40 @@ def cp_gather_after_forward(x: Any, forward_batch, stream: Optional[Any] = None) assert strategy is not None if isinstance(x, tuple): - hidden_states, *rest = x - hidden_states = strategy.gather_hidden_states( - hidden_states, forward_batch, stream + gathered = tuple( + ( + strategy.gather_hidden_states(item, forward_batch, stream) + if item is not None + else None + ) + for item in x ) # MiMo's text-only body returns (hidden_states, None); logits expects a tensor. - if len(rest) == 1 and rest[0] is None: - return hidden_states - return (hidden_states, *rest) + if len(gathered) == 2 and gathered[1] is None: + return gathered[0] + return gathered return strategy.gather_hidden_states(x, forward_batch, stream) +def cp_materialize_global_token_order( + x: Any, forward_batch, stream: Optional[Any] = None +): + """Materialize a CP tensor in the global logical token order.""" + if is_cp_v2_active(forward_batch): + strategy = get_cp_strategy() + assert strategy is not None + return strategy.gather_kv_cache(x, forward_batch, stream) + + # TODO(hzh0425): Keep the legacy gather temporarily for CP-v1 compatibility. Remove it + # with the follow-up CP-v1 cleanup. + from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output + + return cp_all_gather_rerange_output( + x, get_parallel().attn_cp_size, forward_batch, stream + ) + + @contextmanager def cp_shard_model_inputs( complete_hidden_states: Any, @@ -293,6 +327,8 @@ __all__ = [ "get_cp_strategy", "is_cp_v2_active", "cp_gather_after_forward", + "cp_materialize_global_token_order", + "cp_round_robin_input_ids_v2", "cp_shard_hidden_states", "cp_shard_model_inputs", "cp_shard_position_ids", diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index 3c716c840..5f0483140 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -374,12 +374,18 @@ class EagerRunner(BaseRunner): hidden_states = cp_gather_after_forward( hidden_states, forward_batch, torch.cuda.current_stream() ) + logits_kwargs = {} + # DSV4 returns (hidden_states, hidden_states_before_norm) from its model body. + if isinstance(hidden_states, tuple): + hidden_states, hidden_states_before_norm = hidden_states + logits_kwargs["hidden_states_before_norm"] = hidden_states_before_norm return model.logits_processor( forward_batch.input_ids, hidden_states, model.lm_head, forward_batch, aux_hidden_states, + **logits_kwargs, ) def _execute_idle( diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 7353e5a6a..7981f26ed 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -59,6 +59,11 @@ from sglang.srt.layers.communicator_dsa_cp import ( dsa_cp_reduce_scatter_hidden_states, ) from sglang.srt.layers.cp.cp_decode_attn_tp import get_cp_decode_attn_tp_ctx +from sglang.srt.layers.cp.utils import ( + cp_materialize_global_token_order, + cp_round_robin_input_ids_v2, + is_cp_v2_active, +) from sglang.srt.layers.dp_attention import ( _tbo_event, attn_tp_all_gather, @@ -1141,9 +1146,8 @@ class MQALayer(MqaAttentionBase): # DSA CP: keep bf16 kv around for the cross-rank all-gather, then # write to the FlashMLA cache after gather. kv = self._compute_kv_bf16(x, positions, qkv_a=qkv_a) - kv = cp_all_gather_rerange_output( + kv = cp_materialize_global_token_order( kv.contiguous(), - self.cp_size, forward_batch, torch.cuda.current_stream(), ) @@ -1192,9 +1196,8 @@ class MQALayer(MqaAttentionBase): # unified_kv + DSA CP: the 2-source prefill path needs the # FULL current-chunk KV (extend source + ring write), so # all-gather the per-rank bf16 KV across the CP group. - kv = cp_all_gather_rerange_output( + kv = cp_materialize_global_token_order( kv.contiguous(), - self.cp_size, forward_batch, torch.cuda.current_stream(), ) @@ -1202,9 +1205,8 @@ class MQALayer(MqaAttentionBase): # NSA CP: keep bf16 kv around for the cross-rank all-gather, then # write to the FlashMLA cache after gather. kv = self._compute_kv_bf16(x_linear, positions, qkv_a=qkv_a) - kv = cp_all_gather_rerange_output( + kv = cp_materialize_global_token_order( kv.contiguous(), - self.cp_size, forward_batch, torch.cuda.current_stream(), ) @@ -2437,8 +2439,13 @@ class DeepseekV4Model(nn.Module): input_embeds: Optional[torch.Tensor], pp_proxy_tensors: Optional[PPProxyTensors] = None, ) -> Union[torch.Tensor, PPProxyTensors]: + cp_v2_active = is_cp_v2_active(forward_batch) + use_prefill_cp = dsa_use_prefill_cp(forward_batch) if self.pp_group.is_first_rank: - hidden_states = self.embed_tokens(input_ids) + if input_embeds is None: + hidden_states = self.embed_tokens(input_ids) + else: + hidden_states = input_embeds hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1) else: assert pp_proxy_tensors is not None @@ -2462,11 +2469,16 @@ class DeepseekV4Model(nn.Module): else: input_ids_global = input_ids - if dsa_use_prefill_cp(forward_batch): - if self.pp_group.is_first_rank: - hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) - positions = cp_split_and_rebuild_position(forward_batch, positions) - input_ids = cp_round_robin_input_ids(input_ids) + if use_prefill_cp: + if cp_v2_active: + input_ids = cp_round_robin_input_ids_v2(input_ids, forward_batch) + else: + if self.pp_group.is_first_rank: + hidden_states = cp_split_and_rebuild_data( + forward_batch, hidden_states + ) + positions = cp_split_and_rebuild_position(forward_batch, positions) + input_ids = cp_round_robin_input_ids(input_ids) input_ids_global = input_ids # Reset Compressor's per-step freqs_cis cache from any previous step. @@ -2474,7 +2486,7 @@ class DeepseekV4Model(nn.Module): if hasattr(forward_batch, _attr): delattr(forward_batch, _attr) capture_dspark = self.dspark_layers_to_capture is not None - if capture_dspark and dsa_use_prefill_cp(forward_batch): + if capture_dspark and use_prefill_cp: raise NotImplementedError( "DSpark aux hidden-state capture is not supported together with " "DeepSeek-V4 prefill context parallelism (attn_cp_size > 1). Disable one " @@ -2529,7 +2541,7 @@ class DeepseekV4Model(nn.Module): ) # CP all-gather only on the last PP rank; PP IPC carries CP-split tensors. - if self.pp_group.is_last_rank and dsa_use_prefill_cp(forward_batch): + if self.pp_group.is_last_rank and use_prefill_cp and not cp_v2_active: hidden_states = cp_all_gather_rerange_output( hidden_states, self.cp_size, diff --git a/python/sglang/srt/models/deepseek_v4_nextn.py b/python/sglang/srt/models/deepseek_v4_nextn.py index f8d039524..94abf6e5d 100644 --- a/python/sglang/srt/models/deepseek_v4_nextn.py +++ b/python/sglang/srt/models/deepseek_v4_nextn.py @@ -13,6 +13,10 @@ 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 ( + cp_round_robin_input_ids_v2, + is_cp_v2_active, +) from sglang.srt.layers.dp_attention import ( dp_gather_partial, get_global_dp_buffer_len, @@ -115,6 +119,9 @@ class DeepseekV4ModelNextN(nn.Module): self.shared_head = nn.Module() self.shared_head.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + def get_input_embeddings(self) -> nn.Module: + return self.embed_tokens + def hc_head( self, x: torch.Tensor, @@ -137,6 +144,8 @@ class DeepseekV4ModelNextN(nn.Module): forward_batch: ForwardBatch, input_embeds: torch.Tensor = None, ) -> torch.Tensor: + cp_v2_active = is_cp_v2_active(forward_batch) + use_prefill_cp = dsa_use_prefill_cp(forward_batch) if input_embeds is None: hidden_states = self.embed_tokens(input_ids) else: @@ -167,10 +176,13 @@ class DeepseekV4ModelNextN(nn.Module): else: input_ids_global = input_ids - if dsa_use_prefill_cp(forward_batch): - hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) - positions = cp_split_and_rebuild_position(forward_batch, positions) - input_ids = cp_round_robin_input_ids(input_ids) + if use_prefill_cp: + if cp_v2_active: + input_ids = cp_round_robin_input_ids_v2(input_ids, forward_batch) + else: + hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) + positions = cp_split_and_rebuild_position(forward_batch, positions) + input_ids = cp_round_robin_input_ids(input_ids) input_ids_global = input_ids hidden_states, residual, post, comb = self.decoder( @@ -185,7 +197,7 @@ class DeepseekV4ModelNextN(nn.Module): # deferred fused hc_post state. hidden_states = self.decoder.hc_post(hidden_states, residual, post, comb) - if dsa_use_prefill_cp(forward_batch): + if use_prefill_cp and not cp_v2_active: hidden_states = cp_all_gather_rerange_output( hidden_states, self.cp_size, @@ -244,7 +256,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM): positions: torch.Tensor, forward_batch: ForwardBatch, ) -> torch.Tensor: - if self.dsa_enable_prefill_cp: + if self.dsa_enable_prefill_cp and not is_cp_v2_active(forward_batch): if can_dsa_cp_split(len(input_ids), self.cp_size, True, forward_batch): forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( len(input_ids), diff --git a/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py b/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py index cfac77f38..74b7ced61 100644 --- a/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py +++ b/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py @@ -1,7 +1,7 @@ -"""B200 extra CI: DeepSeek-V4-Flash FP4 with attn-CP (DSA prefill CP). +"""B200 extra CI: DeepSeek-V4-Flash FP4 with attn-CP. Balanced recipe (TP=4, DeepEP, EAGLE) plus --attn-cp-size=4 with the -DSA prefill-CP round-robin-split mode. Split out of +DSA prefill-CP interleave strategy. Split out of models_e2e/test_deepseek_v4_flash_fp4_b200.py so the `cp` group covers all context-parallel tests. @@ -29,6 +29,7 @@ SERVER_LAUNCH_TIMEOUT = 3600 DEEPEP_CONFIG = '{"normal_dispatch":{"num_sms":96},"normal_combine":{"num_sms":96}}' _DEEPEP_ENV = { + "SGLANG_ENABLE_CP_V2": "1", "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "1024", # The draft-extend graph pool costs ~4.5 GB here (DeepEP MoE workspace is # captured at full dispatch capacity), which starves the eager prefill @@ -38,6 +39,7 @@ _DEEPEP_ENV = { } _MEGAMOE_ENV = { + "SGLANG_ENABLE_CP_V2": "1", "SGLANG_OPT_DEEPGEMM_MEGA_MOE_NUM_MAX_TOKENS_PER_RANK": "8320", "SGLANG_OPT_DEEPGEMM_MEGA_MOE_USE_FP4_ACTS": "1", "SGLANG_OPT_DEEPGEMM_MEGA_MOE_USE_MXF4_KIND": "1", @@ -81,9 +83,9 @@ class TestDSV4FlashFP4B200Balanced_CP_DeepEP( "1", "--speculative-num-draft-tokens", "2", - "--enable-dsa-prefill-context-parallel", - "--dsa-prefill-cp-mode", - "round-robin-split", + "--enable-prefill-cp", + "--cp-strategy", + "interleave", "--deepep-config", DEEPEP_CONFIG, "--mem-fraction-static", @@ -132,9 +134,9 @@ class TestDSV4FlashFP4B200Balanced_CP_Megamoe( "1", "--speculative-num-draft-tokens", "2", - "--enable-dsa-prefill-context-parallel", - "--dsa-prefill-cp-mode", - "round-robin-split", + "--enable-prefill-cp", + "--cp-strategy", + "interleave", "--deepep-config", DEEPEP_CONFIG, ], @@ -181,12 +183,13 @@ class TestDSV4FlashFP4B200Balanced_CP_NonDeepEP( "1", "--speculative-num-draft-tokens", "2", - "--enable-dsa-prefill-context-parallel", - "--dsa-prefill-cp-mode", - "round-robin-split", + "--enable-prefill-cp", + "--cp-strategy", + "interleave", "--moe-runner-backend", # for fp4 checkpoint "flashinfer_mxfp4", ], + env={"SGLANG_ENABLE_CP_V2": "1"}, ) @classmethod