一句话定义:梯度检查点(Gradient Checkpointing)是一种训练时的显存优化技术——前向传播时只保存少数几个中间结果(检查点),反向传播需要其余中间结果时,从最近的检查点重新前向计算一遍,用额外的计算量换更低的显存占用。
它解决什么问题
训练神经网络做反向传播时,必须用到前向传播中每一层的中间激活值(activation)。层数越多、序列越长、batch 越大,这些激活占的显存就越多,常常成为能不能跑起来的硬瓶颈。
梯度检查点的做法很直接:不再把所有激活都存下来,而是每隔若干层只存一个“检查点”。反向传播走到某一段时,就从该段的检查点重新做一次前向,把它内部的激活临时算出来,用完即弃。
打个比方:追一部长剧,硬盘不够把每集都录下来,于是只记下每五集的剧情梗概(检查点)。需要某集细节时,从最近的梗概处重新看过去。看剧总时间变长了,但硬盘占用小了很多。
和相邻概念的区别
| 技术 | 主要省下的显存 | 代价 |
|---|---|---|
| 梯度检查点 | 中间激活 | 额外前向计算,训练变慢 |
| 梯度累积(Gradient Accumulation) | batch 维度的激活 | 需要更多训练步数 |
| 混合精度(Mixed Precision) | 激活与参数存储 | 数值稳定性要留意 |
| 参数分片 / 卸载(Sharding / Offload) | 参数、优化器状态 | 通信或 IO 开销 |
特别要区分“检查点”这个词的两种用法:模型保存时说的 checkpoint,是把权重存成文件,用于恢复训练或推理;梯度检查点保存的是训练过程中的中间激活,是显存管理手段,不改变模型本身。理论上它不改变训练结果。
主流深度学习框架都提供这项功能,具体 API 名称与用法以官方文档为准。
对从业者和普通人的意义
对 AI 从业者,它把“显存不够”从死路变成可调参数:多花点时间,就能在有限显卡上训练更深的模型、更长的上下文或更大的输入。代价通常是多出约一次前向传播量级的计算,训练明显变慢,所以它更适合显存吃紧、算力尚可的场景,而不是盲目全程开启。
对普通职场人,这是一个通用的工程思路:把全部中间结果缓存下来最省时间但最占空间;只存关键节点、需要时重算,是时间与空间之间的经典权衡。类似的思路也出现在数据库索引、浏览器缓存和编译优化里——先问清楚瓶颈是“算不过来”还是“存不下”,再决定要不要做这种交换。
