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:
co-authored by
jiayisuse
Yinghai Lu
Hao Zhang
Yichao Fu
parent
db1eb48651
commit
87d60a2229
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user