梯度检查点

Gradient Checkpointing

Gradient Checkpointing 梯度检查点(激活重计算)是一种经典的显存-算力折衷技巧:在前向传播时不保存所有层的激活值,只保存关键检查点;反向传播到这一段时重新算一遍激活值,换取 30–50% 的显存下降,代价是训练时长增加约 20%。

详细解释

梯度检查点(Gradient Checkpointing,别名 Activation Checkpointing / Activation Recomputation,激活重计算) 早在 2016 年就出现在训练 ResNet 的场景里,现在已成为所有大模型训练/微调脚本的默认打开选项——因为不打开它,你可能连 7B 模型的 Fine-Tune 都启动不了

训练(Training)时显存占用的三大头

  1. 模型权重(Weights):模型本身,比如 7B BF16 = 14 GB;
  2. 优化器状态(Optimizer States):Adam 需要每个参数存 m 和 v,BF16+FP32 混合后这部分是权重的 2 至 4 倍(7B 大概 56 GB)——这就是为什么 Fine-Tune 比推理显存要多很多。
  3. 激活值(Activations):前向传播时每一层的输出,反向传播要用到它。序列越长、batch-size 越大、模型层数越深,激活越多。7B + 4K 上下文 + batch-size=8 时,激活值能涨到 80 GB+,比权重和优化器加起来还多。

梯度检查点就是专门砍第 3 条(激活)的。思路很朴素:

  • 前向传播:不把每一层的激活值都存下来(那太费了),只存”每 N 层一个”的检查点(Checkpoint)。
  • 反向传播:需要某一层的激活值但发现没存 → 用最近的一个检查点,再跑一小段前向传播”重算”出来

所以本质上是:用”多做一些前向计算”换”少存一堆中间张量”,典型效果是激活显存下降 30–50%,训练时间上升 ~20%(算力换空间)。

和其他显存优化手段的组合关系

优化手段作用对象显存节省倍数是否与梯度检查点叠加
梯度检查点激活值(Activations)基础底座,必开
混合精度(BF16/FP16)权重 + 激活(存储格式)✅ 组合,必开
ZeRO-1/2/3 / FSDP(DeepSpeed)优化器状态/权重分片2× – 8×✅ 组合,DDP 多卡训练标配
QLoRA权重 4-bit + 只训 LoRA权重 4× + 优化器 100×✅ 叠加后单卡 70B 不是梦
FlashAttention-2Attention 层激活/临时矩阵~10×(Attention 部分)✅ 组合,强烈建议
序列并行(Sequence Parallel)激活按序列切分~2×✅ 多卡场景叠加

最强组合(单卡 70B QLoRA 微调):梯度检查点 + BF16 AMP + QLoRA NF4 + Paged Optimizers + FlashAttention-2 → 24G 搞定。

常见坑位

  1. 什么时候关梯度检查点?:当你显存非常充足(4×H100 80G 跑 7B),且追求训练吞吐量,此时可以关。用它换来的 20% 速度提升比省显存更有意义。
  2. “Selective” 梯度检查点:不是所有层都值得重计算。经验上**只对 Transformer Block 做检查点(一层一个)**就够了;对 Embedding/LM Head、RMSNorm 这些小层重算的 overhead 不值得。
  3. reentrant vs non-reentrant:PyTorch 原生 torch.utils.checkpoint.checkpoint 有两个模式;现代(PyTorch 2.1+)推荐 use_reentrant=False,配合 torch.compile 加速效果更好。
  4. LoRA 训练时是否要开?:要!即使是 LoRA,激活值仍然是大头。几乎所有微调脚本默认都是打开的,只有少数老框架要手动调 gradient_checkpointing=True
  5. 唯元智创 企业培训课程里专门有一张《微调显存优化 8 板斧》全景图,按优先级开梯度检查点 + FlashAttention + QLoRA + PagedOptimizers,100 次里 95 次都够用了,剩下再上 ZeRO / 多卡。

常见问题

为什么推理时不用梯度检查点?
推理时没有反向传播,激活值用完就可以立刻丢(不需要留着算梯度),所以本来就不占多少显存。梯度检查点只在训练/微调场景有意义——推理时它只会白白多算前向传播,没有任何好处。别在 vLLM 启动参数里找 gradient_checkpointing 了,没有。
开了梯度检查点显存还是 OOM 怎么办?
按优先级降显存:(1)把 per_device_train_batch_size 从 8 降到 1(梯度累加 gradient_accumulation_steps 调大,等效 batch 不变);(2)切 QLoRA 4-bit + NF4;(3)开启 Paged Optimizers / CPU offload(FSDP CPU offload / DeepSpeed ZeRO-Offload);(4)再不行加 Gradient Accumulation + 切半精度更狠的选项。一般到这一步就不会 OOM 了。
它和 Checkpoint(训练保存权重)是一个东西吗?
完全不是。中文翻译撞了同一个词”Checkpoint”但语境不同:Checkpoint(模型检查点) 指训练过程中定期把权重存到磁盘上,是”保存进度”的意思;Gradient Checkpointing(梯度检查点)特指前向时保留部分激活张量的显存优化技术。别混,社区日常英文常简写成 GC,但中文语境一定要区分。