DiffSynth-Studio 梯度检查点(Gradient Checkpointing)与 Offload 实战指南:从原理到 Pipeline 级细粒度训练优化
2026/9/15 10:44:22 网站建设 项目流程

DiffSynth-Studio 梯度检查点(Gradient Checkpointing)与 Offload 实战指南:从原理到 Pipeline 级细粒度训练优化

【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio

导读

本文围绕 DiffSynth-Studio 的diffsynth.core.gradient模块展开,系统讲解梯度检查点(Gradient Checkpointing)及其 Offload 变体的数学原理、封装接口gradient_checkpoint_forward的使用方式,并结合仓库内 Qwen-Image、WanVideo 等实际 Pipeline 与训练脚本,说明如何在 Block 层级做细粒度梯度检查点、何时应该开启 Offload,帮助你在训练大模型时以可控的计算开销换取更低的显存占用。

一、为什么需要梯度检查点:以 ToyModel 为例理解原理

1.1 一个最小的例子

梯度检查点的核心目标是在训练阶段减少显存(VRAM)占用。我们先看docs/zh/API_Reference/core/gradient.md中给出的最小模型:

import torch class ToyModel(torch.nn.Module): def __init__(self): super().__init__() self.activation = torch.nn.Sigmoid() def forward(self, x): return self.activation(x) model = ToyModel() x = torch.randn((2, 3)) y = model(x)

在这个模型里,输入 $x$ 经过 Sigmoid 激活函数得到输出 $y=\frac{1}{1+e^{-x}}$。

1.2 反向传播中的“空间换时间”与“时间换空间”

假设损失函数为 $\mathcal L$,反向传播时我们得到 $\frac{\partial \mathcal L}{\partial y}$,需要继续计算 $\frac{\partial \mathcal L}{\partial x}$。由 $y(1-y)$ 恰好是 Sigmoid 的导数:

$$\frac{\partial \mathcal L}{\partial x}=\frac{\partial \mathcal L}{\partial y}\cdot\frac{\partial y}{\partial x}=\frac{\partial \mathcal L}{\partial y}\cdot y(1-y)$$

  • 不启用梯度检查点:训练框架在前向传播时保存所有辅助梯度计算的中间变量(例如这里的 $y$),反向传播时直接复用,避免重新计算 exp,计算速度最快,但代价是中间激活值要常驻显存。
  • 启用梯度检查点:中间变量不再保存,只有输入 $x$ 被保留;反向传播经过该层时重新前向计算一遍以得到中间变量,从而显著降低显存峰值,但会引入额外的计算量(典型是约 1 倍左右的前向重算开销),速度变慢。

一句话总结:梯度检查点是用“重算”换“显存”。当模型参数量越来越大、激活值(activation)占据的显存不可忽视时,这项技术几乎成为大模型训练的必选项。

二、核心 API:gradient_checkpoint_forward

2.1 接口与三种开关组合

diffsynth.core.gradient对外暴露的入口只有一个函数gradient_checkpoint_forward,由 diffsynth/core/gradient/init.py 导出,实现在 diffsynth/core/gradient/gradient_checkpoint.py 中。其签名如下:

def gradient_checkpoint_forward( model, use_gradient_checkpointing, use_gradient_checkpointing_offload, *args, **kwargs, ):

文档中的调用示例(可直接运行):

import torch from diffsynth.core.gradient import gradient_checkpoint_forward class ToyModel(torch.nn.Module): def __init__(self): super().__init__() self.activation = torch.nn.Sigmoid() def forward(self, x): return self.activation(x) model = ToyModel() x = torch.randn((2, 3)) y = gradient_checkpoint_forward( model, use_gradient_checkpointing=True, use_gradient_checkpointing_offload=False, x=x, )

两个布尔开关共组合出三种行为:

use_gradient_checkpointinguse_gradient_checkpointing_offload行为
FalseFalse与原始model(*args, **kwargs)完全等价,不引入任何额外行为,可安全集成到现有推理/训练代码中
TrueFalse启用梯度检查点:丢弃中间激活,反向传播时重算
TrueTrue在梯度检查点基础上启用 Offload:所有输入参数(激活值)存储到 CPU 内存,进一步降低显存占用,同时计算速度更慢

注意:use_gradient_checkpointing_offload=True隐含启用梯度检查点(use_gradient_checkpointing分支优先判断 offload)。

