Prhub

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

原始 PR 作者 whybeyoung 合并时间 2026-06-02 10:51 文件变更 3 提交数 12 评论 18 代码增减 +106 / -10

执行摘要

PP 并行 DeepGEMM JIT 编译显著缩短启动时间

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.

该 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 的细节处理。
讨论亮点
  • 环境变量保护: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 建议将 timeForwardModeceil_align 等导入移至模块顶部,避免函数内懒导入,作者接受。

  • _dummy_runhc_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_runglobal_num_tokens_cpu 未设置导致 NSA 后端崩溃的问题。

实现拆解

  1. model_runner.pykernel_warmup 中新增判断:当 SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP 开启、ENABLE_JIT_DEEPGEMM 为真、pp_size > 1 且非 spec 模式时,调用 pp_parallel_deep_gemm_warmup(self)

  2. 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=DECODEEXTEND 调用 model_runner._dummy_run,并考虑 PD 分解场景跳过不合适的模式。

  3. 修改 _dummy_run 方法:新增 forward_mode_override 参数,优先使用它决定 capture_forward_mode,从而允许生成模型也强制进入 EXTEND 路径;将原 not self.is_generation 判断改为 capture_forward_mode == ForwardMode.EXTEND;传递 hc_hidden_sizeDecodeInputBuffers.create 以支持 DeepSeek-v4 mHC;在 PP 且 EXTEND 且 attn_cp_size>1 时切片 pp_proxy_tensors 以匹配 CP 后消息大小;在 forward_info.global_num_tokens_cpu 中填入相应值避免 NSA 后端崩溃。

  4. environ.py 中注册 SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP = EnvBool(False)

  5. 根据 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 模型执行器 modified 7.63
python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py DeepGEMM 封装 modified 7.69
python/sglang/srt/environ.py 环境配置 modified 4.35

关键符号

pp_parallel_deep_gemm_warmup _dummy_run kernel_warmup

关键源码片段

python/sglang/srt/model_executor/model_runner.py data-contract

核心入口:修改 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 core-logic

新增 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,
    )

评论区精华

将 PP 并行 warmup 置于环境变量控制 设计

Fridge003 指出该函数不具有通用性,应受环境变量保护。

结论:作者添加了 SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP,默认关闭。 · 已解决

将 warmup 函数移至 DeepGEMM 模块 设计

Fridge003 建议将 pp_parallel_deep_gemm_warmup 移到 compile_utils.py,因为它是 DeepGEMM warmup 逻辑的一部分。

结论:作者采纳,函数移至 compile_utils.py。 · 已解决

Batch size 选择依据 question

Fridge003 询问 'Why these bs?',作者详细解释了基于 n_sms 和 block_m 的启发式设计。

结论:作者给出了清晰的解释,最终实现保留。 · 已解决

修正 _dummy_run 模式判断 正确性

BBuf 确认 capture_forward_mode 替代 not self.is_generation 是正确的,且原本有潜在 bug。

结论:修改被接受。 · 已解决

风险与影响

  1. 默认关闭:功能由环境变量 opt-in,默认不影响现有部署,风险较低。
  2. 模型兼容性:batch size 推导基于 DeepSeek-V4 的 DSA prefill CP 场景,对其他模型(尤其无 CP 或不同 GEMM 形状)可能产生未预编译的 kernel 或多余编译。但功能 opt-in,用户可选择不启用。
  3. OOM 风险:早期版本曾使用 chunked_prefill_size 等大 batch 导致 OOM,后改为 SM 数量推导的小 batch(最大 n_sms*block_m 可能仍较大,但经过 try/except 移除后错误会直接暴露,需确保显存充足)。
  4. PD 分解场景:已添加 disaggregation_mode 判断,prefill-only 或 decode-only 节点会跳过不适用的 forward 模式,避免 indexer OOM。
  5. 依赖顺序:在 kernel_warmup 中调用,早于 CUDA graph capture,不会干扰后续流程。

对用户:大模型 PP 部署启动速度显著提升(实测 12min→4min),提升用户体验和运维效率。对系统:新增环境变量控制,不影响默认行为;代码主要改动在 warmup 阶段,不涉及在线推理路径。对开发者:_dummy_runforward_mode_override 参数可被未来其他 warmup 或 profiling 场景复用。

opt-in by env variable 模型特定(DeepSeek-V4) batch size 可能不通用 OOM 风险 PD 分解兼容

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论