数据并行是最早、最直观的分布式训练策略。核心思想:把数据切开,每张卡处理一部分,然后同步梯度。但朴素的数据并行有严重的内存冗余问题,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 依赖 |
自检
- DDP 中每张卡保存的优化器状态是一样的吗?(答:是的,完全相同------这就是冗余)
- ZeRO-1 和 DDP 的通信量一样,为什么 ZeRO-1 更好?(答:通信量相同但显存减少 ~4x,因为优化器状态不再冗余)
- ZeRO-3 的通信量为什么乘以层数?(答:每层前向/反向都需要 AllGather 完整参数)
- 训练 70B 模型,4 卡 A100 80GB,至少需要 ZeRO 几?(答:70B × 18 bytes ≈ 1260 GB,4 卡 × 80 GB = 320 GB,需要 ZeRO-3 + Offload,或者更多卡)