引言
当模型参数从百万级涨到十亿、百亿级,单张 GPU 立刻撞上两堵墙:显存装不下(参数 + 梯度 + 优化器状态轻松超过 80GB)和训练太慢(一个 epoch 要跑几周)。分布式训练就是拆掉这两堵墙的工程学——把计算、参数、数据切分到多张卡、多台机器上协同完成一次迭代。
本文按「先会用、再懂原理、最后能排错」的顺序展开:先用 DDP 把数据并行跑通,再理解 all-reduce 这类通信原语为什么是性能瓶颈,然后进入张量并行、流水线并行、ZeRO/FSDP 这些大模型必备的切分术,最后给出多机训练的排错清单。全文代码基于 PyTorch 官方 DDP + DeepSpeed。
前置:单卡训练循环与优化器基础见 https://plumephp.com/ml-deep-learning-advanced/;显存与量化压缩的配合见 https://plumephp.com/ml-model-compression-quantization/;优化器状态占用的原理见 https://plumephp.com/ml-gradient-descent-optimizers/;训练完成后上线见 https://plumephp.com/ml-model-deployment/。
目录
- 1. 为什么需要分布式训练
- 2. 数据并行与 DDP
- 3. 通信原语:all-reduce 与集合通信
- 4. 模型并行与张量并行
- 5. 流水线并行与气泡
- 6. ZeRO 与 FSDP:显存账本重构
- 7. DeepSpeed 实战配置
- 8. 梯度累积、混合精度与排错手册
- 9. 总结
- 延伸阅读
1. 为什么需要分布式训练
1.1 两道墙:显存与时间
训练一个大模型时,显存被四样东西瓜分:模型参数、梯度、优化器状态、激活值。以 Adam 为例,每个参数需要:2 字节权重(fp16)+ 2 字节梯度 + 4 字节一阶动量 + 4 字节二阶动量 + 4 字节 fp32 主权重 = 16 字节/参数。一个 7B 模型仅状态就要 112GB,远超单张 A100 的 80GB。
| 资源 | 单卡瓶颈 | 分布式解法 |
|---|---|---|
| 显存(参数/梯度/优化器) | 模型装不下 | ZeRO / FSDP / 模型并行 |
| 显存(激活值) | batch 开不大 | 梯度检查点 / 序列并行 |
| 计算吞吐 | 迭代太慢 | 数据并行 |
| 单机卡数上限 | 8 卡封顶 | 多机 + 集合通信 |
1.2 四种并行的坐标系
数据并行(Data Parallel) :每卡一份完整模型,切分数据 → 最常用
张量并行(Tensor Parallel) :把单层矩阵切到多卡 → 层内切分
流水线并行(Pipeline Parallel):把不同层放到不同卡 → 层间切分
专家并行(Expert Parallel) :MoE 的不同专家分到不同卡 → 稀疏激活
前三种可组合成 3D 并行,是训练千亿模型的标配:数据并行 × 张量并行 × 流水线并行。
一句话:分布式训练的本质是「用通信换显存和时间」——切得越细,单卡负担越轻,但卡之间的通信开销越大,找到平衡点是全部工程的核心。
1.3 什么时候不需要分布式
不是所有任务都值得上多卡:模型能塞进单卡时单卡调好超参收益更高;卡间带宽低(无 NVLink、跨机房)时数据并行同步可能吃掉全部加速比;batch 太小时每卡样本不足,梯度噪声大、收敛变差。分布式是「不得已而为之」的手段,不是性能银弹。
2. 数据并行与 DDP
2.1 DP 与 DDP 的区别
PyTorch 有两代数据并行实现,理解差异很重要:
| 维度 | DataParallel (DP) | DistributedParallel (DDP) |
|---|---|---|
| 进程模型 | 单进程多线程 | 多进程(每卡一进程) |
| 通信方式 | 主卡 gather/scatter | all-reduce 梯度 |
| GIL 影响 | 严重 | 无 |
| 多机支持 | 不支持 | 原生支持 |
| 推荐度 | 已废弃 | 生产标准 |
DP 的所有前向都在主卡汇总,主卡成瓶颈;DDP 每个进程独立跑完整前向反向,只在反向结束时同步梯度,效率高得多,因此是生产标准。
2.2 DDP 的核心机制:梯度 all-reduce
DDP 的关键设计是反向传播时就通信,而不是等全部梯度算完再同步。它把梯度按 bucket 分桶,某个 bucket 的梯度算完立刻触发 all-reduce,与后续层的反向计算重叠(overlap),把通信藏进计算里。
import os
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data.distributed import DistributedSampler
def setup():
dist.init_process_group(backend="nccl")
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
return local_rank
def main():
local_rank = setup()
model = MyModel().to(local_rank)
model = DDP(model, device_ids=[local_rank])
dataset = MyDataset()
sampler = DistributedSampler(dataset, shuffle=True)
loader = torch.utils.data.DataLoader(
dataset, batch_size=32, sampler=sampler, num_workers=4
)
for epoch in range(num_epochs):
sampler.set_epoch(epoch) # 关键:每个 epoch 打乱不同
for batch in loader:
loss = model(**batch)
loss.backward()
optimizer.step()
optimizer.zero_grad(set_to_none=True)
if __name__ == "__main__":
main()
启动方式用 torchrun:
torchrun --nproc_per_node=8 --nnodes=2 --node_rank=0 \
--master_addr=10.0.0.1 --master_port=29500 train.py
2.3 学习率与 batch size 的线性缩放
DDP 让全局 batch = 单卡 batch × 卡数。全局 batch 变大后梯度噪声降低,需要相应放大学习率:线性缩放 lr_new = lr_base × (global_batch / base_batch),或用平方根缩放。实践建议先做 warmup(前 5% 步数线性升温),全局 batch 超过 8K 后通常要换 LAMB 或带自适应裁剪的优化器。
一句话:DDP = 每卡完整模型 + 反向重叠的梯度 all-reduce;它解决的是「训练慢」,不解决「模型装不下」——显存问题要交给 ZeRO/FSDP。
3. 通信原语:all-reduce 与集合通信
3.1 集合通信五件套
分布式训练的通信全部由 NCCL 提供的集合操作构成:
| 原语 | 语义 | 典型用途 |
|---|---|---|
| broadcast | 一对多 | 初始化参数 |
| all-reduce | 多对多归约(求和/平均) | 梯度同步 |
| reduce-scatter | 归约后分片 | ZeRO 梯度分片 |
| all-gather | 收集各卡分片 | ZeRO 参数重建 |
| all-to-all | 全交换 | MoE / 序列并行 |
DDP 用 all-reduce;ZeRO 用 reduce-scatter + all-gather 组合替代 all-reduce,这是理解 ZeRO 通信量的关键。
3.2 Ring All-Reduce 的直觉
all-reduce 的高效实现是 Ring 算法:N 张卡组成环,分两阶段——reduce-scatter(每卡负责一段的部分和,沿环传递累加)+ all-gather(把各卡拥有的完整段沿环广播)。通信量是 2 × (N-1)/N × 数据量,与卡数近似无关,这是它能扩展的核心原因。判断通信是否成为瓶颈看计算通信比:卡内有 NVLink(A100 600GB/s)时数据并行扩展性良好;仅 PCIe(32GB/s)时通信可能吃掉 30%+ 时间;跨机器走以太网/RoCE 则必须靠 bucket 重叠和张量并行缓解。
# 测量实际可用带宽(NCCL 自带测试)
./build/all_reduce_perf -b 8 -e 4G -f 2 -g 8
一句话:all-reduce 是数据并行的通信核心,Ring 算法让它与卡数弱相关;但跨机带宽远低于 NVLink,所以大模型要在张量并行(卡内高带宽)和流水线并行(低通信)之间做取舍。
4. 模型并行与张量并行
4.1 模型并行:按层切
最朴素的模型并行是把不同层放到不同卡,问题是同一时刻只有一张卡在算,其余卡空转,利用率极低。它只在「模型单卡绝对装不下」时使用,实践中已被流水线并行取代。
4.2 张量并行:按矩阵切
张量并行(TP,Megatron-LM 提出)把单个矩阵乘法切到多卡。以 Transformer 的 MLP 为例,第一个线性层按列切(Column Parallel),第二个按行切(Row Parallel):
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):
super().__init__()
self.out_per_rank = out_features // world_size
self.weight = nn.Parameter(
torch.empty(self.out_per_rank, in_features)
)
def forward(self, x):
# x 在每卡相同,输出是完整输出的一段
return torch.nn.functional.linear(x, self.weight)
class RowParallelLinear(nn.Module):
"""按输入维度切分:每卡持有一部分输入,输出 all-reduce 求和"""
def __init__(self, in_features, out_features, world_size):
super().__init__()
self.in_per_rank = in_features // world_size
self.weight = nn.Parameter(
torch.empty(out_features, self.in_per_rank)
)
def forward(self, x):
partial = torch.nn.functional.linear(x, self.weight)
dist.all_reduce(partial) # 各卡部分和相加
return partial
列切 + 行切成对出现,中间无需通信(每卡算自己的分片),只在 Row Parallel 后做一次 all-reduce,这是 TP 通信量小的原因。代价是每层都通信,所以 TP 通常限制在同一节点内(走 NVLink)。Multi-Head Attention 天然适合 TP:按注意力头切分,8 个头分到 4 张卡每卡算 2 个,QKV 投影按列切、输出投影按行切,结构与 MLP 一致——这也是 TP 并行度常设为注意力头数约数的原因。
一句话:张量并行是「层内切矩阵」,通信频繁但量小,必须跑在 NVLink 上;它的价值是让单层就能跨卡,从而突破单卡装不下大层的限制。
5. 流水线并行与气泡
5.1 按阶段切分与微批
流水线并行(PP)把模型按层切成若干 stage,每卡一个 stage,数据切成多个**微批(micro-batch)**依次流过形成流水线:
时间 →
卡0: [F1][F2][F3][F4] [B4][B3][B2][B1]
卡1: [F1][F2][F3][F4] [B4][B3][B2][B1]
卡2: [F1][F2][F3][F4] ...
5.2 气泡率与调度策略
流水线的致命伤是 气泡(bubble):填充和排空阶段有卡在空转。气泡占比近似 (P-1) / (M + P - 1)(P 为 stage 数、M 为微批数),微批越多气泡越小。
| 调度 | 特点 | 适用 |
|---|---|---|
| GPipe | 全前向再全反向,显存占用高 | 简单场景 |
| 1F1B | 前向反向交替,显存省 | 主流默认 |
| Interleaved 1F1B | 每卡持多段,气泡更小 | 追求效率 |
# PyTorch 原生流水线(torch.distributed.pipeline.sync)
from torch.distributed.pipeline.sync import Pipe
model = nn.Sequential(layer1, layer2, layer3, layer4)
model = Pipe(model, chunks=8, checkpoint="except_last")
output = model(input) # 自动切微批、调度 1F1B
5.3 3D 并行的组合逻辑
真实的大模型训练是三种并行的乘积 总卡数 = TP × PP × DP,例如 64 卡 = 8(TP,节点内 NVLink)× 2(PP,跨节点)× 4(DP)。分配原则:TP 放在节点内(通信最密),PP 跨节点(通信最稀),DP 用剩余卡。
一句话:流水线并行用「微批 + 1F1B」把层间切分做成了流水线,气泡是它的固有成本;3D 并行的排布口诀是 TP 在内、PP 在外、DP 补足。
6. ZeRO 与 FSDP:显存账本重构
6.1 显存账本的重新分配
ZeRO(Zero Redundancy Optimizer)的洞察是:数据并行里每张卡都存了一份完整的参数、梯度、优化器状态,这是巨大的冗余。ZeRO 把它们按卡分片,用时再 all-gather 回来。
| 阶段 | 切分内容 | 显存节省(N 卡) | 通信量 |
|---|---|---|---|
| ZeRO-1 | 优化器状态 | ~4x | 与 DDP 相同 |
| ZeRO-2 | + 梯度 | ~8x | 略增 |
| ZeRO-3 | + 参数 | ~N倍 | 显著增加 |
6.2 FSDP:PyTorch 原生实现
FSDP(Fully Sharded Data Parallel)是 PyTorch 对 ZeRO-3 的原生实现,核心是参数分片 + 用时 all-gather + 用完即弃:
import torch
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
mp_policy = MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.bfloat16,
buffer_dtype=torch.bfloat16,
)
model = FSDP(
model,
sharding_strategy=ShardingStrategy.FULL_SHARD, # 对应 ZeRO-3
mixed_precision=mp_policy,
device_id=torch.cuda.current_device(),
use_orig_params=True,
)
ShardingStrategy 的三个选项对应不同 ZeRO 级别:FULL_SHARD(ZeRO-3)、SHARD_GRAD_OP(ZeRO-2)、NO_SHARD(DDP)。
6.3 分片带来的通信代价
ZeRO-3 每层前向要 all-gather 参数、反向要 all-gather + reduce-scatter,通信量约为 DDP 的 1.5 倍,带宽不足时 ZeRO-3 反而比 DDP 慢。经验法则:单节点内、模型装得下用 ZeRO-1/2 或纯 DDP;跨节点、模型装不下才用 ZeRO-3 + 大 bucket + 通信重叠。
一句话:ZeRO/FSDP 用「分片换显存」,把数据并行的冗余存储彻底消除;但分片越彻底、通信越多,ZeRO-3 的收益要建立在足够带宽之上。
7. DeepSpeed 实战配置
7.1 一个可用的 ZeRO-3 配置
DeepSpeed 把分布式策略抽象成一份 JSON,改配置就能切换并行方案:
{
"train_batch_size": 256,
"gradient_accumulation_steps": 4,
"fp16": {
"enabled": true,
"loss_scale": 0,
"initial_scale_power": 16
},
"zero_optimization": {
"stage": 3,
"offload_optimizer": { "device": "cpu", "pin_memory": true },
"offload_param": { "device": "cpu", "pin_memory": true },
"overlap_comm": true,
"contiguous_gradients": true,
"stage3_gather_16bit_weights_on_model_save": true,
"sub_group_size": 1e9,
"reduce_bucket_size": "auto"
},
"gradient_clipping": 1.0,
"steps_per_print": 100
}
offload_optimizer 和 offload_param 把状态/参数卸到 CPU 内存,用带宽换显存,是「单卡跑大模型」的常用手段。
7.2 启动脚本与配置切换
deepspeed --num_gpus=8 --num_nodes=2 train.py --deepspeed ds_config_zero3.json \
--per_device_train_batch_size 8 --gradient_accumulation_steps 4 --bf16
在 HuggingFace Trainer 中,只需把 JSON 路径传进 deepspeed 参数,其余不用改代码:
from transformers import TrainingArguments
args = TrainingArguments(
output_dir="./out",
per_device_train_batch_size=8,
gradient_accumulation_steps=4,
bf16=True,
deepspeed="ds_config_zero3.json",
gradient_checkpointing=True,
)
7.3 三种并行在 DeepSpeed 中的组合
DeepSpeed 通过 zero_optimization.stage 控制数据并行侧,通过外部参数控制 TP/PP:
deepspeed --num_gpus=64 train.py --deepspeed ds_config.json \
--tensor_model_parallel_size 8 --pipeline_model_parallel_size 2 --zero_stage 1
注意:用了 TP/PP 后 ZeRO 一般只开到 stage 1,因为参数已经切分,再叠 stage 3 会重复分片、通信爆炸。
一句话:DeepSpeed 的价值是把「并行策略」变成配置文件,从 ZeRO-1 到 3D 并行靠改 JSON 就能切换;但配置之间不是叠加越多越好,TP/PP 与 ZeRO-3 通常互斥。
8. 梯度累积、混合精度与排错手册
8.1 梯度累积:小显存换大 batch
当显存开不大 batch 时,用梯度累积模拟大 batch:
accum_steps = 4
for i, batch in enumerate(loader):
loss = model(**batch) / accum_steps
loss.backward()
if (i + 1) % accum_steps == 0:
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
optimizer.zero_grad(set_to_none=True)
三个坑:除以 accum_steps 做平均(否则梯度放大)、只在真正 step 时裁剪梯度、与 DDP 的 no_sync 配合避免每个微批都同步(微批内不同步,最后一个微批才同步):
with model.no_sync(): # 微批内不同步梯度
loss.backward()
loss.backward() # 最后一个微批才同步
8.2 混合精度:fp16 vs bf16
| 精度 | 动态范围 | 是否需 loss scaling | 硬件要求 |
|---|---|---|---|
| fp16 | 窄,易溢出 | 需要 | V100+ |
| bf16 | 宽,与 fp32 同 | 不需要 | A100+ |
优先 bf16:省掉 loss scaling 的调试烦恼,数值更稳;fp16 则必须配 GradScaler。
8.3 多机排错清单
| 现象 | 根因 | 处理 |
|---|---|---|
| 卡在 init_process_group | 网络/端口不通 | 检查 master_addr、防火墙、NCCL_IB |
| NCCL timeout | 某卡慢或掉队 | 调大 NCCL_TIMEOUT,定位慢卡 |
| 显存 OOM 但单卡够 | ZeRO 分片不均 | 检查 batch 能否整除、调 sub_group_size |
| loss 不降 / 变 NaN | 学习率未缩放、loss scale 溢出 | 加 warmup、切 bf16、降 lr |
| 多机比单机还慢 | 通信未重叠 | 开 overlap_comm、调大 bucket |
| 结果不可复现 | 随机种子未按 rank 设置 | seed = base_seed + rank |
常用环境变量:
export NCCL_DEBUG=INFO # 打印通信细节
export NCCL_IB_DISABLE=1 # 无 InfiniBand 时禁用
export NCCL_SOCKET_IFNAME=eth0 # 指定网卡
8.4 一个诊断流程
1. 单卡能否跑通?→ 不能则先修模型代码
2. 单机 8 卡能否跑通?→ 不能则查 NCCL/拓扑
3. 多机能否跑通?→ 查网络与时钟同步
4. 加速比是否线性?→ 不线性则 profile 通信占比;收敛是否与单卡一致?对比前 100 步 loss
一句话:梯度累积、混合精度是「单卡榨显存」的工具,多机排错的顺序永远是「单卡 → 单机 → 多机」,逐层排除。
9. 总结
9.1 策略选择路线
模型单卡装得下 + 想更快 → DDP 数据并行
模型单卡装不下 + 有 NVLink → 张量并行(节点内)
模型很多层 + 想省通信 → 流水线并行(跨节点)
模型装不下 + 无高带宽 → ZeRO-1/2 + CPU offload
千亿参数 + 大规模集群 → 3D 并行(TP×PP×DP)+ ZeRO-1
9.2 关键决策点
| 问题 | 选择 |
|---|---|
| 只是训练慢,模型装得下 | DDP,别引入复杂度 |
| 单层就装不下 | 张量并行,限制在节点内 |
| 层数多、跨机通信贵 | 流水线并行 + 多微批 |
| 显存紧张、带宽尚可 | ZeRO-2 |
| 显存极紧、带宽充足 | ZeRO-3 + offload |
| MoE 稀疏模型 | 专家并行 + all-to-all |
9.3 一句话心法
分布式训练没有银弹,只有「显存—通信—复杂度」的三角权衡——先问清瓶颈是显存还是时间,再选切分方式;切分越细、通信越贵,能单卡解决就别上多卡。
延伸阅读
- https://plumephp.com/ml-deep-learning-advanced/ — 单卡训练循环、优化器与正则化基础
- https://plumephp.com/ml-gradient-descent-optimizers/ — 优化器状态与显存占用的原理
- https://plumephp.com/ml-model-compression-quantization/ — 量化与压缩如何减轻显存压力
- https://plumephp.com/ml-model-deployment/ — 训练完成后的推理服务与版本管理
- AI/ML 专题 — 大规模训练与系统工程深度文章
- PyTorch 分布式文档
- DeepSpeed 官方教程
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。