训练大模型时,瓶颈往往不是算力而是显存:优化器状态、梯度、激活值三者叠加,让单卡很快 OOM。分布式训练要同时回答两个问题——怎么把吞吐扩上去(DDP),怎么把显存降下来(FSDP/ZeRO)。本文从显存账本出发,系统讲透 DDP 的梯度同步、ZeRO 的三级切分、FSDP 的分片实现、混合并行拓扑与生产调优。
前置:/ai-collective-communication-nccl/(集合通信与 NCCL)、/ai-cuda-memory-optimization/(显存管理与复用)、/ai-llm-inference-architecture/(并行范式基础)。
目录
- 1. 为什么需要分布式训练:单卡的三重瓶颈
- 2. DDP 数据并行:AllReduce 与梯度同步
- 3. 显存账本:参数、梯度、优化器状态与激活
- 4. ZeRO 三级切分:优化器、梯度与参数
- 5. FSDP:全分片数据并行的实现机制
- 6. 混合并行:TP、PP、DP 的组合拓扑
- 7. 通信与计算重叠:把等待藏起来
- 8. 生产配置与调优清单
- 9. 常见故障与踩坑
- 10. 速查表与一句话记忆
- 延伸阅读
1. 为什么需要分布式训练:单卡的三重瓶颈
单卡训练的瓶颈有三层,逐层浮现。
三层瓶颈(按出现的先后顺序):
□ 第一层:显存装不下模型
- 70B 参数 FP16 = 140 GB 权重
- 加上优化器状态(Adam 需 2 份动量)→ 数倍膨胀
□ 第二层:算力不够,训练太慢
- 单卡 FLOPS 有限,token/s 上不去
- 想要月级训练压缩到周级
□ 第三层:数据吞吐不够
- 数据加载、预处理成为瓶颈
- 需要更多 worker 并行读取
对应的解法:
□ 显存 → 切分(ZeRO / FSDP / TP)
□ 算力 → 并行(DP / TP / PP)
□ 数据 → 更多 DataLoader worker + 流式读取
并行的四个维度:
□ DP(Data Parallel):每个卡一份完整模型,切数据
□ TP(Tensor Parallel):把单层矩阵切到多卡
□ PP(Pipeline Parallel):把不同层放到不同卡
□ EP(Expert Parallel):MoE 专家分散(见 MoE 篇)
选择顺序:
□ 先 DP 扩吞吐(最省事)
□ 显存不够再上 ZeRO/FSDP
□ 单层都放不下才上 TP
□ 模型太深放不下才上 PP
工程要点:分布式训练的第一个问题永远是「显存不够」——权重、优化器状态、梯度、激活四类占用叠加。DP 解决吞吐、ZeRO/FSDP 解决显存、TP 解决单层过大、PP 解决模型过深。先 DP,再按需叠加。
2. DDP 数据并行:AllReduce 与梯度同步
DDP 是最基础的并行方式:每卡一份完整模型,各算各的梯度,然后同步。
DDP 的工作循环(每个 step):
□ 每卡读不同的 mini-batch(DistributedSampler 保证不重叠)
□ 各自前向 + 反向,得到本地梯度
□ AllReduce 梯度:所有卡求平均 → 每卡拿到全局梯度
□ 各自用相同梯度更新 → 参数保持一致
关键点:
□ 模型副本完全相同(初始权重广播一致)
□ 梯度 AllReduce 后各卡一致 → 参数永不漂移
□ 无需中心参数服务器(去中心化 ring allreduce)
# PyTorch DDP 最小示例
import torch, torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
dist.init_process_group(backend="nccl")
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
model = MyModel().cuda()
model = DDP(model, device_ids=[local_rank])
for batch in loader: # DistributedSampler 保证分片
loss = model(batch).loss
loss.backward() # 反向累积本地梯度
optimizer.step() # DDP 已在 backward 中触发 AllReduce
optimizer.zero_grad(set_to_none=True)
Ring AllReduce 的通信量:
□ 每卡发送 2×(N-1)/N × 梯度大小
□ N 很大时趋近 2× 梯度大小 → 与卡数无关
□ 所以 DDP 的扩展性极好(带宽受限而非延迟受限)
□ 70B 模型梯度 140 GB → 每 step 传 ~280 GB → 网络是硬瓶颈
工程要点:DDP 的核心是「反向传播时自动触发梯度 AllReduce」,用 ring allreduce 让通信量与卡数解耦。扩展性好但受网络带宽限制——梯度越大、卡越多,AllReduce 越慢,这时需要梯度分桶、通信重叠或 ZeRO 来降本。
3. 显存账本:参数、梯度、优化器状态与激活
理解显存构成是优化显存的第一步。
以 Adam 训练 7B 模型(FP16 权重 + FP32 主权重)为例:
□ 参数(FP16) : 7B × 2 bytes = 14 GB
□ 梯度(FP16) : 7B × 2 bytes = 14 GB
□ 优化器状态(Adam) :
- FP32 主权重 : 7B × 4 = 28 GB
- 一阶动量 m : 7B × 4 = 28 GB
- 二阶动量 v : 7B × 4 = 28 GB
□ 合计(不含激活) : 14 + 14 + 28×3 = 112 GB
□ 激活值(随 batch/序列长度): 数 GB ~ 数十 GB
→ 单卡 80 GB 根本装不下 → 必须切分
关键洞察:
□ 优化器状态是「最大头」(占 75%)
□ 参数和梯度只占 25%
□ 激活是动态的、随 batch 变化
→ 切分优化器状态收益最高 → 这正是 ZeRO 的切入点
显存优化手段优先级:
□ 1. 混合精度(FP16/BF16 前向反向)→ 直接砍一半激活
□ 2. 梯度检查点(activation checkpointing)→ 用算力换激活显存
□ 3. ZeRO/FSDP 切分 → 砍优化器状态与参数
□ 4. 梯度累积 → 用小 batch 模拟大 batch
□ 5. 8-bit 优化器(bitsandbytes)→ 优化器状态再砍 4x
工程要点:Adam 训练中优化器状态占显存 75%,参数与梯度只占 25%。因此显存优化的杠杆顺序是——混合精度、梯度检查点、ZeRO 切分优化器状态、8-bit 优化器。先算清账本,再选手段,避免盲目上并行。
4. ZeRO 三级切分:优化器、梯度与参数
ZeRO(Zero Redundancy Optimizer)把「每卡都存一份」的冗余彻底切掉。
ZeRO 的三个级别(逐级加深切分):
□ ZeRO-1:切分优化器状态
- 每卡只存 1/N 的优化器状态
- 更新时各卡负责自己那部分参数
- 显存节省约 4x(优化器是大头)
□ ZeRO-2:再切分梯度
- 每卡只保留自己负责参数的梯度
- 反向时用 reduce-scatter 而非 all-reduce
- 显存节省约 8x
□ ZeRO-3:再切分参数本身
- 每卡只存 1/N 的参数
- 前向/反向时需要 all-gather 出完整参数
- 显存节省随卡数线性增长(~Nx)
ZeRO-3 的一层前向:
□ 参数分片在各卡上(本卡只有 1/N)
□ 计算到某层时:all-gather 该层完整参数
□ 算完立刻释放非本卡的参数分片
□ 反向同理,需要重新 all-gather
→ 通信量增加,但显存大幅下降
通信量对比(单步):
□ DDP : 1× all-reduce ≈ 2× 梯度
□ ZeRO-1 : 1× reduce-scatter + 1× all-gather ≈ 2× 梯度
□ ZeRO-2 : 同 ZeRO-1(省显存不增通信)
□ ZeRO-3 : 额外 2× 参数通信 → 通信量约为 DDP 的 1.5~2 倍
→ 省显存换通信,网络差时 ZeRO-3 会变慢
工程要点:ZeRO 三级递进——1 切优化器状态、2 加切梯度、3 加切参数。前两级几乎不增通信量却省 8x 显存,是性价比最高的选择;ZeRO-3 显存随卡数线性下降,但通信量翻倍,适合「模型实在放不下」且网络充足的场景。
5. FSDP:全分片数据并行的实现机制
FSDP(Fully Sharded Data Parallel)是 ZeRO-3 思想在 PyTorch 的原生实现。
FSDP 的核心动作(per-layer):
□ 分片:参数、梯度、优化器状态都按卡数分片
□ 前向:进入某层前 all-gather 完整参数 → 计算 → 释放
□ 反向:再次 all-gather 参数 → 算梯度 → reduce-scatter
□ 更新:各卡只更新自己分片内的参数
与 DDP 的差异:
□ DDP:每卡全量模型副本,通信一次 all-reduce
□ FSDP:每卡 1/N 模型,通信 all-gather + reduce-scatter 多次
□ FSDP 显存随卡数下降,代价是更多通信
# PyTorch FSDP 配置示例
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(),
limit_all_gathers=True, # 限制并发 all-gather,省显存
use_orig_params=True, # 保留原始参数视图,便于优化器分组
)
分片策略选择:
□ FULL_SHARD : 参数+梯度+优化器全切(ZeRO-3)
□ SHARD_GRAD_OP : 只切梯度+优化器(ZeRO-2)
□ NO_SHARD : 等价 DDP(不切)
□ HYBRID_SHARD : 节点内全切 + 节点间复制(省跨机通信)
自动包装(auto_wrap_policy):
□ 按 Transformer Block 粒度分片 → 一次 all-gather 一层
□ 粒度太大 → 峰值显存高;太小 → 通信次数多
工程要点:FSDP 用「按层 all-gather 参数、算完即释放」实现显存随卡数线性下降。关键是选对分片策略(FULL_SHARD 省显存、HYBRID_SHARD 省跨机带宽)和自动包装粒度(按 Transformer Block)。use_orig_params 与 limit_all_gathers 是两个容易被忽略但影响很大的开关。
6. 混合并行:TP、PP、DP 的组合拓扑
超大模型单靠一种并行不够,需要三维组合。
三种并行的分工:
□ DP:复制完整模型,切数据 → 扩吞吐,通信在梯度
□ TP:切单层矩阵 → 切显存也切算力,通信频繁(每层 2 次)
□ PP:切不同层 → 切显存,通信少(层间点对点)
组合原则:
□ TP 放节点内(NVLink 高带宽,扛得住每层通信)
□ PP 跨节点(点对点通信少,适合慢网络)
□ DP 最外层(梯度同步可重叠)
→ 典型拓扑:DP × PP × TP
3D 并行示例(Megatron 风格):
□ 8 卡节点 × 4 节点 = 32 卡
□ TP = 8(节点内 NVLink)
□ PP = 2(跨 2 个节点)
□ DP = 2(复制两份)
□ 3×2×... 需满足 TP×PP×DP = 总卡数
→ 8 × 2 × 2 = 32 ✓
通信量排序(同规模下):
□ TP:每层 all-reduce ×2 → 最频繁 → 必须节点内
□ PP:层间 p2p → 中等 → 可跨机
□ DP:每步一次 all-reduce → 最少 → 最外层
□ EP:all-to-all → 视专家分布 → 见 MoE 篇
工程要点:混合并行的口诀是「TP 节点内、PP 跨节点、DP 最外层」——把通信最频繁的 TP 放在带宽最高的 NVLink 内,把通信最少的 DP 放最外层。三维拓扑必须满足 TP×PP×DP 等于总卡数,且要与网络拓扑对齐,否则通信会成为瓶颈。
7. 通信与计算重叠:把等待藏起来
分布式训练的加速核心不是「算得快」而是「等得少」。
通信与计算重叠的三种手段:
□ 梯度分桶 + 异步 AllReduce:
- 反向算完一层就立刻启动该层梯度的 AllReduce
- 通信与下一层的反向计算重叠
□ FSDP 的 prefetch:
- 提前 all-gather 下一层的参数
- 当前层计算时,下一层参数已在路上
□ PP 的 micro-batch 流水:
- 把 batch 切成 micro-batch
- 前向/反向交错,填满流水线气泡
# DDP 梯度分桶(bucket_cap_mb 控制桶大小)
model = DDP(model, bucket_cap_mb=25, gradient_as_bucket_view=True)
# □ 桶太大 → 首个桶通信启动晚,重叠少
# □ 桶太小 → 通信次数多,启动开销占比高
# □ 经验值 25 MB 是常见甜点
PP 流水线气泡:
□ 无 micro-batch:GPU 大量空闲(气泡)
□ 有 micro-batch:前向/反向交错填充
□ 气泡比例 ≈ (PP-1) / (micro_batch + PP - 1)
□ micro_batch 越多,气泡越小,但激活显存越大
□ 1F1B(一次前向一次反向)是经典调度
工程要点:分布式训练的实际加速来自「通信隐藏」——DDP 的梯度分桶、FSDP 的参数 prefetch、PP 的 micro-batch 流水,本质都是让通信与计算重叠。bucket_cap_mb 调优、prefetch 开关、micro-batch 数量是三个最常调的重叠参数。
8. 生产配置与调优清单
把分布式训练落到生产,以下是可复用的配置经验。
配置决策树:
□ 模型能放单卡 → 直接用,不上并行
□ 放不下但单层放得下 → FSDP FULL_SHARD 或 ZeRO-3
□ 显存还紧张 → 加梯度检查点 + 8-bit 优化器
□ 单层放不下 → 加 TP(节点内)
□ 层数太多放不下 → 加 PP
□ 跨机带宽差 → FSDP 用 HYBRID_SHARD
关键调优参数:
□ bucket_cap_mb : DDP 梯度桶大小(默认 25)
□ gradient_accumulation: 小 batch 模拟大 batch
□ activation_checkpoint: 用算力换激活显存
□ cpu_offload : 优化器状态放 CPU(慢但省显存)
□ prefetch / limit_all_gathers: FSDP 通信调度
□ NCCL 环境变量 : NCCL_IB_DISABLE / NCCL_SOCKET_IFNAME
吞吐诊断顺序:
□ 1. 看 GPU 利用率:低 → 数据加载或通信是瓶颈
□ 2. 看通信占比:>30% → 考虑降通信(分片策略/重叠)
□ 3. 看显存峰值:接近上限 → 加检查点或降 batch
□ 4. 看数据加载:DataLoader worker 是否打满
工程要点:生产配置按「模型能不能放单卡 → 单层能不能放下 → 网络好不好」三问决策。DDP 的 bucket_cap_mb、FSDP 的分片策略与 prefetch、梯度累积与激活检查点是最高频的调优旋钮;诊断顺序永远是先看 GPU 利用率再看通信占比。
9. 常见故障与踩坑
分布式训练的坑大多与「多卡一致性」和「通信」有关。
高频踩坑:
□ 忘记 set_device → 所有进程挤在 GPU 0,其余卡空闲
□ DistributedSampler 未设 shuffle=False(验证集)→ 评估结果乱
□ BatchNorm 在多卡下统计不一致 → 换 SyncBatchNorm 或 LayerNorm
□ 学习率未随全局 batch 缩放 → 训练不收敛
□ 梯度未做梯度裁剪 → 多卡下梯度范数暴涨
□ FSDP 中 use_orig_params=False → 优化器参数分组失效
□ NCCL 超时(默认 30 分钟)→ 大模型单步超时需调 NCCL_TIMEOUT
□ 卡间负载不均(PP 气泡)→ 部分卡空转
故障定位工具:
□ torch.distributed 的 debug 模式 → 打印每步通信
□ NCCL_DEBUG=INFO → 看通信拓扑与带宽
□ nvidia-smi dmon → 实时看每卡利用率
□ torch profiler → 定位计算/通信时间占比
□ 检查 all-reduce 是否真的发生(梯度是否被同步)
收敛性问题的排查:
□ 单卡小规模能收敛吗?(先排除数据/模型问题)
□ 梯度范数是否正常?(多卡下易爆炸)
□ 学习率缩放了吗?(linear / sqrt 两种规则)
□ 随机种子是否同步?(影响 dropout、数据顺序)
工程要点:分布式训练的故障集中在一致性(set_device、Sampler、BN、种子)与通信(NCCL 超时、拓扑、带宽)。收敛异常先退回单卡验证,再排查学习率缩放、梯度范数与随机种子同步。NCCL_DEBUG 与 profiler 是定位通信问题的两把利器。
10. 速查表与一句话记忆
| 问题 | 一句话答案 |
|---|---|
| 为什么要分布式 | 单卡显存/算力/数据三重瓶颈 |
| DDP 做什么 | 每卡全量模型、切数据、梯度 AllReduce |
| 显存大头是谁 | 优化器状态占约 75%(Adam) |
| ZeRO-1/2/3 区别 | 切优化器状态 / 加切梯度 / 加切参数 |
| FSDP 是什么 | ZeRO-3 的 PyTorch 原生实现,按层 all-gather |
| 分片策略怎么选 | 省显存 FULL_SHARD,省跨机带宽 HYBRID_SHARD |
| 混合并行口诀 | TP 节点内、PP 跨节点、DP 最外层 |
| 怎么加速 | 通信与计算重叠(分桶、prefetch、流水) |
| 常见坑 | set_device、Sampler、BN、学习率缩放、NCCL 超时 |
| 诊断顺序 | 先看 GPU 利用率,再看通信占比 |
一句话记忆:分布式训练 = DDP 扩吞吐(梯度 AllReduce)+ ZeRO/FSDP 省显存(优化器状态是最大头,三级切分)+ 混合并行(TP 节点内、PP 跨节点、DP 最外层)+ 通信重叠(分桶、prefetch、micro-batch)——显存账本先算清,故障先查一致性与 NCCL。
延伸阅读
- /ai-collective-communication-nccl/ — 集合通信原语与 NCCL 调优
- /ai-cuda-memory-optimization/ — 显存管理与碎片治理
- /ai-llm-inference-architecture/ — TP/PP/EP 并行范式基础
- /ai-distributed-inference-gpu-cluster/ — 分布式推理与集群拓扑
- /ai-performance-tuning-checklist/ — 性能调优通用清单
- 高性能计算专题 — 多卡通信与拓扑优化
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。