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_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,
|
||||||
|
|||||||
Reference in New Issue
Block a user