Prhub

#27436 [diffusion] Enable breakable CUDA graph (BCG) for diffusion DiTs

原始 PR 作者 BBuf 合并时间 2026-07-08 14:45 文件变更 31 提交数 86 评论 14 代码增减 +2952 / -503

执行摘要

为扩散 DiT 模型启用 Breakable CUDA Graph

扩散DiT模型的前向中包含动态注意力等无法被单一CUDA Graph捕获的操作,导致整个前向必须运行在Eager模式。Breakable CUDA Graph允许将图安全的操作(MLP、调制层等)捕获为多个图段,中间用Eager断点衔接,从而在保留动态灵活性同时最大化CUDA图收益。PR Body提供了Qwen-Image和Z-Image等模型的Profiler对比,直观显示BCG模式下GPU执行时间缩短。

该PR是扩散推理性能优化的关键里程碑,核心架构设计(BCG Runner + Prompt Padding)具有较高参考价值。建议团队阅读 runner.pybreakable_cuda_graph.pyprompt_padding.py 理解整体机制。未来添加新模型时务必参考已有padder实现,并在CI中添加对应BCG回归测试。

讨论亮点

Review中三个关键问题由gemini-code-assist[bot]提出:

  • 流同步竞态条件(Critical):图捕获在独立流上进行,回放时若当前流未等待输入拷贝完成,可能读取垃圾数据。BBuf通过 _replay 中的 event.synchronize() 解决。
  • _clone_output 不支持 dict/ModelOutput(High):若模型返回dict或ModelOutput,_clone_output直接返回引用导致输出被后续步骤覆写。BBuf添加了对dict和ModelOutput的克隆支持。
  • _weak_ref_if_tensor 缺少dict支持(Medium):中间张量不能弱引用,阻止共享mem pool回收。BBuf添加了dict分支。
    此外,Oasis-Git建议runner继承base、capture在初始化时完成、模型辅助函数分离到独立文件,这些后续被采纳。mickqian关注的 getattr 使用和初始化位置也已优化。

实现拆解

  1. BCG原语重定位:将 python/sglang/srt/model_executor/runner_backend_utils/breakable_cuda_graph/ 下的核心原语(BreakableCUDAGraphBreakableCUDAGraphCaptureeager_on_graph等)迁移到新的共享包 python/sglang/srt/breakable_cuda_graph/,并在原位置保留向后兼容的重新导出模块。新增 get_current_replay_token 用于扩散replay token机制。
  2. 扩散BCG Runner:新增 python/sglang/multimodal_gen/runtime/breakable_cuda_graph/runner.py,定义 BaseBreakableCUDAGraphRunnerDiffusionBreakableCUDAGraphRunner。Runner封装目标模块,处理签名提取(基于kwargs的形状和dtype)、图捕获、回放和输出克隆。签名不匹配时自动回落Eager。支持CPU输入到capture device的拷贝。
  3. Prompt Padding框架:新增 python/sglang/multimodal_gen/runtime/breakable_cuda_graph/prompt_padding.py,提供 first_tensorselect_text_bucketpad_tensor_dimpad_nested_dim等工具。通过将不同长度的文本条件填充到预定义桶,使得同一resolution的图可以跨prompt复用。桶选择采用最小适配原则,超出最大桶时退回Eager。
  4. Model-Specific Padders:为每个支持模型在 breakable_cuda_graph/model_padders/ 下实现专属pad函数:pad_qwen_prompt_kwargs(Qwen-Image)、pad_zimage_prompt_kwargs(Z-Image,需重建caption频率和mask)、pad_ideogram_prompt_kwargs(Ideogram,需解析indicator标记和动态分段mask),以及GLM-Image等模型的padder。
  5. 集成与服务参数:修改 python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py 中的 _predict_noise,插入 _maybe_get_bcg_runner_bcg_run 调用链。新增 server_args.py 中的 --enable-breakable-cuda-graph--bcg-text-buckets 参数,以及模型支持验证函数。Warmup阶段强制使用server-based warmup以捕获所有timestep和CFG分支。
  6. 测试与CI:新增 test_diffusion_bcg_service_validation.py 等测试文件,对支持模型添加BCG模式的端到端验证。调整部分CI配置以适应BCG引入的内存和精度阈值变化。
