渐进式去冗余,从优化器状态到参数的三级分片

配套代码


上一章我们看到 DDP 的内存问题:为了保证训练一致性(通过 All-Reduce 同步梯度),每个 GPU 都需要存储完整的模型状态。4 个 GPU 就是 4 份完整副本(参数、梯度、优化器状态)。ZeRO(Zero Redundancy Optimizer)的核心思想很直接:既然最终状态是一致的,那就每个 GPU 只存一部分,需要的时候再通信取回。

训练状态的冗余分析

先量化一下 DDP 的浪费。以混合精度 + Adam 为例,𝑁 N 个 GPU 训练一个 Φ Φ 参数的模型,每个 GPU 需要存储:

合计 16 Φ 16Φ bytes,其中优化器状态占了 75%。

𝑁 N 个 GPU 就是 𝑁 N 倍冗余:全局存储 16 𝑁 Φ 16NΦ bytes,但实际只需要 16 Φ 16Φ bytes。ZeRO 的三个 Stage 就是按从大到小的顺序,依次消除这些冗余。

ZeRO Stage 1:分片优化器状态

Stage 1 只做一件事:把优化器状态均分到 𝑁 N 个 GPU 上

每个参数有一个”owner” rank,只有 owner 存储该参数的 Adam 状态(fp32 参数副本、一阶矩 𝑚 m、二阶矩 𝑣 v)。训练流程变为:

  1. 前向传播:和 DDP 一样,各自独立计算
  2. 反向传播:计算梯度后,通过 reduce(不是 all_reduce)发送到 owner rank
  3. 参数更新:owner rank 更新参数,然后 broadcast 广播更新后的参数

参数分配策略

首先要决定每个参数的 owner。最简单的方式是轮询(round-robin):

systems/distributed_training/02_zero1.py