训练神经网络时,显存里装的不只有模型权重。前向计算产生的大量中间张量要留到反向传播,用来计算梯度。模型越深、序列越长,这些激活越容易把显存占满。激活检查点会少留一部分中间结果,等反向传播走到那里时再算一次。
我把 2016 年的低内存训练论文和 PyTorch 当前文档对了一遍。这项方法交换的是计算与显存。它和保存到磁盘、用于断点续训的模型检查点同名,处理的事情并不相同。前者管理一次训练步骤里的中间张量,后者保存模型、优化器和训练进度。
前向时少存,反向时重算
普通训练会让前向产生的张量继续留在计算图里。激活检查点只保存选定区段的输入,区段内部的中间张量可以释放。反向传播需要它们时,框架重新调用那段前向计算,重建所需结果,再继续算梯度。
最早的系统研究把一个有 n 层的网络分段,给出约为平方根 n 的内存方案,每个小批次多做一次前向计算。论文在一千层残差网络上把内存从 48GB 降到 7GB,运行时间增加约百分之三十。那是特定网络和实现留下的比例,今天训练语言模型不能直接照搬。
检查点放得太少,省不了多少显存。每一层都重算,又会增加大量计算。PyTorch后来加入选择性激活检查点,可以指定哪些操作保存、哪些操作重算。矩阵乘法很贵时可以优先保存它的结果,把计算较轻、占内存较多的操作留给重算。实际分段要看模型结构和性能记录。
随机操作还会影响正确性。区段里若有 dropout,重算时使用不同随机数,反向看到的函数就变了。PyTorch默认保存和恢复 CPU 及一种设备类型的随机状态,尽量让重算与原前向一致,这也会增加开销。函数若依赖会变化的全局状态,或在区段内部把张量移到未记录的新设备,文档警告可能出现错误梯度。
当前 PyTorch 提供重入与非重入两种实现,并建议显式选择非重入版本。非重入实现能在需要的张量重建完以后提前停止,也支持更多反向调用方式和嵌套结构。版本差异会改变默认行为,训练脚本升级后应重新跑一次梯度和损失对照。
验收要同时记录时间与结果
我会先用很小的固定批次跑两遍,一遍关闭检查点,一遍开启。对比损失、梯度和更新后的参数,确认数值在允许误差内。随后记录峰值显存、单步时间和每秒样本数,再逐步扩大批量或序列长度。省出的显存若能多放有效样本,训练总时间可能值得;单纯把同一批次跑慢,则要重新选择区段。
激活检查点能让装不下的训练先跑起来。它没有减少这次训练要学的参数,也不会免费消除显存限制。重算发生在哪里、慢了多少、梯度是否一致,需要留下一份可复现的对照。