Speed up DeepGEMM JIT warmup with per-PP-rank parallel compile (#26567)

This commit is contained in:
ybyang
2026-06-01 19:51:27 -07:00
committed by GitHub
parent 0574d2b8a5
commit 9fe8b72912
3 changed files with 106 additions and 10 deletions
+1
View File
@@ -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)
@@ -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,
)
@@ -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,