LuAITools.com
提交工具
🧠AI
Gradient Checkpointing

梯度检查点

训练大模型时用「多算一遍」换「少存一点」:不把中间激活全部存下,只在需要时重算,显存占用大幅下降。

梯度检查点是什么?

训练神经网络时,反向传播需要用到正向计算时留下的一堆中间结果(激活值)。模型越大,这些中间结果越多,显存很快就不够用。梯度检查点(Gradient Checkpointing)的思路很朴素:不把这些中间结果全存起来,而是只保留几个「检查点」,等到反向传播要用的时候,再从头重算一遍。

它为什么能省显存?

显存和计算,是两种可以互换的资源
要么你花显存把中间结果存下来,要么你花算力在需要时重新算出来。梯度检查点选择了后者——用多一点的计算量,换来显存占用的大幅下降。
只存关键节点
它不会一个不落地记录每一步,只在某些位置存下「快照」。中间缺的部分,需要时就从最近的快照重算。

代价是什么?

多花一些算力
因为要重算,整体计算量大约会多出 20%~30%,训练速度也会慢一点。
但换来的收益很值
同样的显卡,能装下更大的模型或更大的批量(batch size),往往能换来更好的训练效果和更快的收敛。

什么时候该用它?

当你训练大模型时「显存不够、模型装不下」,梯度检查点就是性价比很高的一招。它常和混合精度训练、梯度累积等技术一起用,是工程上「挤显存」的常规手段。

一句话记住:梯度检查点就是用「多算一遍」换「少存一点」,让显存不再卡住大模型。

评论