优化器与显存开销
这一篇回答一个问题:为什么 ZeRO 第一刀就砍向优化器状态?
答案是一笔账 —— 优化器状态占了模型静态显存的 75%。
每个参数要为优化器付多少显存
前提是混合精度训练:前向反向用 BF16/FP16 提速,优化器内部一律 FP32。
按「每个参数额外存几份 FP32」数账(FP32 每数 4 字节):
| 优化器 | 额外状态 | 每参数 | 7B 模型的优化器开销 |
|---|---|---|---|
| SGD | 无 | 4 B(仅 FP32 副本) | 约 28 GB |
| Momentum SGD | 一阶动量 | 8 B | 约 56 GB |
| Adam / AdamW | 一阶 + 二阶动量 | 12 B | 约 84 GB |
AdamW 每参数额外 12 字节。7B 模型光优化器状态就是 84 GB —— 已经超过一张 80 GB 的 A100。
完整账本
| 显存构成 | 大小 | 占比 |
|---|---|---|
| 优化器状态(FP32 主权重 + | 84 GB | 75% |
| BF16 梯度 | 14 GB | 12.5% |
| BF16 参数 | 14 GB | 12.5% |
这张比例表就是 ZeRO 的设计动机。 ZeRO-1 第一刀砍优化器状态,因为它是占比最大的一块,切它的性价比最高。
而在 DDP 里,每张卡都各存一份完整的 84 GB —— 冗余到了极点。这正是 07-ZeRO 显存优化系列 要解决的问题。
这里的 112 GB 只是**静态显存**(参数 + 梯度 + 优化器状态),不含激活值。激活随 batch / 序列长度 / 深度增长,长序列下常常才是大头(见 [[10-混合精度与显存优化]])。
为什么优化器状态必须用 FP32
一个自然的疑问:既然前向反向都用 BF16 了,优化器为什么不用 BF16 省一半显存?
答案藏在「更新量太小」这件事上:
| 环节 | 数字 |
|---|---|
| 每步参数更新量 | 约等于 学习率 × 梯度,量级常在 |
| BF16 的尾数 | 只有 7 位 |
| 后果 | 表示不了这么小的变化 —— 加到较大的参数上等于没加(大数吃小数) |
而且训练要跑成千上万步,每步的舍入误差会不断累积 —— 用低精度存参数最终会导致训练发散。
所以混合精度训练里,优化器必须保存一份 FP32 的「主权重(master weights)」:每步用 FP32 完成更新,再把结果转成 BF16 供下一次前向使用。
这份副本正是那 28 GB 的来源 —— 省不掉。机制的完整解释(
的判据)见 01-数值计算与精度。
「大数吃小数」:低精度存不下这么小的更新量
设某参数 x ≈ 1.0,这一步的更新量 δ ≈ 1e-5
FP32(尾数 23 位):
x 附近可表示的最小间距 ≈ 1e-7
⇒ x + δ 表示得出来 ⇒ 更新生效
BF16(尾数 7 位):
x 附近可表示的最小间距 ≈ 1e-2
⇒ fl(x + δ) = x —— 更新被直接舍掉
x x+δ(想要更新到的地方)
●────────┃
│ │
└── 更新量小于可表示间距 ⇒ 舍入回 x,这一步等于没做
⇒ 训练要跑成千上万步,每步的舍入误差不断累积 ⇒ 用低精度存参数最终会发散
⇒ 所以混合精度训练必须留一份 FP32 主权重:每步在 FP32 上更新,再转 BF16 供前向大 Batch 优化器:LARS 与 LAMB
前面讲的是「优化器占多少显存」。这一节换个视角:分布式本身会给优化器出难题。
数据并行有个直接副作用 —— 有效 Batch Size 被线性放大。
大 Batch 为什么训练不稳
三层原因,递进:
- Batch 越大,梯度是越多样本的平均,方差更小、方向更干净 → 理论上可以配更大学习率
- 但直接线性放大学习率,训练往往一上来就震荡甚至发散
- 更麻烦的是,模型里不同层的参数尺度差异极大(Embedding 层和 LayerNorm 层根本不是一个量级)—— 用一个全局统一的学习率伺候所有层,怎么调都顾此失彼
核心矛盾在第三条:该用逐层的、自适应的学习率,而不是一刀切。
上 LARS / LAMB 之前,先试两招更简单的
- 线性缩放规则 —— batch 放大
倍,学习率也放大 倍 - 学习率 warmup —— 前几百到几千步让学习率从很小线性爬到目标值,避开初期剧烈震荡
这两招能覆盖大多数中等规模场景。 只有 batch 大到它们也失效时,才轮到逐层自适应的 LARS / LAMB。
LARS:给每层单独定步长
直觉:参数本身很大、梯度很小的层,说明它更新得太保守,可以放开步子;反过来梯度相对参数很大的层,得收着点防跑飞。
最早为 ResNet 的大 Batch 训练设计。
LAMB:把 LARS 的思想搬进 Adam
Adam + 逐层信赖域缩放:
- 每层的更新量先由 Adam 正常算出(已带自适应学习率)
- 再乘一个逐层的信赖域缩放因子
,把该层的更新幅度约束在合理范围内
最著名的战绩:把 BERT 预训练的 batch size 拉到 65536 还能稳定收敛,训练时间从 3 天压到 76 分钟。
| 优化器 | 基础 | 缩放粒度 | 主要战场 |
|---|---|---|---|
| LARS | SGD + Momentum | 逐层(参数范数 / 梯度范数) | ResNet 等 CV 大 batch |
| LAMB | Adam | 逐层(信赖域 × Adam 更新) | BERT / Transformer 大 batch |
什么时候该换 LAMB:数据并行卡数很多(64+)、有效 batch 极大、AdamW 已经压不住训练不稳时。
卡数不多、batch 适中的常规场景,老老实实用 AdamW —— 别为了用而用。这是一条与 08-FlashAttention 结尾「生产环境用库」同类的判据:先试简单的,不到边界不引入复杂度。
这条线索的回扣
把这一篇与前面几篇串起来:
优化器状态占静态显存 75%(84 GB / 112 GB)
↓
DDP 每卡冗余存储这一份
↓
ZeRO-1 第一刀切它 ← 性价比最高
↓
但它必须用 FP32(否则更新被舍掉)
↓
所以这份 12Ψ 省不掉,只能切开「必须用 FP32」与「占比 75%」这两条合起来,才完整解释了 ZeRO 的设计顺序 —— 只看到占比会以为随便切哪块都行,看到 FP32 的约束才知道这块既不能降精度、又必须动它。
从「优化器占 75%」到「ZeRO 只能切、不能省」:
优化器状态占静态显存 75%(84 GB / 112 GB)
│
▼
DDP 里每张卡都冗余存这一份(冗余到了极点)
│
▼
ZeRO-1 第一刀就是它 —— 占比最大,切它的性价比最高
│
▼
但它必须用 FP32(否则更新被舍掉、训练发散)
│
▼
这份 12Ψ 既不能降精度、又必须动它 ⇒ 只能切开,不能省掉
⇒ 只看到「占 75%」会以为随便切哪一块都行;
再加上「必须 FP32」这条约束,才知道它是唯一动得了的那一块相关
- 02-分布式训练总论与显存账本 ——
的完整账本 - 07-ZeRO 显存优化系列 —— 切分优化器状态的三个阶段
- 01-数值计算与精度 —— 为什么低精度更新会被舍掉
- 06-数据并行:DP、DDP 与 FSDP —— 有效 batch 随卡数放大
- 10-混合精度与显存优化 —— 混合精度流程与 master weight 的位置
参考
- https://caomaolufei.github.io/AIInfraGuide/guides/模块三-分布式训练/31-优化器原理与显存开销分析
- https://arxiv.org/abs/1412.6980
- https://arxiv.org/abs/1904.00962
- https://arxiv.org/abs/1708.03888
- https://arxiv.org/abs/1910.02054
YJ