梯度累积
为什么需要梯度累积

Step1: 核心机制
- 一个完整batch被切成
K个micro-batch - 每个
micro-batch各自前向 + 反转,梯度在.grad中累加。 - 攒够
k次后统一执行一次optimizer.step(), 再zero_grad()清零。
Step2: 数学等价性


Step3: 代码实现框架

代码测试输出

完整学习参考
激活检查点与激活卸载
痛点: 显存压力不只来自参数
- 前向传播为了反向传播,必须把每层中间输出(激活值)一直保存到显存里;层级越深、序列越长、batch越大,激活值占比越高。
- 显存占用等于 O(L X B X S X D), L-层级数,B-batch size, S-序列长度, D-隐藏维度。
- 问题:能不能少存激活,仍然完成反向传播?
Step1: 两条优化路线
- Activation Checkpointing(重计算换显存):不保存所有层激活,每隔几层存一个"检查点(Checkpoint)"
- Activation Offload(搬运换显存): 把部分激活临时搬到CPU或者其他存储层级。
- 代价对比: checkpointing多花约20% ~ 30%时间,节省成倍甚至数倍显存;本质是"时间换空间"

Step2: 显存节省分析

Step3: 代码实现部分

工程优化要点

完整学习参考
总结
-
梯度累积:决定batch能开多大,训练稳不稳
-
激活检查点与卸载: 决定显存怎么省