【infra之路】06-数据并行-DDP与ZeRO三阶段

数据并行是最早、最直观的分布式训练策略。核心思想:把数据切开,每张卡处理一部分,然后同步梯度。但朴素的数据并行有严重的内存冗余问题,ZeRO(Zero Redundancy Optimizer)通过三个阶段逐步消除冗余,是 DeepSpeed 和 FSDP 的核心。


回顾:单卡训练的显存去哪了?

在讲分布式之前,先搞清楚单卡训练的显存构成。以 7B 参数模型、FP16 混合精度训练为例:

复制代码
┌─────────────────────────────────────────────────────────┐
│                    单卡显存占用                            │
├─────────────────┬───────────────┬───────────────────────┤
│ 模型参数         │ 7B × 2B      │ = 14 GB               │
│ 梯度             │ 7B × 2B      │ = 14 GB               │
│ 优化器状态        │ 7B × 12B     │ = 84 GB               │
│  (Adam: FP32副本  │  7B × 4B = 28 GB                    │
│   一阶动量 m      │  7B × 4B = 28 GB                    │
│   二阶动量 v      │  7B × 4B = 28 GB)                   │
│ 激活值           │ 变动          │ ~10-50 GB             │
├─────────────────┼───────────────┼───────────────────────┤
│ 合计             │               │ ~122-162 GB           │
└─────────────────┴───────────────┴───────────────────────┘

关键观察:优化器状态占了总显存的 68% (84 / 122)。这是 ZeRO 重点优化的目标。

想知道为什么优化器占用了这么多显存可以点这里

朴素数据并行(Data Parallelism, DP)

工作原理

复制代码
           训练数据
          /   |   \
         /    |    \
    ┌────┐ ┌────┐ ┌────┐
    │GPU0│ │GPU1│ │GPU2│    每张卡持有完整模型副本
    │完整│ │完整│ │完整│    每张卡处理不同的数据 batch
    │模型│ │模型│ │模型│
    └──┬─┘ └──┬─┘ └──┬─┘
       │      │      │
    梯度0  梯度1  梯度2     各自计算梯度
       │      │      │
       └──────┼──────┘
              ▼
         AllReduce           所有梯度求平均
              │
       ┌──────┼──────┐
       ▼      ▼      ▼
    更新0   更新1   更新2     各自用平均梯度更新模型
    (相同)  (相同)  (相同)    更新后模型仍然一致

每张 GPU 持有完整的模型副本,处理不同的数据子集,计算完梯度后做 AllReduce 求平均,然后各自更新参数。因为初始参数相同、梯度相同、学习率相同,所以更新后的参数也相同------保持同步。

PyTorch DDP(Distributed Data Parallel)

DDP 是 DP 的生产级实现,比旧版 DataParallel(单进程多线程)更稳定、更高效(多进程)。

python 复制代码
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP

# 初始化(torchrun 自动设置环境变量)
dist.init_process_group(backend='nccl')
local_rank = int(os.environ['LOCAL_RANK'])
torch.cuda.set_device(local_rank)

# 包装模型
model = MyModel().to(local_rank)
model = DDP(model, device_ids=[local_rank])

# 训练循环和单卡完全一样
for batch in dataloader:
    loss = model(batch).loss
    loss.backward()          # DDP 自动在 backward 中插入 AllReduce
    optimizer.step()
    optimizer.zero_grad()

DDP 的核心优化:梯度桶化(Gradient Bucketing)。它不会等所有梯度都算完再 AllReduce,而是把梯度分成若干 bucket,边算边同步------和反向传播的计算重叠,隐藏通信延迟。

复制代码
反向传播:  [layer_n 梯度] [layer_n-1 梯度] [layer_n-2 梯度] ...
AllReduce:               [bucket 1     ]  [bucket 2     ] ...
                         ↑ 计算和通信重叠(overlap)

DDP 的致命问题:显存冗余

DDP 中每张卡都保存完整的模型参数、梯度、优化器状态。

复制代码
4 卡 DDP 训练 7B 模型:

每张卡: 14 GB (参数) + 14 GB (梯度) + 84 GB (优化器) = 112 GB
4 张卡总计: 112 GB × 4 = 448 GB

但其中真正不同的只有梯度(因为数据不同)!
参数和优化器状态在每张卡上完全相同 → 3/4 是浪费的冗余

