数据并行:DP、DDP 与 FSDP
数据并行瞄准的是训练太慢 —— 数据量太大,单卡一个 epoch 要跑很久。它的做法是「多找几位老师同时批卷,但每人手里都得有一份完整的评分标准」。
关键约束:每卡都要装得下完整模型。 所以它只解决「跑不完」,不解决「装不下」—— 这条边界,以及后来 FSDP 如何跨过它,就是本篇的主线。
数学基础:它是严格等价的
数据并行不是近似技巧,在同步 SGD 下与单卡训练严格等价。 理解这点才能明白为什么梯度必须「取平均」。
全局 batch 大小
把
对
精确等于单卡在全局 batch 上的梯度。 所以「 卡数据并行」与「单卡跑一个 大小的 batch」,每一步的参数更新在数学上一模一样。
两个直接推论:
| 推论 | 说明 |
|---|---|
| 梯度要 AllReduce 求平均,不是求和 | 本地损失已在 local batch 内除过 |
| 有效 batch 变大了 | Effective Batch = 单卡 local batch × |
等价性有三个前提
同步梯度、各卡 local batch 大小相同、同一份初始参数。任一条被破坏(最后一个 batch 不满、异步更新、参数不一致),等价性就不成立 —— 而这正是下面 DP 与 DDP 的差别所在。
三代演进
第一代:DP(DataParallel)—— 主卡中心化
单进程多线程,一个进程管所有 GPU,反复做「分发—收集」:
GPU0(主卡)持有模型与输入
→ Scatter 输入到各卡
→ 各卡复制模型副本 + 前向
→ Gather 输出回主卡
→ 主卡算 loss + 反向
→ 各卡梯度汇总回主卡求和
→ 主卡更新参数,下次迭代再复制三个致命缺陷:
| 缺陷 | 原因 |
|---|---|
| GIL 限制 | 单进程多线程,Python 全局解释器锁让多线程无法真正并行调度 |
| 负载不均 | 主卡额外承担 Scatter/Gather、loss 计算、梯度汇总 |
| 通信低效 | 每步都重新复制模型;数据走主卡中转,容易挤在 PCIe 上而非 NVLink |
典型症状是「主卡 OOM,其他卡显存还很空」 —— 因为输出汇总与 loss 计算都堆在主卡。
DP 的病根是「单进程 + 主卡中心化」。DDP 的所有改进,本质都是在拆掉这两个前提。
第一代 DP:单进程多线程 + 主卡中心化
主卡 GPU0 其他卡
┌────────────────────┐
│ 持有模型与输入 │
│ Scatter 输入 ─────┼──────▶ 各卡复制模型副本 + 前向
│ │◀────── Gather 输出
│ 算 loss + 反向 │
│ 汇总梯度并求和 ◀──┼─────── 各卡梯度回传
│ 更新参数 │
│ 下次迭代再复制 ───┼──────▶ 模型副本重新分发
└────────────────────┘
三个后果:
· GIL 让单进程里的多线程调度不起来
· 主卡额外承担 Scatter/Gather、loss 计算、梯度汇总
· 数据都走主卡中转 ⇒ 容易挤在 PCIe 上而不是 NVLink
⇒ 典型症状是「主卡 OOM,其他卡显存还很空」第二代:DDP —— 多进程、去中心化
每块 GPU 一个独立进程,各持有完整模型副本,彼此地位对等,没有主卡,梯度同步走去中心化的 AllReduce。
Bucket 机制:把通信藏进计算
最朴素的做法是「等整个反向算完,再统一发起一次 AllReduce」—— 但通信时 GPU 在干等。
DDP 的洞察:反向是从最后一层往前逐层算的,靠后的层梯度先就绪。既然如此,为什么要等全部算完?
于是 DDP 把梯度按反向计算顺序打包成若干 Bucket(默认约 25 MB):
- 某个 Bucket 内所有梯度就绪 → 立即发起 AllReduce
- 与此同时,更靠前的层继续算反向
通信与计算就这样重叠起来。
model = DDP(model, device_ids=[local_rank], bucket_cap_mb=25)bucket_cap_mb 是双向权衡:太小 → 通信次数多、每次的固定开销占比高;太大 → 要等更久才凑满一桶,重叠效果差。默认值对多数模型够用。
未参与前向的参数会让 AllReduce 死锁
如果模型里有没被走到的参数(条件分支),它们的梯度永远不就绪,对应 Bucket 永远等不满,AllReduce 卡住 → deadlock。
解法是 find_unused_parameters=True,让 DDP 主动标记未使用参数。但它有额外开销,能避免则避免。
超大规模会失效
DDP 在中等规模下几乎线性加速,但红利在超大规模失效:
- 卡数增长 → 协调开销与网络需求显著增长,通信逐渐盖不住计算
- 512+ GPU 时通信开始受限于「环延迟(ring latency)」 —— 信号绕 Ring 传播一圈的时间,DP 通信无法再被完全重叠
到这一点就该转向其他并行维度(TP / PP / ZeRO),而不是继续堆 DP。
DDP 的根本局限
每卡仍要装下完整的「参数 + 梯度 + 优化器状态」
DDP 加卡只能摊薄计算时间,不能摊薄单卡显存。 7B 模型的
GB,单张 80 GB 的 H100 直接装不下 —— 无论加多少卡,DDP 都救不了这个数字。
第三代:FSDP —— 把那 也切开
参数分片,按需 AllGather。
前向:AllGather 拼出当前 unit 的完整参数 → 算完 → 释放
反向:AllGather 再次拼参数 → 算梯度 → ReduceScatter 同步并分片 → 释放四种分片策略:一个「显存 vs 通信」的旋钮
| 策略 | 分片内容 | 显存 | 通信 | 场景 |
|---|---|---|---|---|
FULL_SHARD | 参数 + 梯度 + 优化器(≈ ZeRO-3) | 最高 | 最高 | 大模型,显存紧张 |
SHARD_GRAD_OP | 梯度 + 优化器,参数不分片(≈ ZeRO-2) | 中 | 中 | 中等模型,想省通信 |
HYBRID_SHARD | 机内 FULL_SHARD + 机间数据并行 | 高 | 机间较低 | 多机大模型首选 |
NO_SHARD | 不分片(等价 DDP) | 无 | 最低 | 调试对照 |
HYBRID_SHARD的洞察很值得记:机内有高带宽 NVLink,适合通信密集的 FULL_SHARD;机间只有较慢的 IB,就退化成通信量小的数据并行。把「高频通信」关在机内,「低频通信」才跨机 —— 这与 13-3D 并行与混合并行策略 的通信域划分原则是同一条判据。
FSDP2
PyTorch 正在推进新一代 API fully_shard。与 FSDP1 的「整个模块打包分片」不同,FSDP2 是 per-parameter 分片,底层基于 DTensor:
- 分片粒度到单个参数,避免整块打包带来的显存/通信浪费
- 基于 DTensor,与 TP / SP 组合时接口统一 —— 这是搭建 2D/3D 并行的基础
- 更清晰的初始化与 checkpoint 语义
显存账本:三代对比
(后三项合称优化器状态,共
| 方案 | 参数 | 梯度 | 优化器状态 | 单卡合计 | |
|---|---|---|---|---|---|
| DDP | 约 112 GB | ||||
| SHARD_GRAD_OP(ZeRO-2) | 约 26 GB | ||||
| FULL_SHARD(ZeRO-3) | 约 14 GB |
三条结论:
- DDP 的显存与卡数无关 —— 加卡不减负
- 优化器状态是最大头(
,占 75%)→ ZeRO-1/2 只切它就能省掉大半 - FULL_SHARD 随卡数线性下降 —— 理论上能训任意大的模型(只要卡够多)
这张表不算激活值
激活值随 batch / 序列长度 / 深度增长,长序列训练时往往才是显存大头。FSDP 分的是模型状态,对激活值无能为力 —— 那要靠激活重计算、SP / CP 等手段(见 10-混合精度与显存优化、11-长序列训练与上下文并行)。
选型时先估模型状态,再给激活与碎片留 1.2–1.5 倍余量。
通信量:DDP vs FSDP
FSDP 用更多通信换更少显存。
| 每步单卡通信量 | 拆解 | |
|---|---|---|
| DDP | 一次 AllReduce | |
| FSDP | 前向 AllGather |
为什么 AllReduce 是
因为 AllReduce = ReduceScatter + AllGather(见 03-集合通信与 NCCL)—— 它本身就是两个
| 维度 | DDP | FSDP(FULL_SHARD) |
|---|---|---|
| 核心思路 | 每卡完整模型,梯度 AllReduce | 参数分片,按需 AllGather |
| 每卡通信量/步 | ||
| 单卡显存 | ||
| 重叠手段 | Bucket 机制 | prefetch 预取下一层参数 |
DDP 与 FSDP 是一组清晰的权衡对偶:DDP 省通信费显存,FSDP 省显存费通信。没有免费午餐 —— 选哪个取决于你的瓶颈是显存还是通信。
而且
的额外通信能否被掩盖,很依赖网络带宽与 prefetch 效果:高带宽机内(NVLink / NVSwitch)基本能重叠掉;低带宽跨机场景则会暴露成瓶颈 —— 这正是 HYBRID_SHARD存在的意义。
三代方案落在同一张「显存 vs 通信」的坐标上:
单卡显存
16Ψ ┤ DDP(每步 2Ψ 通信) ← 加卡不减显存
│ ·
8Ψ ┤ ·
4Ψ ┤ · SHARD_GRAD_OP(≈ ZeRO-2,2Ψ + 14Ψ/N)
2Ψ ┤ · FULL_SHARD(≈ ZeRO-3,每步 3Ψ 通信)
└──────┬────────┬──────────▶ 每卡每步通信量
2Ψ 3Ψ
⇒ 越往右下走,越是用通信换显存 —— 没有免费午餐
⇒ HYBRID_SHARD 的取法:机内用 FULL_SHARD(带宽高,撑得起 3Ψ)、
机间退化成 DP(带宽低,只跑 2Ψ)—— 把高频通信关在机内选型
问题只有一个:参数 + 梯度 + 优化器状态 + 激活值,单卡装得下吗?
装得下 → DDP(最简单、通信最省)
装不下 → 先试 SHARD_GRAD_OP(通信更省)
还不够 → FULL_SHARD
多机跨节点 → HYBRID_SHARD判据是「装不下的到底是哪一项」(完整分支见 10-混合精度与显存优化 的全景表):
| 装不下的是 | 该上 |
|---|---|
| 优化器状态 | ZeRO-1 |
| 梯度 | ZeRO-2 / SHARD_GRAD_OP |
| 参数 | ZeRO-3 / FULL_SHARD / TP |
| 激活值 | 激活重计算 / SP / CP |
相关
- 02-分布式训练总论与显存账本 ——
账本与五大策略全景 - 07-ZeRO 显存优化系列 —— FSDP 四策略与 ZeRO 阶段的对应
- 03-集合通信与 NCCL —— AllReduce / AllGather / ReduceScatter 的通信量
- 03-多卡互联与集群网络 ——
HYBRID_SHARD的带宽依据 01-PyTorch 框架与训练循环——no_sync()与梯度的累加语义- 13-3D 并行与混合并行策略 —— DP 在混合并行里的位置
参考
- https://caomaolufei.github.io/AIInfraGuide/guides/模块三-分布式训练/41-数据并行详解
- https://pytorch.org/docs/stable/notes/ddp.html
- https://pytorch.org/docs/stable/fsdp.html
- https://arxiv.org/abs/1910.02054
- https://pytorch.org/docs/stable/distributed.fsdp.fully_shard.html
YJ