From 9fe8b7291212d3f8eeda8c37214b8621336ab2fe Mon Sep 17 00:00:00 2001 From: ybyang <10629930+whybeyoung@users.noreply.github.com> Date: Tue, 2 Jun 2026 10:51:27 +0800 Subject: [PATCH] Speed up DeepGEMM JIT warmup with per-PP-rank parallel compile (#26567) --- python/sglang/srt/environ.py | 1 + .../layers/deep_gemm_wrapper/compile_utils.py | 64 ++++++++++++++++++- .../sglang/srt/model_executor/model_runner.py | 51 ++++++++++++--- 3 files changed, 106 insertions(+), 10 deletions(-) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 799db26cc..9ed722192 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -464,6 +464,7 @@ class Envs: SGLANG_DG_USE_NVRTC = EnvBool(False) SGLANG_USE_DEEPGEMM_BMM = EnvBool(False) SGLANG_DEEPGEMM_SANITY_CHECK = EnvBool(False) + SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP = EnvBool(False) # DeepSeek MHA Optimization SGLANG_CHUNKED_PREFIX_CACHE_THRESHOLD = EnvInt(8192) 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 f4ae6663c..5a6a49bfb 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py @@ -1,5 +1,6 @@ import logging import os +import time from contextlib import contextmanager, nullcontext from enum import IntEnum, auto from typing import Dict, List, Tuple @@ -13,8 +14,9 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( ) from sglang.srt.environ import envs from sglang.srt.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM +from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.server_args import ServerArgs -from sglang.srt.utils import ceil_div, get_available_gpu_memory, is_musa +from sglang.srt.utils import ceil_align, ceil_div, get_available_gpu_memory, is_musa logger = logging.getLogger(__name__) @@ -450,3 +452,63 @@ def _deep_gemm_execution_hook( if m > 0: _maybe_compile_deep_gemm_one_type_all(kernel_type, n, k, num_groups) yield + + +def pp_parallel_deep_gemm_warmup(model_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. + """ + # 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 + # to attn_cp_size for DSA prefill CP's seq_len % cp_size == 0 assert. + n_sms = torch.cuda.get_device_properties(model_runner.device).multi_processor_count + block_m = 64 + cp = max(model_runner.attn_cp_size, 1) + batch_sizes = sorted( + { + ceil_align(bs, cp) + for bs in ( + 1, + 2 * block_m, + max(n_sms // 8, 2) * block_m, + max(n_sms // 4, 4) * block_m, + n_sms * block_m, + ) + } + ) + + # In PD, prefill-only nodes never decode (indexer would OOM at large + # bs) and decode-only nodes never extend. + disagg_mode = model_runner.server_args.disaggregation_mode + run_decode = model_runner.is_generation and disagg_mode != "prefill" + run_extend = disagg_mode != "decode" + + logger.info( + "PP-parallel DeepGEMM warmup start " + "(pp_rank=%d, tp_rank=%d, batch_sizes=%s, disagg=%s).", + model_runner.pp_rank, + model_runner.tp_rank, + batch_sizes, + disagg_mode, + ) + + t0 = time.perf_counter() + with torch.inference_mode(): + for bs in batch_sizes: + if run_decode: + model_runner._dummy_run( + batch_size=bs, forward_mode_override=ForwardMode.DECODE + ) + if run_extend: + model_runner._dummy_run( + batch_size=bs, forward_mode_override=ForwardMode.EXTEND + ) + + logger.info( + "PP-parallel DeepGEMM warmup done in %.2fs (pp_rank=%d).", + time.perf_counter() - t0, + model_runner.pp_rank, + ) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 8d59c848d..1345846f6 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -2336,16 +2336,25 @@ class ModelRunner(ModelRunnerKVCacheMixin): return attn_backend_wrapper(self, full_attention_backend) def kernel_warmup(self): - """ - Warmup and tune kernels before cuda graph capture. - Currently only doing FlashInfer autotune. - """ + """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. @@ -2473,9 +2482,20 @@ class ModelRunner(ModelRunnerKVCacheMixin): / 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): - """Run a dummy forward pass for warmup/profiling.""" - if self.is_generation: + 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 @@ -2546,11 +2566,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): 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), ) buffers.num_token_non_padded[...] = num_tokens # For extend mode - if not self.is_generation: + 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 @@ -2572,8 +2593,16 @@ class ModelRunner(ModelRunnerKVCacheMixin): 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[:num_tokens] for k, v in buffers.pp_proxy_tensors.items()} + {k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()} ) if require_mlp_tp_gather_: @@ -2592,6 +2621,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) ) global_dp_buffer_len = num_tokens * self.server_args.dp_size + global_num_tokens_cpu = [num_tokens] * self.server_args.dp_size elif require_attn_tp_gather(self.server_args): buffers.global_num_tokens_gpu.copy_( torch.tensor( @@ -2608,8 +2638,10 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) ) global_dp_buffer_len = num_tokens + global_num_tokens_cpu = [num_tokens] else: global_dp_buffer_len = None + global_num_tokens_cpu = None def get_spec_info(): spec_info = None @@ -2698,6 +2730,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): 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,