这就是 ZeRO 要解决的问题。


ZeRO:零冗余优化器

ZeRO 的核心思想:把冗余的数据切分到不同 GPU 上,每个 GPU 只保存 1/N 的状态

ZeRO 分三个阶段,逐步切分更多内容:

复制代码
                    切分内容                    显存节省
                ┌──────────────┐           ┌──────────┐
DDP (无切分)     │ 参数 梯度 优化 │           │ 0%       │
                ├──────────────┤           ├──────────┤
ZeRO-1          │ 参数 梯度 [优化]│  →  切分  │ ~4x      │
                ├──────────────┤           ├──────────┤
ZeRO-2          │ 参数 [梯度][优化]│ →  切分  │ ~8x      │
                ├──────────────┤           ├──────────┤
ZeRO-3          │[参数][梯度][优化]│ → 全切分 │ ~N x     │
                └──────────────┘           └──────────┘

ZeRO Stage 1:切分优化器状态

原理:Adam 的优化器状态(FP32 参数副本、一阶动量 m、二阶动量 v)占 84 GB,是显存的大头。把它平均分成 N 份,每个 GPU 只保存 1/N。

复制代码
                GPU 0          GPU 1          GPU 2          GPU 3
                ┌──────┐      ┌──────┐      ┌──────┐      ┌──────┐
模型参数 (FP16)  │ 完整  │      │ 完整  │      │ 完整  │      │ 完整  │  ← 未切分
                ├──────┤      ├──────┤      ├──────┤      ├──────┤
梯度 (FP16)     │ 完整  │      │ 完整  │      │ 完整  │      │ 完整  │  ← 未切分
                ├──────┤      ├──────┤      ├──────┤      ├──────┤
优化器状态       │ 1/4  │      │ 1/4  │      │ 1/4  │      │ 1/4  │  ← 切分!
(FP32参数+m+v)  │ 21GB │      │ 21GB │      │ 21GB │      │ 21GB │
                └──────┘      └──────┘      └──────┘      └──────┘

训练流程变化

复制代码
1. 前向传播:和 DDP 一样(每张卡有完整参数)
2. 反向传播:和 DDP 一样(每张卡算出完整梯度)
3. 梯度同步:ReduceScatter(不是 AllReduce!)
   → 每个 GPU 只保留自己负责那 1/N 参数的梯度
4. 参数更新:每个 GPU 只更新自己负责的 1/N 参数
5. 参数同步:AllGather(把更新后的参数收集给所有人)

DP/DDP 用 AllReduce:
  GPU 0: [全部梯度] ─┐
  GPU 1: [全部梯度] ─┼─ AllReduce → 每个 GPU 拿到 [全部平均梯度]
  GPU 2: [全部梯度] ─┤             然后各自更新全部参数
  GPU 3: [全部梯度] ─┘

ZeRO-1 用 ReduceScatter + AllGather:
  GPU 0: [全部梯度] ─┐
  GPU 1: [全部梯度] ─┼─ ReduceScatter → GPU 0 拿 [梯度 1/4]
  GPU 2: [全部梯度] ─┤                  GPU 1 拿 [梯度 2/4]
  GPU 3: [全部梯度] ─┘                  GPU 2 拿 [梯度 3/4]
                                        GPU 3 拿 [梯度 4/4]
                     各自更新自己负责的 1/4 参数
                        ↓
                     AllGather → 每个 GPU 拿到更新后的完整参数

通信量分析

步骤 DDP (AllReduce) ZeRO-1 (ReduceScatter + AllGather)
通信量 2 × 参数量 参数量 + 参数量 = 2 × 参数量
结论 通信量完全相同! 但显存省了 ~4x

这是 ZeRO 最精妙的地方------Stage 1 在不增加通信量的情况下,把优化器状态切分了。

显存计算(4 卡,7B 模型,FP16)

复制代码
每张卡显存 = 参数(完整) + 梯度(完整) + 优化器状态(1/4)
           = 14 GB + 14 GB + 21 GB
           = 49 GB

对比 DDP: 14 + 14 + 84 = 112 GB → 节省 56%

ZeRO Stage 2:切分优化器状态 + 梯度

原理:既然梯度在 ReduceScatter 之后每个 GPU 只需要 1/N,那就只保存 1/N 的梯度。

