第 13 题:PyTorch的checkpoint机制(梯度检查点)如何节省显存?计算开销在哪里?
题目
PyTorch的checkpoint机制(梯度检查点)如何节省显存?计算开销在哪里?
完整讲解
一、为什么需要 checkpoint?
前向时,为了 backward 能算梯度,通常要把中间激活都存下来,显存占用 ≈ 与层数/序列长度成正比。大模型或长序列时,激活显存很容易成为瓶颈。梯度检查点(gradient checkpointing) 的思路是:前向时只存一部分激活(如每隔几层存一次),其余不存;反向时用到某段中间激活时,从最近的一个检查点重新算一遍那段前向,得到激活再算梯度。这样用重复计算换显存。
二、显存怎么省下来的?
- 不 checkpoint:前向每一层的激活都保留,backward 时直接用;显存 ∝ 层数 × 每层激活大小。
- Checkpoint:只保留「检查点」处的激活(如每 k 层一个),中间层的激活算完就丢;backward 到某层需要激活时,从上一个检查点重算到该层,得到激活后再继续反向。所以显存从「所有层」变成「检查点数量 × 单点激活 + 重算时的临时激活」,通常能降到原来的 1/√k 量级(k 为分段长度)或更好,取决于分段策略。
三、计算开销在哪里?
- 多算一遍(或半遍)前向:反向时每两检查点之间的层要重新 forward 一次才能拿到激活,再 backward。所以总计算量 ≈ 1 次完整前向 + 1 次完整反向 + 若干次「分段前向」;分段越细,重算次数越多,训练变慢。
- 典型 trade-off:若每 2 层一个检查点,显存约可减半,但 backward 时大约多算 0.5 次前向;若每 √L 层一个,显存可降到 O(√L) 量级,多算量也在 O(√L) 量级(L=层数)。所以 checkpoint 是「用时间换显存」。
四、PyTorch 里怎么用?
torch.utils.checkpoint.checkpoint(fn, *args, **kwargs):把 fn 当成一个「块」:前向时只算一遍、不存中间激活(只保留 fn 的输入若需要);backward 时用 args 重新调用 fn 得到中间结果再反传。常用于把某一层或几层包起来:out = checkpoint(block, x)。
checkpoint_sequential:对 nn.Sequential 的多个子模块分段做 checkpoint,每段一个检查点。
- 自定义:大模型里常按「每 N 层」或「每个 transformer block」做 checkpoint,在 forward 里手动调用
checkpoint(block, x)。
五、面试可怎么说?
- 省显存:不存所有中间激活,只存检查点;反向时从检查点重算得到激活再反传,显存从 O(层数) 降到 O(检查点数) 量级。
- 开销:反向时要多算若干次「分段前向」,训练变慢;是典型的用算力换显存。
- 用法:
checkpoint(fn, *args) 把 fn 包成一段,或对 Sequential 用 checkpoint_sequential,或在大模型里按 block 包。
面试要点
- Checkpoint = 前向少存激活、只存检查点;反向时从检查点重算前向再反传,从而省显存。
- 开销 = 反向阶段多算若干次分段前向,训练变慢;用计算换显存。
- PyTorch:checkpoint(fn, *args)、checkpoint_sequential;大模型常按 block 或每 N 层包。
记忆要点
- 省显存:只存部分激活(检查点),其余反向时重算;显存从 O(L) 到 O(检查点数)。
- 开销:反向时多算分段前向,总计算量增加;trade-off 算力换显存。
- 使用:checkpoint(fn,…)、checkpoint_sequential 或按 block 包装。