From 856b0dc74b774cfbacc4b4d01d43224f000de14a Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Fri, 19 Jun 2026 13:13:13 -0700 Subject: [PATCH] refactor(runner): move kernel warmup into the shared runner lifecycle (warmup()) (#28739) Co-authored-by: Claude Opus 4.8 (1M context) --- .../layers/deep_gemm_wrapper/compile_utils.py | 11 +- .../sglang/srt/model_executor/model_runner.py | 518 ++------------- .../srt/model_executor/runner/__init__.py | 2 +- .../srt/model_executor/runner/base_runner.py | 590 +++++++++++++++++- .../runner/decode_cuda_graph_runner.py | 145 +---- .../srt/model_executor/runner/eager_runner.py | 2 + .../runner/prefill_cuda_graph_runner.py | 3 + .../attention_methods/dense_attention.py | 5 + .../attention_methods/dsa_attention.py | 1 + .../attention_methods/dsv4_attention.py | 1 + .../attention_methods/dual_chunk_attention.py | 1 + .../attention_methods/gdn_attention.py | 1 + .../attention_methods/kda_attention.py | 1 + .../attention_methods/lightning_attention.py | 1 + .../attention_methods/mamba2_attention.py | 1 + .../attention_methods/mla_attention.py | 1 + 16 files changed, 678 insertions(+), 606 deletions(-) diff --git a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py index e46f29980..11183bc9a 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py @@ -416,12 +416,17 @@ def _deep_gemm_execution_hook( 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 DeepGEMM JIT compiles in parallel instead of serially via the warmup /generate flowing through the pipeline. Opt-in via 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 # cover the brackets real /generate hits (smallest decode shape, # 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(): for bs in batch_sizes: if run_decode: - model_runner._dummy_run( + runner._dummy_run( batch_size=bs, forward_mode_override=ForwardMode.DECODE ) if run_extend: - model_runner._dummy_run( + runner._dummy_run( batch_size=bs, forward_mode_override=ForwardMode.EXTEND ) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 020d71b35..f3fe4daa0 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -18,7 +18,6 @@ from __future__ import annotations import contextlib import datetime import gc -import hashlib import inspect import logging import os @@ -27,7 +26,6 @@ import threading import time from collections import defaultdict from dataclasses import dataclass -from pathlib import Path from typing import Any, Callable, List, Optional, Tuple, Union import torch @@ -35,7 +33,6 @@ import torch.distributed as dist from torch import nn 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 ( BailingHybridConfig, FalconH1Config, @@ -129,12 +126,8 @@ from sglang.srt.layers.cp.utils import ( get_cp_strategy, ) from sglang.srt.layers.dp_attention import ( - DpPaddingMode, get_attention_tp_group, - get_attention_tp_size, initialize_dp_attention, - set_dp_buffer_len, - set_is_extend_in_batch, ) from sglang.srt.layers.logits_processor import LogitsProcessorOutput 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, ) from sglang.srt.model_executor.forward_batch_info import ( - CaptureHiddenMode, ForwardBatch, ForwardMode, PPProxyTensors, @@ -175,9 +167,6 @@ from sglang.srt.model_executor.runner import ( EagerRunner, 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.remote_instance_weight_loader_utils import ( RemoteInstanceWeightLoaderBackend, @@ -193,10 +182,7 @@ from sglang.srt.server_args import ( get_global_server_args, set_global_server_args_for_scheduler, ) -from sglang.srt.speculative.spec_info import ( - SpeculativeAlgorithm, - create_dummy_verify_input, -) +from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.state_capturer.base import TopkCaptureOutput from sglang.srt.state_capturer.indexer_topk import ( create_indexer_capturer, @@ -213,7 +199,6 @@ from sglang.srt.utils import ( broadcast_pyobj, cpu_has_amx_support, dynamic_import, - empty_context, enable_show_time_cost, get_available_gpu_memory, get_bool_env_var, @@ -224,14 +209,11 @@ from sglang.srt.utils import ( is_npu, log_info_on_rank0, monkey_patch_p2p_access_check, - require_attn_tp_gather, require_gathered_buffer, - require_mlp_tp_gather, reserve_rope_cache_for_long_sequences, set_cuda_arch, 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.nvtx_pytorch_hooks import PytHooks from sglang.srt.utils.nvtx_utils import profile_range @@ -889,22 +871,14 @@ class ModelRunner(ModelRunnerKVCacheMixin): # runs with aux hidden state capture enabled. self.init_aux_hidden_state_capture() - # The eager (no-cuda-graph) phase runner. Always built: it serves both - # the fully-disabled case (decode/prefill runners point at it) and the - # eager fallback when a cuda-graph runner can't run a batch. - self.eager_runner = EagerRunner(self) - + # Device-specific attention-backend init. cg_supported gates decode + # cuda-graph capture (out-of-tree / unknown platforms may not support it). + cg_supported = True if self.device == "cuda" or self.device == "musa": self.init_cublas() 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": self.init_attention_backend() - if not disable_cuda_graph: - self.init_decode_cuda_graph() elif self.device == "npu": self.init_attention_backend() # 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().cpu_group, ) - if not disable_cuda_graph: - self.init_decode_cuda_graph() elif current_platform.is_out_of_tree(): self.init_attention_backend() - if current_platform.support_cuda_graph() and not disable_cuda_graph: - self.init_decode_cuda_graph() - else: - self.decode_cuda_graph_runner = None - self.graph_mem_usage = 0 + cg_supported = current_platform.support_cuda_graph() + else: + self.init_attention_backend() + cg_supported = False + + # 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: self.decode_cuda_graph_runner = None self.graph_mem_usage = 0 - self.init_attention_backend() - if disable_cuda_graph: # Decode cuda graph disabled: route eager decode through the # EagerRunner (the dispatch gate isinstance(..., EagerRunner) keeps # _forward_raw off any replay branch). 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: register_forward_hooks(self.model, server_args.forward_hooks) - self.init_prefill_cuda_graph() - self.prealloc_symmetric_memory_pool() 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) 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): self.use_ngram_embedding = self.model_config.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: 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 if not hasattr(self.model, "model"): logger.warning( diff --git a/python/sglang/srt/model_executor/runner/__init__.py b/python/sglang/srt/model_executor/runner/__init__.py index 0422f4457..746eed72a 100644 --- a/python/sglang/srt/model_executor/runner/__init__.py +++ b/python/sglang/srt/model_executor/runner/__init__.py @@ -7,7 +7,7 @@ BaseCudaGraphBackend chosen via cuda_graph_config. Public API: - 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). - BaseCudaGraphRunner — abstract cuda-graph base; bucket padding + capture-loop scaffolding on top of BaseRunner. diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 74e7f1b7a..6b72420b7 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -15,22 +15,176 @@ from __future__ import annotations +import datetime +import hashlib +import inspect import logging 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 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.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: - from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.model_runner import ModelRunner 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): def __init__(self, model_runner: ModelRunner) -> None: self.model_runner = model_runner @@ -43,6 +197,438 @@ class BaseRunner(ABC): self.attn_tp_rank = get_parallel().attn_tp_rank 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 def can_run_graph(self, forward_batch: ForwardBatch) -> bool: ... diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index da4b877cb..52b12509d 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -43,7 +43,6 @@ from sglang.srt.distributed.parallel_state import ( set_pdmux_status, ) 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.dp_attention import ( DpPaddingMode, @@ -62,7 +61,6 @@ from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, ForwardBatch, ForwardMode, - NgramEmbeddingInfo, PPProxyTensors, compute_local_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): """Decode-phase CUDA graph runner. @@ -775,6 +642,18 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): return forward_batch, attn_backend, pp_proxy_tensors 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() if self.enable_profile_cuda_graph: profile_context = self._init_profile_context_and_memory_record() diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index 4755f9ad3..456c5ac23 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -116,6 +116,8 @@ class EagerRunner(BaseRunner): ), 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: # Eager never runs a cuda graph; callers dispatch on isinstance(..., diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index 3da471de5..dba9899e7 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -587,6 +587,9 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): return forward_batch, self.model_runner.attn_backend 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 graph_capture() as graph_capture_context: self.stream = graph_capture_context.stream diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py index 8da926f89..fb51ed4f5 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py @@ -325,6 +325,11 @@ class MockModelRunner(ModelRunner): self.pp_size = 1 self.is_draft_worker = False 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 = ( max(case.input_lens) if case.forward_mode.is_target_verify() diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py index 51230f641..f31a3d328 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py @@ -306,6 +306,7 @@ class DSAMockModelRunner(ModelRunner): self.page_size = case.page_size self.model_config = model_config self.tp_size = 1 + self._kernel_warmed_up = True self.dp_size = 1 self.pp_size = 1 self.server_args = make_mock_server_args( diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py index 5a6ce04cd..9ab3aaa41 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py @@ -412,6 +412,7 @@ class MockDSV4ModelRunner: self.sliding_window_size = DSV4_SWA_WINDOW self.use_mla_backend = True self.is_draft_worker = False + self._kernel_warmed_up = True @property def hybrid_gdn_config(self): diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py index 5a896abdb..55290bf9b 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py @@ -321,6 +321,7 @@ class DualChunkMockModelRunner(ModelRunner): self.page_size = case.page_size self.model_config = model_config self.tp_size = 1 + self._kernel_warmed_up = True self.dp_size = 1 self.pp_size = 1 self.server_args = make_mock_server_args( diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py index e8ce45411..f7e8e5b42 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py @@ -304,6 +304,7 @@ class MockGDNModelRunner(ModelRunner): self.sliding_window_size = None self.use_mla_backend = False self.is_draft_worker = False + self._kernel_warmed_up = True @property def hybrid_gdn_config(self): diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py index ff6faa324..70328bd2c 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py @@ -310,6 +310,7 @@ class MockKDAModelRunner(ModelRunner): self.sliding_window_size = None self.use_mla_backend = False self.is_draft_worker = False + self._kernel_warmed_up = True @property def hybrid_gdn_config(self): diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py index 23e729ace..72910077d 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py @@ -319,6 +319,7 @@ class MockLightningModelRunner(ModelRunner): self.sliding_window_size = None self.use_mla_backend = False self.is_draft_worker = False + self._kernel_warmed_up = True @property def hybrid_gdn_config(self): diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py index 11ff4e9c7..998acc31f 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py @@ -454,6 +454,7 @@ class MockMamba2ModelRunner(ModelRunner): self.sliding_window_size = None self.use_mla_backend = False self.is_draft_worker = False + self._kernel_warmed_up = True @property def hybrid_gdn_config(self): diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py index 890be7273..f3db98e5b 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py @@ -306,6 +306,7 @@ class MockMLAModelRunner(ModelRunner): self.sliding_window_size = None self.use_mla_backend = True self.is_draft_worker = False + self._kernel_warmed_up = True @property def hybrid_gdn_config(self):