复制代码
                GPU 0          GPU 1          GPU 2          GPU 3
                ┌──────┐      ┌──────┐      ┌──────┐      ┌──────┐
模型参数 (FP16) │ 完整  │      │ 完整  │      │ 完整  │     │ 完整  │  ← 未切分
                ├──────┤      ├──────┤      ├──────┤      ├──────┤
梯度 (FP16)     │ 1/4  │      │ 1/4  │      │ 1/4  │      │ 1/4  │  ← 也切分了!
                ├──────┤      ├──────┤      ├──────┤      ├──────┤
优化器状态       │ 1/4  │      │ 1/4  │      │ 1/4  │      │ 1/4  │  ← 切分
                └──────┘      └──────┘      └──────┘      └──────┘

关键实现细节:反向传播时,梯度仍然是完整计算的(因为参数是完整的)。但每算完一层的梯度,就立刻做 ReduceScatter,把这一层的梯度分发到负责的 GPU,释放掉其他 GPU 上的副本。这样峰值显存更低。

复制代码
反向传播过程:

Layer N:  计算梯度 → ReduceScatter → 只保留自己负责的 1/N 梯度 → 释放其余
Layer N-1: 计算梯度 → ReduceScatter → 只保留自己负责的 1/N 梯度 → 释放其余
...
Layer 0:  计算梯度 → ReduceScatter → 只保留自己负责的 1/N 梯度 → 释放其余

通信量:和 Stage 1 相同(ReduceScatter 的通信量不变)。

显存计算(4 卡,7B 模型,FP16)

复制代码
每张卡显存 = 参数(完整) + 梯度(1/4) + 优化器状态(1/4)
           = 14 GB + 3.5 GB + 21 GB
           = 38.5 GB

对比 DDP: 112 GB → 节省 66%
对比 Stage 1: 49 GB → 再省 21%

ZeRO Stage 3:切分优化器状态 + 梯度 + 参数

原理:参数也切成 N 份,每个 GPU 只保存 1/N 的参数。这是最激进的切分。

复制代码
                GPU 0          GPU 1          GPU 2          GPU 3
                ┌──────┐      ┌──────┐      ┌──────┐      ┌──────┐
模型参数 (FP16)  │ 1/4  │      │ 1/4  │      │ 1/4  │      │ 1/4  │  ← 全切分了!
                ├──────┤      ├──────┤      ├──────┤      ├──────┤
梯度 (FP16)     │ 1/4  │      │ 1/4  │      │ 1/4  │      │ 1/4  │
                ├──────┤      ├──────┤      ├──────┤      ├──────┤
优化器状态       │ 1/4  │      │ 1/4  │      │ 1/4  │      │ 1/4  │
                └──────┘      └──────┘      └──────┘      └──────┘

问题:前向/反向传播时需要完整参数,但每个 GPU 只有 1/N。怎么办?

解决方案:按需 AllGather。

复制代码
前向传播每一层:
  1. AllGather 该层的参数(从 4 个 GPU 收集完整参数)
  2. 用完整参数做前向计算
  3. 丢弃非自己负责的参数(释放显存)

反向传播每一层:
  1. AllGather 该层的参数
  2. 用完整参数做反向计算
  3. ReduceScatter 梯度(分发给负责的 GPU)
  4. 丢弃非自己负责的参数

Layer 3:  AllGather → 前向 → 丢弃  |  AllGather → 反向 → ReduceScatter → 丢弃
Layer 2:  AllGather → 前向 → 丢弃  |  AllGather → 反向 → ReduceScatter → 丢弃
Layer 1:  AllGather → 前向 → 丢弃  |  AllGather → 反向 → ReduceScatter → 丢弃
Layer 0:  AllGather → 前向 → 丢弃  |  AllGather → 反向 → ReduceScatter → 丢弃

通信量分析

阶段 通信量
前向传播 每层 AllGather 参数量 → 总计 参数量 × 层数
反向传播 每层 AllGather + ReduceScatter → 总计 2 × 参数量 × 层数
参数更新 ReduceScatter 梯度 + AllGather 参数 → 2 × 参数量
总计 ~3 × 参数量 × 层数

Stage 3 的通信量远大于 Stage 1/2(乘以层数),这就是代价。

显存计算(4 卡,7B 模型,FP16)

