好问题,这是理解 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(不含临时参数)
└────────────────┘