Extract cuda-graph setup into a module (#31168)
This commit is contained in:
@@ -19,7 +19,6 @@ import contextlib
|
||||
import inspect
|
||||
import logging
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Union
|
||||
|
||||
@@ -40,9 +39,6 @@ from sglang.srt.distributed import (
|
||||
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
||||
maybe_init_shared_mooncake_transfer_engine,
|
||||
)
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
prealloc_symmetric_memory_pool,
|
||||
)
|
||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||
from sglang.srt.dllm.config import DllmConfig
|
||||
from sglang.srt.elastic_ep.elastic_ep import (
|
||||
@@ -68,8 +64,6 @@ from sglang.srt.eplb.expert_location import (
|
||||
set_global_expert_location_metadata,
|
||||
)
|
||||
from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
|
||||
from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner
|
||||
from sglang.srt.hardware_backend.xpu.graph_runner.xpu_graph_runner import XPUGraphRunner
|
||||
from sglang.srt.kv_canary.api import install_canary
|
||||
from sglang.srt.kv_canary.runner.canary_manager import context_tuple
|
||||
from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env
|
||||
@@ -91,11 +85,7 @@ from sglang.srt.mem_cache.kv_cache_configurator import (
|
||||
KVCacheConfigurator,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
||||
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
||||
from sglang.srt.model_executor.cuda_graph_config import (
|
||||
Backend,
|
||||
Phase,
|
||||
check_cuda_graph_backend,
|
||||
cuda_graph_fully_disabled,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import (
|
||||
@@ -107,14 +97,17 @@ from sglang.srt.model_executor.forward_context import (
|
||||
forward_context,
|
||||
has_forward_context,
|
||||
)
|
||||
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 import misc_utils
|
||||
from sglang.srt.model_executor.model_runner_components.attention_backend_setup import (
|
||||
build_attention_backends,
|
||||
configure_aux_hidden_state_capture,
|
||||
get_attention_backend,
|
||||
)
|
||||
from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import (
|
||||
capture_cuda_graphs,
|
||||
capture_decode_graph,
|
||||
capture_prefill_graph,
|
||||
)
|
||||
from sglang.srt.model_executor.model_runner_components.kv_pool_runtime import (
|
||||
compute_post_capture_kv_resize,
|
||||
is_post_capture_kv_active,
|
||||
@@ -122,7 +115,6 @@ from sglang.srt.model_executor.model_runner_components.kv_pool_runtime import (
|
||||
from sglang.srt.model_executor.model_runner_components.layer_setup import (
|
||||
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 (
|
||||
@@ -160,12 +152,10 @@ from sglang.srt.model_executor.model_runner_components.weight_updater import (
|
||||
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
||||
from sglang.srt.model_executor.runner import (
|
||||
EagerRunner,
|
||||
PrefillCudaGraphRunner,
|
||||
get_batch_sizes_to_capture,
|
||||
)
|
||||
from sglang.srt.model_loader.utils import resolve_language_model
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_flags, get_server_args
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.server_args import ( # noqa: F401 (re-export)
|
||||
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS,
|
||||
@@ -194,7 +184,6 @@ from sglang.srt.utils import (
|
||||
get_available_gpu_memory,
|
||||
is_host_cpu_arm64,
|
||||
is_npu,
|
||||
log_info_on_rank0,
|
||||
numa_utils,
|
||||
require_gathered_buffer,
|
||||
reserve_rope_cache_for_long_sequences,
|
||||
@@ -745,58 +734,13 @@ class ModelRunner:
|
||||
self.decode_attention_backend_str = backends.decode_attention_backend_str
|
||||
|
||||
def init_cuda_graphs(self, capture_decode_cuda_graph: bool = True):
|
||||
"""Capture cuda graphs. Requires init_attention_backends() to have run.
|
||||
|
||||
Spec draft runners pass capture_decode_cuda_graph=False
|
||||
because they capture their own decode-style graphs separately.
|
||||
"""
|
||||
|
||||
self.graph_shared_output = GraphSharedOutput.create_for_model_runner(self)
|
||||
|
||||
# The eager (no-cuda-graph) phase runner, built AFTER the attention
|
||||
# backend so its __init__ can warm up kernels (run-once) and allocate the
|
||||
# fixed-max static buffer — both before the cuda-graph runners, so that
|
||||
# buffer is canonical in the shared pool and the cg runners coalesce onto
|
||||
# it. Always built: it serves both the fully-disabled case (decode/prefill
|
||||
# runners point at it) and the eager fallback when a cg runner can't run a
|
||||
# batch.
|
||||
self.eager_runner = EagerRunner(self)
|
||||
|
||||
# cuda-graph capture: prefill before decode, so both coalesce onto the
|
||||
# eager buffer allocated above. (init_prefill_cuda_graph routes prefill
|
||||
# to the eager runner when the prefill graph is disabled.)
|
||||
self.init_prefill_cuda_graph()
|
||||
|
||||
self.decode_cuda_graph_runner = None
|
||||
self.graph_mem_usage = 0
|
||||
|
||||
if capture_decode_cuda_graph:
|
||||
if self.device in ("cuda", "musa", "cpu", "npu", "xpu"):
|
||||
self.init_decode_cuda_graph()
|
||||
elif (
|
||||
current_platform.is_out_of_tree()
|
||||
and current_platform.support_cuda_graph()
|
||||
):
|
||||
self.init_decode_cuda_graph()
|
||||
else:
|
||||
self.decode_cuda_graph_runner = self.eager_runner
|
||||
|
||||
# Register forward hooks AFTER cuda-graph capture so their tensor ops are
|
||||
# not traced into any captured graph — capture stays hook-free and hooks
|
||||
# fire only on the eager forward path (capture replay never runs Python
|
||||
# hooks anyway).
|
||||
if self.server_args.forward_hooks:
|
||||
register_forward_hooks(self.model, self.server_args.forward_hooks)
|
||||
|
||||
prealloc_symmetric_memory_pool(
|
||||
is_draft_worker=self.is_draft_worker,
|
||||
enable_symm_mem=self.server_args.enable_symm_mem,
|
||||
device=self.device,
|
||||
forward_stream=self.forward_stream,
|
||||
capture = capture_cuda_graphs(
|
||||
model_runner=self, capture_decode_cuda_graph=capture_decode_cuda_graph
|
||||
)
|
||||
|
||||
if self.canary_manager is not None and not self.is_draft_worker:
|
||||
self.canary_manager.mark_init_finished()
|
||||
self.eager_runner = capture.eager_runner
|
||||
self.prefill_cuda_graph_runner = capture.prefill_runner
|
||||
self.decode_cuda_graph_runner = capture.decode.runner
|
||||
self.graph_mem_usage = capture.decode.graph_mem_usage
|
||||
|
||||
def init_routed_experts_capturer(self):
|
||||
if self.is_draft_worker:
|
||||
@@ -1093,196 +1037,18 @@ class ModelRunner:
|
||||
)
|
||||
|
||||
def init_decode_cuda_graph(self):
|
||||
"""Capture device graphs."""
|
||||
self.decode_cuda_graph_runner = None
|
||||
self.graph_mem_usage = 0
|
||||
|
||||
if not self.is_generation:
|
||||
# TODO: Currently, cuda graph only captures decode steps, which only exists for generation models
|
||||
return
|
||||
|
||||
if self.server_args.model_impl.lower() == ModelImpl.MINDSPORE:
|
||||
return
|
||||
|
||||
if self.device != "cpu" and check_cuda_graph_backend(
|
||||
Phase.DECODE, Backend.DISABLED
|
||||
):
|
||||
return
|
||||
|
||||
if self.device == "cpu" and not get_flags().capture.enable_torch_compile:
|
||||
return
|
||||
|
||||
tic = time.perf_counter()
|
||||
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||
graph_backend = defaultdict(
|
||||
lambda: f"{current_platform.device_name} graph",
|
||||
{
|
||||
"cuda": "CUDA graph",
|
||||
"musa": "CUDA graph",
|
||||
"cpu": "CPU graph",
|
||||
"npu": "NPU graph",
|
||||
"xpu": "XPU graph",
|
||||
},
|
||||
)
|
||||
role = "draft" if self.is_draft_worker else "target"
|
||||
if self.spec_algorithm.is_speculative():
|
||||
capture_name = f"{role} verify"
|
||||
num_tokens_per_req = (
|
||||
self.spec_algorithm.get_num_tokens_per_req_for_target_verify(
|
||||
self.server_args.speculative_num_draft_tokens,
|
||||
self.is_draft_worker,
|
||||
)
|
||||
)
|
||||
else:
|
||||
capture_name = f"{role} decode"
|
||||
num_tokens_per_req = 1
|
||||
capture_bs, _ = get_batch_sizes_to_capture(self, num_tokens_per_req)
|
||||
decode_backend = self.server_args.cuda_graph_config.decode.backend
|
||||
logger.info(
|
||||
f"Capture {capture_name} {graph_backend[self.device]} begin. "
|
||||
f"backend={decode_backend}, num_tokens_per_req={num_tokens_per_req}, "
|
||||
f"bs={capture_bs}, avail mem={before_mem:.2f} GB"
|
||||
)
|
||||
|
||||
if current_platform.is_out_of_tree():
|
||||
GraphRunnerCls = current_platform.get_graph_runner_cls()
|
||||
self.decode_cuda_graph_runner = GraphRunnerCls(self)
|
||||
else:
|
||||
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
||||
DecodeCudaGraphRunner,
|
||||
)
|
||||
|
||||
graph_runners = defaultdict(
|
||||
lambda: DecodeCudaGraphRunner,
|
||||
{
|
||||
"cpu": CPUGraphRunner,
|
||||
"npu": NPUGraphRunner,
|
||||
"xpu": XPUGraphRunner,
|
||||
},
|
||||
)
|
||||
self.decode_cuda_graph_runner = graph_runners[self.device](self)
|
||||
|
||||
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||
self.graph_mem_usage = before_mem - after_mem
|
||||
logger.info(
|
||||
f"Capture {capture_name} {graph_backend[self.device]} end. "
|
||||
f"elapsed={time.perf_counter() - tic:.2f} s, "
|
||||
f"mem usage={self.graph_mem_usage:.2f} GB, avail mem={after_mem:.2f} GB."
|
||||
)
|
||||
capture = capture_decode_graph(model_runner=self)
|
||||
self.decode_cuda_graph_runner = capture.runner
|
||||
self.graph_mem_usage = capture.graph_mem_usage
|
||||
|
||||
def init_prefill_cuda_graph(self, force_for_draft_worker: bool = False):
|
||||
"""Initialize prefill CUDA graph runner."""
|
||||
self.prefill_cuda_graph_runner = None
|
||||
|
||||
if check_cuda_graph_backend(Phase.PREFILL, Backend.DISABLED):
|
||||
logger.info(
|
||||
"Disable prefill CUDA graph because cuda_graph_config "
|
||||
"resolved prefill.backend='disabled' (e.g. via "
|
||||
"--cuda-graph-backend-prefill=disabled or auto-disable rules)."
|
||||
)
|
||||
# Prefill cuda graph disabled: route eager prefill through the
|
||||
# EagerRunner (its can_run_graph returns False, so _forward_raw's
|
||||
# extend branch falls through to the eager path).
|
||||
if not self.is_draft_worker:
|
||||
self.prefill_cuda_graph_runner = self.eager_runner
|
||||
return
|
||||
|
||||
# Draft models skip here during __init__; the eagle worker calls
|
||||
# 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
|
||||
|
||||
# Skip prefill CG for EAGLE target on tc_piecewise: that backend
|
||||
# captures CaptureHiddenMode.NULL while runtime requests FULL, so
|
||||
# the captured graph is dead, and capturing it perturbs FP4 /
|
||||
# TRTLLM-MoE state and corrupts decode replay (see #28386). BCG
|
||||
# captures FULL for EAGLE target in PrefillCudaGraphRunner.__init__
|
||||
# (restored from #25795), so it does NOT need this skip.
|
||||
if (
|
||||
self.spec_algorithm.is_eagle()
|
||||
and not self.is_draft_worker
|
||||
and not self.server_args.enable_return_hidden_states
|
||||
and not check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
|
||||
):
|
||||
logger.info(
|
||||
"Disable prefill CUDA graph for EAGLE target on tc_piecewise "
|
||||
"to avoid FP4/MoE decode-replay corruption (#28386)."
|
||||
)
|
||||
self.prefill_cuda_graph_runner = self.eager_runner
|
||||
return
|
||||
|
||||
# Resolve the decoder once. Some VLM wrappers (for example Kimi-VL)
|
||||
# expose it as ``language_model`` rather than ``model``.
|
||||
try:
|
||||
language_model = resolve_language_model(self.model)
|
||||
except AttributeError:
|
||||
logger.warning(
|
||||
"Disable prefill CUDA graph because the model is not a language model"
|
||||
)
|
||||
return
|
||||
|
||||
# Disable prefill CUDA graph for non capture size
|
||||
if not self.server_args.cuda_graph_config.prefill.bs:
|
||||
logger.warning(
|
||||
"Disable prefill CUDA graph because the capture size is not set"
|
||||
)
|
||||
return
|
||||
|
||||
# Collect attention layers and moe layers from the model. Keep a VLM
|
||||
# wrapper that exposes ``language_model`` unchanged: assigning it to
|
||||
# ``model`` would register a duplicate module alias and duplicate the
|
||||
# model's state-dict namespace.
|
||||
if hasattr(self.model, "model"):
|
||||
self.model.model = language_model
|
||||
|
||||
# Find the module that owns the decoder `layers`. Models wrap it at
|
||||
# varying depths: a direct text model exposes `.layers`, a CausalLM
|
||||
# wraps it as `.model.layers`, and some multimodal models add another
|
||||
# level (e.g. DeepSeek-OCR: OCR wrapper -> Deepseek*ForCausalLM ->
|
||||
# text model -> `.layers`). Descend the `.model` chain until we find it.
|
||||
layer_model = language_model
|
||||
while not hasattr(layer_model, "layers") and hasattr(layer_model, "model"):
|
||||
layer_model = layer_model.model
|
||||
|
||||
if not hasattr(layer_model, "layers"):
|
||||
logger.warning(
|
||||
"Disable prefill CUDA graph because the model does not have a 'layers' attribute"
|
||||
)
|
||||
return
|
||||
|
||||
self.attention_layers, self.moe_layers, self.moe_fusions, self.dsa_indexers = (
|
||||
compute_attention_and_moe_layers(layer_model)
|
||||
)
|
||||
|
||||
if len(self.attention_layers) < self.model_config.num_hidden_layers:
|
||||
# TODO(yuwei): support Non-Standard GQA
|
||||
log_info_on_rank0(
|
||||
logger,
|
||||
"Disable prefill CUDA graph because some layers do not apply Standard GQA",
|
||||
)
|
||||
return
|
||||
|
||||
tic = time.perf_counter()
|
||||
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||
prefill_backend = self.server_args.cuda_graph_config.prefill.backend
|
||||
role = "draft" if self.is_draft_worker else "target"
|
||||
capture_name = f"{role} prefill"
|
||||
capture_num_tokens = sorted(self.server_args.cuda_graph_config.prefill.bs)
|
||||
logger.info(
|
||||
f"Capture {capture_name} CUDA graph begin. "
|
||||
f"backend={prefill_backend}, num_tokens={capture_num_tokens}, "
|
||||
f"avail mem={before_mem:.2f} GB"
|
||||
)
|
||||
|
||||
self.prefill_cuda_graph_runner = PrefillCudaGraphRunner(self)
|
||||
|
||||
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||
mem_usage = before_mem - after_mem
|
||||
logger.info(
|
||||
f"Capture {capture_name} CUDA graph end. "
|
||||
f"elapsed={time.perf_counter() - tic:.2f} s, "
|
||||
f"mem usage={mem_usage:.2f} GB, avail mem={after_mem:.2f} GB."
|
||||
self.prefill_cuda_graph_runner = capture_prefill_graph(
|
||||
model_runner=self,
|
||||
eager_runner=self.eager_runner,
|
||||
force_for_draft_worker=force_for_draft_worker,
|
||||
)
|
||||
|
||||
def init_threads_binding(self):
|
||||
|
||||
@@ -0,0 +1,314 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import msgspec
|
||||
|
||||
from sglang.srt.configs.model_config import ModelImpl
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
prealloc_symmetric_memory_pool,
|
||||
)
|
||||
from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner
|
||||
from sglang.srt.hardware_backend.xpu.graph_runner.xpu_graph_runner import XPUGraphRunner
|
||||
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
||||
from sglang.srt.model_executor.cuda_graph_config import (
|
||||
Backend,
|
||||
Phase,
|
||||
check_cuda_graph_backend,
|
||||
)
|
||||
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,
|
||||
get_batch_sizes_to_capture,
|
||||
)
|
||||
from sglang.srt.model_loader.utils import resolve_language_model
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
from sglang.srt.utils import get_available_gpu_memory, log_info_on_rank0
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.model_executor.runner.base_runner import BaseRunner
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DecodeGraphCapture(msgspec.Struct, frozen=True, kw_only=True):
|
||||
runner: Optional[BaseRunner]
|
||||
graph_mem_usage: float
|
||||
|
||||
|
||||
class CudaGraphsCapture(msgspec.Struct, frozen=True, kw_only=True):
|
||||
eager_runner: EagerRunner
|
||||
prefill_runner: Optional[BaseRunner]
|
||||
decode: DecodeGraphCapture
|
||||
|
||||
|
||||
def capture_cuda_graphs(
|
||||
*, model_runner: ModelRunner, capture_decode_cuda_graph: bool = True
|
||||
) -> CudaGraphsCapture:
|
||||
"""Capture cuda graphs. Requires init_attention_backends() to have run.
|
||||
|
||||
Spec draft runners pass capture_decode_cuda_graph=False
|
||||
because they capture their own decode-style graphs separately.
|
||||
|
||||
"""
|
||||
|
||||
model_runner.graph_shared_output = GraphSharedOutput.create_for_model_runner(
|
||||
model_runner
|
||||
)
|
||||
|
||||
# The eager (no-cuda-graph) phase runner, built AFTER the attention
|
||||
# backend so its __init__ can warm up kernels (run-once) and allocate the
|
||||
# fixed-max static buffer — both before the cuda-graph runners, so that
|
||||
# buffer is canonical in the shared pool and the cg runners coalesce onto
|
||||
# it. Always built: it serves both the fully-disabled case (decode/prefill
|
||||
# runners point at it) and the eager fallback when a cg runner can't run a
|
||||
# batch.
|
||||
eager_runner = EagerRunner(model_runner)
|
||||
|
||||
# cuda-graph capture: prefill before decode, so both coalesce onto the
|
||||
# eager buffer allocated above. (capture_prefill_graph routes prefill
|
||||
# to the eager runner when the prefill graph is disabled.)
|
||||
prefill_runner = capture_prefill_graph(
|
||||
model_runner=model_runner, eager_runner=eager_runner
|
||||
)
|
||||
|
||||
decode = DecodeGraphCapture(runner=None, graph_mem_usage=0)
|
||||
if capture_decode_cuda_graph:
|
||||
if model_runner.device in ("cuda", "musa", "cpu", "npu", "xpu"):
|
||||
decode = capture_decode_graph(model_runner=model_runner)
|
||||
elif (
|
||||
current_platform.is_out_of_tree() and current_platform.support_cuda_graph()
|
||||
):
|
||||
decode = capture_decode_graph(model_runner=model_runner)
|
||||
else:
|
||||
decode = DecodeGraphCapture(runner=eager_runner, graph_mem_usage=0)
|
||||
|
||||
# Register forward hooks AFTER cuda-graph capture so their tensor ops are
|
||||
# not traced into any captured graph — capture stays hook-free and hooks
|
||||
# fire only on the eager forward path (capture replay never runs Python
|
||||
# hooks anyway).
|
||||
if model_runner.server_args.forward_hooks:
|
||||
register_forward_hooks(
|
||||
model_runner.model, model_runner.server_args.forward_hooks
|
||||
)
|
||||
|
||||
prealloc_symmetric_memory_pool(
|
||||
is_draft_worker=model_runner.is_draft_worker,
|
||||
enable_symm_mem=model_runner.server_args.enable_symm_mem,
|
||||
device=model_runner.device,
|
||||
forward_stream=model_runner.forward_stream,
|
||||
)
|
||||
|
||||
if model_runner.canary_manager is not None and not model_runner.is_draft_worker:
|
||||
model_runner.canary_manager.mark_init_finished()
|
||||
|
||||
return CudaGraphsCapture(
|
||||
eager_runner=eager_runner, prefill_runner=prefill_runner, decode=decode
|
||||
)
|
||||
|
||||
|
||||
def capture_prefill_graph(
|
||||
*,
|
||||
model_runner: ModelRunner,
|
||||
eager_runner: EagerRunner,
|
||||
force_for_draft_worker: bool = False,
|
||||
) -> Optional[BaseRunner]:
|
||||
"""Initialize prefill CUDA graph runner."""
|
||||
|
||||
if check_cuda_graph_backend(Phase.PREFILL, Backend.DISABLED):
|
||||
logger.info(
|
||||
"Disable prefill CUDA graph because cuda_graph_config "
|
||||
"resolved prefill.backend='disabled' (e.g. via "
|
||||
"--cuda-graph-backend-prefill=disabled or auto-disable rules)."
|
||||
)
|
||||
# Prefill cuda graph disabled: route eager prefill through the
|
||||
# EagerRunner (its can_run_graph returns False, so _forward_raw's
|
||||
# extend branch falls through to the eager path).
|
||||
if not model_runner.is_draft_worker:
|
||||
return eager_runner
|
||||
return None
|
||||
|
||||
# Draft models skip here during __init__; the eagle worker calls
|
||||
# this method explicitly (force_for_draft_worker=True) after
|
||||
# init_lm_head so graphs capture the final embedding weights.
|
||||
if model_runner.is_draft_worker and not force_for_draft_worker:
|
||||
return None
|
||||
|
||||
# Skip prefill CG for EAGLE target on tc_piecewise: that backend
|
||||
# captures CaptureHiddenMode.NULL while runtime requests FULL, so
|
||||
# the captured graph is dead, and capturing it perturbs FP4 /
|
||||
# TRTLLM-MoE state and corrupts decode replay (see #28386). BCG
|
||||
# captures FULL for EAGLE target in PrefillCudaGraphRunner.__init__
|
||||
# (restored from #25795), so it does NOT need this skip.
|
||||
if (
|
||||
model_runner.spec_algorithm.is_eagle()
|
||||
and not model_runner.is_draft_worker
|
||||
and not model_runner.server_args.enable_return_hidden_states
|
||||
and not check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
|
||||
):
|
||||
logger.info(
|
||||
"Disable prefill CUDA graph for EAGLE target on tc_piecewise "
|
||||
"to avoid FP4/MoE decode-replay corruption (#28386)."
|
||||
)
|
||||
return eager_runner
|
||||
|
||||
# Resolve the decoder once. Some VLM wrappers (for example Kimi-VL)
|
||||
# expose it as ``language_model`` rather than ``model``.
|
||||
try:
|
||||
language_model = resolve_language_model(model_runner.model)
|
||||
except AttributeError:
|
||||
logger.warning(
|
||||
"Disable prefill CUDA graph because the model is not a language model"
|
||||
)
|
||||
return None
|
||||
|
||||
# Disable prefill CUDA graph for non capture size
|
||||
if not model_runner.server_args.cuda_graph_config.prefill.bs:
|
||||
logger.warning("Disable prefill CUDA graph because the capture size is not set")
|
||||
return None
|
||||
|
||||
# Collect attention layers and moe layers from the model. Keep a VLM
|
||||
# wrapper that exposes ``language_model`` unchanged: assigning it to
|
||||
# ``model`` would register a duplicate module alias and duplicate the
|
||||
# model's state-dict namespace.
|
||||
if hasattr(model_runner.model, "model"):
|
||||
model_runner.model.model = language_model
|
||||
|
||||
# Find the module that owns the decoder `layers`. Models wrap it at
|
||||
# varying depths: a direct text model exposes `.layers`, a CausalLM
|
||||
# wraps it as `.model.layers`, and some multimodal models add another
|
||||
# level (e.g. DeepSeek-OCR: OCR wrapper -> Deepseek*ForCausalLM ->
|
||||
# text model -> `.layers`). Descend the `.model` chain until we find it.
|
||||
layer_model = language_model
|
||||
while not hasattr(layer_model, "layers") and hasattr(layer_model, "model"):
|
||||
layer_model = layer_model.model
|
||||
|
||||
if not hasattr(layer_model, "layers"):
|
||||
logger.warning(
|
||||
"Disable prefill CUDA graph because the model does not have a 'layers' attribute"
|
||||
)
|
||||
return None
|
||||
|
||||
(
|
||||
model_runner.attention_layers,
|
||||
model_runner.moe_layers,
|
||||
model_runner.moe_fusions,
|
||||
model_runner.dsa_indexers,
|
||||
) = compute_attention_and_moe_layers(layer_model)
|
||||
|
||||
if len(model_runner.attention_layers) < model_runner.model_config.num_hidden_layers:
|
||||
# TODO(yuwei): support Non-Standard GQA
|
||||
log_info_on_rank0(
|
||||
logger,
|
||||
"Disable prefill CUDA graph because some layers do not apply Standard GQA",
|
||||
)
|
||||
return None
|
||||
|
||||
tic = time.perf_counter()
|
||||
before_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id)
|
||||
prefill_backend = model_runner.server_args.cuda_graph_config.prefill.backend
|
||||
role = "draft" if model_runner.is_draft_worker else "target"
|
||||
capture_name = f"{role} prefill"
|
||||
capture_num_tokens = sorted(model_runner.server_args.cuda_graph_config.prefill.bs)
|
||||
logger.info(
|
||||
f"Capture {capture_name} CUDA graph begin. "
|
||||
f"backend={prefill_backend}, num_tokens={capture_num_tokens}, "
|
||||
f"avail mem={before_mem:.2f} GB"
|
||||
)
|
||||
|
||||
prefill_runner = PrefillCudaGraphRunner(model_runner)
|
||||
|
||||
after_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id)
|
||||
mem_usage = before_mem - after_mem
|
||||
logger.info(
|
||||
f"Capture {capture_name} CUDA graph end. "
|
||||
f"elapsed={time.perf_counter() - tic:.2f} s, "
|
||||
f"mem usage={mem_usage:.2f} GB, avail mem={after_mem:.2f} GB."
|
||||
)
|
||||
return prefill_runner
|
||||
|
||||
|
||||
def capture_decode_graph(*, model_runner: ModelRunner) -> DecodeGraphCapture:
|
||||
"""Capture device graphs."""
|
||||
no_capture = DecodeGraphCapture(runner=None, graph_mem_usage=0)
|
||||
|
||||
if not model_runner.is_generation:
|
||||
# TODO: Currently, cuda graph only captures decode steps, which only exists for generation models
|
||||
return no_capture
|
||||
if model_runner.server_args.model_impl.lower() == ModelImpl.MINDSPORE:
|
||||
return no_capture
|
||||
if model_runner.device != "cpu" and check_cuda_graph_backend(
|
||||
Phase.DECODE, Backend.DISABLED
|
||||
):
|
||||
return no_capture
|
||||
if model_runner.device == "cpu" and not get_flags().capture.enable_torch_compile:
|
||||
return no_capture
|
||||
|
||||
tic = time.perf_counter()
|
||||
before_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id)
|
||||
graph_backend = defaultdict(
|
||||
lambda: f"{current_platform.device_name} graph",
|
||||
{
|
||||
"cuda": "CUDA graph",
|
||||
"musa": "CUDA graph",
|
||||
"cpu": "CPU graph",
|
||||
"npu": "NPU graph",
|
||||
"xpu": "XPU graph",
|
||||
},
|
||||
)
|
||||
role = "draft" if model_runner.is_draft_worker else "target"
|
||||
if model_runner.spec_algorithm.is_speculative():
|
||||
capture_name = f"{role} verify"
|
||||
num_tokens_per_req = (
|
||||
model_runner.spec_algorithm.get_num_tokens_per_req_for_target_verify(
|
||||
model_runner.server_args.speculative_num_draft_tokens,
|
||||
model_runner.is_draft_worker,
|
||||
)
|
||||
)
|
||||
else:
|
||||
capture_name = f"{role} decode"
|
||||
num_tokens_per_req = 1
|
||||
capture_bs, _ = get_batch_sizes_to_capture(model_runner, num_tokens_per_req)
|
||||
decode_backend = model_runner.server_args.cuda_graph_config.decode.backend
|
||||
logger.info(
|
||||
f"Capture {capture_name} {graph_backend[model_runner.device]} begin. "
|
||||
f"backend={decode_backend}, num_tokens_per_req={num_tokens_per_req}, "
|
||||
f"bs={capture_bs}, avail mem={before_mem:.2f} GB"
|
||||
)
|
||||
|
||||
if current_platform.is_out_of_tree():
|
||||
GraphRunnerCls = current_platform.get_graph_runner_cls()
|
||||
runner = GraphRunnerCls(model_runner)
|
||||
else:
|
||||
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
||||
DecodeCudaGraphRunner,
|
||||
)
|
||||
|
||||
graph_runners = defaultdict(
|
||||
lambda: DecodeCudaGraphRunner,
|
||||
{
|
||||
"cpu": CPUGraphRunner,
|
||||
"npu": NPUGraphRunner,
|
||||
"xpu": XPUGraphRunner,
|
||||
},
|
||||
)
|
||||
runner = graph_runners[model_runner.device](model_runner)
|
||||
|
||||
after_mem = get_available_gpu_memory(model_runner.device, model_runner.gpu_id)
|
||||
graph_mem_usage = before_mem - after_mem
|
||||
logger.info(
|
||||
f"Capture {capture_name} {graph_backend[model_runner.device]} end. "
|
||||
f"elapsed={time.perf_counter() - tic:.2f} s, "
|
||||
f"mem usage={graph_mem_usage:.2f} GB, avail mem={after_mem:.2f} GB."
|
||||
)
|
||||
return DecodeGraphCapture(runner=runner, graph_mem_usage=graph_mem_usage)
|
||||
Reference in New Issue
Block a user