Clean up prefill CUDA graph runner (#31654)

This commit is contained in:
Lianmin Zheng
2026-07-20 12:14:21 -07:00
committed by GitHub
parent 4e8eb1457b
commit 54aaedd76d
4 changed files with 358 additions and 316 deletions
@@ -357,8 +357,9 @@ class BaseRunner(ABC):
num_tokens_per_req = 1
if mr.spec_algorithm.is_speculative():
if mr.is_draft_worker:
if not mr.spec_algorithm.supports_target_verify_for_draft():
raise RuntimeError("This should not happen")
assert (
mr.spec_algorithm.supports_target_verify_for_draft()
), "This should not happen"
capture_forward_mode = ForwardMode.TARGET_VERIFY
num_tokens_per_req = mr.decode_num_tokens_per_req()
@@ -384,8 +385,6 @@ class BaseRunner(ABC):
)
)
seq_len_fill_value = mr.attn_backend.get_cuda_graph_seq_len_fill_value()
if get_flags().capture.enable_torch_compile:
set_torch_compile_config()
should_disable_torch_compile = not getattr(
@@ -402,33 +401,34 @@ class BaseRunner(ABC):
# NOTE: aux hidden state capture (eagle3/dflash) is already
# configured by init_aux_hidden_state_capture() in initialize().
require_mlp_tp_gather_ = require_mlp_tp_gather(mr.server_args)
if require_gathered_buffer(mr.server_args):
assert require_mlp_tp_gather_ or require_attn_tp_gather(mr.server_args)
# Token-axis buffer views and counters.
input_ids = buffers.input_ids[:num_tokens]
positions = buffers.positions[:num_tokens]
out_cache_loc = buffers.out_cache_loc[:num_tokens]
mrope_positions = buffers.mrope_positions[:, :num_tokens]
buffers.num_token_non_padded[...] = num_tokens
# Batch-axis buffer views.
req_pool_indices = buffers.req_pool_indices[:batch_size]
seq_lens = buffers.seq_lens[:batch_size]
seq_lens_cpu = buffers.seq_lens_cpu[:batch_size]
# Optional buffer views.
# Eager-reuse drops the logits buffer; only buffer sets that carry one slice it.
next_token_logits_buffer = (
buffers.next_token_logits_buffer[:num_tokens]
if buffers.next_token_logits_buffer is not None
else None
)
mrope_positions = buffers.mrope_positions[:, :num_tokens]
req_pool_indices = buffers.req_pool_indices[:batch_size]
seq_lens = buffers.seq_lens[:batch_size]
seq_lens_cpu = buffers.seq_lens_cpu[:batch_size]
encoder_lens = (
buffers.encoder_lens[:batch_size]
if buffers.encoder_lens is not None
else None
)
buffers.num_token_non_padded[...] = num_tokens
# For extend mode
if capture_forward_mode == ForwardMode.EXTEND:
seq_len_fill_value = mr.attn_backend.get_cuda_graph_seq_len_fill_value()
extend_prefix_lens_cpu = [0] * batch_size
extend_seq_lens_cpu = [seq_len_fill_value] * batch_size
extend_num_tokens = num_tokens
@@ -462,9 +462,15 @@ class BaseRunner(ABC):
{k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()}
)
# TP-gather requirements for global token metadata.
require_mlp_tp_gather_ = require_mlp_tp_gather(mr.server_args)
require_attn_tp_gather_ = require_attn_tp_gather(mr.server_args)
if require_gathered_buffer(mr.server_args):
assert require_mlp_tp_gather_ or require_attn_tp_gather_
if require_mlp_tp_gather_:
global_num_tokens_cpu = [num_tokens] * mr.server_args.dp_size
elif require_attn_tp_gather(mr.server_args):
elif require_attn_tp_gather_:
global_num_tokens_cpu = [num_tokens]
else:
global_num_tokens_cpu = None
@@ -480,6 +486,7 @@ class BaseRunner(ABC):
global_dp_buffer_len = None
global_num_tokens_cpu = None
# Speculative metadata and hidden-state capture mode.
spec_info = create_dummy_verify_input(
mr.spec_algorithm,
mr.server_args,
@@ -502,6 +509,7 @@ class BaseRunner(ABC):
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL
)
# Optional LoRA metadata.
if mr.server_args.enable_lora:
lora_ids = [None] * batch_size
else:
@@ -540,11 +548,11 @@ class BaseRunner(ABC):
global_forward_mode=capture_forward_mode,
lora_ids=lora_ids,
)
if buffers.ngram_embedding_info is not None:
forward_batch.ngram_embedding_info = buffers.ngram_embedding_info.slice(
batch_size
)
if lora_ids is not None:
mr.lora_manager.prepare_lora_batch(forward_batch)
@@ -552,6 +560,9 @@ class BaseRunner(ABC):
mr.attn_backend.init_forward_metadata(forward_batch)
def run_once():
# Reused dummy batches may carry DP-local lazy caches from a prior
# forward. Clear them, then refresh the process-wide DP buffer and
# MoE mode metadata read by model code during this standalone run.
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len(
global_dp_buffer_len,
@@ -14,22 +14,25 @@
"""PrefillCudaGraphRunner — runs the EXTEND phase under a pluggable backend.
Backend selection comes from cuda_graph_config.prefill:
- "tc_piecewise" — default, TcPiecewiseCudaGraphBackend: torch.compile
wraps the model; per-shape graphs live in
torch.compile's internal cache. Multi-batch supported.
- "breakable" — BreakableCudaGraphBackend: segmented capture (no
torch.compile). Captures with bs=1; rejects multi-req
prefill in can_run_graph.
- "full" — FullCudaGraphBackend: one whole-forward graph per
num_tokens bucket, captured with a fixed number of
request slots (cuda_graph_config.prefill.full_prefill_max_req). Replay
pads num_tokens up to the nearest bucket and pads
the request axis with zero-length sentinel requests;
bs > slots falls back to eager. Attention metadata
follows the decode-style 2-step contract
(init_forward_metadata_out_graph before capture/replay).
- "disabled" handled at the model_runner level — runner not
constructed.
- "breakable" — default on CUDA, BreakableCudaGraphBackend:
segmented capture (no torch.compile). Captures the
transformer body with one request slot, then replays it
with live batch metadata; multi-request prefill is
supported by running attention metadata and the
LM-head/logits tail outside the captured body.
- "full" — FullCudaGraphBackend: one graph per num_tokens bucket
for the captured transformer body. Capture uses
cuda_graph_config.prefill.full_prefill_max_req request
slots (auto-derived when unset); replay pads num_tokens
to the nearest bucket and pads unused request slots with
zero-length sentinels. bs > slots falls back to eager.
Attention metadata is refreshed out-of-graph against the
slot-padded batch before capture/replay.
- "tc_piecewise"TcPiecewiseCudaGraphBackend: torch.compile
wraps the model; per-shape compiled/captured pieces live
in torch.compile's internal cache. Multi-request prefill
is supported.
- "disabled" — handled at the model_runner level; runner not constructed.
"""
from __future__ import annotations
@@ -37,7 +40,7 @@ from __future__ import annotations
import copy
import inspect
import logging
import warnings
from contextlib import contextmanager
from typing import TYPE_CHECKING, Dict, Optional, Union
import torch
@@ -96,16 +99,17 @@ from sglang.srt.runtime_context import get_parallel
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
from sglang.srt.utils import (
get_available_gpu_memory,
get_bool_env_var,
is_hip,
is_npu,
require_attn_tp_gather,
require_gathered_buffer,
require_mlp_tp_gather,
)
from sglang.srt.utils.aiter import maybe_pre_warm_aiter_chip_info
if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.model_executor.model_runner import ModelRunner
# Suppress Dynamo warning about tracing through lru_cache-wrapped functions.
warnings.filterwarnings("ignore", message=".*lru_cache.*", module="torch._dynamo")
logger = logging.getLogger(__name__)
# A replay executes every padded token in its capture bucket. Sparse bucket
@@ -113,9 +117,6 @@ logger = logging.getLogger(__name__)
# model work than an exact-shape eager forward.
_MAX_PREFILL_CUDA_GRAPH_PADDING_FACTOR = 2
_is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
def prefill_failure_msg(backend_name: str) -> str:
"""Render PREFILL_CUDA_GRAPH_CAPTURE_FAILED_MSG with a backend-specific
@@ -137,9 +138,9 @@ def prefill_failure_msg(backend_name: str) -> str:
)
# Names of the static prefill input tensors a Breakable-backed prefill
# runner owns. Each is a 1-D int64 tensor of length max_bs; captured
# Breakable segments read from these stable addresses.
# Static prefill input tensors owned by captured-body backends (Breakable and
# Full). Each is a 1-D int64 tensor of length max_bs; captured graphs read
# these stable addresses during replay.
_PREFILL_STATIC_FIELDS = (
"seq_lens",
"extend_seq_lens",
@@ -149,10 +150,6 @@ _PREFILL_STATIC_FIELDS = (
"orig_seq_lens",
)
if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.model_executor.model_runner import ModelRunner
class PrefillCudaGraphRunner(BaseCudaGraphRunner):
"""Prefill-phase CUDA graph runner.
@@ -165,51 +162,48 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
def __init__(self, model_runner: ModelRunner):
super().__init__(model_runner)
# --- core state ------------------------------------------------
self.quant_config = getattr(self.model_runner.model, "quant_config", None)
self.is_multimodal = model_runner.is_multimodal
# --- model flags ----------------------------------------------
self.quant_config = getattr(model_runner.model, "quant_config", None)
self.is_multimodal = model_runner.model_config.is_multimodal
# Classification/reward forwards branch on return_pooled_hidden_states;
# capture must use the same flag value as replay for those models.
self.capture_return_pooled_hidden_states = not model_runner.is_generation
# --- bucket sizes ---------------------------------------------
# --- prefill graph config -------------------------------------
prefill_config = model_runner.server_args.cuda_graph_config.prefill
self.prefill_backend_name = prefill_config.backend
# bs in prefill carries the captured shape (token count for
# tc_piecewise) — one shape knob per phase.
capture_tokens = model_runner.server_args.cuda_graph_config.prefill.bs
capture_tokens = prefill_config.bs
assert capture_tokens is not None, "cuda_graph_config[prefill].bs is not set"
self.capture_num_tokens = sorted(capture_tokens)
self.max_num_tokens = (
max(self.capture_num_tokens) if self.capture_num_tokens else 8192
)
assert self.capture_num_tokens, "cuda_graph_config[prefill].bs is empty"
# --- runner bounds --------------------------------------------
self.max_num_tokens = max(self.capture_num_tokens)
self.max_bs = model_runner.req_to_token_pool.size
# --- capture modes --------------------------------------------
self.capture_forward_mode = ForwardMode.EXTEND
self.capture_hidden_mode = CaptureHiddenMode.NULL
# If returning hidden states is enabled, or if speculative prefill
# needs aux hidden states (DFLASH), capture the FULL variant up front.
# Ported from main #27468.
if (
# Hidden-state capture mode cases:
# - Breakable EAGLE draft: LAST.
# - Breakable EAGLE target: FULL.
# - Return-hidden-states or DFLASH: FULL.
# - Otherwise: NULL.
is_breakable_eagle = (
self.prefill_backend_name == Backend.BREAKABLE
and model_runner.spec_algorithm.is_eagle()
)
needs_full_hidden_states = (
model_runner.server_args.enable_return_hidden_states
or model_runner.spec_algorithm.is_dflash_family()
):
self.capture_hidden_mode = CaptureHiddenMode.FULL
# EAGLE captures FULL hidden states for the target and LAST for the
# draft (can_run_graph rejects on mismatch). BCG only; tc_piecewise
# EAGLE is routed to eager in ModelRunner.init_prefill_cuda_graph.
_cg_cfg = model_runner.server_args.cuda_graph_config
_prefill_backend_name = (
_cg_cfg.prefill.backend if _cg_cfg is not None else Backend.TC_PIECEWISE
)
self.prefill_backend_name = _prefill_backend_name
if (
_prefill_backend_name == Backend.BREAKABLE
and model_runner.spec_algorithm.is_eagle()
):
self.capture_hidden_mode = (
CaptureHiddenMode.LAST
if model_runner.is_draft_worker
else CaptureHiddenMode.FULL
)
if is_breakable_eagle and model_runner.is_draft_worker:
self.capture_hidden_mode = CaptureHiddenMode.LAST
elif is_breakable_eagle or needs_full_hidden_states:
self.capture_hidden_mode = CaptureHiddenMode.FULL
else:
self.capture_hidden_mode = CaptureHiddenMode.NULL
self.mamba_track_enabled = self._is_mamba_track_enabled()
@@ -255,25 +249,11 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
# --- backend ---------------------------------------------------
# When the backend is Breakable, captured segments need stable
# tensor addresses, so we own a set of static int64 buffers here
# and rebind them into capture-time dummy inputs / replay-time
# serving inputs below. Other backends don't need this.
# Initialize the slot to None BEFORE constructing the backend:
# TcPiecewise runs its compile pass during __init__ which calls
# _run_dummy_forward -> capture_prepare, and capture_prepare reads
# self._prefill_static_buffers and self.static_draft_hidden_states.
# self.layer_model has the same ordering requirement: _run_forward
# checks `self.layer_model is not None` to decide whether to call
# the inner stack or outer model.forward, and that check fires
# inside TcPiecewise's _run_compile_pass before backend resolution
# returns.
# TcPiecewise resolves by running a compile pass that calls back into
# capture_prepare / _run_forward, so these fields must exist first.
self._prefill_static_buffers: Optional[Dict[str, torch.Tensor]] = None
self.static_draft_hidden_states: Optional[torch.Tensor] = None
self.layer_model = None
# Set before resolve_prefill_backend — TcPiecewise's _run_compile_pass
# calls back into capture_prepare which reads this. Full overrides below
# once the backend type is known.
self._capture_req_slots = 1
# Same rationale: _run_compile_pass runs a dummy _run_forward before
# resolve_prefill_backend returns, and that forward reads
@@ -281,35 +261,42 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
# default False; the assignment below sets the real value once the
# backend type is known.
self._is_full_backend = False
# TcPiecewise does its compile pass during backend construction.
# Wrap only that path with the prefill CUDA graph failure hint.
try:
self.backend = resolve_prefill_backend(self)
except RuntimeError as e:
if _prefill_backend_name in (Backend.TC_PIECEWISE, Backend.BREAKABLE):
raise Exception(
f"Capture prefill CUDA graph failed: {e}\n"
f"{prefill_failure_msg(_prefill_backend_name)}"
)
raise
if self.prefill_backend_name != Backend.TC_PIECEWISE:
raise
raise RuntimeError(
f"Capture prefill CUDA graph failed: {e}\n"
f"{prefill_failure_msg(self.prefill_backend_name)}"
) from e
self._is_full_backend = isinstance(self.backend, FullCudaGraphBackend)
if self._is_full_backend:
max_req = (
model_runner.server_args.cuda_graph_config.prefill.full_prefill_max_req
)
max_req = prefill_config.full_prefill_max_req
if max_req is None:
# Auto: scale request slots with the chunked prefill size.
max_req = max(model_runner.server_args.chunked_prefill_size // 512, 1)
self._capture_req_slots = min(max_req, self.max_bs)
self._full_cg_seq_lens_cpu = (
torch.zeros((self._capture_req_slots,), dtype=torch.int64, device="cpu")
if self._is_full_backend
else None
)
if isinstance(self.backend, (BreakableCudaGraphBackend, FullCudaGraphBackend)):
self._full_cg_seq_lens_cpu = torch.zeros(
(self._capture_req_slots,), dtype=torch.int64, device="cpu"
)
with torch.device(self.device):
self._prefill_static_buffers = {
name: torch.zeros((self.max_bs,), dtype=torch.int64)
for name in _PREFILL_STATIC_FIELDS
}
elif isinstance(self.backend, BreakableCudaGraphBackend):
self._full_cg_seq_lens_cpu = None
with torch.device(self.device):
self._prefill_static_buffers = {
name: torch.zeros((self.max_bs,), dtype=torch.int64)
for name in _PREFILL_STATIC_FIELDS
}
else:
self._full_cg_seq_lens_cpu = None
# Static hidden_states buffer giving the captured graph a stable
# address; load_batch refreshes it from live spec_info at replay.
@@ -368,8 +355,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
)
# --- aiter chip info pre-warming (AMD) -------------------------
if _use_aiter:
self._pre_warm_aiter_chip_info()
maybe_pre_warm_aiter_chip_info()
# --- capture --------------------------------------------------
self.device_module.synchronize()
@@ -400,11 +386,15 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
self.model_runner.model_config.vocab_size, rows=rows
)
def _uses_eager_prefill_tail(self) -> bool:
return self.prefill_backend_name in (Backend.BREAKABLE, Backend.FULL)
def _prefill_logits_buffer_rows(self, forward_batch: ForwardBatch) -> int:
if not forward_batch.return_logprob:
return forward_batch.batch_size
if not isinstance(self.backend, BreakableCudaGraphBackend):
return forward_batch.batch_size
assert (
self._uses_eager_prefill_tail()
), "Prefill return_logprob requires an eager logits tail."
global_num_tokens = forward_batch.global_num_tokens_for_logprob_cpu
if global_num_tokens is not None:
@@ -433,40 +423,47 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
buf.copy_(local)
return buf
_aiter_chip_info_cached = False
def _get_layer_model_positions(self, forward_batch: ForwardBatch) -> torch.Tensor:
"""Mirror outer multimodal wrappers when BCG captures layer_model directly."""
if forward_batch.mrope_positions is None:
return forward_batch.positions
@classmethod
def _pre_warm_aiter_chip_info(cls):
"""Pre-populate aiter chip info env vars before CUDA graph capture.
model = self.model_runner.model
if getattr(model, "is_mrope_enabled", False):
return forward_batch.mrope_positions
aiter's get_cu_num_custom_op and get_gfx_custom_op call
subprocess.run(rocminfo) to query GPU info. During CUDA graph capture
the GPU context is locked, so rocminfo hangs indefinitely. Pre-calling
them here caches the results as environment variables so the subprocess
is never invoked during capture. Only runs once per process.
"""
if cls._aiter_chip_info_cached:
return
cls._aiter_chip_info_cached = True
language_model = getattr(model, "language_model", None)
if getattr(language_model, "is_mrope_enabled", False):
return forward_batch.mrope_positions
import os
return forward_batch.positions
try:
from aiter.jit.utils.chip_info import get_cu_num, get_gfx
if not os.environ.get("CU_NUM"):
cu_num = get_cu_num()
os.environ["CU_NUM"] = str(cu_num)
logger.info(f"Pre-warmed aiter CU_NUM={cu_num}")
if not os.environ.get("GPU_ARCHS"):
gfx = get_gfx()
os.environ["GPU_ARCHS"] = gfx
logger.info(f"Pre-warmed aiter GPU_ARCHS={gfx}")
except ImportError:
pass
except Exception as e:
logger.warning(f"Failed to pre-warm aiter chip info: {e}")
@contextmanager
def _prefill_forward_context(
self,
forward_batch: ForwardBatch,
*,
num_tokens: Optional[int] = None,
raw_num_tokens: Optional[int] = None,
):
with (
forward_context(
ForwardContext(attn_backend=self.model_runner.attn_backend)
),
set_tc_piecewise_forward_context(
forward_batch,
self.attention_layers,
self.quant_config,
self.moe_layers,
self.moe_fusions,
dsa_indexers=self.dsa_indexers,
mha_companion_layers=self.mha_companion_layers,
num_tokens=num_tokens,
raw_num_tokens=raw_num_tokens,
full_graph=self._is_full_backend,
),
):
yield
@torch.no_grad()
def _run_forward(self, forward_batch: ForwardBatch, num_tokens: int):
@@ -496,26 +493,9 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
)
set_is_extend_in_batch(False)
with (
forward_context(
ForwardContext(attn_backend=self.model_runner.attn_backend)
),
set_tc_piecewise_forward_context(
forward_batch,
self.attention_layers,
self.quant_config,
self.moe_layers,
self.moe_fusions,
dsa_indexers=self.dsa_indexers,
mha_companion_layers=self.mha_companion_layers,
# FULL backend: the whole transformer body is captured in one
# graph (no eager-break seams), so fusion gates keyed on BCG's
# split execution may fuse. (Both BCG and Full set layer_model
# -- the backend type is the discriminator, not layer_model.)
full_graph=self._is_full_backend,
),
):
if self.layer_model is not None:
with self._prefill_forward_context(forward_batch):
if self._uses_eager_prefill_tail():
# BCG / Full: capture the transformer body only.
positions = self._get_layer_model_positions(forward_batch)
return self.layer_model.forward(
forward_batch.input_ids,
@@ -523,27 +503,13 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
forward_batch,
forward_batch.input_embeds,
)
# tc_piecewise: compile/capture the outer model.forward path.
return self.model_runner.model.forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
)
def _get_layer_model_positions(self, forward_batch: ForwardBatch) -> torch.Tensor:
"""Mirror outer multimodal wrappers when BCG captures layer_model directly."""
if forward_batch.mrope_positions is None:
return forward_batch.positions
model = self.model_runner.model
if getattr(model, "is_mrope_enabled", False):
return forward_batch.mrope_positions
language_model = getattr(model, "language_model", None)
if getattr(language_model, "is_mrope_enabled", False):
return forward_batch.mrope_positions
return forward_batch.positions
def _run_dummy_forward(self, num_tokens: int) -> None:
"""Build a dummy ForwardBatch at this shape, init attn metadata,
run forward once. Used by TcPiecewiseCudaGraphBackend.prepare
@@ -732,15 +698,8 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
):
return False
num_tokens = len(forward_batch.input_ids)
if forward_batch.return_logprob and not isinstance(
self.backend, BreakableCudaGraphBackend
):
for start_len, seq_len in zip(
forward_batch.extend_logprob_start_lens_cpu,
forward_batch.extend_seq_lens_cpu,
):
if start_len is not None and start_len < seq_len:
return False
if forward_batch.return_logprob and not self._uses_eager_prefill_tail():
return False
if num_tokens > self.max_num_tokens:
return False
padded_num_tokens = self._pad_to_bucket(num_tokens, self.capture_num_tokens)
@@ -750,10 +709,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
# captured shape. The factor above only rejects replays whose padded
# model work is disproportionate to the useful token count.
#
# Multi-req replay is supported by BCG via the layer_model.forward
# monkey-patch in replay(): the captured bs=1 graph runs the
# transformer stack, then the outer model.forward runs
# logits_processor eagerly on top with live multi-req metadata.
# Multi-req replay is supported by body-capture backends via the
# layer_model.forward monkey-patch in replay(): the captured graph runs
# the transformer stack, then the outer model.forward runs
# logits_processor eagerly on top with live request metadata.
return True
def _build_capture_spec_info(self, num_tokens: int):
@@ -1087,6 +1046,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
or forward_batch.return_pooled_hidden_states
),
)
if self._is_full_backend:
forward_batch.next_token_logits_buffer = (
static_forward_batch.next_token_logits_buffer
)
if (
isinstance(self.backend, BreakableCudaGraphBackend)
@@ -1138,6 +1101,122 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
self._static_num_tokens = static_num_tokens
return static_forward_batch
def _execute_body_capture(
self,
forward_batch: ForwardBatch,
static_forward_batch: ForwardBatch,
static_num_tokens: int,
raw_num_tokens: int,
**kwargs,
):
# BCG / Full: replay the captured body, run the LM head +
# logits_processor eagerly.
shape_key = ShapeKey(size=self._static_num_tokens)
full_path = self._is_full_backend
static_n = self._static_num_tokens
ie_idx = self._input_embeds_arg_idx
def replay_layer_forward(*args, **layer_kwargs):
# The captured body graph reads activations from the static
# input_embeds slot. The outer model.forward (run eagerly)
# passes the live embeddings into layer_model.forward as the
# 4th positional arg (or input_embeds kwarg): for multimodal
# batches these are the composed text+vision embeds, for
# text-only batches they are get_input_embeddings()(input_ids).
# Copy them into the slot before replay so the graph sees the
# current request's embeddings (mirrors main's BCG closure).
if self.buffer_registry.has_slot("input_embeds"):
ie = layer_kwargs.get("input_embeds")
if ie is None and ie_idx is not None and len(args) > ie_idx:
ie = args[ie_idx]
if ie is not None:
self.buffer_registry.get_slot("input_embeds").slice_for(
1, static_n
)[: ie.shape[0]].copy_(ie)
hs = self.backend.replay(shape_key, static_forward_batch, **kwargs)
return hs[:raw_num_tokens] if full_path else hs
original_layer_forward = self.layer_model.forward
self.layer_model.forward = replay_layer_forward
# For Full, run the eager tail against the raw user-facing batch so it
# uses real request metadata instead of padded slots. BCG has no
# request-slot padding, so static_forward_batch is already the serving batch.
tail_batch = forward_batch if full_path else static_forward_batch
try:
with self._prefill_forward_context(
static_forward_batch,
num_tokens=static_num_tokens,
raw_num_tokens=raw_num_tokens,
):
return self.model_runner.model.forward(
tail_batch.input_ids,
tail_batch.positions,
tail_batch,
**kwargs,
)
finally:
self.layer_model.forward = original_layer_forward
def _execute_tc_piecewise(
self,
static_forward_batch: ForwardBatch,
static_num_tokens: int,
raw_num_tokens: int,
**kwargs,
):
with self._prefill_forward_context(
static_forward_batch,
num_tokens=static_num_tokens,
raw_num_tokens=raw_num_tokens,
):
return self.backend.replay(
ShapeKey(size=self._static_num_tokens),
static_forward_batch,
**kwargs,
)
def _trim_logits_output(
self, output: LogitsProcessorOutput
) -> LogitsProcessorOutput:
# Preserve mm_input_embeds for speculative decoding.
mm_input_embeds = None
if (
self.model_runner.spec_algorithm.is_speculative()
and output.mm_input_embeds is not None
):
mm_input_embeds = output.mm_input_embeds[: self.raw_num_tokens]
logits_rows = self.raw_bs if self._is_full_backend else self.raw_num_tokens
return LogitsProcessorOutput(
next_token_logits=(
output.next_token_logits[:logits_rows]
if output.next_token_logits is not None
else None
),
hidden_states=(
output.hidden_states[: self.raw_num_tokens]
if output.hidden_states is not None
else None
),
input_token_logprobs=output.input_token_logprobs,
input_top_logprobs_val=output.input_top_logprobs_val,
input_top_logprobs_idx=output.input_top_logprobs_idx,
input_token_ids_logprobs_val=output.input_token_ids_logprobs_val,
input_token_ids_logprobs_idx=output.input_token_ids_logprobs_idx,
mm_input_embeds=mm_input_embeds,
)
def _finalize_execute_output(
self, output
) -> Union[LogitsProcessorOutput, PPProxyTensors, EmbeddingPoolerOutput]:
if isinstance(output, LogitsProcessorOutput):
return self._trim_logits_output(output)
if isinstance(output, EmbeddingPoolerOutput):
return output
assert isinstance(output, PPProxyTensors)
raise NotImplementedError(
"PPProxyTensors is not supported in PrefillCudaGraphRunner yet."
)
def execute(
self, forward_batch: ForwardBatch, **kwargs
) -> Union[LogitsProcessorOutput, PPProxyTensors, EmbeddingPoolerOutput]:
@@ -1146,130 +1225,19 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
static_num_tokens = len(static_forward_batch.input_ids)
raw_num_tokens = self.raw_num_tokens
if self.layer_model is not None:
# BCG / Full: replay the captured body, run the LM head +
# logits_processor eagerly. For Full, slice hidden_states to
# raw_num_tokens and pass the raw forward_batch so the eager
# tail runs at raw_bs.
shape_key = ShapeKey(size=self._static_num_tokens)
full_path = self._is_full_backend
static_n = self._static_num_tokens
ie_idx = self._input_embeds_arg_idx
def replay_layer_forward(*args, **layer_kwargs):
# The captured BCG graph reads activations from the static
# input_embeds slot. The outer model.forward (run eagerly)
# passes the live embeddings into layer_model.forward as the
# 4th positional arg (or input_embeds kwarg): for multimodal
# batches these are the composed text+vision embeds, for
# text-only batches they are get_input_embeddings()(input_ids).
# Copy them into the slot before replay so the graph sees the
# current request's embeddings (mirrors main's BCG closure).
if self.buffer_registry.has_slot("input_embeds"):
ie = layer_kwargs.get("input_embeds")
if ie is None and ie_idx is not None and len(args) > ie_idx:
ie = args[ie_idx]
if ie is not None:
self.buffer_registry.get_slot("input_embeds").slice_for(
1, static_n
)[: ie.shape[0]].copy_(ie)
hs = self.backend.replay(shape_key, static_forward_batch, **kwargs)
return hs[:raw_num_tokens] if full_path else hs
original_layer_forward = self.layer_model.forward
self.layer_model.forward = replay_layer_forward
# For Full, run the eager LM head + logits_processor against
# the raw user-facing batch so the tail's req_slots-sized work and
# buffers collapse to real bs. For BCG, static_forward_batch
# IS the raw batch (bs=1 has no padding), so we keep the
# existing call unchanged.
tail_batch = forward_batch if full_path else static_forward_batch
try:
with (
forward_context(
ForwardContext(attn_backend=self.model_runner.attn_backend)
),
set_tc_piecewise_forward_context(
static_forward_batch,
self.attention_layers,
self.quant_config,
self.moe_layers,
self.moe_fusions,
dsa_indexers=self.dsa_indexers,
mha_companion_layers=self.mha_companion_layers,
num_tokens=static_num_tokens,
raw_num_tokens=raw_num_tokens,
full_graph=full_path,
),
):
output = self.model_runner.model.forward(
tail_batch.input_ids,
tail_batch.positions,
tail_batch,
**kwargs,
)
finally:
self.layer_model.forward = original_layer_forward
if self._uses_eager_prefill_tail():
output = self._execute_body_capture(
forward_batch,
static_forward_batch,
static_num_tokens,
raw_num_tokens,
**kwargs,
)
else:
# TC_PIECEWISE path. backend.replay calls the compiled
# outer model.forward directly (torch.compile handles
# multi-req via bs-invariant FX-traced kernels). Full/BCG use
# the captured-body path above; only tc_piecewise reaches here.
with (
forward_context(
ForwardContext(attn_backend=self.model_runner.attn_backend)
),
set_tc_piecewise_forward_context(
static_forward_batch,
self.attention_layers,
self.quant_config,
self.moe_layers,
self.moe_fusions,
dsa_indexers=self.dsa_indexers,
mha_companion_layers=self.mha_companion_layers,
num_tokens=static_num_tokens,
raw_num_tokens=raw_num_tokens,
),
):
output = self.backend.replay(
ShapeKey(size=self._static_num_tokens),
static_forward_batch,
**kwargs,
)
if isinstance(output, LogitsProcessorOutput):
# Preserve mm_input_embeds for speculative decoding.
mm_input_embeds = None
if (
self.model_runner.spec_algorithm.is_speculative()
and output.mm_input_embeds is not None
):
mm_input_embeds = output.mm_input_embeds[: self.raw_num_tokens]
logits_rows = (
self.raw_bs if self._is_full_backend else self.raw_num_tokens
)
return LogitsProcessorOutput(
next_token_logits=(
output.next_token_logits[:logits_rows]
if output.next_token_logits is not None
else None
),
hidden_states=(
output.hidden_states[: self.raw_num_tokens]
if output.hidden_states is not None
else None
),
input_token_logprobs=output.input_token_logprobs,
input_top_logprobs_val=output.input_top_logprobs_val,
input_top_logprobs_idx=output.input_top_logprobs_idx,
input_token_ids_logprobs_val=output.input_token_ids_logprobs_val,
input_token_ids_logprobs_idx=output.input_token_ids_logprobs_idx,
mm_input_embeds=mm_input_embeds,
)
elif isinstance(output, EmbeddingPoolerOutput):
return output
else:
assert isinstance(output, PPProxyTensors)
raise NotImplementedError(
"PPProxyTensors is not supported in PrefillCudaGraphRunner yet."
output = self._execute_tc_piecewise(
static_forward_batch,
static_num_tokens,
raw_num_tokens,
**kwargs,
)
return self._finalize_execute_output(output)
@@ -22,6 +22,7 @@ _compiled_fn reused for every shape.
from __future__ import annotations
import warnings
from contextlib import contextmanager
from typing import TYPE_CHECKING, Any, Callable, Optional
@@ -63,6 +64,10 @@ if TYPE_CHECKING:
_VALID_COMPILERS = ("eager", "inductor")
def _suppress_lru_cache_dynamo_warning() -> None:
warnings.filterwarnings("ignore", message=".*lru_cache.*", module="torch._dynamo")
def _toggle_multi_platform_ops(
model: torch.nn.Module, *, reverse: bool, num_tokens: int
) -> None:
@@ -95,6 +100,7 @@ class TcPiecewiseCudaGraphBackend(BaseCudaGraphBackend):
self._language_model: torch.nn.Module = getattr(
model_runner.model, "language_model", model_runner.model
)
_suppress_lru_cache_dynamo_warning()
self._run_compile_pass(cuda_graph_runner)
# model_runner.model.forward is the wrapper that builds LogitsProcessorOutput.
# The compiled trampoline is dispatched internally by it.
+57
View File
@@ -0,0 +1,57 @@
# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""AITER utility helpers."""
import logging
import os
from sglang.srt.utils.common import get_bool_env_var, is_hip
logger = logging.getLogger(__name__)
_aiter_chip_info_cached = False
def maybe_pre_warm_aiter_chip_info() -> None:
"""Pre-populate AITER chip info env vars before CUDA graph capture.
AITER chip info probes can shell out to rocminfo. During graph capture the
GPU context is locked, so pre-caching CU_NUM/GPU_ARCHS avoids a hang.
"""
if not (get_bool_env_var("SGLANG_USE_AITER") and is_hip()):
return
global _aiter_chip_info_cached
if _aiter_chip_info_cached:
return
_aiter_chip_info_cached = True
try:
from aiter.jit.utils.chip_info import get_cu_num, get_gfx
except ImportError:
logger.debug("Skip aiter chip info pre-warm because aiter is unavailable")
return
try:
if not os.environ.get("CU_NUM"):
cu_num = get_cu_num()
os.environ["CU_NUM"] = str(cu_num)
logger.info("Pre-warmed aiter CU_NUM=%s", cu_num)
if not os.environ.get("GPU_ARCHS"):
gfx = get_gfx()
os.environ["GPU_ARCHS"] = gfx
logger.info("Pre-warmed aiter GPU_ARCHS=%s", gfx)
except Exception as e:
logger.warning("Failed to pre-warm aiter chip info: %s", e)