Improve CUDA graph and speculative execution output handling (#37329)

Co-authored-by: jiayisuse <jiayisuse@fb.com>
Co-authored-by: Yinghai Lu <yinghai@meta.com>
Co-authored-by: Hao Zhang <zhisbug@users.noreply.github.com>
Co-authored-by: Yichao Fu <yichaofu@meta.com>
This commit is contained in:
Lianmin Zheng
2026-09-02 17:25:27 -07:00
committed by GitHub
co-authored by jiayisuse Yinghai Lu Hao Zhang Yichao Fu
parent db1eb48651
commit 87d60a2229
8 changed files with 173 additions and 14 deletions
@@ -42,9 +42,10 @@ class AuxHiddenStatePacker:
def finalize(self) -> torch.Tensor: def finalize(self) -> torch.Tensor:
"""Return the packed buffer; callers guard the empty case on ``len()``.""" """Return the packed buffer; callers guard the empty case on ``len()``."""
assert ( if self._buffer is None or self._idx != self._num_captures:
self._buffer is not None and self._idx == self._num_captures raise RuntimeError(
), f"captured {self._idx} of {self._num_captures} aux hidden states" f"captured {self._idx} of {self._num_captures} aux hidden states"
)
return self._buffer return self._buffer
@@ -55,4 +56,6 @@ AuxHiddenStateAccumulator = Union[List[torch.Tensor], AuxHiddenStatePacker]
def pack_aux_hidden_states(aux_hidden_states: AuxHiddenStates) -> torch.Tensor: def pack_aux_hidden_states(aux_hidden_states: AuxHiddenStates) -> torch.Tensor:
if isinstance(aux_hidden_states, torch.Tensor): if isinstance(aux_hidden_states, torch.Tensor):
return aux_hidden_states return aux_hidden_states
if len(aux_hidden_states) == 1:
return aux_hidden_states[0]
return torch.cat(aux_hidden_states, dim=-1) return torch.cat(aux_hidden_states, dim=-1)
@@ -124,8 +124,10 @@ from sglang.srt.model_executor.model_runner_components.kv_pool_runtime import (
is_post_capture_kv_active, is_post_capture_kv_active,
) )
from sglang.srt.model_executor.model_runner_components.layer_setup import ( from sglang.srt.model_executor.model_runner_components.layer_setup import (
AttentionAndMoeLayers,
ModelLayerInfo, ModelLayerInfo,
adjust_hybrid_swa_layer_ids, adjust_hybrid_swa_layer_ids,
compute_attention_and_moe_layers,
resolve_layer_indices, resolve_layer_indices,
) )
from sglang.srt.model_executor.model_runner_components.load_model_utils import ( from sglang.srt.model_executor.model_runner_components.load_model_utils import (
@@ -1435,6 +1437,10 @@ class ModelRunner:
return DecodeCudaGraphRunner 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): def init_decode_cuda_graph(self):
self.decode_cuda_graph_runner = None self.decode_cuda_graph_runner = None
capture = capture_decode_graph(model_runner=self) capture = capture_decode_graph(model_runner=self)
@@ -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.graph_shared_output import GraphSharedOutput
from sglang.srt.model_executor.hook_manager import register_forward_hooks 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 ( from sglang.srt.model_executor.runner import (
EagerRunner, EagerRunner,
PrefillCudaGraphRunner, PrefillCudaGraphRunner,
@@ -444,7 +441,7 @@ def capture_prefill_graph(
model_runner.moe_fusions, model_runner.moe_fusions,
model_runner.dsa_indexers, model_runner.dsa_indexers,
model_runner.mha_companion_layers, 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.attention_layers,
model_runner.mha_companion_layers, model_runner.mha_companion_layers,
+8 -1
View File
@@ -484,6 +484,8 @@ _DSPARK_SKIPPED_WEIGHT_PREFIXES = ("lm_head.", "rotary_emb.")
class DSparkDraftMixin: class DSparkDraftMixin:
supports_pre_gather_target_hidden_projection = True
def __init__(self, config, quant_config=None, prefix: str = "") -> None: def __init__(self, config, quant_config=None, prefix: str = "") -> None:
super().__init__(config=config, quant_config=quant_config, prefix=prefix) super().__init__(config=config, quant_config=quant_config, prefix=prefix)
self._fused_kv_write_cache = None self._fused_kv_write_cache = None
@@ -737,8 +739,13 @@ class DSparkDraftMixin:
cache_loc: torch.Tensor, cache_loc: torch.Tensor,
cache_loc_2d: Optional[torch.Tensor] = None, cache_loc_2d: Optional[torch.Tensor] = None,
commit_lens: Optional[torch.Tensor] = None, commit_lens: Optional[torch.Tensor] = None,
target_hidden_is_projected: bool = False,
) -> None: ) -> 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) bundle = self._fused_kv_write_bundle(pool)
if bundle is not None: if bundle is not None:
@@ -42,6 +42,7 @@ class TargetHiddenKvInjector:
commit_lens: Optional[torch.Tensor] = None, commit_lens: Optional[torch.Tensor] = None,
state_slot: Optional[torch.Tensor] = None, state_slot: Optional[torch.Tensor] = None,
final_pos: Optional[torch.Tensor] = None, final_pos: Optional[torch.Tensor] = None,
target_hidden_is_projected: bool = False,
) -> None: ) -> None:
if target_hidden is None or target_hidden.numel() == 0: if target_hidden is None or target_hidden.numel() == 0:
return return
@@ -71,6 +72,11 @@ class TargetHiddenKvInjector:
pool = self.draft_model_runner.token_to_kv_pool pool = self.draft_model_runner.token_to_kv_pool
if hasattr(pool, "set_swa_key_buffer_radix_fused_norm_rope"): 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( self._inject_mla(
pool=pool, pool=pool,
target_hidden=target_hidden, target_hidden=target_hidden,
@@ -91,6 +97,7 @@ class TargetHiddenKvInjector:
cache_loc=cache_loc, cache_loc=cache_loc,
cache_loc_2d=cache_loc_2d, cache_loc_2d=cache_loc_2d,
commit_lens=commit_lens, commit_lens=commit_lens,
target_hidden_is_projected=target_hidden_is_projected,
) )
def _inject_mla( def _inject_mla(
@@ -1,7 +1,7 @@
import logging import logging
from contextlib import nullcontext from contextlib import nullcontext
from dataclasses import replace from dataclasses import replace
from typing import Optional from typing import Callable, Optional, Protocol, runtime_checkable
import torch 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.cuda_graph_config import Backend
from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode, CaptureHiddenMode,
ForwardMode,
compute_position, compute_position,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
@@ -88,6 +89,41 @@ logger = logging.getLogger(__name__)
_is_npu = is_npu() _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): class DSparkWorkerV2(BaseSpecWorker):
def __init__( def __init__(
@@ -227,6 +263,7 @@ class DSparkWorkerV2(BaseSpecWorker):
), ),
lm_head=lm_head, lm_head=lm_head,
) )
self._target_hidden_projection_enabled = False
self._verify_planner = DSparkVerifyPlanner( self._verify_planner = DSparkVerifyPlanner(
draft_model=self.draft_model, draft_model=self.draft_model,
@@ -377,6 +414,16 @@ class DSparkWorkerV2(BaseSpecWorker):
def init_attention_backends(self): def init_attention_backends(self):
with self._draft_context(): with self._draft_context():
self._draft_worker.init_attention_backends() 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._need_mamba_verify_commit = mambaish_config(
self.model_runner.model_config self.model_runner.model_config
) is not None and hasattr( ) is not None and hasattr(
@@ -487,6 +534,14 @@ class DSparkWorkerV2(BaseSpecWorker):
batch_output = self.target_worker.forward_batch_generation( batch_output = self.target_worker.forward_batch_generation(
batch, capture_hidden_mode=CaptureHiddenMode.FULL 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 logits_output = batch_output.logits_output
next_token_ids = batch_output.next_token_ids next_token_ids = batch_output.next_token_ids
self._tp_sync.sync(SpecTpSyncSite.DSPARK_TARGET, next_token_ids) self._tp_sync.sync(SpecTpSyncSite.DSPARK_TARGET, next_token_ids)
@@ -542,6 +597,7 @@ class DSparkWorkerV2(BaseSpecWorker):
positions=positions, positions=positions,
state_slot=state_slot, state_slot=state_slot,
final_pos=final_pos, final_pos=final_pos,
target_hidden_is_projected=target_hidden_is_projected,
) )
# Avoid copying large hidden-state buffers to CPU in overlap scheduling. # Avoid copying large hidden-state buffers to CPU in overlap scheduling.
logits_output.hidden_states = None logits_output.hidden_states = None
@@ -121,6 +121,13 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
model_config=SimpleNamespace(context_len=8192, num_hidden_layers=1), model_config=SimpleNamespace(context_len=8192, num_hidden_layers=1),
layer_info=SimpleNamespace(start_layer=0, end_layer=1), layer_info=SimpleNamespace(start_layer=0, end_layer=1),
req_to_token_pool=SimpleNamespace(size=1), req_to_token_pool=SimpleNamespace(size=1),
get_cuda_graph_layers=lambda _layer_model: (
[object()],
[],
[],
[],
[None],
),
) )
language_model = SimpleNamespace(layers=[object()]) language_model = SimpleNamespace(layers=[object()])
@@ -129,11 +136,6 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
patch.object( patch.object(
graph_setup, "resolve_language_model", return_value=language_model graph_setup, "resolve_language_model", return_value=language_model
), ),
patch.object(
graph_setup,
"compute_attention_and_moe_layers",
return_value=([object()], [], [], [], [None]),
),
patch.object( patch.object(
graph_setup, graph_setup,
"get_available_gpu_memory", "get_available_gpu_memory",
@@ -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()