【infra之路】详解优化器状态为什么会占用显存

好问题,这是理解 ZeRO 为什么有效的前提。让我从 Adam 优化器的数学原理讲起。

可以看一下---神经网络中优化器的作用

Adam 为什么要保存这么多状态?

先回忆梯度下降的本质:参数沿着梯度方向更新 。但"直接用梯度更新"有很多问题------梯度噪声大、不同参数尺度不同、容易陷入局部最优。Adam 通过维护参数的历史梯度统计信息来解决这些问题。

Adam 的更新公式

对每个参数 θ i \theta_i θi,Adam 维护三个东西:

复制代码
第 1 个:FP32 参数副本 θ_i
  → 因为训练用 FP16,但优化器计算需要 FP32 精度,否则会溢出/下溢

第 2 个:一阶动量 m_i(梯度的指数移动平均)
  m_i = β₁ × m_i + (1 - β₁) × g_i
  → 记录"梯度大致朝哪个方向",平滑噪声

第 3 个:二阶动量 v_i(梯度平方的指数移动平均)
  v_i = β₂ × v_i + (1 - β₂) × g_i²
  → 记录"梯度的变化幅度",自动调节学习率

更新参数时:

复制代码
m̂_i = m_i / (1 - β₁ᵗ)     ← 偏差修正
v̂_i = v_i / (1 - β₂ᵗ)     ← 偏差修正
θ_i = θ_i - lr × m̂_i / (√v̂_i + ε)

显存怎么算的

每个参数需要保存 4 个数值

状态 精度 每参数字节 7B 模型总计
FP16 参数(模型本身) FP16 2 bytes 14 GB
FP16 梯度 FP16 2 bytes 14 GB
FP32 参数副本 FP32 4 bytes 28 GB
一阶动量 m FP32 4 bytes 28 GB
二阶动量 v FP32 4 bytes 28 GB
合计 16 bytes/参数 112 GB

其中优化器状态(FP32 副本 + m + v)= 12 bytes/参数 × 7B = 84 GB

这就是 84 GB 的来源:3 个 FP32 变量 × 4 bytes × 70 亿参数

为什么需要 FP32 副本?

训练时模型参数用 FP16(省显存、加速计算),但 FP16 的范围太小(最大值 65504,精度只有 ~3 位有效数字)。优化器的更新量通常很小(比如 1e-5),如果用 FP16 做 θ - lr × m̂ / √v̂

复制代码
FP16:  θ = 1.234
       lr × m̂ / √v̂ = 0.00001
       θ_new = 1.234 - 0.00001 = 1.234  ← 更新被截断了!完全没变化

FP32:  θ = 1.234567
       lr × m̂ / √v̂ = 0.00001
       θ_new = 1.234557  ← 更新被正确保留

所以 Adam 内部必须用 FP32 做计算和保存状态,更新完后再截断回 FP16 给模型用。

不同优化器的显存对比

优化器 每参数额外状态 额外 bytes/参数 7B 模型额外显存
SGD(无动量) 0 0 GB
SGD + Momentum m(动量) 4 (FP32) 28 GB
AdaGrad v(梯度平方累积) 4 (FP32) 28 GB
Adam FP32副本 + m + v 12 (FP32) 84 GB
AdamW 同 Adam 12 (FP32) 84 GB
8-bit Adam FP32副本 + m(8bit) + v(8bit) 4+1+1 = 6 42 GB
Adafactor 分解的 m, v(行列分解) ~2 ~14 GB

Adam 是显存消耗最大的主流优化器。这就是为什么大模型训练领域有这么多显存优化技术的原因。

显存优化的几种思路

1. 8-bit Adam(bitsandbytes)

把 m 和 v 从 FP32 量化到 INT8(1 byte),节省 6 bytes/参数:

复制代码
原始 Adam: FP32副本(4) + m(4) + v(4) = 12 bytes/参数
8-bit Adam: FP32副本(4) + m_INT8(1) + v_INT8(1) = 6 bytes/参数 → 节省 50%

原理:动量值 m 和 v 的分布比较集中,可以用动态缩放量化到 INT8 而不损失太多精度。

2. Adafactor

利用 Transformer 的结构特点------参数大多是矩阵(m×n),把 m 和 v 分解成行向量和列向量

复制代码
标准 Adam:  v 是 m×n 矩阵 → m×n 个值
Adafactor:  v 分解为 row_factor(m) + col_factor(n) → m+n 个值

例:一个 4096×4096 的权重矩阵
  Adam:     4096 × 4096 = 16,777,216 个 v 值
  Adafactor: 4096 + 4096 = 8,192 个值 → 节省 2000 倍!

代价:精度有损失,某些模型收敛变慢。

3. 去掉 FP32 副本(纯 FP16 优化)

一些研究尝试完全用 FP16 做优化器计算(如 BF16 优化器),省掉 4 bytes/参数的 FP32 副本。BF16 因为指数位和 FP32 一样(8 bit),范围够大,可以在某些场景替代 FP32。

4. ZeRO 切分(本课文内容)

不改优化器本身,而是把状态分散到多卡上。4 卡 ZeRO-1/2/3 时,每卡只保存 1/4 的优化器状态 → 84/4 = 21 GB。

一张图总结

复制代码
7B 模型训练显存构成(FP16 混合精度 + Adam):

                    112 GB 总计
  ┌──────────────────────────────────────────────┐
  │                                              │
  │  模型参数 14GB  │  梯度 14GB  │   优化器状态 84GB      │
  │  (FP16)        │  (FP16)    │                       │
  │  ████          │  ████      │  ████████████████████ │
  │                │            │  FP32副本  m    v     │
  │                │            │   28GB    28GB  28GB  │
  └──────────────────────────────────────────────┘
                      ↑                    ↑
                  真正需要的            68% 是冗余的
                  (不同数据产生          (每张卡都一样)
                   不同梯度)

ZeRO-1 切分优化器状态后(4 卡):
  ┌───────────────────────────────┐
  │  参数 14GB │ 梯度 14GB │ 优化器 21GB │  ← 每卡 49 GB
  └───────────────────────────────┘

ZeRO-3 全切分后(4 卡):
  ┌────────────────┐
  │ 参数 3.5 │ 梯度 3.5 │ 优化器 21GB │  ← 每卡 28 GB(不含临时参数)
  └────────────────┘