Clean up prefill CUDA graph runner (#31654)
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user