diff --git a/python/sglang/srt/layers/aux_hidden_states.py b/python/sglang/srt/layers/aux_hidden_states.py index ac4aff91b..2843abf09 100644 --- a/python/sglang/srt/layers/aux_hidden_states.py +++ b/python/sglang/srt/layers/aux_hidden_states.py @@ -42,9 +42,10 @@ class AuxHiddenStatePacker: def finalize(self) -> torch.Tensor: """Return the packed buffer; callers guard the empty case on ``len()``.""" - assert ( - self._buffer is not None and self._idx == self._num_captures - ), f"captured {self._idx} of {self._num_captures} aux hidden states" + if self._buffer is None or self._idx != self._num_captures: + raise RuntimeError( + f"captured {self._idx} of {self._num_captures} aux hidden states" + ) return self._buffer @@ -55,4 +56,6 @@ AuxHiddenStateAccumulator = Union[List[torch.Tensor], AuxHiddenStatePacker] def pack_aux_hidden_states(aux_hidden_states: AuxHiddenStates) -> torch.Tensor: if isinstance(aux_hidden_states, torch.Tensor): return aux_hidden_states + if len(aux_hidden_states) == 1: + return aux_hidden_states[0] return torch.cat(aux_hidden_states, dim=-1) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 95b4118b1..d6d74ebb8 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -124,8 +124,10 @@ from sglang.srt.model_executor.model_runner_components.kv_pool_runtime import ( is_post_capture_kv_active, ) from sglang.srt.model_executor.model_runner_components.layer_setup import ( + AttentionAndMoeLayers, ModelLayerInfo, adjust_hybrid_swa_layer_ids, + compute_attention_and_moe_layers, resolve_layer_indices, ) from sglang.srt.model_executor.model_runner_components.load_model_utils import ( @@ -1435,6 +1437,10 @@ class ModelRunner: return DecodeCudaGraphRunner + def get_cuda_graph_layers(self, layer_model) -> AttentionAndMoeLayers: + """Return the model layers used by prefill CUDA graph execution.""" + return compute_attention_and_moe_layers(layer_model) + def init_decode_cuda_graph(self): self.decode_cuda_graph_runner = None capture = capture_decode_graph(model_runner=self) diff --git a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py index 96324643f..5cb1b72b1 100644 --- a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py @@ -31,9 +31,6 @@ from sglang.srt.model_executor.graph_memory_usage import ( ) from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput from sglang.srt.model_executor.hook_manager import register_forward_hooks -from sglang.srt.model_executor.model_runner_components.layer_setup import ( - compute_attention_and_moe_layers, -) from sglang.srt.model_executor.runner import ( EagerRunner, PrefillCudaGraphRunner, @@ -444,7 +441,7 @@ def capture_prefill_graph( model_runner.moe_fusions, model_runner.dsa_indexers, model_runner.mha_companion_layers, - ) = compute_attention_and_moe_layers(layer_model) + ) = model_runner.get_cuda_graph_layers(layer_model) ( model_runner.attention_layers, model_runner.mha_companion_layers, diff --git a/python/sglang/srt/models/dspark.py b/python/sglang/srt/models/dspark.py index 7c94d1c80..be3930638 100644 --- a/python/sglang/srt/models/dspark.py +++ b/python/sglang/srt/models/dspark.py @@ -484,6 +484,8 @@ _DSPARK_SKIPPED_WEIGHT_PREFIXES = ("lm_head.", "rotary_emb.") class DSparkDraftMixin: + supports_pre_gather_target_hidden_projection = True + def __init__(self, config, quant_config=None, prefix: str = "") -> None: super().__init__(config=config, quant_config=quant_config, prefix=prefix) self._fused_kv_write_cache = None @@ -737,8 +739,13 @@ class DSparkDraftMixin: cache_loc: torch.Tensor, cache_loc_2d: Optional[torch.Tensor] = None, commit_lens: Optional[torch.Tensor] = None, + target_hidden_is_projected: bool = False, ) -> None: - ctx_hidden = self.project_target_hidden(target_hidden) + ctx_hidden = ( + target_hidden + if target_hidden_is_projected + else self.project_target_hidden(target_hidden) + ) bundle = self._fused_kv_write_bundle(pool) if bundle is not None: diff --git a/python/sglang/srt/speculative/dspark_components/dspark_kv_inject.py b/python/sglang/srt/speculative/dspark_components/dspark_kv_inject.py index baaf55eb7..c283ecff3 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_kv_inject.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_kv_inject.py @@ -42,6 +42,7 @@ class TargetHiddenKvInjector: commit_lens: Optional[torch.Tensor] = None, state_slot: Optional[torch.Tensor] = None, final_pos: Optional[torch.Tensor] = None, + target_hidden_is_projected: bool = False, ) -> None: if target_hidden is None or target_hidden.numel() == 0: return @@ -71,6 +72,11 @@ class TargetHiddenKvInjector: pool = self.draft_model_runner.token_to_kv_pool if hasattr(pool, "set_swa_key_buffer_radix_fused_norm_rope"): + if target_hidden_is_projected: + raise RuntimeError( + "Pre-gather target-hidden projection is not supported by the " + "DSpark MLA KV injection path." + ) self._inject_mla( pool=pool, target_hidden=target_hidden, @@ -91,6 +97,7 @@ class TargetHiddenKvInjector: cache_loc=cache_loc, cache_loc_2d=cache_loc_2d, commit_lens=commit_lens, + target_hidden_is_projected=target_hidden_is_projected, ) def _inject_mla( diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index a06d457c6..4a5e68408 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -1,7 +1,7 @@ import logging from contextlib import nullcontext from dataclasses import replace -from typing import Optional +from typing import Callable, Optional, Protocol, runtime_checkable import torch @@ -19,6 +19,7 @@ from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.model_executor.cuda_graph_config import Backend from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, + ForwardMode, compute_position, ) from sglang.srt.runtime_context import ( @@ -88,6 +89,41 @@ logger = logging.getLogger(__name__) _is_npu = is_npu() +@runtime_checkable +class _SupportsDSparkTargetHiddenProjection(Protocol): + def set_dspark_target_hidden_projector( + self, + projector: Callable[[torch.Tensor], torch.Tensor], + *, + num_context_features: int, + ) -> bool: ... + + def should_project_dspark_target_hidden( + self, + *, + forward_mode: ForwardMode, + capture_hidden_mode: CaptureHiddenMode, + ) -> bool: ... + + +def _configure_target_hidden_projection( + *, target_model, draft_model, is_deepseek_v4_draft: bool +) -> bool: + """Install the optional token-major projection before the target SP gather.""" + if is_deepseek_v4_draft: + return False + if not draft_model.supports_pre_gather_target_hidden_projection: + return False + if not isinstance(target_model, _SupportsDSparkTargetHiddenProjection): + return False + return bool( + target_model.set_dspark_target_hidden_projector( + draft_model.project_target_hidden, + num_context_features=int(draft_model.num_context_features), + ) + ) + + class DSparkWorkerV2(BaseSpecWorker): def __init__( @@ -227,6 +263,7 @@ class DSparkWorkerV2(BaseSpecWorker): ), lm_head=lm_head, ) + self._target_hidden_projection_enabled = False self._verify_planner = DSparkVerifyPlanner( draft_model=self.draft_model, @@ -377,6 +414,16 @@ class DSparkWorkerV2(BaseSpecWorker): def init_attention_backends(self): with self._draft_context(): self._draft_worker.init_attention_backends() + self._target_hidden_projection_enabled = _configure_target_hidden_projection( + target_model=self.target_worker.model_runner.model, + draft_model=self.draft_model, + is_deepseek_v4_draft=self._draft_is_moe, + ) + if self._target_hidden_projection_enabled and self.ps.tp_rank == 0: + logger.info( + "DSpark prefill target-hidden projection runs before " + "sequence-parallel gather." + ) self._need_mamba_verify_commit = mambaish_config( self.model_runner.model_config ) is not None and hasattr( @@ -487,6 +534,14 @@ class DSparkWorkerV2(BaseSpecWorker): batch_output = self.target_worker.forward_batch_generation( batch, capture_hidden_mode=CaptureHiddenMode.FULL ) + # BCG replay skips model-side Python, so re-evaluate the same pure predicate. + target_hidden_is_projected = ( + self._target_hidden_projection_enabled + and self.target_worker.model_runner.model.should_project_dspark_target_hidden( + forward_mode=batch.forward_mode, + capture_hidden_mode=CaptureHiddenMode.FULL, + ) + ) logits_output = batch_output.logits_output next_token_ids = batch_output.next_token_ids self._tp_sync.sync(SpecTpSyncSite.DSPARK_TARGET, next_token_ids) @@ -542,6 +597,7 @@ class DSparkWorkerV2(BaseSpecWorker): positions=positions, state_slot=state_slot, final_pos=final_pos, + target_hidden_is_projected=target_hidden_is_projected, ) # Avoid copying large hidden-state buffers to CPU in overlap scheduling. logits_output.hidden_states = None diff --git a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py index 7ede4fe7e..d3c3bb1c1 100644 --- a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py +++ b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py @@ -121,6 +121,13 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase): model_config=SimpleNamespace(context_len=8192, num_hidden_layers=1), layer_info=SimpleNamespace(start_layer=0, end_layer=1), req_to_token_pool=SimpleNamespace(size=1), + get_cuda_graph_layers=lambda _layer_model: ( + [object()], + [], + [], + [], + [None], + ), ) language_model = SimpleNamespace(layers=[object()]) @@ -129,11 +136,6 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase): patch.object( graph_setup, "resolve_language_model", return_value=language_model ), - patch.object( - graph_setup, - "compute_attention_and_moe_layers", - return_value=([object()], [], [], [], [None]), - ), patch.object( graph_setup, "get_available_gpu_memory", diff --git a/test/registered/unit/spec/test_dspark_target_hidden_projection.py b/test/registered/unit/spec/test_dspark_target_hidden_projection.py new file mode 100644 index 000000000..3cc4c2089 --- /dev/null +++ b/test/registered/unit/spec/test_dspark_target_hidden_projection.py @@ -0,0 +1,81 @@ +import unittest +from types import MethodType, SimpleNamespace +from unittest import mock + +import torch + +from sglang.srt.layers.aux_hidden_states import pack_aux_hidden_states +from sglang.srt.models.dspark import DSparkDraftMixin +from sglang.srt.speculative.dspark_components.dspark_kv_inject import ( + TargetHiddenKvInjector, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +class _Attention: + num_kv_heads = 1 + head_dim = 3 + + def __init__(self) -> None: + self.attn = SimpleNamespace(k_scale=0.5, v_scale=0.25) + self.input = None + + def kv_proj_only(self, hidden_states: torch.Tensor): + self.input = hidden_states + return hidden_states + 1, hidden_states + 2 + + def apply_k_norm(self, hidden_states: torch.Tensor) -> torch.Tensor: + return hidden_states + 3 + + def apply_k_rope( + self, _positions: torch.Tensor, hidden_states: torch.Tensor + ) -> torch.Tensor: + return hidden_states + 4 + + +class DSparkTargetHiddenProjectionTest(CustomTestCase): + def test_single_aux_hidden_state_is_returned_without_copy(self) -> None: + hidden_states = torch.empty(2, 3) + + self.assertIs(hidden_states, pack_aux_hidden_states([hidden_states])) + + def test_preprojected_hidden_is_not_projected_again(self) -> None: + attention = _Attention() + draft_model = SimpleNamespace( + layers=[SimpleNamespace(self_attn=attention)], + project_target_hidden=mock.Mock( + side_effect=AssertionError("projection must not run twice") + ), + _fused_kv_write_bundle=lambda _pool: None, + _stacked_ctx_kv_params=lambda: None, + ) + draft_model.write_target_hidden_kv = MethodType( + DSparkDraftMixin.write_target_hidden_kv, draft_model + ) + pool = SimpleNamespace(set_kv_buffer=mock.Mock()) + injector = TargetHiddenKvInjector( + draft_model=draft_model, + draft_model_runner=SimpleNamespace(token_to_kv_pool=pool), + model_runner=SimpleNamespace(device=torch.device("cpu")), + device=torch.device("cpu"), + verify_num_draft_tokens=2, + block_pos_offsets=torch.arange(2), + ) + projected_hidden = torch.arange(6, dtype=torch.float32).reshape(2, 3) + + injector.inject_target_hidden( + target_hidden=projected_hidden, + cache_loc=torch.arange(2), + positions=torch.arange(2), + target_hidden_is_projected=True, + ) + + draft_model.project_target_hidden.assert_not_called() + self.assertIs(attention.input, projected_hidden) + + +if __name__ == "__main__": + unittest.main()