refactor(runner): move kernel warmup into the shared runner lifecycle (warmup()) (#28739)

Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-06-19 13:13:13 -07:00
committed by GitHub
co-authored by Claude Opus 4.8
parent d705a91de1
commit 856b0dc74b
16 changed files with 678 additions and 606 deletions
@@ -416,12 +416,17 @@ def _deep_gemm_execution_hook(
yield yield
def pp_parallel_deep_gemm_warmup(model_runner) -> None: def pp_parallel_deep_gemm_warmup(runner) -> None:
"""Run per-PP-rank dummy DECODE+EXTEND forwards so each rank's """Run per-PP-rank dummy DECODE+EXTEND forwards so each rank's
DeepGEMM JIT compiles in parallel instead of serially via the warmup DeepGEMM JIT compiles in parallel instead of serially via the warmup
/generate flowing through the pipeline. Opt-in via /generate flowing through the pipeline. Opt-in via
SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP. SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP.
Driven from BaseRunner.warmup(), which passes the runner; the dummy
forwards go through runner._dummy_run (the autotune/dummy-run machinery now
lives on BaseRunner). ModelRunner state is read via runner.model_runner.
""" """
model_runner = runner.model_runner
# n_splits ~= n_sms / ceil(bs/block_m) with block_m=64; sweep 5 bs to # n_splits ~= n_sms / ceil(bs/block_m) with block_m=64; sweep 5 bs to
# cover the brackets real /generate hits (smallest decode shape, # cover the brackets real /generate hits (smallest decode shape,
# mid-low, two mid, and n_splits=1 for ~5K+ token prefill). Ceil-align # mid-low, two mid, and n_splits=1 for ~5K+ token prefill). Ceil-align
@@ -466,11 +471,11 @@ def pp_parallel_deep_gemm_warmup(model_runner) -> None:
with torch.inference_mode(): with torch.inference_mode():
for bs in batch_sizes: for bs in batch_sizes:
if run_decode: if run_decode:
model_runner._dummy_run( runner._dummy_run(
batch_size=bs, forward_mode_override=ForwardMode.DECODE batch_size=bs, forward_mode_override=ForwardMode.DECODE
) )
if run_extend: if run_extend:
model_runner._dummy_run( runner._dummy_run(
batch_size=bs, forward_mode_override=ForwardMode.EXTEND batch_size=bs, forward_mode_override=ForwardMode.EXTEND
) )
+51 -467
View File
@@ -18,7 +18,6 @@ from __future__ import annotations
import contextlib import contextlib
import datetime import datetime
import gc import gc
import hashlib
import inspect import inspect
import logging import logging
import os import os
@@ -27,7 +26,6 @@ import threading
import time import time
from collections import defaultdict from collections import defaultdict
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable, List, Optional, Tuple, Union from typing import Any, Callable, List, Optional, Tuple, Union
import torch import torch
@@ -35,7 +33,6 @@ import torch.distributed as dist
from torch import nn from torch import nn
from sglang.jit_kernel.ngram_embedding import update_token_table_decode from sglang.jit_kernel.ngram_embedding import update_token_table_decode
from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config
from sglang.srt.configs import ( from sglang.srt.configs import (
BailingHybridConfig, BailingHybridConfig,
FalconH1Config, FalconH1Config,
@@ -129,12 +126,8 @@ from sglang.srt.layers.cp.utils import (
get_cp_strategy, get_cp_strategy,
) )
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
DpPaddingMode,
get_attention_tp_group, get_attention_tp_group,
get_attention_tp_size,
initialize_dp_attention, initialize_dp_attention,
set_dp_buffer_len,
set_is_extend_in_batch,
) )
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.moe.hash_topk import HashTopK from sglang.srt.layers.moe.hash_topk import HashTopK
@@ -156,7 +149,6 @@ from sglang.srt.model_executor.cuda_graph_config import (
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 (
CaptureHiddenMode,
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
PPProxyTensors, PPProxyTensors,
@@ -175,9 +167,6 @@ from sglang.srt.model_executor.runner import (
EagerRunner, EagerRunner,
PrefillCudaGraphRunner, PrefillCudaGraphRunner,
) )
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
_allocate_decode_buffers,
)
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
RemoteInstanceWeightLoaderBackend, RemoteInstanceWeightLoaderBackend,
@@ -193,10 +182,7 @@ from sglang.srt.server_args import (
get_global_server_args, get_global_server_args,
set_global_server_args_for_scheduler, set_global_server_args_for_scheduler,
) )
from sglang.srt.speculative.spec_info import ( from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
SpeculativeAlgorithm,
create_dummy_verify_input,
)
from sglang.srt.state_capturer.base import TopkCaptureOutput from sglang.srt.state_capturer.base import TopkCaptureOutput
from sglang.srt.state_capturer.indexer_topk import ( from sglang.srt.state_capturer.indexer_topk import (
create_indexer_capturer, create_indexer_capturer,
@@ -213,7 +199,6 @@ from sglang.srt.utils import (
broadcast_pyobj, broadcast_pyobj,
cpu_has_amx_support, cpu_has_amx_support,
dynamic_import, dynamic_import,
empty_context,
enable_show_time_cost, enable_show_time_cost,
get_available_gpu_memory, get_available_gpu_memory,
get_bool_env_var, get_bool_env_var,
@@ -224,14 +209,11 @@ from sglang.srt.utils import (
is_npu, is_npu,
log_info_on_rank0, log_info_on_rank0,
monkey_patch_p2p_access_check, monkey_patch_p2p_access_check,
require_attn_tp_gather,
require_gathered_buffer, require_gathered_buffer,
require_mlp_tp_gather,
reserve_rope_cache_for_long_sequences, reserve_rope_cache_for_long_sequences,
set_cuda_arch, set_cuda_arch,
slow_rank_detector, slow_rank_detector,
) )
from sglang.srt.utils.common import ceil_align, require_mlp_sync
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
from sglang.srt.utils.nvtx_utils import profile_range from sglang.srt.utils.nvtx_utils import profile_range
@@ -889,22 +871,14 @@ class ModelRunner(ModelRunnerKVCacheMixin):
# runs with aux hidden state capture enabled. # runs with aux hidden state capture enabled.
self.init_aux_hidden_state_capture() self.init_aux_hidden_state_capture()
# The eager (no-cuda-graph) phase runner. Always built: it serves both # Device-specific attention-backend init. cg_supported gates decode
# the fully-disabled case (decode/prefill runners point at it) and the # cuda-graph capture (out-of-tree / unknown platforms may not support it).
# eager fallback when a cuda-graph runner can't run a batch. cg_supported = True
self.eager_runner = EagerRunner(self)
if self.device == "cuda" or self.device == "musa": if self.device == "cuda" or self.device == "musa":
self.init_cublas() self.init_cublas()
self.init_attention_backend() self.init_attention_backend()
self.kernel_warmup()
self._pre_initialize_flashinfer_allreduce_workspace()
if not disable_cuda_graph:
self.init_decode_cuda_graph()
elif self.device == "cpu": elif self.device == "cpu":
self.init_attention_backend() self.init_attention_backend()
if not disable_cuda_graph:
self.init_decode_cuda_graph()
elif self.device == "npu": elif self.device == "npu":
self.init_attention_backend() self.init_attention_backend()
# lazy init for zbal with mix mode (before graph capture when enable_cuda_graph) # lazy init for zbal with mix mode (before graph capture when enable_cuda_graph)
@@ -918,32 +892,44 @@ class ModelRunner(ModelRunnerKVCacheMixin):
get_world_group().world_size, get_world_group().world_size,
get_world_group().cpu_group, get_world_group().cpu_group,
) )
if not disable_cuda_graph:
self.init_decode_cuda_graph()
elif current_platform.is_out_of_tree(): elif current_platform.is_out_of_tree():
self.init_attention_backend() self.init_attention_backend()
if current_platform.support_cuda_graph() and not disable_cuda_graph: cg_supported = current_platform.support_cuda_graph()
self.init_decode_cuda_graph() else:
else: self.init_attention_backend()
self.decode_cuda_graph_runner = None cg_supported = False
self.graph_mem_usage = 0
# 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()
if not disable_cuda_graph and cg_supported:
self.init_decode_cuda_graph()
else: else:
self.decode_cuda_graph_runner = None self.decode_cuda_graph_runner = None
self.graph_mem_usage = 0 self.graph_mem_usage = 0
self.init_attention_backend()
if disable_cuda_graph: if disable_cuda_graph:
# Decode cuda graph disabled: route eager decode through the # Decode cuda graph disabled: route eager decode through the
# EagerRunner (the dispatch gate isinstance(..., EagerRunner) keeps # EagerRunner (the dispatch gate isinstance(..., EagerRunner) keeps
# _forward_raw off any replay branch). # _forward_raw off any replay branch).
self.decode_cuda_graph_runner = self.eager_runner self.decode_cuda_graph_runner = self.eager_runner
self.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 server_args.forward_hooks: if server_args.forward_hooks:
register_forward_hooks(self.model, server_args.forward_hooks) register_forward_hooks(self.model, server_args.forward_hooks)
self.init_prefill_cuda_graph()
self.prealloc_symmetric_memory_pool() self.prealloc_symmetric_memory_pool()
if self.canary_manager is not None and not self.is_draft_worker: if self.canary_manager is not None and not self.is_draft_worker:
@@ -2490,431 +2476,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
full_attention_backend = ATTENTION_BACKENDS[backend_str](self) full_attention_backend = ATTENTION_BACKENDS[backend_str](self)
return attn_backend_wrapper(self, full_attention_backend) return attn_backend_wrapper(self, full_attention_backend)
def kernel_warmup(self):
"""Warmup and tune kernels before cuda graph capture."""
if self.device != "cuda":
return
if self._should_run_flashinfer_autotune():
self._flashinfer_autotune()
if (
envs.SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP.get()
and deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM
and self.pp_size > 1
and not self.spec_algorithm.is_speculative()
):
from sglang.srt.layers.deep_gemm_wrapper.compile_utils import (
pp_parallel_deep_gemm_warmup,
)
pp_parallel_deep_gemm_warmup(self)
def _pre_initialize_flashinfer_allreduce_workspace(self):
"""Pre-initialize flashinfer allreduce fusion workspaces.
Must run before CUDA graph capture to avoid collective operations
(broadcasts, barriers) inside the graph capture context, which can
deadlock with custom_all_reduce.register_graph_buffers.
"""
if self.server_args.flashinfer_allreduce_fusion_backend is None:
return
from sglang.srt.layers.communicator import FUSE_ALLREDUCE_MAX_BATCH_SIZE
from sglang.srt.layers.flashinfer_comm_fusion import pre_initialize_workspaces
pre_initialize_workspaces(
max_token_num=FUSE_ALLREDUCE_MAX_BATCH_SIZE,
hidden_dim=self.model_config.hidden_size,
dtype=self.dtype,
)
def _should_run_flashinfer_autotune(self) -> bool:
"""Check if flashinfer autotune should be run."""
if self.server_args.disable_flashinfer_autotune:
return False
# CuteDSL v1 (cutedsl runner + deepep a2a) bypasses MoeRunner and must not
# be autotuned -- its _dummy_run would dispatch more tokens per rank than
# SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK, tripping a DeepEP assert.
# Read server_args directly to avoid depending on initialize_moe_config()
# having already populated the MoE backend globals.
if (
self.server_args.moe_runner_backend == "flashinfer_cutedsl"
and self.server_args.moe_a2a_backend == "deepep"
):
return False
backend_str = self.server_args.moe_runner_backend
# TODO smor- support other cases for flashinfer autotune, such as, mamba backend
moe_needs_autotune = backend_str in [
"flashinfer_trtllm",
"flashinfer_trtllm_routed",
"flashinfer_mxfp4",
"flashinfer_cutedsl",
"flashinfer_cutlass",
]
from sglang.srt.layers.quantization.fp4_utils import (
get_fp4_gemm_runner_backend,
)
model_uses_fp4 = self.model_config.quantization in (
"modelopt_fp4",
"modelopt_mixed",
)
fp4_gemm_needs_autotune = model_uses_fp4 and (
get_fp4_gemm_runner_backend().is_flashinfer_cutlass()
or get_fp4_gemm_runner_backend().is_flashinfer_cutedsl()
)
from sglang.srt.layers.quantization.fp8_utils import (
get_fp8_gemm_runner_backend,
)
from sglang.srt.utils import is_sm100_supported
model_uses_modelopt_fp8 = self.model_config.quantization in (
"modelopt",
"modelopt_fp8",
"modelopt_mixed",
)
fp8_gemm_needs_autotune = (
get_fp8_gemm_runner_backend().is_flashinfer_cutlass()
or (model_uses_modelopt_fp8 and is_sm100_supported())
)
if not (
moe_needs_autotune or fp4_gemm_needs_autotune or fp8_gemm_needs_autotune
):
return False
major, _ = torch.cuda.get_device_capability()
if major < 9:
return False
if self.spec_algorithm.is_speculative():
return not self.is_draft_worker
return True
def _flashinfer_autotune(self):
"""Run flashinfer autotune."""
from flashinfer.autotuner import autotune
from sglang.srt.layers.logits_processor import autotune_dummy_run_mode
cache_path = self._flashinfer_autotune_cache_path()
if envs.SGLANG_FLASHINFER_AUTOTUNE_CACHE.get():
autotune_cache = cache_path
logger.info("Running FlashInfer autotune with cache: %s", autotune_cache)
else:
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
runs_dir = cache_path.parent / "runs"
runs_dir.mkdir(parents=True, exist_ok=True)
autotune_cache = (
runs_dir / f"{cache_path.stem}.{timestamp}{cache_path.suffix}"
)
logger.info(
"Running FlashInfer autotune (cache reuse DISABLED via "
"SGLANG_FLASHINFER_AUTOTUNE_CACHE=0); writing fresh result to: %s",
autotune_cache,
)
# Run warmup on the non-default stream to avoid NCCL 2.29+ cudaMemcpyBatchAsync
# calls on default stream (unsupported by CUDA) when --enable-symm-mem is used.
self.forward_stream.wait_stream(torch.cuda.current_stream())
with torch.get_device_module(self.device).stream(self.forward_stream):
with (
torch.inference_mode(),
autotune(True, cache=str(autotune_cache)),
autotune_dummy_run_mode(),
):
self._dummy_run(batch_size=self.req_to_token_pool.size)
torch.cuda.current_stream().wait_stream(self.forward_stream)
logger.info("FlashInfer autotune completed.")
def _flashinfer_autotune_cache_path(self) -> Path:
import flashinfer
major, minor = torch.cuda.get_device_capability(self.device)
arch = f"sm{major}{minor}"
flashinfer_version = getattr(flashinfer, "__version__", "unknown")
server_args = self.server_args
model_key = "|".join(
[
str(server_args.model_path),
str(self.dtype),
str(server_args.quantization),
str(server_args.moe_runner_backend),
str(self.tp_size),
str(self.pp_size),
str(self.dp_size),
str(self.moe_ep_size),
str(self.model_config.hf_config.__class__.__name__),
]
)
cache_key = hashlib.sha256(model_key.encode()).hexdigest()[:16]
cache_dir = (
Path(envs.SGLANG_CACHE_DIR.get())
/ "flashinfer"
/ "autotune"
/ flashinfer_version
/ arch
/ cache_key
)
cache_dir.mkdir(parents=True, exist_ok=True)
return (
cache_dir
/ f"rank_tp{self.tp_rank}_pp{self.pp_rank}_dp{self.dp_rank or 0}.json"
)
def _dummy_run(
self,
batch_size: int,
run_ctx=None,
forward_mode_override: Optional[ForwardMode] = None,
):
"""Run a dummy forward pass for warmup/profiling.
forward_mode_override forces EXTEND/DECODE regardless of
is_generation (used by the PP-parallel DeepGEMM warmup).
"""
if forward_mode_override is not None:
capture_forward_mode = forward_mode_override
elif self.is_generation:
capture_forward_mode = ForwardMode.DECODE
else:
capture_forward_mode = ForwardMode.EXTEND
capture_hidden_mode = CaptureHiddenMode.NULL
num_tokens_per_bs = 1
if self.spec_algorithm.is_speculative():
if self.is_draft_worker:
if not self.spec_algorithm.supports_target_verify_for_draft():
raise RuntimeError("This should not happen")
capture_forward_mode = ForwardMode.TARGET_VERIFY
num_tokens_per_bs = (
self.spec_algorithm.get_num_tokens_per_bs_for_target_verify(
self.server_args.speculative_num_draft_tokens, self.is_draft_worker
)
)
if self.server_args.enable_return_hidden_states:
capture_hidden_mode = CaptureHiddenMode.FULL
num_tokens = batch_size * num_tokens_per_bs
# Keep warmup aligned with scheduler MLP-sync padding.
if require_mlp_sync(self.server_args):
attn_tp_size = get_attention_tp_size()
if attn_tp_size > 1 and num_tokens % attn_tp_size != 0:
num_tokens = ceil_align(num_tokens, attn_tp_size)
batch_size = num_tokens // num_tokens_per_bs
seq_len_fill_value = self.attn_backend.get_cuda_graph_seq_len_fill_value()
if self.server_args.enable_torch_compile:
set_torch_compile_config()
should_disable_torch_compile = not getattr(
self.model, "_can_torch_compile", True
)
if should_disable_torch_compile:
log_info_on_rank0(
logger,
"Transformers backend model reports it is not torch.compile "
"compatible (e.g. dynamic rope scaling). Disabling torch.compile.",
)
self.server_args.enable_torch_compile = False
# 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(self.server_args)
if require_gathered_buffer(self.server_args):
assert require_mlp_tp_gather_ or require_attn_tp_gather(self.server_args)
buffers = _allocate_decode_buffers(
device=self.device,
max_bs=batch_size,
max_num_token=num_tokens,
hidden_size=self.model_config.hidden_size,
vocab_size=self.model_config.vocab_size,
dtype=self.model_config.dtype,
dp_size=self.server_args.dp_size,
pp_size=self.server_args.pp_size,
is_encoder_decoder=self.model_config.is_encoder_decoder,
require_mlp_tp_gather=require_mlp_tp_gather_,
seq_len_fill_value=seq_len_fill_value,
encoder_len_fill_value=(
getattr(self.model_config.hf_config, "max_source_positions", 0)
if self.model_config.is_encoder_decoder
else 0
),
num_tokens_per_bs=num_tokens_per_bs,
cache_loc_dtype=torch.int64,
enable_mamba_track=False,
hc_hidden_size=getattr(self.model_config, "hc_hidden_size", None),
pp_proxy_topk_size=self.get_pp_proxy_topk_size(),
)
buffers.num_token_non_padded[...] = num_tokens
# For extend mode
if capture_forward_mode == ForwardMode.EXTEND:
extend_prefix_lens_cpu = [0] * batch_size
extend_seq_lens_cpu = [seq_len_fill_value] * batch_size
extend_num_tokens = num_tokens
extend_seq_lens = torch.full(
(batch_size,), seq_len_fill_value, dtype=torch.int32, device=self.device
)
extend_prefix_lens = torch.zeros(
(batch_size,), dtype=torch.int32, device=self.device
)
extend_start_loc = torch.arange(
0, num_tokens, num_tokens_per_bs, dtype=torch.int32, device=self.device
)
else:
extend_prefix_lens_cpu = None
extend_seq_lens_cpu = None
extend_num_tokens = None
extend_seq_lens = None
extend_prefix_lens = None
extend_start_loc = None
if self.server_args.pp_size > 1:
# PP0 already cp-split hidden_states before send.
pp_hidden_tokens = num_tokens
if (
capture_forward_mode == ForwardMode.EXTEND
and self.pp_rank != 0
and self.attn_cp_size > 1
):
pp_hidden_tokens = num_tokens // self.attn_cp_size
pp_proxy_tensors = PPProxyTensors(
{k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()}
)
if require_mlp_tp_gather_:
global_num_tokens_cpu = [num_tokens] * self.server_args.dp_size
elif require_attn_tp_gather(self.server_args):
global_num_tokens_cpu = [num_tokens]
else:
global_num_tokens_cpu = None
if global_num_tokens_cpu is not None:
global_dp_buffer_len = sum(global_num_tokens_cpu)
num_tokens_tensor = torch.tensor(
global_num_tokens_cpu, dtype=torch.int32, device=self.device
)
buffers.global_num_tokens_gpu.copy_(num_tokens_tensor)
buffers.global_num_tokens_for_logprob_gpu.copy_(num_tokens_tensor)
else:
global_dp_buffer_len = None
global_num_tokens_cpu = None
spec_info = create_dummy_verify_input(
self.spec_algorithm,
self.server_args,
buffers.custom_mask,
num_tokens_per_bs,
self.is_draft_worker,
)
if spec_info is not None and (
self.spec_algorithm.is_eagle() or self.spec_algorithm.is_standalone()
):
# MTP models (e.g. deepseek_nextn) read spec_info.hidden_states
# during forward; provide a dummy so warmup doesn't crash.
spec_info.hidden_states = torch.zeros(
(num_tokens, self.model_config.hidden_size),
dtype=self.dtype,
device=self.device,
)
if capture_hidden_mode != CaptureHiddenMode.FULL:
capture_hidden_mode = (
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL
)
if self.server_args.enable_lora:
lora_ids = [None] * batch_size
else:
lora_ids = None
forward_batch = ForwardBatch(
forward_mode=capture_forward_mode,
batch_size=batch_size,
input_ids=buffers.input_ids,
req_pool_indices=buffers.req_pool_indices,
seq_lens=buffers.seq_lens,
seq_lens_cpu=buffers.seq_lens_cpu,
next_token_logits_buffer=buffers.next_token_logits_buffer,
orig_seq_lens=buffers.seq_lens,
out_cache_loc=buffers.out_cache_loc,
seq_lens_sum=buffers.seq_lens.sum().item(),
encoder_lens=buffers.encoder_lens,
return_logprob=False,
positions=buffers.positions,
extend_num_tokens=extend_num_tokens,
extend_seq_lens=extend_seq_lens,
extend_prefix_lens=extend_prefix_lens,
extend_start_loc=extend_start_loc,
extend_prefix_lens_cpu=extend_prefix_lens_cpu,
extend_seq_lens_cpu=extend_seq_lens_cpu,
global_num_tokens_gpu=buffers.global_num_tokens_gpu,
global_num_tokens_cpu=global_num_tokens_cpu,
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
global_dp_buffer_len=global_dp_buffer_len,
mrope_positions=buffers.mrope_positions,
spec_algorithm=self.spec_algorithm,
spec_info=spec_info,
capture_hidden_mode=capture_hidden_mode,
num_token_non_padded=buffers.num_token_non_padded,
global_forward_mode=capture_forward_mode,
lora_ids=lora_ids,
)
if lora_ids is not None:
self.lora_manager.prepare_lora_batch(forward_batch)
self.attn_backend.init_forward_metadata(forward_batch)
def run_once():
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len(
global_dp_buffer_len,
num_tokens,
forward_batch.dp_padding_mode.is_max_len(),
global_num_tokens_cpu,
)
set_is_extend_in_batch(False)
kwargs = {}
if (
self.server_args.pp_size > 1
and "pp_proxy_tensors"
in inspect.signature(self.model.forward).parameters
):
kwargs["pp_proxy_tensors"] = PPProxyTensors(
{k: v.clone() for k, v in pp_proxy_tensors.tensors.items()}
)
if not self.is_generation:
kwargs["get_embedding"] = True
logits_output_or_pp_proxy_tensors = self.model.forward(
buffers.input_ids,
forward_batch.positions,
forward_batch,
**kwargs,
)
return logits_output_or_pp_proxy_tensors
torch.get_device_module(self.device).synchronize()
self.tp_group.barrier()
with forward_context(ForwardContext(attn_backend=self.attn_backend)):
with torch.inference_mode(), run_ctx or empty_context():
run_once()
def maybe_init_ngram_embedding(self): def maybe_init_ngram_embedding(self):
self.use_ngram_embedding = self.model_config.use_ngram_embedding self.use_ngram_embedding = self.model_config.use_ngram_embedding
if self.use_ngram_embedding: if self.use_ngram_embedding:
@@ -3038,6 +2599,29 @@ class ModelRunner(ModelRunnerKVCacheMixin):
if self.is_draft_worker and not force_for_draft_worker: if self.is_draft_worker and not force_for_draft_worker:
return return
# EAGLE-family target worker: the prefill graph captures
# CaptureHiddenMode.NULL, but target prefill needs FULL hidden states to
# feed the draft, so can_run_graph always rejects it and prefill runs
# eagerly. The graph is therefore never used; capturing it (now before
# the decode graph) can perturb backend state on FP4 / TRTLLM-MoE paths
# and corrupt decode replay, so skip its capture and route prefill
# through the eager runner — the runtime path either way. With
# enable_return_hidden_states the prefill graph is FULL and usable, so
# only skip when it would capture NULL.
if (
self.spec_algorithm.is_eagle()
and not self.is_draft_worker
and not self.server_args.enable_return_hidden_states
):
logger.info(
"Disable prefill CUDA graph for the EAGLE target worker: target "
"prefill needs FULL hidden states but the prefill graph captures "
"NULL, so the graph is unused; skipping its capture keeps decode "
"graph capture clean."
)
self.prefill_cuda_graph_runner = self.eager_runner
return
# Disable piecewise CUDA graph for non-language models # Disable piecewise CUDA graph for non-language models
if not hasattr(self.model, "model"): if not hasattr(self.model, "model"):
logger.warning( logger.warning(
@@ -7,7 +7,7 @@ BaseCudaGraphBackend chosen via cuda_graph_config.
Public API: Public API:
- BaseRunner — minimal abstract base shared by the cuda-graph runners - BaseRunner — minimal abstract base shared by the cuda-graph runners
and the eager runner (shared __init__ + abstract and the eager runner (shared __init__ + warmup + abstract
can_run_graph/load_batch/execute). can_run_graph/load_batch/execute).
- BaseCudaGraphRunner — abstract cuda-graph base; bucket padding + - BaseCudaGraphRunner — abstract cuda-graph base; bucket padding +
capture-loop scaffolding on top of BaseRunner. capture-loop scaffolding on top of BaseRunner.
@@ -15,22 +15,176 @@
from __future__ import annotations from __future__ import annotations
import datetime
import hashlib
import inspect
import logging import logging
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any from pathlib import Path
from types import SimpleNamespace
from typing import TYPE_CHECKING, Any, Optional
import torch import torch
from sglang.srt.batch_overlap.two_batch_overlap import TboCudaGraphRunnerPlugin from sglang.srt.batch_overlap.two_batch_overlap import TboCudaGraphRunnerPlugin
from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config
from sglang.srt.environ import envs
from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.dp_attention import (
DpPaddingMode,
set_dp_buffer_len,
set_is_extend_in_batch,
)
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
ForwardMode,
NgramEmbeddingInfo,
PPProxyTensors,
)
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
from sglang.srt.speculative.spec_info import create_dummy_verify_input
from sglang.srt.utils import (
empty_context,
log_info_on_rank0,
require_attn_tp_gather,
require_gathered_buffer,
require_mlp_tp_gather,
)
from sglang.srt.utils.common import ceil_align, require_mlp_sync
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def _allocate_decode_buffers(
*,
device: torch.device,
max_bs: int,
max_num_token: int,
hidden_size: int,
vocab_size: int,
dtype: torch.dtype,
dp_size: int,
pp_size: int,
is_encoder_decoder: bool,
require_mlp_tp_gather: bool,
seq_len_fill_value: int,
encoder_len_fill_value: int,
num_tokens_per_bs: int,
cache_loc_dtype: torch.dtype,
enable_mamba_track: bool,
ne_token_table: Optional[torch.Tensor] = None,
hc_hidden_size: Optional[int] = None,
) -> SimpleNamespace:
"""Allocate the FB-shared decode buffers."""
with torch.device(device):
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
input_embeds = torch.zeros((max_num_token, hidden_size), dtype=dtype)
req_pool_indices = torch.zeros((max_bs,), dtype=torch.int64)
seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int64)
out_cache_loc = torch.zeros((max_num_token,), dtype=cache_loc_dtype)
positions = torch.zeros((max_num_token,), dtype=torch.int64)
mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64)
num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
custom_mask = torch.ones(
(max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_bs,
dtype=torch.bool,
)
next_token_logits_buffer = torch.zeros(
(max_num_token, vocab_size),
dtype=torch.float,
)
mamba_track_indices = (
torch.zeros((max_bs,), dtype=torch.int64) if enable_mamba_track else None
)
mamba_track_mask = (
torch.zeros((max_bs,), dtype=torch.bool) if enable_mamba_track else None
)
if pp_size > 1:
# mHC (e.g. DSV4) flattens residual into hidden_states (size = hc_hidden_size).
is_mhc = hc_hidden_size is not None
hs = hc_hidden_size if is_mhc else hidden_size
pp_proxy_tensors = {
"hidden_states": torch.zeros((max_bs, hs), dtype=dtype),
}
if not is_mhc:
pp_proxy_tensors["residual"] = torch.zeros(
(max_bs, hidden_size), dtype=dtype
)
else:
pp_proxy_tensors = None
if is_encoder_decoder:
encoder_lens = torch.full(
(max_bs,), encoder_len_fill_value, dtype=torch.int32
)
else:
encoder_lens = None
if require_mlp_tp_gather:
global_num_tokens_gpu = torch.zeros((dp_size,), dtype=torch.int32)
global_num_tokens_for_logprob_gpu = torch.zeros(
(dp_size,), dtype=torch.int32
)
else:
global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
global_num_tokens_for_logprob_gpu = torch.zeros((1,), dtype=torch.int32)
ngram_embedding_info = (
NgramEmbeddingInfo(
token_table=ne_token_table,
column_starts=torch.zeros([max_bs], dtype=torch.int32),
req_lens=torch.ones([max_bs], dtype=torch.int32),
out_column_starts=torch.zeros([max_bs], dtype=torch.int32),
out_req_lens=torch.ones([max_bs], dtype=torch.int32),
)
if ne_token_table is not None
else None
)
if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get():
rids_int = torch.zeros((max_bs,), dtype=torch.int64)
bootstrap_room_ids_int = torch.full((max_bs,), -1, dtype=torch.int64)
else:
rids_int = None
bootstrap_room_ids_int = None
seq_lens_cpu = torch.full(
(max_bs,),
seq_len_fill_value,
dtype=torch.int64,
device="cpu",
)
return SimpleNamespace(
input_ids=input_ids,
input_embeds=input_embeds,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
out_cache_loc=out_cache_loc,
positions=positions,
mrope_positions=mrope_positions,
num_token_non_padded=num_token_non_padded,
custom_mask=custom_mask,
next_token_logits_buffer=next_token_logits_buffer,
mamba_track_indices=mamba_track_indices,
mamba_track_mask=mamba_track_mask,
encoder_lens=encoder_lens,
global_num_tokens_gpu=global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
pp_proxy_tensors=pp_proxy_tensors,
ngram_embedding_info=ngram_embedding_info,
rids_int=rids_int,
bootstrap_room_ids_int=bootstrap_room_ids_int,
)
class BaseRunner(ABC): class BaseRunner(ABC):
def __init__(self, model_runner: ModelRunner) -> None: def __init__(self, model_runner: ModelRunner) -> None:
self.model_runner = model_runner self.model_runner = model_runner
@@ -43,6 +197,438 @@ class BaseRunner(ABC):
self.attn_tp_rank = get_parallel().attn_tp_rank self.attn_tp_rank = get_parallel().attn_tp_rank
self.tbo_plugin = TboCudaGraphRunnerPlugin() self.tbo_plugin = TboCudaGraphRunnerPlugin()
def warmup(self) -> None:
"""Run kernel warmup + autotune once, gated by mr._kernel_warmed_up."""
mr = self.model_runner
if getattr(mr, "_kernel_warmed_up", False):
return
mr._kernel_warmed_up = True
if mr.device != "cuda":
return
self._pre_initialize_flashinfer_allreduce_workspace()
if self._should_run_flashinfer_autotune():
self._flashinfer_autotune()
if (
envs.SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP.get()
and deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM
and mr.pp_size > 1
and not mr.spec_algorithm.is_speculative()
):
from sglang.srt.layers.deep_gemm_wrapper.compile_utils import (
pp_parallel_deep_gemm_warmup,
)
pp_parallel_deep_gemm_warmup(self)
def _pre_initialize_flashinfer_allreduce_workspace(self):
"""Allocate flashinfer allreduce workspaces; must run before CG capture
to keep broadcasts/barriers outside the capture context (else deadlock
with custom_all_reduce.register_graph_buffers).
"""
mr = self.model_runner
if mr.server_args.flashinfer_allreduce_fusion_backend is None:
return
from sglang.srt.layers.communicator import FUSE_ALLREDUCE_MAX_BATCH_SIZE
from sglang.srt.layers.flashinfer_comm_fusion import pre_initialize_workspaces
pre_initialize_workspaces(
max_token_num=FUSE_ALLREDUCE_MAX_BATCH_SIZE,
hidden_dim=mr.model_config.hidden_size,
dtype=mr.dtype,
)
def _should_run_flashinfer_autotune(self) -> bool:
"""Check if flashinfer autotune should be run."""
mr = self.model_runner
if mr.server_args.disable_flashinfer_autotune:
return False
# CuteDSL v1 (cutedsl runner + deepep a2a) bypasses MoeRunner and must not
# be autotuned -- its _dummy_run would dispatch more tokens per rank than
# SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK, tripping a DeepEP assert.
# Read server_args directly to avoid depending on initialize_moe_config()
# having already populated the MoE backend globals.
if (
mr.server_args.moe_runner_backend == "flashinfer_cutedsl"
and mr.server_args.moe_a2a_backend == "deepep"
):
return False
backend_str = mr.server_args.moe_runner_backend
# TODO smor- support other cases for flashinfer autotune, such as, mamba backend
moe_needs_autotune = backend_str in [
"flashinfer_trtllm",
"flashinfer_trtllm_routed",
"flashinfer_mxfp4",
"flashinfer_cutedsl",
"flashinfer_cutlass",
]
from sglang.srt.layers.quantization.fp4_utils import (
get_fp4_gemm_runner_backend,
)
model_uses_fp4 = mr.model_config.quantization in (
"modelopt_fp4",
"modelopt_mixed",
)
fp4_gemm_needs_autotune = model_uses_fp4 and (
get_fp4_gemm_runner_backend().is_flashinfer_cutlass()
or get_fp4_gemm_runner_backend().is_flashinfer_cutedsl()
)
from sglang.srt.layers.quantization.fp8_utils import (
get_fp8_gemm_runner_backend,
)
from sglang.srt.utils import is_sm100_supported
model_uses_modelopt_fp8 = mr.model_config.quantization in (
"modelopt",
"modelopt_fp8",
"modelopt_mixed",
)
fp8_gemm_needs_autotune = (
get_fp8_gemm_runner_backend().is_flashinfer_cutlass()
or (model_uses_modelopt_fp8 and is_sm100_supported())
)
if not (
moe_needs_autotune or fp4_gemm_needs_autotune or fp8_gemm_needs_autotune
):
return False
major, _ = torch.cuda.get_device_capability()
if major < 9:
return False
if mr.spec_algorithm.is_speculative():
return not mr.is_draft_worker
return True
def _flashinfer_autotune(self):
"""Run flashinfer autotune."""
from flashinfer.autotuner import autotune
from sglang.srt.layers.logits_processor import autotune_dummy_run_mode
mr = self.model_runner
cache_path = self._flashinfer_autotune_cache_path()
if envs.SGLANG_FLASHINFER_AUTOTUNE_CACHE.get():
autotune_cache = cache_path
logger.info("Running FlashInfer autotune with cache: %s", autotune_cache)
else:
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
runs_dir = cache_path.parent / "runs"
runs_dir.mkdir(parents=True, exist_ok=True)
autotune_cache = (
runs_dir / f"{cache_path.stem}.{timestamp}{cache_path.suffix}"
)
logger.info(
"Running FlashInfer autotune (cache reuse DISABLED via "
"SGLANG_FLASHINFER_AUTOTUNE_CACHE=0); writing fresh result to: %s",
autotune_cache,
)
# Run warmup on the non-default stream to avoid NCCL 2.29+ cudaMemcpyBatchAsync
# calls on default stream (unsupported by CUDA) when --enable-symm-mem is used.
mr.forward_stream.wait_stream(torch.cuda.current_stream())
with torch.get_device_module(mr.device).stream(mr.forward_stream):
with (
torch.inference_mode(),
autotune(True, cache=str(autotune_cache)),
autotune_dummy_run_mode(),
):
self._dummy_run(batch_size=mr.req_to_token_pool.size)
torch.cuda.current_stream().wait_stream(mr.forward_stream)
logger.info("FlashInfer autotune completed.")
def _flashinfer_autotune_cache_path(self) -> Path:
import flashinfer
mr = self.model_runner
major, minor = torch.cuda.get_device_capability(mr.device)
arch = f"sm{major}{minor}"
flashinfer_version = getattr(flashinfer, "__version__", "unknown")
server_args = mr.server_args
model_key = "|".join(
[
str(server_args.model_path),
str(mr.dtype),
str(server_args.quantization),
str(server_args.moe_runner_backend),
str(mr.tp_size),
str(mr.pp_size),
str(mr.dp_size),
str(mr.moe_ep_size),
str(mr.model_config.hf_config.__class__.__name__),
]
)
cache_key = hashlib.sha256(model_key.encode()).hexdigest()[:16]
cache_dir = (
Path(envs.SGLANG_CACHE_DIR.get())
/ "flashinfer"
/ "autotune"
/ flashinfer_version
/ arch
/ cache_key
)
cache_dir.mkdir(parents=True, exist_ok=True)
return (
cache_dir / f"rank_tp{mr.tp_rank}_pp{mr.pp_rank}_dp{mr.dp_rank or 0}.json"
)
def _dummy_run(
self,
batch_size: int,
run_ctx=None,
forward_mode_override: Optional[ForwardMode] = None,
):
"""Run a dummy forward pass for warmup/profiling.
forward_mode_override forces EXTEND/DECODE regardless of
is_generation (used by the PP-parallel DeepGEMM warmup).
"""
mr = self.model_runner
if forward_mode_override is not None:
capture_forward_mode = forward_mode_override
elif mr.is_generation:
capture_forward_mode = ForwardMode.DECODE
else:
capture_forward_mode = ForwardMode.EXTEND
capture_hidden_mode = CaptureHiddenMode.NULL
num_tokens_per_bs = 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")
capture_forward_mode = ForwardMode.TARGET_VERIFY
num_tokens_per_bs = (
mr.spec_algorithm.get_num_tokens_per_bs_for_target_verify(
mr.server_args.speculative_num_draft_tokens, mr.is_draft_worker
)
)
if mr.server_args.enable_return_hidden_states:
capture_hidden_mode = CaptureHiddenMode.FULL
num_tokens = batch_size * num_tokens_per_bs
# Keep warmup aligned with scheduler MLP-sync padding.
if require_mlp_sync(mr.server_args):
attn_tp_size = get_parallel().attn_tp_size
if attn_tp_size > 1 and num_tokens % attn_tp_size != 0:
num_tokens = ceil_align(num_tokens, attn_tp_size)
batch_size = num_tokens // num_tokens_per_bs
seq_len_fill_value = mr.attn_backend.get_cuda_graph_seq_len_fill_value()
if mr.server_args.enable_torch_compile:
set_torch_compile_config()
should_disable_torch_compile = not getattr(
mr.model, "_can_torch_compile", True
)
if should_disable_torch_compile:
log_info_on_rank0(
logger,
"Transformers backend model reports it is not torch.compile "
"compatible (e.g. dynamic rope scaling). Disabling torch.compile.",
)
mr.server_args.enable_torch_compile = False
# 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)
buffers = _allocate_decode_buffers(
device=mr.device,
max_bs=batch_size,
max_num_token=num_tokens,
hidden_size=mr.model_config.hidden_size,
vocab_size=mr.model_config.vocab_size,
dtype=mr.model_config.dtype,
dp_size=mr.server_args.dp_size,
pp_size=mr.server_args.pp_size,
is_encoder_decoder=mr.model_config.is_encoder_decoder,
require_mlp_tp_gather=require_mlp_tp_gather_,
seq_len_fill_value=seq_len_fill_value,
encoder_len_fill_value=(
getattr(mr.model_config.hf_config, "max_source_positions", 0)
if mr.model_config.is_encoder_decoder
else 0
),
num_tokens_per_bs=num_tokens_per_bs,
cache_loc_dtype=torch.int64,
enable_mamba_track=False,
hc_hidden_size=getattr(mr.model_config, "hc_hidden_size", None),
)
buffers.num_token_non_padded[...] = num_tokens
# For extend mode
if capture_forward_mode == ForwardMode.EXTEND:
extend_prefix_lens_cpu = [0] * batch_size
extend_seq_lens_cpu = [seq_len_fill_value] * batch_size
extend_num_tokens = num_tokens
extend_seq_lens = torch.full(
(batch_size,), seq_len_fill_value, dtype=torch.int32, device=mr.device
)
extend_prefix_lens = torch.zeros(
(batch_size,), dtype=torch.int32, device=mr.device
)
extend_start_loc = torch.arange(
0, num_tokens, num_tokens_per_bs, dtype=torch.int32, device=mr.device
)
else:
extend_prefix_lens_cpu = None
extend_seq_lens_cpu = None
extend_num_tokens = None
extend_seq_lens = None
extend_prefix_lens = None
extend_start_loc = None
if mr.server_args.pp_size > 1:
# PP0 already cp-split hidden_states before send.
pp_hidden_tokens = num_tokens
if (
capture_forward_mode == ForwardMode.EXTEND
and mr.pp_rank != 0
and mr.attn_cp_size > 1
):
pp_hidden_tokens = num_tokens // mr.attn_cp_size
pp_proxy_tensors = PPProxyTensors(
{k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()}
)
if require_mlp_tp_gather_:
global_num_tokens_cpu = [num_tokens] * mr.server_args.dp_size
elif require_attn_tp_gather(mr.server_args):
global_num_tokens_cpu = [num_tokens]
else:
global_num_tokens_cpu = None
if global_num_tokens_cpu is not None:
global_dp_buffer_len = sum(global_num_tokens_cpu)
num_tokens_tensor = torch.tensor(
global_num_tokens_cpu, dtype=torch.int32, device=mr.device
)
buffers.global_num_tokens_gpu.copy_(num_tokens_tensor)
buffers.global_num_tokens_for_logprob_gpu.copy_(num_tokens_tensor)
else:
global_dp_buffer_len = None
global_num_tokens_cpu = None
spec_info = create_dummy_verify_input(
mr.spec_algorithm,
mr.server_args,
buffers.custom_mask,
num_tokens_per_bs,
mr.is_draft_worker,
)
if spec_info is not None and (
mr.spec_algorithm.is_eagle() or mr.spec_algorithm.is_standalone()
):
# MTP models (e.g. deepseek_nextn) read spec_info.hidden_states
# during forward; provide a dummy so warmup doesn't crash.
spec_info.hidden_states = torch.zeros(
(num_tokens, mr.model_config.hidden_size),
dtype=mr.dtype,
device=mr.device,
)
if capture_hidden_mode != CaptureHiddenMode.FULL:
capture_hidden_mode = (
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL
)
if mr.server_args.enable_lora:
lora_ids = [None] * batch_size
else:
lora_ids = None
forward_batch = ForwardBatch(
forward_mode=capture_forward_mode,
batch_size=batch_size,
input_ids=buffers.input_ids,
req_pool_indices=buffers.req_pool_indices,
seq_lens=buffers.seq_lens,
seq_lens_cpu=buffers.seq_lens_cpu,
next_token_logits_buffer=buffers.next_token_logits_buffer,
orig_seq_lens=buffers.seq_lens,
out_cache_loc=buffers.out_cache_loc,
seq_lens_sum=buffers.seq_lens.sum().item(),
encoder_lens=buffers.encoder_lens,
return_logprob=False,
positions=buffers.positions,
extend_num_tokens=extend_num_tokens,
extend_seq_lens=extend_seq_lens,
extend_prefix_lens=extend_prefix_lens,
extend_start_loc=extend_start_loc,
extend_prefix_lens_cpu=extend_prefix_lens_cpu,
extend_seq_lens_cpu=extend_seq_lens_cpu,
global_num_tokens_gpu=buffers.global_num_tokens_gpu,
global_num_tokens_cpu=global_num_tokens_cpu,
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
global_dp_buffer_len=global_dp_buffer_len,
mrope_positions=buffers.mrope_positions,
spec_algorithm=mr.spec_algorithm,
spec_info=spec_info,
capture_hidden_mode=capture_hidden_mode,
num_token_non_padded=buffers.num_token_non_padded,
global_forward_mode=capture_forward_mode,
lora_ids=lora_ids,
)
if lora_ids is not None:
mr.lora_manager.prepare_lora_batch(forward_batch)
mr.attn_backend.init_forward_metadata(forward_batch)
def run_once():
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len(
global_dp_buffer_len,
num_tokens,
forward_batch.dp_padding_mode.is_max_len(),
global_num_tokens_cpu,
)
set_is_extend_in_batch(False)
kwargs = {}
if (
mr.server_args.pp_size > 1
and "pp_proxy_tensors" in inspect.signature(mr.model.forward).parameters
):
kwargs["pp_proxy_tensors"] = PPProxyTensors(
{k: v.clone() for k, v in pp_proxy_tensors.tensors.items()}
)
if not mr.is_generation:
kwargs["get_embedding"] = True
logits_output_or_pp_proxy_tensors = mr.model.forward(
buffers.input_ids,
forward_batch.positions,
forward_batch,
**kwargs,
)
return logits_output_or_pp_proxy_tensors
torch.get_device_module(mr.device).synchronize()
mr.tp_group.barrier()
with forward_context(ForwardContext(attn_backend=mr.attn_backend)):
with torch.inference_mode(), run_ctx or empty_context():
run_once()
@abstractmethod @abstractmethod
def can_run_graph(self, forward_batch: ForwardBatch) -> bool: ... def can_run_graph(self, forward_batch: ForwardBatch) -> bool: ...
@@ -43,7 +43,6 @@ from sglang.srt.distributed.parallel_state import (
set_pdmux_status, set_pdmux_status,
) )
from sglang.srt.dllm.config import DllmConfig from sglang.srt.dllm.config import DllmConfig
from sglang.srt.environ import envs
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
DpPaddingMode, DpPaddingMode,
@@ -62,7 +61,6 @@ from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode, CaptureHiddenMode,
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
NgramEmbeddingInfo,
PPProxyTensors, PPProxyTensors,
compute_local_num_token_non_padded, compute_local_num_token_non_padded,
enable_num_token_non_padded, enable_num_token_non_padded,
@@ -163,137 +161,6 @@ def build_replay_fb_view(
) )
def _allocate_decode_buffers(
*,
device: torch.device,
max_bs: int,
max_num_token: int,
hidden_size: int,
vocab_size: int,
dtype: torch.dtype,
dp_size: int,
pp_size: int,
is_encoder_decoder: bool,
require_mlp_tp_gather: bool,
seq_len_fill_value: int,
encoder_len_fill_value: int,
num_tokens_per_bs: int,
cache_loc_dtype: torch.dtype,
enable_mamba_track: bool,
ne_token_table: Optional[torch.Tensor] = None,
hc_hidden_size: Optional[int] = None,
pp_proxy_topk_size: Optional[int] = None,
) -> SimpleNamespace:
"""Allocate the FB-shared decode buffers as a namespace adopted by
``build_decode_registry(source=...)``."""
with torch.device(device):
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
input_embeds = torch.zeros((max_num_token, hidden_size), dtype=dtype)
req_pool_indices = torch.zeros((max_bs,), dtype=torch.int64)
seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int64)
out_cache_loc = torch.zeros((max_num_token,), dtype=cache_loc_dtype)
positions = torch.zeros((max_num_token,), dtype=torch.int64)
mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64)
num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
custom_mask = torch.ones(
(max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_bs,
dtype=torch.bool,
)
next_token_logits_buffer = torch.zeros(
(max_num_token, vocab_size),
dtype=torch.float,
)
mamba_track_indices = (
torch.zeros((max_bs,), dtype=torch.int64) if enable_mamba_track else None
)
mamba_track_mask = (
torch.zeros((max_bs,), dtype=torch.bool) if enable_mamba_track else None
)
if pp_size > 1:
# mHC (e.g. DSV4) flattens residual into hidden_states (size = hc_hidden_size).
is_mhc = hc_hidden_size is not None
hs = hc_hidden_size if is_mhc else hidden_size
pp_proxy_tensors = {
"hidden_states": torch.zeros((max_bs, hs), dtype=dtype),
}
if not is_mhc:
pp_proxy_tensors["residual"] = torch.zeros(
(max_bs, hidden_size), dtype=dtype
)
if pp_proxy_topk_size is not None:
pp_proxy_tensors["topk_indices"] = torch.zeros(
(max_num_token, pp_proxy_topk_size), dtype=torch.int32
)
else:
pp_proxy_tensors = None
if is_encoder_decoder:
encoder_lens = torch.full(
(max_bs,), encoder_len_fill_value, dtype=torch.int32
)
else:
encoder_lens = None
if require_mlp_tp_gather:
global_num_tokens_gpu = torch.zeros((dp_size,), dtype=torch.int32)
global_num_tokens_for_logprob_gpu = torch.zeros(
(dp_size,), dtype=torch.int32
)
else:
global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
global_num_tokens_for_logprob_gpu = torch.zeros((1,), dtype=torch.int32)
ngram_embedding_info = (
NgramEmbeddingInfo(
token_table=ne_token_table,
column_starts=torch.zeros([max_bs], dtype=torch.int32),
req_lens=torch.ones([max_bs], dtype=torch.int32),
out_column_starts=torch.zeros([max_bs], dtype=torch.int32),
out_req_lens=torch.ones([max_bs], dtype=torch.int32),
)
if ne_token_table is not None
else None
)
if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get():
rids_int = torch.zeros((max_bs,), dtype=torch.int64)
bootstrap_room_ids_int = torch.full((max_bs,), -1, dtype=torch.int64)
else:
rids_int = None
bootstrap_room_ids_int = None
seq_lens_cpu = torch.full(
(max_bs,),
seq_len_fill_value,
dtype=torch.int64,
device="cpu",
)
return SimpleNamespace(
input_ids=input_ids,
input_embeds=input_embeds,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
out_cache_loc=out_cache_loc,
positions=positions,
mrope_positions=mrope_positions,
num_token_non_padded=num_token_non_padded,
custom_mask=custom_mask,
next_token_logits_buffer=next_token_logits_buffer,
mamba_track_indices=mamba_track_indices,
mamba_track_mask=mamba_track_mask,
encoder_lens=encoder_lens,
global_num_tokens_gpu=global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
pp_proxy_tensors=pp_proxy_tensors,
ngram_embedding_info=ngram_embedding_info,
rids_int=rids_int,
bootstrap_room_ids_int=bootstrap_room_ids_int,
)
class DecodeCudaGraphRunner(BaseCudaGraphRunner): class DecodeCudaGraphRunner(BaseCudaGraphRunner):
"""Decode-phase CUDA graph runner. """Decode-phase CUDA graph runner.
@@ -775,6 +642,18 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
return forward_batch, attn_backend, pp_proxy_tensors return forward_batch, attn_backend, pp_proxy_tensors
def capture(self) -> None: def capture(self) -> None:
# Warm up + autotune kernels once before capture (run-once across the
# decode + prefill runners; see BaseRunner.warmup).
self.warmup()
# warmup() may disable torch.compile for a model whose _can_torch_compile
# is False; recompute the compile bucket so capture matches.
if self.enable_torch_compile and not (
self.model_runner.server_args.enable_torch_compile
):
self.enable_torch_compile = False
_, self.compile_bs = get_batch_sizes_to_capture(
self.model_runner, self.num_tokens_per_bs
)
profile_context = empty_context() profile_context = empty_context()
if self.enable_profile_cuda_graph: if self.enable_profile_cuda_graph:
profile_context = self._init_profile_context_and_memory_record() profile_context = self._init_profile_context_and_memory_record()
@@ -116,6 +116,8 @@ class EagerRunner(BaseRunner):
), ),
dp_size=sa.dp_size, dp_size=sa.dp_size,
) )
# Eager has no capture step, so warm up here (run-once via mr._kernel_warmed_up).
self.warmup()
def can_run_graph(self, forward_batch: ForwardBatch) -> bool: def can_run_graph(self, forward_batch: ForwardBatch) -> bool:
# Eager never runs a cuda graph; callers dispatch on isinstance(..., # Eager never runs a cuda graph; callers dispatch on isinstance(...,
@@ -587,6 +587,9 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
return forward_batch, self.model_runner.attn_backend return forward_batch, self.model_runner.attn_backend
def capture(self) -> None: def capture(self) -> None:
# Warm up + autotune kernels once before capture (run-once across the
# decode + prefill runners; see BaseRunner.warmup).
self.warmup()
with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc): with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc):
with graph_capture() as graph_capture_context: with graph_capture() as graph_capture_context:
self.stream = graph_capture_context.stream self.stream = graph_capture_context.stream
@@ -325,6 +325,11 @@ class MockModelRunner(ModelRunner):
self.pp_size = 1 self.pp_size = 1
self.is_draft_worker = False self.is_draft_worker = False
self.spec_algorithm = SpeculativeAlgorithm.NONE self.spec_algorithm = SpeculativeAlgorithm.NONE
# The runner lifecycle warms up kernels in capture() / first execute()
# via BaseRunner.warmup(); this mock never calls init_backends and has no
# real kernels to warm up, so mark it done (warmup becomes a no-op for
# the runner-mode attention tests that drive capture directly).
self._kernel_warmed_up = True
speculative_num_draft_tokens = ( speculative_num_draft_tokens = (
max(case.input_lens) max(case.input_lens)
if case.forward_mode.is_target_verify() if case.forward_mode.is_target_verify()
@@ -306,6 +306,7 @@ class DSAMockModelRunner(ModelRunner):
self.page_size = case.page_size self.page_size = case.page_size
self.model_config = model_config self.model_config = model_config
self.tp_size = 1 self.tp_size = 1
self._kernel_warmed_up = True
self.dp_size = 1 self.dp_size = 1
self.pp_size = 1 self.pp_size = 1
self.server_args = make_mock_server_args( self.server_args = make_mock_server_args(
@@ -412,6 +412,7 @@ class MockDSV4ModelRunner:
self.sliding_window_size = DSV4_SWA_WINDOW self.sliding_window_size = DSV4_SWA_WINDOW
self.use_mla_backend = True self.use_mla_backend = True
self.is_draft_worker = False self.is_draft_worker = False
self._kernel_warmed_up = True
@property @property
def hybrid_gdn_config(self): def hybrid_gdn_config(self):
@@ -321,6 +321,7 @@ class DualChunkMockModelRunner(ModelRunner):
self.page_size = case.page_size self.page_size = case.page_size
self.model_config = model_config self.model_config = model_config
self.tp_size = 1 self.tp_size = 1
self._kernel_warmed_up = True
self.dp_size = 1 self.dp_size = 1
self.pp_size = 1 self.pp_size = 1
self.server_args = make_mock_server_args( self.server_args = make_mock_server_args(
@@ -304,6 +304,7 @@ class MockGDNModelRunner(ModelRunner):
self.sliding_window_size = None self.sliding_window_size = None
self.use_mla_backend = False self.use_mla_backend = False
self.is_draft_worker = False self.is_draft_worker = False
self._kernel_warmed_up = True
@property @property
def hybrid_gdn_config(self): def hybrid_gdn_config(self):
@@ -310,6 +310,7 @@ class MockKDAModelRunner(ModelRunner):
self.sliding_window_size = None self.sliding_window_size = None
self.use_mla_backend = False self.use_mla_backend = False
self.is_draft_worker = False self.is_draft_worker = False
self._kernel_warmed_up = True
@property @property
def hybrid_gdn_config(self): def hybrid_gdn_config(self):
@@ -319,6 +319,7 @@ class MockLightningModelRunner(ModelRunner):
self.sliding_window_size = None self.sliding_window_size = None
self.use_mla_backend = False self.use_mla_backend = False
self.is_draft_worker = False self.is_draft_worker = False
self._kernel_warmed_up = True
@property @property
def hybrid_gdn_config(self): def hybrid_gdn_config(self):
@@ -454,6 +454,7 @@ class MockMamba2ModelRunner(ModelRunner):
self.sliding_window_size = None self.sliding_window_size = None
self.use_mla_backend = False self.use_mla_backend = False
self.is_draft_worker = False self.is_draft_worker = False
self._kernel_warmed_up = True
@property @property
def hybrid_gdn_config(self): def hybrid_gdn_config(self):
@@ -306,6 +306,7 @@ class MockMLAModelRunner(ModelRunner):
self.sliding_window_size = None self.sliding_window_size = None
self.use_mla_backend = True self.use_mla_backend = True
self.is_draft_worker = False self.is_draft_worker = False
self._kernel_warmed_up = True
@property @property
def hybrid_gdn_config(self): def hybrid_gdn_config(self):