Skip to content

优化器与显存开销 ​

标签
AI/infra/显存管理
AI/infra/优化器
字数
2159 字
阅读时间
9 分钟

这一篇回答一个问题:为什么 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。

完整账本 ​

BF16 参数=2×7×109=14 GBBF16 梯度=2×7×109=14 GB优化器状态=12×7×109=84 GB(FP32 主权重 28+m 28+v 28)合计≈112 GB
显存构成大小占比
优化器状态(FP32 主权重 + m + v)84 GB75%
BF16 梯度14 GB12.5%
BF16 参数14 GB12.5%

这张比例表就是 ZeRO 的设计动机。 ZeRO-1 第一刀砍优化器状态,因为它是占比最大的一块,切它的性价比最高。

而在 DDP 里,每张卡都各存一份完整的 84 GB —— 冗余到了极点。这正是 07-ZeRO 显存优化系列 要解决的问题。

这里的 112 GB 只是**静态显存**(参数 + 梯度 + 优化器状态),不含激活值。激活随 batch / 序列长度 / 深度增长,长序列下常常才是大头(见 [[10-混合精度与显存优化]])。

为什么优化器状态必须用 FP32 ​

一个自然的疑问:既然前向反向都用 BF16 了,优化器为什么不用 BF16 省一半显存?

答案藏在「更新量太小」这件事上:

环节数字
每步参数更新量约等于 学习率 × 梯度,量级常在 10−4 甚至更小
BF16 的尾数只有 7 位
后果表示不了这么小的变化 —— 加到较大的参数上等于没加(大数吃小数)

而且训练要跑成千上万步,每步的舍入误差会不断累积 —— 用低精度存参数最终会导致训练发散。

所以混合精度训练里,优化器必须保存一份 FP32 的「主权重(master weights)」:每步用 FP32 完成更新,再把结果转成 BF16 供下一次前向使用。

这份副本正是那 28 GB 的来源 —— 省不掉。机制的完整解释(fl(x+δ)=x 的判据)见 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 被线性放大。N 张卡就是单卡的 N 倍,卡到几十上百张时有效 batch 轻松到几千甚至几万。

大 Batch 为什么训练不稳 ​

三层原因,递进:

  1. Batch 越大,梯度是越多样本的平均,方差更小、方向更干净 → 理论上可以配更大学习率
  2. 但直接线性放大学习率,训练往往一上来就震荡甚至发散
  3. 更麻烦的是,模型里不同层的参数尺度差异极大(Embedding 层和 LayerNorm 层根本不是一个量级)—— 用一个全局统一的学习率伺候所有层,怎么调都顾此失彼

核心矛盾在第三条:该用逐层的、自适应的学习率,而不是一刀切。

上 LARS / LAMB 之前,先试两招更简单的

  1. 线性缩放规则 —— batch 放大 k 倍,学习率也放大 k 倍
  2. 学习率 warmup —— 前几百到几千步让学习率从很小线性爬到目标值,避开初期剧烈震荡

这两招能覆盖大多数中等规模场景。 只有 batch 大到它们也失效时,才轮到逐层自适应的 LARS / LAMB。

LARS:给每层单独定步长 ​

λl=‖θl‖‖gl‖

直觉:参数本身很大、梯度很小的层,说明它更新得太保守,可以放开步子;反过来梯度相对参数很大的层,得收着点防跑飞。

最早为 ResNet 的大 Batch 训练设计。

LAMB:把 LARS 的思想搬进 Adam ​

Adam + 逐层信赖域缩放:

  1. 每层的更新量先由 Adam 正常算出(已带自适应学习率)
  2. 再乘一个逐层的信赖域缩放因子 ϕ(‖θl‖)/‖rl‖,把该层的更新幅度约束在合理范围内

最著名的战绩:把 BERT 预训练的 batch size 拉到 65536 还能稳定收敛,训练时间从 3 天压到 76 分钟。

优化器基础缩放粒度主要战场
LARSSGD + Momentum逐层(参数范数 / 梯度范数)ResNet 等 CV 大 batch
LAMBAdam逐层(信赖域 × 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」这条约束,才知道它是唯一动得了的那一块

相关 ​

参考 ​

贡献者 ​

文件历史 ​