文件 模块 状态 重要度
python/sglang/multimodal_gen/runtime/breakable_cuda_graph/model_padders/zimage.py Z-Image 衬垫 added 9.35
python/sglang/multimodal_gen/runtime/breakable_cuda_graph/runner.py BCG 运行器 added 9.25
python/sglang/srt/breakable_cuda_graph/breakable_cuda_graph.py BCG 原语 added 9.25
python/sglang/multimodal_gen/runtime/breakable_cuda_graph/prompt_padding.py Prompt 衬垫 added 9.25
python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py 去噪阶段 modified 8.95
python/sglang/multimodal_gen/runtime/server_args.py 服务参数 modified 8.72

关键符号

BreakableCUDAGraph BreakableCUDAGraphCapture eager_on_graph break_graph BaseBreakableCUDAGraphRunner DiffusionBreakableCUDAGraphRunner pad_qwen_prompt_kwargs pad_zimage_prompt_kwargs pad_ideogram_prompt_kwargs pad_glm_image_prompt_kwargs select_text_bucket _bcg_run _maybe_get_bcg_runner

关键源码片段

python/sglang/multimodal_gen/runtime/breakable_cuda_graph/model_padders/zimage.py data-contract

Z-Image 模型专属 padder,展示最复杂的 caption 频率重建和 mask 处理逻辑,是理解 BCG 文本填充的典型范例。

# python/sglang/multimodal_gen/runtime/breakable_cuda_graph/model_padders/zimage.py
# 展示 pad_zimage_prompt_kwargs 核心逻辑def pad_zimage_prompt_kwargs(
    call_kwargs: dict, current_model: Any, buckets: tuple[int, ...]
) -> dict:
    # 从 encoder_hidden_states 中提取第一个 caption 张量
    caption = _first_caption_tensor(call_kwargs.get("encoder_hidden_states"))
    if caption is None:
        return call_kwargs
​
    seq = _caption_seq_len(caption)
    cap_freq = None
    freqs_cis = call_kwargs.get("freqs_cis")
    if isinstance(freqs_cis, (tuple, list)) and len(freqs_cis) == 2:
        cap_freq = bcg_utils.first_tensor(freqs_cis[0])
    cap_freq_len = int(cap_freq.shape[0]) if torch.is_tensor(cap_freq) else seq
​
    # 选择最小的适配桶,超出最大桶则回落 eager
    bucket = bcg_utils.select_text_bucket(max(seq, cap_freq_len), buckets)
    if bucket is None:
        return call_kwargs
​
    # 只保留 BCG 相关的关键字
    out = {
        key: value
        for key, value in call_kwargs.items()
        if key in {
            "hidden_states", "timestep", "guidance",
            "encoder_hidden_states", "encoder_attention_mask",
            "freqs_cis",
        }
    }
​
    # 填充 encoder_hidden_states 到 bucket 长度
    out["encoder_hidden_states"] = _pad_caption(
        out["encoder_hidden_states"], target=bucket
    )
    # 填充频率信息(可能重建)
    out["freqs_cis"] = _pad_caption_freqs(
        out.get("freqs_cis"), current_model, target=bucket
    )
    # 生成并填充 mask
    if "encoder_attention_mask" in out:
        out["encoder_attention_mask"] = _caption_mask(
            out, caption=caption, seq=seq, bucket=bucket
        )
    else:
        out["encoder_attention_mask"] = _caption_mask(
            call_kwargs, caption=caption, seq=seq, bucket=bucket
        )
    return out
python/sglang/multimodal_gen/runtime/breakable_cuda_graph/runner.py dependency-wiring

BCG Runner 核心文件,定义了 BaseBreakableCUDAGraphRunner 和 DiffusionBreakableCUDAGraphRunner,实现图签名匹配、捕获与回放全套逻辑。

# python/sglang/multimodal_gen/runtime/breakable_cuda_graph/runner.py
# 展示 BaseBreakableCUDAGraphRunner 的 capture 和 __call__ 核心逻辑class BaseBreakableCUDAGraphRunner:
    def capture(self, *args, **kwargs) -> None:
        """为当前签名捕获 CUDA 图段,离线广播;服务期间不再调用。"""
        sig = _signature_kwargs(kwargs)
        if sig in self._graph_cache:
            return # 已经捕获过
        with BreakableCUDAGraphCapture(self._graph, device=self._device):
            output = self._module(*args, **kwargs)
        self._graph_cache[sig] = (output, copy.deepcopy(self._graph.segments))
