Extract cuda-graph setup into a module (#31168)
This commit is contained in:
@@ -19,7 +19,6 @@ import contextlib
|
|||||||
import inspect
|
import inspect
|
||||||
import logging
|
import logging
|
||||||
import time
|
import time
|
||||||
from collections import defaultdict
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Optional, Union
|
from typing import Optional, Union
|
||||||
|
|
||||||
@@ -40,9 +39,6 @@ from sglang.srt.distributed import (
|
|||||||
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
||||||
maybe_init_shared_mooncake_transfer_engine,
|
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.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.dllm.config import DllmConfig
|
from sglang.srt.dllm.config import DllmConfig
|
||||||
from sglang.srt.elastic_ep.elastic_ep import (
|
from sglang.srt.elastic_ep.elastic_ep import (
|
||||||
@@ -68,8 +64,6 @@ from sglang.srt.eplb.expert_location import (
|
|||||||
set_global_expert_location_metadata,
|
set_global_expert_location_metadata,
|
||||||
)
|
)
|
||||||
from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
|
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.api import install_canary
|
||||||
from sglang.srt.kv_canary.runner.canary_manager import context_tuple
|
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
|
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,
|
KVCacheConfigurator,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
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 (
|
from sglang.srt.model_executor.cuda_graph_config import (
|
||||||
Backend,
|
|
||||||
Phase,
|
|
||||||
check_cuda_graph_backend,
|
|
||||||
cuda_graph_fully_disabled,
|
cuda_graph_fully_disabled,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import (
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
@@ -107,14 +97,17 @@ from sglang.srt.model_executor.forward_context import (
|
|||||||
forward_context,
|
forward_context,
|
||||||
has_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 import misc_utils
|
||||||
from sglang.srt.model_executor.model_runner_components.attention_backend_setup import (
|
from sglang.srt.model_executor.model_runner_components.attention_backend_setup import (
|
||||||
build_attention_backends,
|
build_attention_backends,
|
||||||
configure_aux_hidden_state_capture,
|
configure_aux_hidden_state_capture,
|
||||||
get_attention_backend,
|
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 (
|
from sglang.srt.model_executor.model_runner_components.kv_pool_runtime import (
|
||||||
compute_post_capture_kv_resize,
|
compute_post_capture_kv_resize,
|
||||||
is_post_capture_kv_active,
|
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 (
|
from sglang.srt.model_executor.model_runner_components.layer_setup import (
|
||||||
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 (
|
||||||
@@ -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.pool_configurator import MemoryPoolConfig
|
||||||
from sglang.srt.model_executor.runner import (
|
from sglang.srt.model_executor.runner import (
|
||||||
EagerRunner,
|
EagerRunner,
|
||||||
PrefillCudaGraphRunner,
|
|
||||||
get_batch_sizes_to_capture,
|
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.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.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
from sglang.srt.server_args import ( # noqa: F401 (re-export)
|
from sglang.srt.server_args import ( # noqa: F401 (re-export)
|
||||||
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS,
|
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS,
|
||||||
@@ -194,7 +184,6 @@ from sglang.srt.utils import (
|
|||||||
get_available_gpu_memory,
|
get_available_gpu_memory,
|
||||||
is_host_cpu_arm64,
|
is_host_cpu_arm64,
|
||||||
is_npu,
|
is_npu,
|
||||||
log_info_on_rank0,
|
|
||||||
numa_utils,
|
numa_utils,
|
||||||
require_gathered_buffer,
|
require_gathered_buffer,
|
||||||
reserve_rope_cache_for_long_sequences,
|
reserve_rope_cache_for_long_sequences,
|
||||||
@@ -745,58 +734,13 @@ class ModelRunner:
|
|||||||
self.decode_attention_backend_str = backends.decode_attention_backend_str
|
self.decode_attention_backend_str = backends.decode_attention_backend_str
|
||||||
|
|
||||||
def init_cuda_graphs(self, capture_decode_cuda_graph: bool = True):
|
def init_cuda_graphs(self, capture_decode_cuda_graph: bool = True):
|
||||||
"""Capture cuda graphs. Requires init_attention_backends() to have run.
|
capture = capture_cuda_graphs(
|
||||||
|
model_runner=self, capture_decode_cuda_graph=capture_decode_cuda_graph
|
||||||
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,
|
|
||||||
)
|
)
|
||||||
|
self.eager_runner = capture.eager_runner
|
||||||
if self.canary_manager is not None and not self.is_draft_worker:
|
self.prefill_cuda_graph_runner = capture.prefill_runner
|
||||||
self.canary_manager.mark_init_finished()
|
self.decode_cuda_graph_runner = capture.decode.runner
|
||||||
|
self.graph_mem_usage = capture.decode.graph_mem_usage
|
||||||
|
|
||||||
def init_routed_experts_capturer(self):
|
def init_routed_experts_capturer(self):
|
||||||
if self.is_draft_worker:
|
if self.is_draft_worker:
|
||||||
@@ -1093,196 +1037,18 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def init_decode_cuda_graph(self):
|
def init_decode_cuda_graph(self):
|
||||||
"""Capture device graphs."""
|
|
||||||
self.decode_cuda_graph_runner = None
|
self.decode_cuda_graph_runner = None
|
||||||
self.graph_mem_usage = 0
|
self.graph_mem_usage = 0
|
||||||
|
capture = capture_decode_graph(model_runner=self)
|
||||||
if not self.is_generation:
|
self.decode_cuda_graph_runner = capture.runner
|
||||||
# TODO: Currently, cuda graph only captures decode steps, which only exists for generation models
|
self.graph_mem_usage = capture.graph_mem_usage
|
||||||
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."
|
|
||||||
)
|
|
||||||
|
|
||||||
def init_prefill_cuda_graph(self, force_for_draft_worker: bool = False):
|
def init_prefill_cuda_graph(self, force_for_draft_worker: bool = False):
|
||||||
"""Initialize prefill CUDA graph runner."""
|
|
||||||
self.prefill_cuda_graph_runner = None
|
self.prefill_cuda_graph_runner = None
|
||||||
|
self.prefill_cuda_graph_runner = capture_prefill_graph(
|
||||||
if check_cuda_graph_backend(Phase.PREFILL, Backend.DISABLED):
|
model_runner=self,
|
||||||
logger.info(
|
eager_runner=self.eager_runner,
|
||||||
"Disable prefill CUDA graph because cuda_graph_config "
|
force_for_draft_worker=force_for_draft_worker,
|
||||||
"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."
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_threads_binding(self):
|
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