优化器状态到底占多少显存?Adam 训一个 7B 模型,光状态就要 60GB+

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

AI technology illustration

我最早也踩过这个坑。明明模型只占十几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_avgexp_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

(0)
上一篇 2026年8月28日
下一篇 2026年8月28日

相关推荐