use_reentrant=False 粒度
Gradient checkpointing 用重算 forward 换显存——训练大模型或高分辨率输入时 OOM 的常用解法,wall-clock 涨 20–40% 要算进预算。

1. OOM 时:用重算换显存
训 ViT-L 或 4K 语义分割,batch size 已经压到 1 仍 OOM——gradient checkpointing 是标准解法。机制很直接:forward 时不保存中间 activation,backward 时重算该段的 forward,用计算换内存。典型收益是 peak memory 降 30–50%,代价是 wall-clock 涨 20–40%。预算里必须算进 throughput 损失,不是「免费午餐」;若项目 deadline 紧,要先测 slowdown 是否可接受,再决定是否同时上 gradient accumulation 撑有效 batch。
OOM 排查顺序:先确认是 activation 还是 optimizer state 占主导——前者 checkpoint 有效,后者需 ZeRO/FSDP 或换 AdamW 为更省内存的 optimizer 变体。activation 显存大致与「保存的 tensor 数量 × 单 tensor 体积 × batch」成正比;深 transformer 每层 self-attention 的 与中间 FFN 是主要占用。checkpoint 在 segment 边界不保存中间态,backward 时从 segment 输入重新 forward 一遍,时间近似翻倍该段 forward,但 memory 只保留 segment 边界上的少量 tensor。optimizer state(Adam 的 m、v)与参数量同级,checkpoint 不碰这部分——若 OOM 来自 optimizer,应优先减 batch、用 8-bit optimizer 或 sharding。
2. PyTorch API 与推荐用法
from torch.utils.checkpoint import checkpoint
def forward(self, x):
return checkpoint(self.block, x, use_reentrant=False)PyTorch 2.0+ 推荐 use_reentrant=False——与 custom autograd、DDP、部分 hook 兼容更好。旧版 reentrant=True 与某些 in-place op 或 custom Function 冲突,backward 报错或 silent wrong grad。HuggingFace 一行开启:model.gradient_checkpointing_enable(),内部按 transformer block 粒度 checkpoint,不必手写每 block;启用后 model.config.use_cache 须 False,否则 KV cache 与 checkpoint 语义冲突。
use_reentrant=False 走非重入 autograd 路径,与 torch.compile、activation checkpoint 嵌套、部分 register_hook 组合更稳。手写时常见模式是把 整段 TransformerBlock 或 Bottleneck 包进 checkpoint(fn, *args),而不是拆成 attention 与 FFN 两次 checkpoint——两次边界会多一次重算开销。HuggingFace 的 gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False}) 应对齐 PyTorch 2 默认。若模型含 past_key_values(生成缓存),训练阶段必须关 cache;推理 export 时再单独开 cache,勿把 training checkpoint wrapper 带进 deploy graph。
3. 粒度:按 block,不按 layer
粒度决定收益与代价的平衡点。按 transformer block 或 ResNet stage checkpoint 是 sweet spot;对每个 Conv/Linear 都 checkpoint 会 slowdown 40%+ 而 memory 收益递减。经验法则:checkpoint 段内的 activation 占全网 10–25% 时 ROI 最高。checkpoint 块内禁止 in-place op(如 x += ...、F.relu_(x)),否则 backward 重算时状态不一致。
与 activation checkpointing 的边界:checkpoint 只省 activation memory,不省参数和 optimizer state。eval/inference 不需要 checkpoint——只在 training forward 包一层,export ONNX 前去掉 checkpoint wrapper。划分 segment 时可按 显存占比 粗算:对 ViT,12 层均分 checkpoint 每 2 层一段,通常比 12 段各 checkpoint 一层更省 wall-clock;对 U-Net,encoder 每个 stage 与 decoder 对称 block 各一段,避免在 concat/skip 边界 checkpoint——重算时 skip tensor 仍须保存,边界选错会省不了 memory。PyTorch 还提供 checkpoint_sequential 把 nn.Sequential 切成多段,适合 ResNet-style stack。
4. 与 AMP、DDP 的交点
checkpoint 重算的 forward 仍在 autocast 内,一般无额外问题。若 custom Function 有 fp16 敏感 op,单独测 NaN。DDP 下每 rank 独立 checkpoint,无特殊配置。多卡时 effective batch 变大,有时不用 checkpoint 也能训——先测单卡 checkpoint vs 多卡无 checkpoint 的 wall-clock 与 memory,选 Pareto 优解。
与 gradient accumulation 组合时,checkpoint 在 每个 micro-batch 的 forward 生效,accumulation 不改变 segment 重算次数——总 slowdown 与「无 accumulation、同等 global batch」接近。FSDP 与 checkpoint 可叠加:sharding 省参数/optimizer,checkpoint 省 activation;两者都开时先 profile 哪项是瓶颈,避免重复优化。若 backward 出现 RuntimeError: Trying to backward through the graph a second time,常见原因是 reentrant checkpoint 与 retain_graph 冲突,优先改 use_reentrant=False 并检查是否在 checkpoint 段外重复使用同一 graph。
5. 案例:ViT-B 4K 分割从 OOM 到 batch=2
某 4K 语义分割,ViT-B encoder + U-Net decoder,单卡 24GB batch=1 OOM。按 encoder 每 stage、decoder 每 block 包 checkpoint 后 batch=2 可训,peak memory 从 23GB 降到 14GB,throughput 降 28%。若对每个 patch embed 单独 checkpoint,throughput 再降 15% 而 memory 仅多降 1GB——粒度太细,收益递减。
该案例的验收不只 OOM 消失:固定 seed 对比 checkpoint 前后 mIoU 应 parity(数值误差在 1e-4 量级内);同时记录 samples/sec 与 max_memory_allocated,写入 run card,避免后续同事为再省 1GB 把粒度切到 per-layer。4K 输入下 decoder 上采样层 activation 体积大,checkpoint 重点放在 decoder 中段往往比只 checkpoint encoder 更划算——encoder 若已冻结,其 activation 仍占显存,冻结不等于不存 activations。
6. 验收
torch.cuda.max_memory_allocated()对比开/关 checkpoint 的 peak memory,目标降 30%+。- 同一 batch 测 throughput(samples/sec),确认 slowdown 在可接受范围(通常 20–40%)。
- 固定 seed,对比 checkpoint 前后 val metric——应 parity,不应因重算引入数值 drift。
- 若用 HuggingFace,确认 block 粒度符合预期,不是 per-layer。
- 故意在 checkpoint 块内加 in-place op,确认能复现报错——证明约束被遵守。
Gradient checkpointing 的 use_reentrant=False 和 block 粒度是落地时的两个关键开关。改一项时固定 seed 与 batch,memory 与 throughput 同测,才分得清是粒度太细还是 block 划分不合理。
相关
也可以看看
- ·3 分钟阅读
grad_norm spike 与 misclass grid
只看 train loss 不够——要 log lr、grad norm、val slice、GPU 利用率和样例可视化,异常通常是 loss 曲线先以外的信号。
- ·5 分钟阅读
beta=0.999 与 eval 权重
训练时对参数做指数滑动平均,验证和部署常用 EMA 权重——曲线更顺,有时泛化更好,但要与 optimizer 步和 checkpoint 策略一致。
- ·3 分钟阅读
梯度裁剪:爆炸刹车
clip_grad_norm 能止住 NaN 连锁,但根因常是 LR 过大、loss scale 或 bad batch——裁剪是买时间,不是根治。
johan's blog