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_DG_USE_NVRTC = EnvBool(False)
SGLANG_USE_DEEPGEMM_BMM = EnvBool(False) SGLANG_USE_DEEPGEMM_BMM = EnvBool(False)
SGLANG_DEEPGEMM_SANITY_CHECK = EnvBool(False) SGLANG_DEEPGEMM_SANITY_CHECK = EnvBool(False)
SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP = EnvBool(False)
# DeepSeek MHA Optimization # DeepSeek MHA Optimization
SGLANG_CHUNKED_PREFIX_CACHE_THRESHOLD = EnvInt(8192) SGLANG_CHUNKED_PREFIX_CACHE_THRESHOLD = EnvInt(8192)
@@ -1,5 +1,6 @@
import logging import logging
import os import os
import time
from contextlib import contextmanager, nullcontext from contextlib import contextmanager, nullcontext
from enum import IntEnum, auto from enum import IntEnum, auto
from typing import Dict, List, Tuple 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.environ import envs
from sglang.srt.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM 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.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__) logger = logging.getLogger(__name__)
@@ -450,3 +452,63 @@ def _deep_gemm_execution_hook(
if m > 0: if m > 0:
_maybe_compile_deep_gemm_one_type_all(kernel_type, n, k, num_groups) _maybe_compile_deep_gemm_one_type_all(kernel_type, n, k, num_groups)
yield 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) return attn_backend_wrapper(self, full_attention_backend)
def kernel_warmup(self): def kernel_warmup(self):
""" """Warmup and tune kernels before cuda graph capture."""
Warmup and tune kernels before cuda graph capture.
Currently only doing FlashInfer autotune.
"""
if self.device != "cuda": if self.device != "cuda":
return return
if self._should_run_flashinfer_autotune(): if self._should_run_flashinfer_autotune():
self._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): def _pre_initialize_flashinfer_allreduce_workspace(self):
"""Pre-initialize flashinfer allreduce fusion workspaces. """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" / 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): def _dummy_run(
"""Run a dummy forward pass for warmup/profiling.""" self,
if self.is_generation: 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 capture_forward_mode = ForwardMode.DECODE
else: else:
capture_forward_mode = ForwardMode.EXTEND capture_forward_mode = ForwardMode.EXTEND
@@ -2546,11 +2566,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
num_tokens_per_bs=num_tokens_per_bs, num_tokens_per_bs=num_tokens_per_bs,
cache_loc_dtype=torch.int64, cache_loc_dtype=torch.int64,
enable_mamba_track=False, enable_mamba_track=False,
hc_hidden_size=getattr(self.model_config, "hc_hidden_size", None),
) )
buffers.num_token_non_padded[...] = num_tokens buffers.num_token_non_padded[...] = num_tokens
# For extend mode # For extend mode
if not self.is_generation: if capture_forward_mode == ForwardMode.EXTEND:
extend_prefix_lens_cpu = [0] * batch_size extend_prefix_lens_cpu = [0] * batch_size
extend_seq_lens_cpu = [seq_len_fill_value] * batch_size extend_seq_lens_cpu = [seq_len_fill_value] * batch_size
extend_num_tokens = num_tokens extend_num_tokens = num_tokens
@@ -2572,8 +2593,16 @@ class ModelRunner(ModelRunnerKVCacheMixin):
extend_start_loc = None extend_start_loc = None
if self.server_args.pp_size > 1: 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( 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_: if require_mlp_tp_gather_:
@@ -2592,6 +2621,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
) )
) )
global_dp_buffer_len = num_tokens * self.server_args.dp_size 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): elif require_attn_tp_gather(self.server_args):
buffers.global_num_tokens_gpu.copy_( buffers.global_num_tokens_gpu.copy_(
torch.tensor( torch.tensor(
@@ -2608,8 +2638,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
) )
) )
global_dp_buffer_len = num_tokens global_dp_buffer_len = num_tokens
global_num_tokens_cpu = [num_tokens]
else: else:
global_dp_buffer_len = None global_dp_buffer_len = None
global_num_tokens_cpu = None
def get_spec_info(): def get_spec_info():
spec_info = None spec_info = None
@@ -2698,6 +2730,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
extend_prefix_lens_cpu=extend_prefix_lens_cpu, extend_prefix_lens_cpu=extend_prefix_lens_cpu,
extend_seq_lens_cpu=extend_seq_lens_cpu, extend_seq_lens_cpu=extend_seq_lens_cpu,
global_num_tokens_gpu=buffers.global_num_tokens_gpu, 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, global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(), dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
global_dp_buffer_len=global_dp_buffer_len, global_dp_buffer_len=global_dp_buffer_len,