复制代码
每张卡显存 = 参数(1/4) + 梯度(1/4) + 优化器状态(1/4) + 临时完整参数(前向/反向)
           = 3.5 GB + 3.5 GB + 21 GB + 14 GB (临时)
           ≈ 42 GB (峰值,含临时参数)

但参数和状态部分随 GPU 数量线性下降!
8 卡时: 1.75 + 1.75 + 10.5 + 14 = 28 GB
16 卡时: 0.875 + 0.875 + 5.25 + 14 = 21 GB

注意:临时参数(14 GB)不随 GPU 数量变化,这是 Stage 3 的显存下限。可以通过 参数预取(prefetch)梯度检查点(gradient checkpointing) 来优化。


ZeRO 三阶段完整对比

复制代码
显存占用(7B 模型,FP16,4 卡):

DDP:      ████████████████████████████████████████████████████ 112 GB/卡
ZeRO-1:   █████████████████████████ 49 GB/卡
ZeRO-2:   ███████████████████ 38.5 GB/卡
ZeRO-3:   ███████████████████ 42 GB/卡 (峰值,含临时参数)

通信量(每步):

DDP:      ████████████████ 2M (AllReduce)
ZeRO-1:   ████████████████ 2M (ReduceScatter + AllGather)
ZeRO-2:   ████████████████ 2M (逐层 ReduceScatter)
ZeRO-3:   ██████████████████████████████████████████ ~3M × 层数
DDP ZeRO-1 ZeRO-2 ZeRO-3
切分内容 优化器状态 优化器+梯度 优化器+梯度+参数
通信量 2M 2M 2M 3M × 层数
通信原语 AllReduce ReduceScatter + AllGather 逐层 ReduceScatter 逐层 AllGather + ReduceScatter
显存节省 ~4x ~8x ~N x(随 GPU 数线性)
适用场景 模型能放进单卡 模型参数能放进单卡 同左,梯度也放不下时 模型参数都放不下时
实现难度

DeepSpeed 实战代码

基本使用(ZeRO-2)

python 复制代码
# ds_config.json
{
    "train_batch_size": 32,
    "gradient_accumulation_steps": 4,
    "fp16": { "enabled": true },
    "zero_optimization": {
        "stage": 2,
        "offload_optimizer": {
            "device": "none"  // 或 "cpu" 启用 ZeRO-Offload
        },
        "allgather_partitions": true,
        "allgather_bucket_size": 5e8,
        "reduce_scatter": true,
        "reduce_bucket_size": 5e8,
        "overlap_comm": true  // 通信和计算重叠
    }
}
python 复制代码
# train.py
import deepspeed
import argparse

# DeepSpeed 会自动解析 config
parser = argparse.ArgumentParser()
parser = deepspeed.add_config_arguments(parser)
args = parser.parse_args()

model = MyModel()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)

# 一行代码启用 ZeRO
model, optimizer, _, _ = deepspeed.initialize(
    args=args,
    model=model,
    optimizer=optimizer,
    model_parameters=model.parameters(),
    config="ds_config.json"
)

# 训练循环和单卡几乎一样
for batch in dataloader:
    loss = model(batch).loss
    model.backward(loss)       # 自动处理梯度 + ZeRO 切分
    model.step()               # 自动处理参数更新 + AllGather

启动:

bash 复制代码
deepspeed --num_gpus=4 train.py --deepspeed_config ds_config.json
# 或多机:
deepspeed --num_gpus=4 --num_nodes=2 --hostfile hostfile train.py

ZeRO-3 配置

json 复制代码
{
    "zero_optimization": {
        "stage": 3,
        "stage3_max_live_parameters": 1e9,
        "stage3_max_reuse_distance": 1e9,
        "stage3_prefetch_bucket_size": 5e7,
        "stage3_param_persistence_threshold": 1e5,
        "reduce_bucket_size": 5e8,
        "sub_group_size": 1e9,
        "overlap_comm": true,
        "offload_optimizer": {
            "device": "cpu",
            "pin_memory": true
        },
        "offload_param": {
            "device": "cpu",
            "pin_memory": true
        }
    }
}

ZeRO-Offload:把显存压力转移到 CPU 内存

复制代码
正常 ZeRO-3:
  GPU 显存: [参数 1/N] + [梯度 1/N] + [优化器 1/N] + [临时参数]

