分布式训练 — 从单卡到千卡集群
#分布式训练 · #DeepSpeed · #ZeRO · #FSDP · #数据并行 · #模型并行 · #流水线并行
当单卡放不下模型或训练太慢时,分布式训练是必经之路。本文从数据并行出发,逐步深入到 ZeRO、3D 并行和 FSDP 实战。
为什么需要分布式训练?
mermaid
graph TD
subgraph "单卡训练困境"
SINGLE["单张 GPU<br/>80GB A100"] --> Q1{"模型大小?"}
Q1 -->|"< 80GB"| OK["可以训练<br/>但慢"]
Q1 -->|"> 80GB"| FAIL["❌ OOM<br/>无法训练"]
end
subgraph "分布式训练解决"
DIST["多张 GPU 协作"] --> DP["数据并行<br/>每卡一份模型"]
DIST --> MP["模型并行<br/>每卡一部分层"]
DIST --> ZERO["ZeRO 优化<br/>切分优化器状态"]
end
style FAIL fill:#e74c3c,color:#fff
style DIST fill:#2ecc71,color:#fff| 模型 | 参数量 | FP16 权重 | 梯度 | 优化器状态(Adam) | 总需求 | 能单卡训练? |
|---|---|---|---|---|---|---|
| Qwen4-7B | 7B | 14 GB | 14 GB | 28 GB | 56 GB | ✅ A100 80GB |
| Qwen4-72B | 72B | 144 GB | 144 GB | 288 GB | 576 GB | ❌ 需要 8×A100 |
| LLaMA 4 Scout | 109B×MoE | 218 GB | — | — | >1TB | ❌ 需要多机 |
注意:上面是"最小需求"。实际训练还需要激活值存储(Activation Memory),通常是权重的 2-5 倍。所以 7B 模型用 LoRA 才能在 24GB 消费级显卡上跑。
数据并行 (Data Parallel, DP)
原理
每张 GPU 持有完整模型副本,但处理不同 batch 的数据。梯度在所有 GPU 间平均后统一更新。
mermaid
graph TD
subgraph "数据并行流程"
DATA["训练数据<br/>Batch = 256"] --> SPLIT["切分为 4 份"]
SPLIT --> GPU0["GPU 0<br/>batch 64"]
SPLIT --> GPU1["GPU 1<br/>batch 64"]
SPLIT --> GPU2["GPU 2<br/>batch 64"]
SPLIT --> GPU3["GPU 3<br/>batch 64"]
GPU0 --> G0["梯度₀"]
GPU1 --> G1["梯度₁"]
GPU2 --> G2["梯度₂"]
GPU3 --> G3["梯度₃"]
G0 --> ALLREDUCE["AllReduce<br/>平均所有梯度"]
G1 --> ALLREDUCE
G2 --> ALLREDUCE
G3 --> ALLREDUCE
ALLREDUCE --> UPDATE["每张卡用平均梯度更新参数<br/>保证模型副本一致"]
end
style ALLREDUCE fill:#e74c3c,color:#fffPyTorch DDP 实现
python
"""
PyTorch DistributedDataParallel (DDP) 完整示例
"""
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler
def setup_distributed(rank: int, world_size: int):
"""初始化分布式环境"""
dist.init_process_group(
backend="nccl", # NVIDIA NCCL 通信库
init_method="tcp://localhost:12355",
rank=rank,
world_size=world_size,
)
torch.cuda.set_device(rank)
def cleanup():
dist.destroy_process_group()
def train_ddp(rank: int, world_size: int):
"""单进程的训练入口"""
setup_distributed(rank, world_size)
# 模型 + DDP 包装
model = nn.Linear(1024, 10).cuda(rank)
model = DDP(model, device_ids=[rank])
# 数据集 + 分布式采样器
dataset = torch.randn(10000, 1024) # 示例数据
sampler = DistributedSampler(
dataset,
num_replicas=world_size,
rank=rank,
shuffle=True,
)
dataloader = DataLoader(dataset, batch_size=64, sampler=sampler)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
criterion = nn.MSELoss()
for epoch in range(10):
sampler.set_epoch(epoch) # 每 epoch 重新 shuffle
for batch in dataloader:
data = batch.cuda(rank)
target = torch.randn(64, 10).cuda(rank)
output = model(data)
loss = criterion(output, target)
optimizer.zero_grad()
loss.backward() # DDP 自动 AllReduce 梯度
optimizer.step()
if rank == 0:
print(f"Epoch {epoch}, Loss: {loss.item():.4f}")
cleanup()
# 启动方式 (命令行):
# torchrun --nproc_per_node=4 train.py数据并行的瓶颈
| 问题 | 原因 | 影响 |
|---|---|---|
| 显存冗余 | 每张卡存完整优化器状态 | 训练大模型受单卡显存限制 |
| 通信开销 | AllReduce 传输整个梯度 | 卡数越多通信占比越大 |
| 负载不均 | 最后一张卡可能数据少 | 计算完成时间不同步 |
通信量量化
数据并行的通信瓶颈可以用公式精确定量。对于参数量为
| 通信原语 | 单步通信量 (每卡) | 说明 |
|---|---|---|
| DDP AllReduce | 梯度总大小 × 2(Ring发送+接收),N 为卡数 | |
| ZeRO-3 AllGather | 前向时收集他卡参数 | |
| ZeRO-3 ReduceScatter | 反向时归约梯度 |
实例:7B 模型 FP16 训练,
= 14 GB,64 张卡下:
- DDP 每步通信量 ≈
≈ 27.6 GB/卡——这就是为什么纯数据并行扩展性差 - ZeRO-3 每步(前向+反向)≈
≈ 27.6 GB/卡——与 DDP 总量相同,但显存优势巨大
python
"""
通信量计算器
"""
def communication_cost(
params_gb: float, # 参数量 (GB, FP16)
num_gpus: int, # GPU 数量
strategy: str = "ddp", # ddp | zero3
) -> dict:
"""计算单步通信量"""
phi = params_gb * 1e9 # 转为 bytes
factor = (num_gpus - 1) / num_gpus
if strategy == "ddp":
# Ring AllReduce: 2× 数据量
per_gpu = 2 * phi * factor
elif strategy == "zero3":
# AllGather(fwd) + ReduceScatter(bwd)
per_gpu = 2 * phi * factor
# 假设 NVLink 600 GB/s 双向带宽
bw = 600e9
time_ms = (per_gpu / bw) * 1000
return {
"per_gpu_gb": per_gpu / 1e9,
"total_gb": per_gpu * num_gpus / 1e9,
"time_ms": time_ms,
"time_percent": time_ms / 100, # 假设总步时 ~100ms
}
# 7B 模型, 8 卡
cost = communication_cost(14, 8, "ddp")
print(f"DDP: {cost['per_gpu_gb']:.1f} GB/卡, "
f"通信占比 ~{cost['time_percent']:.0f}%")
# 7B 模型, 64 卡 — 通信占比爆炸
cost = communication_cost(14, 64, "ddp")
print(f"DDP 64卡: {cost['per_gpu_gb']:.1f} GB/卡, "
f"通信占比 ~{cost['time_percent']:.0f}% ← 严重瓶颈")| 卡数 | DDP 通信量/卡 | 通信耗时 (NVLink) | 通信占比 |
|---|---|---|---|
| 4 | ~21 GB | ~35 ms | ~35% |
| 8 | ~24.5 GB | ~41 ms | ~41% |
| 16 | ~26.3 GB | ~44 ms | ~44% |
| 64 | ~27.6 GB | ~46 ms | ~46% |
关键结论:DDP 通信量随卡数增长趋于饱和(
),但通信占比始终很高。这就是为什么大模型训练必须引入 TP/PP 来减少 DP 维度。
ZeRO — 优化器状态切分
ZeRO 的三个阶段
DeepSpeed ZeRO (Zero Redundancy Optimizer) 通过切分优化器状态、梯度和参数,逐步消除数据并行中的显存冗余。
mermaid
graph TD
subgraph "ZeRO 三个阶段"
BASE["基础数据并行<br/>每卡完整副本"] --> Z1["ZeRO-1<br/>切分优化器状态<br/>(β₁, β₂, momentum)"]
Z1 --> Z2["ZeRO-2<br/>+ 切分梯度"]
Z2 --> Z3["ZeRO-3<br/>+ 切分模型参数<br/>几乎线性扩展"]
end
style BASE fill:#f39c12,color:#fff
style Z3 fill:#2ecc71,color:#fff显存节省对比
| 阶段 | 优化器状态 | 梯度 | 参数 | 单卡 7B 显存 | 能训多大 |
|---|---|---|---|---|---|
| 无并行 | 完整 | 完整 | 完整 | 56 GB | ~10B |
| ZeRO-1 | 切分 | 完整 | 完整 | 42 GB | ~14B |
| ZeRO-2 | 切分 | 切分 | 完整 | 28 GB | ~20B |
| ZeRO-3 | 切分 | 切分 | 切分 | 14 GB | ~100B+ |
DeepSpeed ZeRO 配置
json
{
"train_batch_size": 64,
"gradient_accumulation_steps": 4,
"optimizer": {
"type": "AdamW",
"params": {
"lr": 1e-4,
"betas": [0.9, 0.95],
"weight_decay": 0.1
}
},
"zero_optimization": {
"stage": 3,
"offload_optimizer": {
"device": "cpu",
"pin_memory": true
},
"offload_param": {
"device": "cpu",
"pin_memory": true
},
"overlap_comm": true,
"contiguous_gradients": true,
"reduce_bucket_size": 5e7,
"stage3_prefetch_bucket_size": 5e7,
"stage3_param_persistence_threshold": 1e5,
"sub_group_size": 1e9
},
"fp16": {
"enabled": true
}
}bash
# DeepSpeed 启动命令
deepspeed --num_gpus=8 train.py \
--deepspeed ds_config_zero3.json \
--model_name Qwen/Qwen4-7B \
--per_device_train_batch_size 8 \
--gradient_accumulation_steps 4模型并行 (Model Parallelism)
Tensor Parallelism (TP) — 层内切分
将一个 Transformer 层内的权重矩阵沿列/行切分到多张 GPU。
mermaid
graph TD
subgraph "Tensor Parallelism: 列切分 (Column-wise)"
X["输入 x<br/>(B, L, d_model)"] --> G0["GPU 0: A[:, :d/2]<br/>y₀ = x · A₀"]
X --> G1["GPU 1: A[:, d/2:]<br/>y₁ = x · A₁"]
G0 --> CAT["拼接: y = [y₀ | y₁]"]
G1 --> CAT
end
subgraph "Tensor Parallelism: 行切分 (Row-wise)"
X2["输入 x"] --> SPLIT["切分: x₀, x₁"]
SPLIT --> G0R["GPU 0: y₀ = x₀ · B₀"]
SPLIT --> G1R["GPU 1: y₁ = x₁ · B₁"]
G0R --> REDUCE["AllReduce: y = y₀ + y₁"]
G1R --> REDUCE
end
style CAT fill:#3498db,color:#fff
style REDUCE fill:#e74c3c,color:#fff| 切分方式 | 前向传播 | 反向传播 | 通信量 |
|---|---|---|---|
| 列切分 | 输入复制,输出拼接 | 梯度 AllReduce | fwd: AllGather, bwd: ReduceScatter |
| 行切分 | 输入切分,输出求和 | 梯度切分 | fwd: AllReduce, bwd: AllGather |
python
"""
Tensor Parallelism 的列切分实现
"""
import torch
import torch.nn as nn
import torch.distributed as dist
class ColumnParallelLinear(nn.Module):
"""列切分线性层"""
def __init__(self, in_features, out_features, world_size, rank):
super().__init__()
self.rank = rank
self.world_size = world_size
assert out_features % world_size == 0
local_out = out_features // world_size
self.weight = nn.Parameter(torch.randn(local_out, in_features) * 0.02)
def forward(self, x):
# 每张卡计算: y_local = x @ W_local^T
y_local = x @ self.weight.T # (B, seq, local_out)
# AllGather: 收集所有卡的输出并拼接
y_list = [torch.zeros_like(y_local) for _ in range(self.world_size)]
dist.all_gather(y_list, y_local)
y = torch.cat(y_list, dim=-1) # (B, seq, out_features)
return yPipeline Parallelism (PP) — 层间切分
将模型按层切分,每张 GPU 负责连续的若干层。流水线中,micro-batch 在不同 GPU 间流动。
mermaid
graph LR
subgraph "Pipeline: 4 GPU, 4 micro-batches"
MB1["MB1"] --> L0["GPU0<br/>层1-8"]
L0 --> L1["GPU1<br/>层9-16"]
L1 --> L2["GPU2<br/>层17-24"]
L2 --> L3["GPU3<br/>层25-32"]
L3 --> LOSS["Loss"]
end
NOTE["关键: micro-batch 实现流水线并行<br/>避免 GPU 空闲等待"]3D 并行 — 组合使用
mermaid
graph TD
subgraph "3D 并行架构"
NODE1["节点 1 (8 GPU)"] --> |"DP: 数据副本"| NODE2["节点 2 (8 GPU)"]
subgraph "单节点内"
G0["GPU 0"] --> |"TP: 层内切分"| G1["GPU 1"]
G0 --> |"PP: 层间切分"| G4["GPU 4-7"]
G4 --> G5["GPU 5"]
G4 --> G6["GPU 6"]
G4 --> G7["GPU 7"]
end
end
style NODE1 fill:#3498db,color:#fff
style G0 fill:#e74c3c,color:#fff
style G4 fill:#2ecc71,color:#fff| 并行策略 | 拆分维度 | 适用规模 | 通信模式 |
|---|---|---|---|
| DP | 数据 batch | 2-16 卡 | AllReduce 梯度 |
| TP | 单层权重矩阵 | 2-8 卡 (单机) | AllReduce/AllGather 激活 |
| PP | 模型层 | 2-N 卡 | P2P 发送激活/梯度 |
| ZeRO-3 | 参数/梯度/优化器 | N 卡 | AllGather + ReduceScatter |
生产环境典型配置:TP=2 + PP=4 + DP=16 = 128 GPU。先用 PP 和 TP 满足单模型副本的显存需求,再用 DP 加速。
FSDP — PyTorch 原生 ZeRO-3
FSDP (Fully Sharded Data Parallel) 是 PyTorch 官方实现的 ZeRO-3 等效方案。
python
"""
PyTorch FSDP 完整示例
"""
import torch
import torch.nn as nn
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
MixedPrecision,
ShardingStrategy,
BackwardPrefetch,
)
from torch.distributed.fsdp.wrap import (
transformer_auto_wrap_policy,
size_based_auto_wrap_policy,
)
def create_fsdp_model(model, rank):
"""包装模型为 FSDP"""
# 混合精度配置
mixed_precision = MixedPrecision(
param_dtype=torch.bfloat16, # 参数用 BF16
reduce_dtype=torch.float32, # 梯度归约用 FP32
buffer_dtype=torch.float32,
)
# 自动包装策略: 按 Transformer 层包装
auto_wrap_policy = transformer_auto_wrap_policy
model = FSDP(
model,
auto_wrap_policy=auto_wrap_policy,
mixed_precision=mixed_precision,
sharding_strategy=ShardingStrategy.FULL_SHARD,
backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
device_id=rank,
cpu_offload=False,
)
return model
# 启动: torchrun --nproc_per_node=8 train_fsdp.py
def train_fsdp():
import torch.distributed as dist
dist.init_process_group(backend="nccl")
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
# 创建模型并 FSDP 包装
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen4-7B",
torch_dtype=torch.bfloat16,
)
model = create_fsdp_model(model, local_rank)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
for epoch in range(3):
for batch in dataloader:
# FSDP 自动处理前向的参数 AllGather 和
# 反向的梯度 ReduceScatter
loss = model(**batch).loss
loss.backward()
optimizer.step()
optimizer.zero_grad()
dist.destroy_process_group()通信原语速查
| 原语 | 操作 | 用途 |
|---|---|---|
| AllReduce | 所有卡求和/平均,结果分发到所有卡 | 梯度同步 |
| AllGather | 收集所有卡的数据,拼接后分发 | TP 列切分的输出拼接 |
| ReduceScatter | 先求和再按卡切分分发 | ZeRO-3 梯度归约 |
| Broadcast | 一张卡的数据复制到所有卡 | 模型参数初始化 |
| P2P (Send/Recv) | 点对点传输 | Pipeline Parallelism |
分布式训练 Checklist
| 检查项 | 说明 |
|---|---|
| ✅ NCCL 版本 ≥ 2.18 | nccl-tests 验证带宽 |
| ✅ NVLink/NVSwitch 可用 | nvidia-smi topo -m 查看拓扑 |
| ✅ 梯度累积对齐 | global_batch = per_device × grad_accum × num_gpus |
| ✅ 学习率缩放 | lr 按 sqrt 或 linear 缩放,warmup 比例不变 |
| ✅ 随机种子同步 | 所有 rank 使用相同 seed |
| ✅ Checkpoint 一致性 | 保存/加载时处理分片状态的聚合 |
参考
- ZeRO: Memory Optimizations Toward Training Trillion Parameter Models — Rajbhandari et al., 2019
- DeepSpeed — Microsoft
- PyTorch FSDP — Meta
- Megatron-LM: Training Multi-Billion Parameter Language Models — Shoeybi et al., 2019
- Efficient Large-Scale Language Model Training on GPU Clusters — Narayanan et al., 2021
登录后即可发表评论 👇