Skip to content

DeepSpeed ZeRO 优化器状态分片

DeepSpeed ZeRO(Zero Redundancy Optimizer)通过消除数据并行中的冗余存储,将显存占用从 O(模型大小) 降低到 O(模型大小/N),使训练更大模型成为可能。

DeepSpeed ZeRO架构

冗余分析

标准数据并行中,每个 GPU 存储完整副本:

组件单 GPU 显存N 个 GPU 总显存冗余
模型参数ΨN×ΨN-1 倍冗余
梯度ΨN×ΨN-1 倍冗余
优化器状态 (Adam)2N×ΨN-1 倍冗余
总计4N×Ψ-

ZeRO 逐步消除这些冗余。

ZeRO 阶段

ZeRO-1:优化器状态分片

将优化器状态(Adam 的 m 和 v)分片到各 GPU:

python
# DeepSpeed ZeRO-1 配置
ds_config = {
    "zero_optimization": {
        "stage": 1,
    },
    "train_batch_size": 128,
    "gradient_accumulation_steps": 4,
}
  • 显存:从 4Ψ 降至 2Ψ + 2Ψ/N
  • 通信:与标准 DDP 相同(梯度 AllReduce)

ZeRO-2:梯度分片

在 ZeRO-1 基础上,将梯度也分片存储:

python
ds_config = {
    "zero_optimization": {
        "stage": 2,
        "offload_optimizer": {
            "device": "none",  # 不使用 CPU Offload
        },
    },
}
  • 显存:从 4Ψ 降至 2Ψ/N + Ψ
  • 通信:梯度 ReduceScatter(比 AllReduce 通信量减半)

ZeRO-3:参数分片

将模型参数也分片,仅在计算时聚合:

python
ds_config = {
    "zero_optimization": {
        "stage": 3,
        "overlap_comm": True,       # 通信与计算重叠
        "contiguous_gradients": True,
        "reduce_bucket_size": 5e8,
    },
}
  • 显存:从 4Ψ 降至 4Ψ/N(理想情况)
  • 通信:参数 AllGather + 梯度 ReduceScatter

显存对比

配置7B 模型 (FP16)13B 模型70B 模型
DDP28GB52GB280GB
ZeRO-1 (8GPU)16GB30GB161GB
ZeRO-2 (8GPU)10GB18GB98GB
ZeRO-3 (8GPU)4GB7GB35GB

阶段选择

  • ZeRO-1:最小改动,适合模型接近单 GPU 极限
  • ZeRO-2:最佳平衡,适合大多数场景
  • ZeRO-3:模型远超单 GPU 容量时使用

ZeRO-3 性能

ZeRO-3 的参数 AllGather 增加通信量,小模型上性能损失可达 20-30%。仅在模型确实无法放入单 GPU 时使用 ZeRO-3。

相关资源

最近更新