ZeRO-Offload:
  GPU 显存: [临时参数] + [计算中的激活值]
  CPU 内存: [参数 1/N] + [梯度 1/N] + [优化器 1/N]  ← 移到 CPU!

  CPU 做 Adam 更新 → 通过 PCIe 传回 GPU

适用场景:GPU 显存实在不够时(比如单卡训练 7B 模型),用 CPU 内存换 GPU 显存。代价是 PCIe 带宽 (~32 GB/s) 远慢于 GPU 显存带宽 (~2 TB/s),训练速度会降 2-5 倍。


PyTorch FSDP(Fully Sharded Data Parallel)

FSDP 是 PyTorch 原生实现的 ZeRO-3,不需要 DeepSpeed 依赖。FSDP2(PyTorch 2.x)基于 DTensor,更加灵活。

python 复制代码
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
from functools import partial

# 定义哪些模块要被独立切分(通常是 Transformer Layer)
auto_wrap_policy = partial(
    transformer_auto_wrap_policy,
    transformer_layer_cls={TransformerBlock}
)

# 包装模型
model = FSDP(
    model,
    auto_wrap_policy=auto_wrap_policy,
    sharding_strategy=ShardingStrategy.FULL_SHARD,  # = ZeRO-3
    # sharding_strategy=ShardingStrategy.SHARD_GRAD_OP,  # = ZeRO-2
    # sharding_strategy=ShardingStrategy.NO_SHARD,  # = DDP
    mixed_precision=MixedPrecision(
        param_dtype=torch.float16,
        reduce_dtype=torch.float16,
        buffer_dtype=torch.float16,
    ),
    device_id=torch.cuda.current_device(),
)

# 训练循环和 DDP 完全一样
for batch in dataloader:
    loss = model(batch).loss
    loss.backward()    # FSDP 自动处理 AllGather + ReduceScatter
    optimizer.step()
    optimizer.zero_grad()

FSDP vs DeepSpeed 对比

FSDP DeepSpeed
维护方 PyTorch 官方 Microsoft
集成度 PyTorch 原生,无需额外依赖 需要 pip install deepspeed
ZeRO 阶段 Stage 2 (SHARD_GRAD_OP), Stage 3 (FULL_SHARD) Stage 1/2/3 都支持
ZeRO-Offload 有限支持 完整支持
MoE 支持
灵活性 与 PyTorch 生态无缝集成 功能更全,但耦合度高

三阶段通信量总结(N 个 GPU,模型参数量 M)

阶段 前向通信 反向通信 更新通信 每步总通信
DDP 0 AllReduce 2M 0 2M
ZeRO-1 0 ReduceScatter M AllGather M 2M
ZeRO-2 0 逐层 ReduceScatter M AllGather M 2M
ZeRO-3 逐层 AllGather M×L 逐层 AllGather+ReduceScatter 2M×L AllGather M 3M×L + M

L = 模型层数。可以看到 Stage 1 和 Stage 2 的通信量和 DDP 完全相同,但显存大幅减少------这就是为什么它们是首选方案。


本课小结

概念 要点
DDP 完整模型副本 + 梯度 AllReduce,简单但有显存冗余
ZeRO-1 切分优化器状态,省 ~4x 显存,通信量不变
ZeRO-2 +切分梯度,省 ~8x 显存,通信量不变
ZeRO-3 +切分参数,显存随 GPU 数线性下降,但通信量 ×层数
ZeRO-Offload 把优化器/参数移到 CPU 内存,用 PCIe 带宽换显存
FSDP PyTorch 原生 ZeRO-3 实现,无需 DeepSpeed 依赖

自检

  1. DDP 中每张卡保存的优化器状态是一样的吗?(答:是的,完全相同------这就是冗余)
  2. ZeRO-1 和 DDP 的通信量一样,为什么 ZeRO-1 更好?(答:通信量相同但显存减少 ~4x,因为优化器状态不再冗余)
  3. ZeRO-3 的通信量为什么乘以层数?(答:每层前向/反向都需要 AllGather 完整参数)
  4. 训练 70B 模型,4 卡 A100 80GB,至少需要 ZeRO 几?(答:70B × 18 bytes ≈ 1260 GB,4 卡 × 80 GB = 320 GB,需要 ZeRO-3 + Offload,或者更多卡)