执行摘要
- 一句话:PP 并行 DeepGEMM JIT 编译显著缩短启动时间
- 推荐动作:该 PR 设计清晰、讨论充分,值得精读。重点关注:
- 如何将 PP 串行编译变为并行:利用
kernel_warmup 在每个 rank 独立执行 dummy forward。
- Batch size 的推导方法:基于 GPU SM 数和 block_m 的启发式,覆盖 n_splits 范围。
- 对
_dummy_run 的重构:引入 forward_mode_override,使生成模型也能触发 EXTEND 路径,并修复了原逻辑对 is_generation 的依赖。
- 兼容 PD 分解和 DSA prefill CP 的细节处理。
功能与动机
Today the startup warmup /generate request flows through PP stages serially: stage k can only start its DeepGEMM JIT compile after stage k-1 finishes. For large MoE models like DeepSeek-V4 this dominates startup time.
实现拆解
-
在 model_runner.py 的 kernel_warmup 中新增判断:当 SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP 开启、ENABLE_JIT_DEEPGEMM 为真、pp_size > 1 且非 spec 模式时,调用 pp_parallel_deep_gemm_warmup(self)。
-
在 compile_utils.py 中新增 pp_parallel_deep_gemm_warmup 函数:根据 GPU SM 数量和 block_m=64 推导 5 个代表性 batch_size(覆盖解码小形状到数千 token 预填),ceil-align 到 attn_cp_size;然后依次以 forward_mode_override=DECODE 和 EXTEND 调用 model_runner._dummy_run,并考虑 PD 分解场景跳过不合适的模式。
-
修改 _dummy_run 方法:新增 forward_mode_override 参数,优先使用它决定 capture_forward_mode,从而允许生成模型也强制进入 EXTEND 路径;将原 not self.is_generation 判断改为 capture_forward_mode == ForwardMode.EXTEND;传递 hc_hidden_size 到 DecodeInputBuffers.create 以支持 DeepSeek-v4 mHC;在 PP 且 EXTEND 且 attn_cp_size>1 时切片 pp_proxy_tensors 以匹配 CP 后消息大小;在 forward_info.global_num_tokens_cpu 中填入相应值避免 NSA 后端崩溃。
-
在 environ.py 中注册 SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP = EnvBool(False)。
-
根据 review 反馈调整:将函数体从 model_runner.py 移至 compile_utils.py;将 import 语句提到文件顶部;移除 try/except 让错误直接暴露;添加 n_sms*block_m 大 batch 以确保 n_splits=1 也被编译;用 disaggregation_mode 跳过不适用节点。
关键文件:
python/sglang/srt/model_executor/model_runner.py(模块 模型执行器;类别 source;类型 data-contract;符号 _dummy_run, kernel_warmup): 核心入口:修改 kernel_warmup 添加并行 warmup 调用;重构 _dummy_run 增加 forward_mode_override 和重要修复
python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py(模块 DeepGEMM 封装;类别 source;类型 core-logic;符号 pp_parallel_deep_gemm_warmup, _deep_gemm_execution_hook): 新增 pp_parallel_deep_gemm_warmup 函数,实现并行 dummpy forward 的 batch size 推导和调度
python/sglang/srt/environ.py(模块 环境配置;类别 source;类型 configuration): 注册环境变量 SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP,作为 opt-in 开关
关键符号:pp_parallel_deep_gemm_warmup, _dummy_run, kernel_warmup
关键源码片段
python/sglang/srt/model_executor/model_runner.py
核心入口:修改 kernel_warmup 添加并行 warmup 调用;重构 _dummy_run 增加 forward_mode_override 和重要修复
# 在 kernel_warmup 中插入 PP 并行 DeepGEMM warmup
# 条件:环境变量启用、JIT DeepGEMM 可用、pp_size > 1、非 speculative 模式
def kernel_warmup(self):
"""Warmup and tune kernels before cuda graph capture."""
if self.device != "cuda":
return
if self._should_run_flashinfer_autotune():
self._flashinfer_autotune()
# PP 并行 DeepGEMM warmup:每个 rank 独立触发 JIT 编译
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)
# _dummy_run 增加 forward_mode_override 参数
# 允许强制以 DECODE 或 EXTEND 模式运行 dummy forward,
# 从而让生成模型也能触发 EXTEND 路径的 JIT 编译。
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
# ... 后续逻辑(省略,关键修改如上)
python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py
新增 pp_parallel_deep_gemm_warmup 函数,实现并行 dummpy forward 的 batch size 推导和调度
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.
n_sms = torch.cuda.get_device_properties(
model_runner.device
).multi_processor_count
block_m = 64
cp = max(model_runner.attn_cp_size, 1)
# ceil-align to attn_cp_size for DSA prefill CP
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, # n_splits=1,覆盖大预填
)
}
)
# 考虑 PD 分解:prefill-only 节点不跑 DECODE,decode-only 节点不跑 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,
)
评论区精华
-
环境变量保护:Fridge003 指出该功能不应默认启用,理由是对非 DeepSeek 模型未验证且 batch size 选择具有针对性,因此改为由 SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP 控制(默认关闭)。
-
函数放置:Fridge003 建议将 pp_parallel_deep_gemm_warmup 移至 compile_utils.py 作为 DeepGEMM warmup 逻辑的一部分,作者照做。
-
Batch size 选择依据:为何选择这 5 个 batch size?作者解释是为了覆盖 n_splits 的不同 bracket,由 SM 数和 block_m 推导,确保各种形状的 JIT 内核都被预编译。
-
导入位置:BBuf 建议将 time、ForwardMode、ceil_align 等导入移至模块顶部,避免函数内懒导入,作者接受。
-
_dummy_run 中 hc_hidden_size 的传递:Fridge003 询问原因,作者解释 DeepSeek-v4 mHC 需要该参数以正确分配 PP IPC buffer。
-
not self.is_generation 改为 capture_forward_mode == ForwardMode.EXTEND:BBuf 确认该修改是正确的,因为原逻辑在生成模型 + override=EXTEND 时会跳过 extend 设置,属潜在 bug。
-
DSA prefill CP 的 pp_proxy_tensors 切片:Fridge003 建议去掉 DSA 专用判断,改为通用条件,作者实现为仅当 capture_forward_mode == EXTEND and self.pp_rank != 0 and self.attn_cp_size > 1 时切片。
-
全局 num_tokens_cpu 缺失:最后一轮修复了 _dummy_run 中 global_num_tokens_cpu 未设置导致 NSA 后端崩溃的问题。
- 将 PP 并行 warmup 置于环境变量控制 (design): 作者添加了 SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP,默认关闭。
- 将 warmup 函数移至 DeepGEMM 模块 (design): 作者采纳,函数移至 compile_utils.py。
- Batch size 选择依据 (question): 作者给出了清晰的解释,最终实现保留。
- 修正 _dummy_run 模式判断 (correctness): 修改被接受。
风险与影响
- 风险:
- 默认关闭:功能由环境变量 opt-in,默认不影响现有部署,风险较低。
- 模型兼容性:batch size 推导基于 DeepSeek-V4 的 DSA prefill CP 场景,对其他模型(尤其无 CP 或不同 GEMM 形状)可能产生未预编译的 kernel 或多余编译。但功能 opt-in,用户可选择不启用。
- OOM 风险:早期版本曾使用
chunked_prefill_size 等大 batch 导致 OOM,后改为 SM 数量推导的小 batch(最大 n_sms*block_m 可能仍较大,但经过 try/except 移除后错误会直接暴露,需确保显存充足)。
- PD 分解场景:已添加
disaggregation_mode 判断,prefill-only 或 decode-only 节点会跳过不适用的 forward 模式,避免 indexer OOM。
- 依赖顺序:在
kernel_warmup 中调用,早于 CUDA graph capture,不会干扰后续流程。
- 影响:对用户:大模型 PP 部署启动速度显著提升(实测 12min→4min),提升用户体验和运维效率。对系统:新增环境变量控制,不影响默认行为;代码主要改动在 warmup 阶段,不涉及在线推理路径。对开发者:_dummy_run 的 forward_mode_override 参数可被未来其他 warmup 或 profiling 场景复用。
- 风险标记:opt-in by env variable, 模型特定(DeepSeek-V4), batch size 可能不通用, OOM 风险, PD 分解兼容
关联脉络
参与讨论