diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index d11eb406d..563a829f0 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -750,7 +750,14 @@ class DSV4AttnMetadata: if src_val is None and dst_val is None: continue assert dst_val is not None, f"{field_name=} {src_val=} {dst_val=}" - dst_val.copy_(src_val) + shape_mismatch = dst_val.shape != src_val.shape + assert not shape_mismatch or field_name in self._CP_GLOBAL_FIELDS, ( + f"Only CP-global replay metadata may use a shorter live prefix, " + f"got {field_name=} {src_val.shape=} {dst_val.shape=}" + ) + _copy_tensor_allowing_storage_alias( + dst_val, src_val, pad_value=0 if shape_mismatch else None + ) # These fields are safe to replace because captured kernels only need # the current per-replay objects, or the field is produced inside the @@ -1002,6 +1009,27 @@ def _prefill_graph_max_seq_len() -> Optional[int]: return get_exec().graph.cuda_graph_config.prefill.max_seq_len +def _copy_tensor_allowing_storage_alias( + dst: torch.Tensor, src: torch.Tensor, *, pad_value: Optional[int] = None +) -> None: + """Copy replay metadata while preserving capture-stable destination addresses.""" + if dst is src: + return + if dst.untyped_storage().data_ptr() == src.untyped_storage().data_ptr(): + src = src.clone() + if dst.shape == src.shape: + dst.copy_(src) + return + assert ( + pad_value is not None + and dst.ndim == src.ndim + and dst.shape[0] >= src.shape[0] + and dst.shape[1:] == src.shape[1:] + ), f"Cannot copy replay metadata from {src.shape=} to {dst.shape=}" + dst.fill_(pad_value) + dst[: src.shape[0]].copy_(src) + + @dataclass class DSV4Metadata: core_attn_metadata: DSV4AttnMetadata diff --git a/python/sglang/srt/layers/cp/bcg.py b/python/sglang/srt/layers/cp/bcg.py index 6631c9d48..06f234c54 100644 --- a/python/sglang/srt/layers/cp/bcg.py +++ b/python/sglang/srt/layers/cp/bcg.py @@ -23,17 +23,22 @@ import torch from sglang.srt.arg_groups.overrides import ( attention_backends_of, + model_config_of, resolved_view, resolving_view, ) +from sglang.srt.configs.model_config import is_deepseek_v4 from sglang.srt.layers.cp.base import get_cp_strategy +from sglang.srt.layers.cp.interleave import InterleaveCPStrategy from sglang.srt.layers.cp.padding import get_cp_padding_align_size from sglang.srt.layers.cp.utils import ( cp_gather_after_forward, + cp_shard_hidden_states, cp_split_before_forward, prepare_cp_forward, ) from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy +from sglang.srt.layers.logits_processor import LogitsMetadata from sglang.srt.model_executor.forward_batch_info import PPProxyTensors if TYPE_CHECKING: @@ -50,12 +55,18 @@ def supports_prefill_cp_bcg(server_args: ServerArgs) -> bool: cfg = resolving_view(server_args) resolved = resolved_view(server_args) prefill_attention_backend, _ = attention_backends_of(resolved_view(server_args)) + supports_layout = ( + cfg.cp_strategy == "zigzag" and prefill_attention_backend == "trtllm_mha" + ) or ( + cfg.cp_strategy == "interleave" + and prefill_attention_backend == "dsv4" + and is_deepseek_v4(model_config_of(server_args).hf_config) + ) return ( cfg.enable_prefill_cp and cfg.pp_size == 1 and resolved.attn_cp_size == cfg.tp_size - and cfg.cp_strategy == "zigzag" - and prefill_attention_backend == "trtllm_mha" + and supports_layout ) @@ -67,8 +78,12 @@ def enable_cp_bcg_capture(server_args: ServerArgs) -> bool: def filter_prefill_cp_bcg_capture_num_tokens( capture_num_tokens: list[int], server_args: ServerArgs ) -> list[int]: - """Keep only token buckets where the zigzag CP strategy can run.""" - min_num_tokens = resolved_view(server_args).attn_cp_size * 2 + """Keep only token buckets where the configured CP strategy can run.""" + cfg = resolving_view(server_args) + cp_segments_per_token_block = 2 if cfg.cp_strategy == "zigzag" else 1 + min_num_tokens = ( + resolved_view(server_args).attn_cp_size * cp_segments_per_token_block + ) filtered = [size for size in capture_num_tokens if size >= min_num_tokens] if not filtered: raise ValueError( @@ -96,6 +111,8 @@ class PrefillCPBCGInput: input_embeds: torch.Tensor positions: torch.Tensor + input_ids: Optional[torch.Tensor] = None + num_token_non_padded: Optional[torch.Tensor] = None bucket_local_tokens: Dict[int, int] = field(default_factory=dict) live_local_tokens: int = 0 @@ -114,12 +131,22 @@ class PrefillCPBCGInput: (runner.max_num_tokens,), dtype=torch.int64, ), + input_ids=torch.zeros((runner.max_num_tokens,), dtype=torch.int64), + num_token_non_padded=torch.zeros((), dtype=torch.int32), ) def required_local_tokens(self, extend_seq_lens: Any) -> Optional[int]: - """Return the aligned CP-local rows required by a live zigzag layout.""" + """Return the aligned CP-local rows required by the active layout.""" strategy = get_cp_strategy() - if not isinstance(strategy, ZigzagCPStrategy) or extend_seq_lens is None: + if extend_seq_lens is None: + return None + if isinstance(strategy, InterleaveCPStrategy): + logical_tokens = ( + sum(int(length) for length in extend_seq_lens) + strategy.cp_size - 1 + ) // strategy.cp_size + align_size = get_cp_padding_align_size() + return (logical_tokens + align_size - 1) // align_size * align_size + if not isinstance(strategy, ZigzagCPStrategy): return None cp_segment_num = strategy.cp_size * 2 @@ -219,6 +246,7 @@ class PrefillCPBCGInput: raw_tokens = int(forward_batch.extend_num_tokens) global_input_ids = forward_batch.input_ids[:raw_tokens] global_positions = forward_batch.positions[:raw_tokens] + local_input_ids = cp_shard_hidden_states(global_input_ids, forward_batch) global_input_embeds = runner.model_runner.model.get_input_embeddings()( global_input_ids ) @@ -249,12 +277,31 @@ class PrefillCPBCGInput: input_embeds = self.input_embeds[:captured_local_tokens] positions = self.positions[:captured_local_tokens] + assert self.input_ids is not None + input_ids = self.input_ids[:captured_local_tokens] input_embeds.zero_() positions.zero_() + input_ids.zero_() input_embeds[:live_local_tokens].copy_(local_input_embeds) positions[:live_local_tokens].copy_(local_positions) + input_ids[:live_local_tokens].copy_(local_input_ids) forward_batch.input_embeds = input_embeds - forward_batch.positions = positions + forward_batch._cp_positions = positions + # Keep the global input_ids field intact: the runner uses its length to + # select the global capture bucket. The DSV4 body consumes this fixed, + # rank-local view for hash routing and MegaMoE. + forward_batch._cp_input_ids = input_ids + forward_batch.input_ids_global = input_ids + if forward_batch.num_token_non_padded is not None: + assert self.num_token_non_padded is not None + metadata = forward_batch.attn_cp_metadata + logical_tokens = ( + metadata.per_rank_logical_token or metadata.per_rank_actual_token + ) + strategy = get_cp_strategy() + assert strategy is not None + self.num_token_non_padded.fill_(logical_tokens[strategy.cp_rank]) + forward_batch.num_token_non_padded = self.num_token_non_padded self.live_local_tokens = live_local_tokens @@ -307,10 +354,50 @@ def execute_prefill_cp_bcg( static_forward_batch, torch.cuda.current_stream(), ) - return model.logits_processor( - forward_batch.input_ids, + if aux_hidden_states is not None: + if torch.is_tensor(aux_hidden_states): + aux_hidden_states = cp_gather_after_forward( + aux_hidden_states, static_forward_batch, torch.cuda.current_stream() + ) + else: + aux_hidden_states = [ + cp_gather_after_forward( + aux, static_forward_batch, torch.cuda.current_stream() + ) + for aux in aux_hidden_states + ] + hidden_states_before_norm = None + if isinstance(hidden_states, tuple): + assert len(hidden_states) == 2 + hidden_states, hidden_states_before_norm = hidden_states + + input_ids = forward_batch.input_ids + logits_metadata = forward_batch + tail = None + language_model = getattr(model, "model", None) + if ( + capture_aux_hidden_states + and getattr(language_model, "late_layer_start", None) is not None + and forward_batch.forward_mode.is_extend_without_speculative() + ): + tail_metadata = runner.model_runner.attn_backend.tail_forward_metadata + tail = tail_metadata.late_layer_tail + input_ids = tail.rows(input_ids) + logits_metadata = LogitsMetadata.from_forward_batch(forward_batch) + logits_metadata.extend_seq_lens = tail.extend_seq_lens + logits_metadata.extend_seq_lens_cpu = tail.extend_seq_lens_cpu + logits_metadata.extend_logprob_start_lens_cpu = tail.extend_seq_lens_cpu + + output = model.logits_processor( + input_ids, hidden_states, model.lm_head, - forward_batch, + logits_metadata, aux_hidden_states, + hidden_states_before_norm=( + None if aux_hidden_states is not None else hidden_states_before_norm + ), ) + if tail is not None: + output.hidden_states_token_indices = tail.token_indices + return output diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index 48695bc62..d5489394f 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -701,6 +701,9 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): def _get_layer_model_positions(self, forward_batch: ForwardBatch) -> torch.Tensor: """Mirror outer multimodal wrappers when BCG captures layer_model directly.""" + cp_positions = getattr(forward_batch, "_cp_positions", None) + if cp_positions is not None: + return cp_positions if forward_batch.mrope_positions is None: return forward_batch.positions @@ -782,7 +785,9 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): if self._uses_eager_prefill_tail(): # BCG / Full: capture the transformer body only. positions = self._get_layer_model_positions(forward_batch) - input_ids = forward_batch.input_ids + input_ids = getattr( + forward_batch, "_cp_input_ids", forward_batch.input_ids + ) kwargs = _build_layer_model_forward_kwargs( self.layer_model, forward_batch, pp_proxy_tensors ) diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 8bcc19b02..69ae59dcb 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -2039,7 +2039,10 @@ class MQALayer(MqaAttentionBase): if ( forward_batch.forward_mode.is_extend() and is_in_breakable_cuda_graph() - and not getattr(attn_backend, "low_ratio_prefill_graph", False) + and ( + dsa_use_prefill_cp(forward_batch) + or not getattr(attn_backend, "low_ratio_prefill_graph", False) + ) ): bcg_deepseek_v4_low_ratio_sources(self, x, q_lora, positions) else: @@ -4411,11 +4414,18 @@ class DeepseekV4Model(nn.Module): ) if self.engram_hasher is not None: if cp_extend: - # n-gram hashing needs each token's predecessors: hash the whole prompt + # N-gram hashing needs each token's predecessors, so hash the + # whole prompt before selecting this CP rank's interleaved rows. + # The hasher builds request-to-token indices dynamically; keep + # that work at an eager break during breakable graph capture. total = int(forward_batch.attn_cp_metadata.total_seq_lens) - hash_ids = self.engram_hasher( - forward_batch.input_ids[:total], forward_batch - ) + global_input_ids = forward_batch.input_ids[:total] + if is_in_breakable_cuda_graph(): + hash_ids = bcg_deepseek_v4_engram_hash_ids( + self.engram_hasher, global_input_ids + ) + else: + hash_ids = self.engram_hasher(global_input_ids, forward_batch) parallel = get_parallel() hash_ids = hash_ids[parallel.attn_cp_rank :: parallel.attn_cp_size] pad_rows = hidden_states.shape[0] - hash_ids.shape[0]