​
    def __call__(self, *args, **kwargs):
        sig = _signature_kwargs(kwargs)
        entry = self._graph_cache.get(sig)
        if entry is None:
            # 签名未缓存,回退到 eager 模式
            return self._module(*args, **kwargs)
        expected_output, _ = entry
        # 将输入拷贝到静态缓冲区
        for buf, live in zip(self._input_buffers, _flatten_kwargs(kwargs)):
            buf.copy_(live, non_blocking=True)
        # 回放图段
        for seg in self._graph.segments:
            seg.replay()
        # 读取并克隆输出(避免被后续回放覆写)
        return _clone_output(expected_output)
python/sglang/srt/breakable_cuda_graph/breakable_cuda_graph.py dependency-wiring

模型无关的 BCG 原语层,被 LLM 和扩散运行时共享,包含流追踪、图分段、弱引用等关键机制。

# python/sglang/srt/breakable_cuda_graph/breakable_cuda_graph.py
# 展示 BreakableCUDAGraphCapture 的核心生命周期class BreakableCUDAGraphCapture:
    """上下文管理器,进入时启动捕获模式,退出时收集图段。"""
    def __init__(self, graph, device):
        self._graph = graph
        self._device = device
        self._stream = torch.cuda.Stream(device=device)
        self._segments = []
​
    def __enter__(self):
        _current_capture_var.set(self)
        _current_stream_var.set(self._stream)
        _install_wait_stream_hook() # 追踪 fork/join
        self._stream.__enter__()
        return self
​
    def __exit__(self, *exc):
        # 结束当前活跃图段
        if _is_stream_capturing(self._stream):
            seg = torch.cuda.CUDAGraph()
            seg.capture_end()
            self._segments.append(seg)
        _uninstall_wait_stream_hook()
        self._stream.__exit__(*exc)
        _current_capture_var.set(None)
        _current_stream_var.set(None)
        _forked_streams_var.set(None)
​
    @property
    def segments(self):
        return list(self._segments)

评论区精华

流同步竞态条件 正确性

gemini-code-assist[bot] 指出图在独立流上捕获,回放时若当前流未等待输入拷贝完成,可能导致数据不完整。

结论:BBuf 在 _replay 中添加 event.synchronize(),确保输入拷贝完成后再回放。 · 已解决

_clone_output 不支持 dict/ModelOutput 正确性

gemini-code-assist[bot] 指出如果模型返回 dict 或 ModelOutput,_clone_output 直接返回引用导致输出被后续步骤覆写。

结论:BBuf 添加了对 dict 和 ModelOutput 的递归克隆支持。 · 已解决

_weak_ref_if_tensor 缺少 dict 支持 正确性

gemini-code-assist[bot] 指出 _weak_ref_if_tensor 未处理 dict,导致断点间中间张量不能被弱引用,阻止 mem pool 回收。

结论:BBuf 在 breakable_cuda_graph.py 中添加了 dict 分支的弱引用处理。 · 已解决

BCG 架构设计:runner 继承 vs 直接注入 设计

Oasis-Git 建议 runner 应继承 base runner 并实现 capture/replay API,避免直接注入 transformer.forward。

结论:BBuf 重构后采用 BaseBreakableCUDAGraphRunner 和 DiffusionBreakableCUDAGraphRunner 的继承模式,capture 在 warmup 时完成。 · 已解决

模型特定辅助函数位置 设计

mickqian 询问将模型 padder 放在 stages/ 是否合适,建议独立目录。

结论:BBuf 将 padder 移至 breakable_cuda_graph/model_padders/。 · 已解决

风险与影响

  • 正确性风险:流同步处理仍需细粒度验证,尤其多流场景下;桶选择边界(seq等于桶大小时)需确认填充函数行为。
  • 内存风险:每个(resolution, bucket, timestep分支)组合产生独立图段,显存可能显著增长。用户需根据模型调整 --bcg-text-buckets
  • 兼容性风险:频繁 main 合并引入的BCG原语API变化可能影响LLM侧稳定性,需持续跟踪CI。
  • 测试覆盖缺口:多resolution切换、超桶回落Eager等场景缺乏覆盖。
  • 用户:使用支持模型时传递 --enable-breakable-cuda-graph 获得性能提升,代价是warmup时间延长和显存增加。其他模型无影响。
  • 系统:新增 sglang.srt.breakable_cuda_graph 共享包,影响线程内stream同步和contextvar管理;扩散warmup阶段变重。
  • 团队:未来添加新扩散模型需编写padder并加入支持列表;BCG原语的双侧维护通过共享包减轻。
核心路径变更 内存使用增加 流同步竞态风险

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论