Enable breakable CUDA graph for eagle (#25795)
This commit is contained in:
@@ -23,6 +23,7 @@ breaks are inserted eagerly via :func:`eager_on_graph` decorated callables
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import bisect
|
import bisect
|
||||||
|
import inspect
|
||||||
import logging
|
import logging
|
||||||
from typing import TYPE_CHECKING, Union
|
from typing import TYPE_CHECKING, Union
|
||||||
|
|
||||||
@@ -53,6 +54,7 @@ from sglang.srt.model_executor.cuda_graph_runner import (
|
|||||||
set_global_graph_memory_pool,
|
set_global_graph_memory_pool,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import (
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
|
CaptureHiddenMode,
|
||||||
PPProxyTensors,
|
PPProxyTensors,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
||||||
@@ -77,11 +79,6 @@ class BreakableCudaGraphRunner:
|
|||||||
graph capture of the eager kernel stream.
|
graph capture of the eager kernel stream.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# replay_prepare shares its buffer-population logic with the PCG runner —
|
|
||||||
# bind the method here without inheriting. __init__, capture, and replay
|
|
||||||
# diverge enough that inheritance would obscure more than it saves.
|
|
||||||
replay_prepare = PiecewiseCudaGraphRunner.replay_prepare
|
|
||||||
|
|
||||||
def __init__(self, model_runner: ModelRunner):
|
def __init__(self, model_runner: ModelRunner):
|
||||||
self.model_runner = model_runner
|
self.model_runner = model_runner
|
||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
@@ -103,6 +100,18 @@ class BreakableCudaGraphRunner:
|
|||||||
)
|
)
|
||||||
self.max_bs = model_runner.req_to_token_pool.size
|
self.max_bs = model_runner.req_to_token_pool.size
|
||||||
|
|
||||||
|
self.capture_hidden_mode = CaptureHiddenMode.NULL
|
||||||
|
if model_runner.server_args.enable_return_hidden_states:
|
||||||
|
self.capture_hidden_mode = CaptureHiddenMode.FULL
|
||||||
|
if (
|
||||||
|
model_runner.spec_algorithm is not None
|
||||||
|
and model_runner.spec_algorithm.is_eagle()
|
||||||
|
):
|
||||||
|
if model_runner.is_draft_worker:
|
||||||
|
self.capture_hidden_mode = CaptureHiddenMode.LAST
|
||||||
|
else:
|
||||||
|
self.capture_hidden_mode = CaptureHiddenMode.FULL
|
||||||
|
|
||||||
log_info_on_rank0(
|
log_info_on_rank0(
|
||||||
logger,
|
logger,
|
||||||
f"[BCG] Capture num tokens: {self.capture_num_tokens}",
|
f"[BCG] Capture num tokens: {self.capture_num_tokens}",
|
||||||
@@ -133,6 +142,16 @@ class BreakableCudaGraphRunner:
|
|||||||
type(language_model).__name__,
|
type(language_model).__name__,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
self.use_input_embeds = self.is_multimodal
|
||||||
|
if self.use_input_embeds:
|
||||||
|
sig = inspect.signature(self.layer_model.forward)
|
||||||
|
params = list(sig.parameters)
|
||||||
|
if "input_embeds" not in params:
|
||||||
|
raise ValueError(
|
||||||
|
f"layer_model.forward must accept 'input_embeds' for "
|
||||||
|
f"multimodal BCG, got params: {params}"
|
||||||
|
)
|
||||||
|
self._input_embeds_arg_idx = params.index("input_embeds")
|
||||||
|
|
||||||
# Memory pool
|
# Memory pool
|
||||||
if get_global_graph_memory_pool() is None:
|
if get_global_graph_memory_pool() is None:
|
||||||
@@ -178,6 +197,16 @@ class BreakableCudaGraphRunner:
|
|||||||
input_embeds = None
|
input_embeds = None
|
||||||
mrope_positions = None
|
mrope_positions = None
|
||||||
|
|
||||||
|
if model_runner.is_draft_worker:
|
||||||
|
from sglang.srt.speculative.eagle_utils import get_draft_hidden_dim
|
||||||
|
|
||||||
|
hidden_dim = get_draft_hidden_dim(model_runner)
|
||||||
|
self.static_draft_hidden_states = torch.zeros(
|
||||||
|
(self.max_num_tokens, hidden_dim),
|
||||||
|
dtype=model_runner.dtype,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
|
||||||
self.buffers = PrefillInputBuffers(
|
self.buffers = PrefillInputBuffers(
|
||||||
input_ids=input_ids,
|
input_ids=input_ids,
|
||||||
out_cache_loc=out_cache_loc,
|
out_cache_loc=out_cache_loc,
|
||||||
@@ -221,6 +250,7 @@ class BreakableCudaGraphRunner:
|
|||||||
forward_batch.input_ids,
|
forward_batch.input_ids,
|
||||||
forward_batch.positions,
|
forward_batch.positions,
|
||||||
forward_batch,
|
forward_batch,
|
||||||
|
input_embeds=forward_batch.input_embeds,
|
||||||
)
|
)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
@@ -234,11 +264,18 @@ class BreakableCudaGraphRunner:
|
|||||||
"""
|
"""
|
||||||
from sglang.srt.layers.dp_attention import DpPaddingMode
|
from sglang.srt.layers.dp_attention import DpPaddingMode
|
||||||
from sglang.srt.model_executor.forward_batch_info import (
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
CaptureHiddenMode,
|
|
||||||
ForwardBatch,
|
ForwardBatch,
|
||||||
ForwardMode,
|
ForwardMode,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
spec_info = None
|
||||||
|
if self.model_runner.is_draft_worker:
|
||||||
|
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
||||||
|
|
||||||
|
spec_info = EagleDraftInput(
|
||||||
|
hidden_states=self.static_draft_hidden_states[:num_tokens],
|
||||||
|
)
|
||||||
|
|
||||||
buffers = self.buffers
|
buffers = self.buffers
|
||||||
bs = 1
|
bs = 1
|
||||||
with torch.device(self.device):
|
with torch.device(self.device):
|
||||||
@@ -292,8 +329,8 @@ class BreakableCudaGraphRunner:
|
|||||||
buffers.mrope_positions[:, :num_tokens] if self.is_multimodal else None
|
buffers.mrope_positions[:, :num_tokens] if self.is_multimodal else None
|
||||||
),
|
),
|
||||||
spec_algorithm=None,
|
spec_algorithm=None,
|
||||||
spec_info=None,
|
spec_info=spec_info,
|
||||||
capture_hidden_mode=CaptureHiddenMode.NULL,
|
capture_hidden_mode=self.capture_hidden_mode,
|
||||||
num_token_non_padded=None,
|
num_token_non_padded=None,
|
||||||
global_forward_mode=ForwardMode.EXTEND,
|
global_forward_mode=ForwardMode.EXTEND,
|
||||||
lora_ids=None,
|
lora_ids=None,
|
||||||
@@ -341,6 +378,8 @@ class BreakableCudaGraphRunner:
|
|||||||
return False
|
return False
|
||||||
if forward_batch.forward_mode.is_target_verify():
|
if forward_batch.forward_mode.is_target_verify():
|
||||||
return False
|
return False
|
||||||
|
if forward_batch.capture_hidden_mode != self.capture_hidden_mode:
|
||||||
|
return False
|
||||||
if forward_batch.input_embeds is not None:
|
if forward_batch.input_embeds is not None:
|
||||||
return False
|
return False
|
||||||
if forward_batch.replace_embeds is not None:
|
if forward_batch.replace_embeds is not None:
|
||||||
@@ -374,6 +413,18 @@ class BreakableCudaGraphRunner:
|
|||||||
|
|
||||||
return graph, output
|
return graph, output
|
||||||
|
|
||||||
|
def replay_prepare(self, forward_batch, **kwargs):
|
||||||
|
# TODO: fix PiecewiseCudaGraphRunner to support draft workers as well.
|
||||||
|
static_forward_batch = PiecewiseCudaGraphRunner.replay_prepare(
|
||||||
|
self, forward_batch, **kwargs
|
||||||
|
)
|
||||||
|
if self.model_runner.is_draft_worker and forward_batch.spec_info is not None:
|
||||||
|
num_tokens = len(forward_batch.input_ids)
|
||||||
|
self.static_draft_hidden_states[:num_tokens].copy_(
|
||||||
|
forward_batch.spec_info.hidden_states
|
||||||
|
)
|
||||||
|
return static_forward_batch
|
||||||
|
|
||||||
def replay(
|
def replay(
|
||||||
self,
|
self,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
@@ -391,6 +442,22 @@ class BreakableCudaGraphRunner:
|
|||||||
# outer forward the captured hidden_states; logits_processor / pooler
|
# outer forward the captured hidden_states; logits_processor / pooler
|
||||||
# then runs eagerly on top with the live multi-req forward_batch.
|
# then runs eagerly on top with the live multi-req forward_batch.
|
||||||
def replay_layer_forward(*args, **layer_kwargs):
|
def replay_layer_forward(*args, **layer_kwargs):
|
||||||
|
ie = layer_kwargs.get("input_embeds") or (
|
||||||
|
args[self._input_embeds_arg_idx]
|
||||||
|
if self.use_input_embeds and len(args) > self._input_embeds_arg_idx
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
if self.use_input_embeds:
|
||||||
|
if ie is None:
|
||||||
|
raise ValueError("BCG replay expects input_embeds but got None")
|
||||||
|
self.buffers.input_embeds[:static_num_tokens].copy_(
|
||||||
|
ie[:static_num_tokens]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
if ie is not None:
|
||||||
|
raise ValueError(
|
||||||
|
"BCG replay got unexpected input_embeds on non-multimodal model"
|
||||||
|
)
|
||||||
captured_graph.replay()
|
captured_graph.replay()
|
||||||
return captured_hidden
|
return captured_hidden
|
||||||
|
|
||||||
@@ -416,11 +483,7 @@ class BreakableCudaGraphRunner:
|
|||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
self.layer_model.forward = original_layer_forward
|
self.layer_model.forward = original_layer_forward
|
||||||
|
|
||||||
if isinstance(output, LogitsProcessorOutput):
|
if isinstance(output, LogitsProcessorOutput):
|
||||||
# Slice trailing-padding off hidden_states; next_token_logits is
|
|
||||||
# bs-shaped from logits_processor (bs <= raw_num_tokens), so the
|
|
||||||
# slice is a no-op for that field but matches PCG's pattern.
|
|
||||||
return LogitsProcessorOutput(
|
return LogitsProcessorOutput(
|
||||||
next_token_logits=output.next_token_logits[: self.raw_num_tokens],
|
next_token_logits=output.next_token_logits[: self.raw_num_tokens],
|
||||||
hidden_states=(
|
hidden_states=(
|
||||||
|
|||||||
@@ -2787,7 +2787,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
f"mem usage={self.graph_mem_usage:.2f} GB. avail mem={after_mem:.2f} GB."
|
f"mem usage={self.graph_mem_usage:.2f} GB. avail mem={after_mem:.2f} GB."
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_piecewise_cuda_graphs(self):
|
def init_piecewise_cuda_graphs(self, force_for_draft_worker: bool = False):
|
||||||
"""Initialize piecewise CUDA graph runner."""
|
"""Initialize piecewise CUDA graph runner."""
|
||||||
self.piecewise_cuda_graph_runner = None
|
self.piecewise_cuda_graph_runner = None
|
||||||
|
|
||||||
@@ -2797,8 +2797,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
# Draft models use decode CUDA graphs, not PCG
|
# Draft models skip here during __init__; the eagle worker calls
|
||||||
if self.is_draft_worker:
|
# this method explicitly (force_for_draft_worker=True) after
|
||||||
|
# init_lm_head so graphs capture the final embedding weights.
|
||||||
|
if self.is_draft_worker and not force_for_draft_worker:
|
||||||
return
|
return
|
||||||
|
|
||||||
# Disable piecewise CUDA graph for non-language models
|
# Disable piecewise CUDA graph for non-language models
|
||||||
|
|||||||
@@ -3743,11 +3743,7 @@ class ServerArgs:
|
|||||||
self.disaggregation_transfer_backend != "fake"
|
self.disaggregation_transfer_backend != "fake"
|
||||||
), "Prefill server does not support 'fake' as the transfer backend"
|
), "Prefill server does not support 'fake' as the transfer backend"
|
||||||
|
|
||||||
if self.disable_piecewise_cuda_graph:
|
|
||||||
self.disable_cuda_graph = True
|
self.disable_cuda_graph = True
|
||||||
logger.warning(
|
|
||||||
"Cuda graph is disabled for prefill server when piecewise cuda graph is not enabled."
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.disaggregation_mode in ("prefill", "decode"):
|
if self.disaggregation_mode in ("prefill", "decode"):
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
@@ -228,3 +229,20 @@ def verify_tree_greedy_func(
|
|||||||
target_predict=target_predict,
|
target_predict=target_predict,
|
||||||
)
|
)
|
||||||
return predicts, accept_index, accept_token_num
|
return predicts, accept_index, accept_token_num
|
||||||
|
|
||||||
|
|
||||||
|
def get_draft_hidden_dim(model_runner: ModelRunner) -> int:
|
||||||
|
"""Derive the hidden dimension of target hidden states fed to the draft model."""
|
||||||
|
hf_config = model_runner.model_config.hf_config
|
||||||
|
eagle_config = getattr(hf_config, "eagle_config", {})
|
||||||
|
use_aux = eagle_config.get("use_aux_hidden_state", False)
|
||||||
|
spec_algorithm = model_runner.spec_algorithm
|
||||||
|
|
||||||
|
if spec_algorithm is not None and spec_algorithm.is_eagle3() and use_aux:
|
||||||
|
base = getattr(hf_config, "target_hidden_size", None)
|
||||||
|
if base is None:
|
||||||
|
base = model_runner.model_config.hidden_size
|
||||||
|
layer_ids = eagle_config.get("eagle_aux_hidden_state_layer_ids", [])
|
||||||
|
num_aux = max(len(layer_ids), 1)
|
||||||
|
return base * num_aux
|
||||||
|
return model_runner.model_config.spec_hidden_size
|
||||||
|
|||||||
@@ -190,6 +190,10 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
speculative_moe_a2a_backend_context(),
|
speculative_moe_a2a_backend_context(),
|
||||||
):
|
):
|
||||||
self.init_attention_backend()
|
self.init_attention_backend()
|
||||||
|
if server_args.enable_breakable_cuda_graph:
|
||||||
|
self.draft_runner.init_piecewise_cuda_graphs(
|
||||||
|
force_for_draft_worker=True
|
||||||
|
)
|
||||||
self.init_cuda_graphs()
|
self.init_cuda_graphs()
|
||||||
|
|
||||||
self.tree_mask_mode = TreeMaskMode.FULL_MASK
|
self.tree_mask_mode = TreeMaskMode.FULL_MASK
|
||||||
|
|||||||
Reference in New Issue
Block a user