你有一张 80GB 的 A100,想训练一个 7B 参数的模型。把模型权重拷进显存,FP16 格式才 14GB。你怎么算都够——直到你敲下 torch.cuda.OutOfMemoryError。

我最早也踩过这个坑。明明模型只占十几GB,为什么训练时多卡都挤不下?后来我才知道,模型本身只是冰山一角,真正的显存大头是优化器状态。Adam 这个最常用的优化器,在训练过程中不是“用完即走”,它要给每一个参数挂两本“账本”,一直记到训练结束。
在 Adam 原始论文 里,每个参数的更新公式都包含两个额外的量:一阶动量 m 和二阶动量 v。
m_t = β1·m_{t-1} + (1-β1)·g_t
v_t = β2·v_{t-1} + (1-β2)·g_t²
m 是梯度的指数平均(理解成“惯性”),v 是梯度平方的指数平均(理解成“自适应步长”)。每一步参数更新都要读这两个值、再更新这两个值,所以它们必须常驻显存。去看 PyTorch 的 Adam 实现,你会发现每个参数张量旁边都跟着 exp_avg 和 exp_avg_sq 两个 FP32 张量,大小和参数一模一样。
那具体占多大?我们算一下混合精度训练(AMP)下一份参数的完整开销:
| 项目 | 每参数字节 | 7B 模型总计 |
|---|---|---|
| 模型参数(FP16 副本) | 2 | 14 GB |
| 梯度(FP16 副本) | 2 | 14 GB |
| Adam m(FP32) | 4 | 28 GB |
| Adam v(FP32) | 4 | 28 GB |
| FP32 主权重 | 4 | 28 GB |
| 合计 | 16 | 112 GB |
所以标题说的 60GB+ 就是 m 和 v 两项——56GB,四舍五入刚好 60。如果算上 FP32 主权重,光优化器涉及的显存就冲到 84GB,一张 A100(80GB)根本装不下。
我最初特别不解:为什么不能直接把 FP16 参数当主权重,非要另存一份 FP32?后来看混合精度训练的资料才明白:FP16 动态范围只有约 5 个十进制数量级,梯度很小的时候直接变 0,更新公式里的 eps 一加就废。所以必须把主权重和优化器状态保留在 FP32 里,保证参数更新的精度。这也是“最省显存”的混合精度方案依然有这么大开销的根本原因。
那有没有办法省?有,但每一条都有代价。
| 方案 | 状态占用(m+v) | 代价 | 省的是总量吗 |
|---|---|---|---|
| 标准 Adam | 8 字节/参数 | 无 | 否 |
| AdaFactor | ≈2 字节/参数 | 低秩近似,收敛变慢或欠拟合 | 是 |
| 8-bit Adam | 2 字节/参数 | 量化/反量化开销,调参敏感 | 是 |
| ZeRO Stage 1 | 8 字节/参数(分片) | 通信开销 | 按卡分摊 |
AdaFactor 把 v 矩阵分解成低秩因子,m 也做了简化,状态从 8 字节降到约 2 字节。但 AdaFactor 论文 自己也承认,在 transformer 上收敛速度和最终精度都可能打折扣。用在大模型上,你可能要多花不少时间来调到理想的效果。
8-bit Adam 把 m 和 v 量化成 int8,显存直接除以 4。从 8-bit Adam 论文 的实验看,在多数任务上效果和 FP32 相当,但量化范围对学习率极敏感,学习率一大就容易溢出。
ZeRO 分片 的思路最直接:状态总量不变,但按数据并行度切碎,每张卡只存自己负责的那一份。DeepSpeed 官方文档里给过 7.5B 模型的示例,用 ZeRO Stage 1 把优化器状态分到 64 张卡上,每卡状态不到 1GB。代价是每次更新参数都要多一次 all-gather 通信,网络带宽差的集群可能拖慢训练。
你也可以写几行 Python 自己估算一下,避免被别人的“经验值”误导:
def memory_count(n):
return {'params': 2*n, 'grads': 2*n, 'm': 4*n, 'v': 4*n, 'master': 4*n}
m = memory_count(7000000000)
print('params: ' + str(m['params']/1e9) + ' GB')
print('adam state: ' + str((m['m']+m['v'])/1e9) + ' GB')
在我自己的实验里,跑 7B 的 LoRA 微调时,很多人以为显存会小很多。确实,LoRA 只训练低秩旁路,但 Adam 需要保存旁路参数的 m 和 v;如果冻结的基座参数还要过反向传播,梯度同样要占空间。所以 LoRA 能省显存,但省的不只是优化器状态,更主要是激活值。这又是另一本账了。
为什么不能用 FP16 保存 m 和 v?
m 和 v 是梯度的指数平均,数值可能很小也可能很大,FP16 只有 10 位尾数,无法保持累积精度。8-bit 优化器用分位数量化缓解了这个问题,但大多数场景下还是要保留 FP32 主权重来保证收敛。
那现在该怎么办?我的建议是:先算账,再选方案。参数量、batch size、模型并行度代进算法先估一轮,然后根据通信带宽决定用 ZeRO 还是 Himel。如果你真的缺显存,可以考虑换优化器——比如 Lion 只维护一个动量,状态占用减半,但需要更多调参经验。技术边界也很清楚:优化器状态不是唯一的大头,激活值、梯度、通信开销都在抢显存。ZeRO 省了显存但加了通信,8-bit 省了显存但可能牺牲精度。没有免费午餐,你只是在选哪种代价自己更愿意付。
最后补一句:那些宣传“消费级显卡跑 7B 全参训练”的教程,十有八九是把冻结基座参数和全参训练混在一起说。全参训练 7B 的显存底线就在那,数学不会骗人。
原创文章,作者:guanweilu,如若转载,请注明出处:https://guanweilu.cn/article/344.html