2.2 源码级行为解析

阅读 diffsynth/core/gradient/gradient_checkpoint.py 可以看到完整的四路分支:

  1. DeepSpeed 优先分支:当use_gradient_checkpointing=True且环境中已安装deepspeed、且deepspeed.checkpointing.is_configured()为真时,走 DeepSpeed 的激活检查点路径。此时先通过judge_args_requires_grad检查所有输入(args + tuple(kwargs.values()))中是否存在requires_grad=True的 Tensor:
    • 若不存在(纯推理或无梯度输入),直接调用model(*args, **kwargs),不打断原始计算图;
    • 否则调用deepspeed.checkpointing.checkpoint(create_custom_forward_use_reentrant(model), *all_args)
  2. Offload 分支use_gradient_checkpointing_offload=True时,用torch.autograd.graph.save_on_cpu()包裹torch.utils.checkpoint.checkpoint(..., use_reentrant=False),把中间激活保存到 CPU 内存。
  3. 普通梯度检查点分支:仅use_gradient_checkpointing=True,直接调用torch.utils.checkpoint.checkpoint(..., use_reentrant=False)
  4. 兜底分支:两者均为False,等价于model(*args, **kwargs),保证接口对推理流程完全透明。

实现中用到两个辅助函数create_custom_forward(module)create_custom_forward_use_reentrant(module),它们的作用是把torch.nn.Module包装成可被torch.utils.checkpoint.checkpoint接受的普通可调用对象;use_reentrant=False则采用非重入(non-reentrant)实现,与torch.autograd.graph.save_on_cpu()等现代 API 配合更稳定。

从源码结构还可以推断:该模块对 DeepSpeed 的支持是可选且自动发现的(_HAS_DEEPSPEED通过try: import deepspeed判定),未安装 DeepSpeed 时自动回退到 PyTorch 原生实现,不会因缺少依赖而报错。

三、最佳实践一:在model_fn的 Block 层级启用细粒度梯度检查点

3.1 为什么不要对整个模型粗暴开启

对整个模型整体启用梯度检查点时,计算效率与显存占用往往都不是最优的:整个模型被当作一个大的可重算单元,任何一层的反向传播都会触发全模型前向重算。更合理的做法是细粒度地按 Block(Transformer Block / 单个子模块)分别包上梯度检查点,把重算范围控制在单层以内。但若为此在每个模型里手工塞入torch.utils.checkpoint.checkpoint,又会让代码变得繁杂、侵入性强。

3.2model_fn中的标准范式

DiffSynth-Studio 的解决方案是把梯度检查点逻辑收敛到 Pipeline 的model_fn中。以 diffsynth/pipelines/qwen_image.py 中的model_fn_qwen_image为例(定义于 L707 附近),它在forward参数中显式接收两个开关:

def model_fn_qwen_image( dit: QwenImageDiT = None, ... use_gradient_checkpointing=False, use_gradient_checkpointing_offload=False, ... ):

随后在遍历 Transformer Block 的主循环里,对每一个 Block调用gradient_checkpoint_forward(见 diffsynth/pipelines/qwen_image.py):

for block_id, block in enumerate(dit.transformer_blocks): text, image = gradient_checkpoint_forward( block, use_gradient_checkpointing, use_gradient_checkpointing_offload, image=image, text=text, temb=conditioning, image_rotary_emb=image_rotary_emb, attention_mask=attention_mask, enable_fp8_attention=enable_fp8_attention, modulate_index=modulate_index, kv_cache=None if kv_cache is None else kv_cache.get(f"block_{block_id}"), )

这种写法的好处:

  • 不改模型结构代码QwenImageDiT的 Block 实现无需任何改动;
  • 粒度细:每个 Transformer Block 是独立的重算单元,重算成本被限制在单层;
  • 可配置:通过传入use_gradient_checkpointing/use_gradient_checkpointing_offload即可在开启/关闭间切换,训练与推理共用同一套model_fn

3.3 其他 Pipeline 中的应用

