Speed up DeepGEMM JIT warmup with per-PP-rank parallel compile (#26567)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user