同样的模式被大量复用于仓库内各类模型的训练路径。例如 diffsynth/pipelines/wan_video.py 中,WanVideo 在循环 Block 时同样用gradient_checkpoint_forward(block, use_gradient_checkpointing, use_gradient_checkpointing_offload, x, context, t_mod, freqs)包裹;而在启用 VAP(Video-Adaptive Prior?)等额外分支时,则直接组合torch.autograd.graph.save_on_cpu()torch.utils.checkpoint.checkpoint(..., use_reentrant=False)(见 diffsynth/pipelines/wan_video.py)。从源码搜索结果看,gradient_checkpoint_forward已被qwen_imagewan_videoflux_imageflux2_imagez_imagekrea2anima_imagemova_audio_videoqwen_video_edit等多个 Pipeline 以及各 DiT/ControlNet 模型文件(如 diffsynth/models/qwen_image_dit.py、diffsynth/models/flux2_dit.py、diffsynth/models/wan_video_dit.py 等)广泛采用,是 DiffSynth-Studio 训练体系中的统一基础设施。

四、最佳实践二:训练脚本中的开关与参数贯通

4.1 命令行参数到model_fn的传递链

在官方训练脚本 examples/qwen_image/model_training/train.py 中,两个开关从命令行参数一路传入:

use_gradient_checkpointing=args.use_gradient_checkpointing, use_gradient_checkpointing_offload=args.use_gradient_checkpointing_offload,

并在训练封装对象中保存为成员,最终随model_fn的调用参数一起注入(见 examples/qwen_image/model_training/train.py),从而直达model_fn_qwen_image内部的gradient_checkpoint_forward。也就是说:训练时你只需要在启动脚本里加一个 flag,无需触碰任何模型代码

4.2 实际训练脚本用法

仓库内大量训练脚本默认开启了梯度检查点。例如 examples/qwen_image/model_training/full/Qwen-Image.sh、examples/qwen_image/model_training/full/Qwen-Image-Edit.sh、examples/qwen_image/model_training/full/Qwen-Image-2512.sh 等均包含:

--use_gradient_checkpointing \

而在 Blockwise ControlNet 系列脚本(如 examples/qwen_image/model_training/full/Qwen-Image-Blockwise-ControlNet-Canny.sh)中,还可以看到被注释掉的# --use_gradient_checkpointing \,便于在“完整微调”与“ControlNet 分支训练”等不同场景下按需取舍——这说明是否开启梯度检查点完全由启动参数决定,切换成本极低。

4.3 Offload 参数

use_gradient_checkpointing_offload同样作为独立命令行参数存在。根据文档指引,Offload 仅需在激活值占用显存过大的模型(例如视频生成模型)中启用,因为此时把激活搬到 CPU 内存能换来最明显的显存收益,代价是更慢的训练速度。

五、什么时候该用 Offload:选择建议

结合文档的最佳实践与源码实现,给出如下决策参考:

  • 梯度检查点(通常需要开启):随着模型参数量增大,中间激活值的显存开销快速膨胀,梯度检查点已成为必要的训练技术,一般默认开启;
  • 梯度检查点 Offload(按需开启):仅当激活值本身过大(典型如视频生成模型的长序列、多帧 token)导致显存依然吃紧时开启。它把激活保存到 CPU 内存,显存占用进一步下降,但 CPU↔GPU 的搬运与重算使训练更慢;
  • 两者都不开:适用于显存充裕的推理场景或小模型训练,gradient_checkpoint_forward此时退化为直接调用,与原始前向完全一致,因此你也可以放心地把这一封装常驻model_fn中,无需为开关写多套代码。

六、小结

  • diffsynth.core.gradient.gradient_checkpoint_forward是 DiffSynth-Studio 统一封装的梯度检查点入口,用两个布尔参数覆盖“不启用 / 启用 / 启用并 Offload”三种模式,且对 DeepSpeed 自动感知、对纯推理透明。
  • 推荐的启用位置是 Pipeline 的model_fn(如model_fn_qwen_image),在 Block 层级做细粒度检查点,兼顾计算效率与代码整洁。
  • 训练侧只需在启动脚本中加入--use_gradient_checkpointing(必要时加--use_gradient_checkpointing_offload)即可生效,参数会一路贯通到model_fn内的每个 Transformer Block。

相关源码与文档索引:接口实现见 diffsynth/core/gradient/gradient_checkpoint.py,模块导出见 diffsynth/core/gradient/init.py,Qwen-Image 的 Block 级接入见 diffsynth/pipelines/qwen_image.py,WanVideo 的 Offload 组合用法见 diffsynth/pipelines/wan_video.py,训练参数传递见 examples/qwen_image/model_training/train.py。

【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询