📑 本页目录(点开跳转)
10 · 数据并行与 AllReduce
⏱ 26 分钟 | ⭐ 最简单也最常用的并行
🎯 一句话
每张卡放一份完整的模型,各自吃不同的数据,算完梯度后求个平均。 它是唯一一种不改变模型结构的并行方式 —— 所以永远先试它。
🔄 一、它怎么工作
🔑 为什么求平均而不是求和: 因为等效 batch 变成了 N 倍,要保持梯度的尺度不变。 求和的话相当于学习率被放大了 N 倍。
# PyTorch DDP:标准写法
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
dist.init_process_group("nccl")
model = DDP(model.to(local_rank), device_ids=[local_rank])
sampler = torch.utils.data.distributed.DistributedSampler(dataset)
loader = DataLoader(dataset, sampler=sampler, batch_size=micro_bs)
for epoch in range(E):
sampler.set_epoch(epoch) # ⭐ 必须!否则每轮的数据划分完全相同
for batch in loader:
loss = model(batch)
loss.backward() # ⭐ AllReduce 在这里自动发生
opt.step(); opt.zero_grad(set_to_none=True)
⚠️
sampler.set_epoch(epoch)是最常被漏掉的一行: 不加的话每个 epoch 的数据顺序和划分完全一样,等于没有 shuffle。
⚙️ 二、DDP 的两个关键优化
① 梯度桶(bucketing)+ 通信计算重叠
结果对照
对照
时间轴:
计算 ████████████████████████
通信 ▓▓▓▓ ▓▓▓▓ ▓▓▓▓ ▓▓▓▓ ← 藏在计算后面
↑ 桶1 ↑桶2 ↑桶3 ↑桶4
💡 这就是为什么 DDP 比手写 AllReduce 快很多。 桶大小可调:
DDP(model, bucket_cap_mb=25)—— 太小则通信次数多、开销大;太大则重叠机会少。
② 梯度累积时要关掉同步
# 🧩 骨架:`loader` 来自你自己的代码,这一段只看写法
# ❌ 每个 micro batch 都 AllReduce → 通信量翻 N 倍
for i, batch in enumerate(loader):
loss = model(batch) / accum
loss.backward()
# ✅ 只在最后一步同步
for i, batch in enumerate(loader):
is_last = (i + 1) % accum == 0
ctx = nullcontext() if is_last else model.no_sync() # ⭐
with ctx:
(model(batch) / accum).backward()
if is_last:
opt.step(); opt.zero_grad(set_to_none=True)
⚠️ 忘了
no_sync()是一个隐蔽的性能杀手: 累积 8 步就意味着通信量是必要值的 8 倍,而代码看起来完全正常。
🔁 三、Ring AllReduce(第 5 章的展开)
结果对照
对照
⭐ 每张卡收发的数据量 = 2 × (N-1)/N × 数据量 ≈ 2 × 数据量
【和卡数几乎无关】
对比朴素的"都发给 0 号卡":0 号要收发 N-1 倍 💀
📐 为什么是 2×(N−1)/N(想看再点)
梯度总量 $D$,切成 $N$ 份,每份 $D/N$。
- ReduceScatter:转 $N-1$ 步,每步发送 $D/N$ → 共发送 $(N-1)D/N$
- AllGather:同样 $N-1$ 步,每步 $D/N$ → 共 $(N-1)D/N$
合计 $2(N-1)D/N \approx 2D$($N$ 大时)。
💡 关键点:随着 N 增大,每卡的通信量趋于常数 2D,不随 N 增长。 但步数是 $2(N-1)$,所以延迟随 N 线性增长 —— 这是 Ring 的软肋。
⭐ 分层 AllReduce:大规模的标准做法
结果对照
💡 NCCL 会自动根据拓扑选择算法 —— 这也是为什么 第 5 章说要先确认拓扑被正确识别(
NCCL_DEBUG=INFO)。
📉 四、扩展效率:为什么卡越多效率越低
信息关系
算一算
一个具体的例子(7B 模型,BF16 梯度 14GB):
通信量 = 2 × 14 GB = 28 GB/步
机内 NVLink 900GB/s → 0.03 秒
跨机 IB 25GB/s → 1.1 秒 ⭐ 差 35 倍
如果单步计算是 0.5 秒:
· 机内:通信可完全重叠 → 扩展效率 ~95%
· 跨机:通信 1.1 秒 > 计算 0.5 秒 → 【通信成为瓶颈】💀
🔑 这就是为什么大规模训练必须做三件事: ① 梯度压缩/量化(FP16 甚至 FP8 通信) ② 分层 AllReduce ③ 换策略 —— 数据并行到了瓶颈就该上 ZeRO 或模型并行
⚠️ 五、五个真实的坑
| 坑 | 症状 | 解法 |
|---|---|---|
忘了 set_epoch |
每轮数据顺序完全一样 | 加上它 ⭐ |
忘了 no_sync() |
梯度累积时通信量翻 N 倍 | 见上面代码 |
| 各卡 batch 不等长 | AllReduce 挂死(有的卡先退出循环) | 补齐或用 join() 上下文 ⭐ |
| BatchNorm 各算各的 | 小 batch 时统计不准,精度掉 | 换 SyncBatchNorm 或用 LayerNorm |
| 随机性不一致 | 各卡 dropout mask 不同导致结果不可复现 | 统一 seed;数据增强的种子要各卡不同 |
💥 "AllReduce 挂死"是最难排查的一个: 症状是训练突然卡住不动、没有报错、GPU 利用率 100%(在空转等待)。 根因常常是某张卡的数据比别人少一个 batch,提前退出了循环, 其他卡在等一个永远不会来的 AllReduce。
✅ 排查:
export TORCH_NCCL_BLOCKING_WAIT=1会让它超时报错而不是永久挂起。
🔗 和站内其他章的关系
| 相关的地方 | 这里的位置 |
|---|---|
| 第 5 章 Ring AllReduce | 本章是它的展开 |
| 第 5 章 AllReduce=RS+AG | 分层和 ZeRO 都基于它 |
| 第 3 章 显存四大块 | 数据并行一块都没省 ⭐ |
| 第 11 章 ZeRO | 数据并行的省显存版 |
| 第 9 章 梯度累积 | 和 no_sync() 配合 |
| 《机器学习与深度学习基础》09 优化器与学习率 warmup 与学习率量级 | 数据并行把全局 batch 放大了 N 倍,学习率必须跟着改 —— 不然扩到几十卡就发散 ⭐ |
| 《Kaggle竞赛方法论》02 训练策略与正则化 梯度累加 | 那里它是"显存不够时模拟大 batch";在 DDP 里不配 no_sync(),每个累积步都会白同步一次 ⚠️ |
✅ 检查点
- 数据并行怎么工作?为什么梯度要求平均而不是求和?
sampler.set_epoch()不加会怎样?- DDP 的梯度桶做了什么?为什么反向传播的顺序让它成为可能?
- 梯度累积时忘了
no_sync()会怎样? - Ring AllReduce 每卡的通信量是多少?它的软肋是什么?
- 分层 AllReduce 为什么能降到 1/8?
- 为什么卡越多扩展效率越低?
- 数据并行省显存吗?
- "训练突然卡死、无报错、GPU 100%"最可能是什么原因?
👀 答案
- 每张卡放完整模型副本,吃不同数据分片,算完梯度后 AllReduce 求平均,所有卡拿到相同梯度各自 step,参数保持一致。求平均是因为等效 batch 变成 N 倍,求和相当于学习率被放大 N 倍。
- 每个 epoch 的数据顺序和划分完全一样,等于没有 shuffle。
- 攒够一个桶(默认 25MB)就立刻开始 AllReduce,让通信和后面层的计算重叠。可能是因为反向传播是从后往前的,最后一层的梯度最先算出来,不用等全部算完。
- 通信量翻 N 倍(累积 8 步就是 8 倍必要通信量),而代码看起来完全正常——隐蔽的性能杀手。
- ≈ 2 × 数据量,和卡数几乎无关。软肋是步数为 2(N−1),延迟随卡数线性增长。
- 先在机内用 NVLink 做一次归约,只让各机的代表跨机通信(数据量变成 1/8),最后机内广播回去。
- 因为卡数增加时计算时间下降(每卡数据变少)但通信量基本不变,所以通信占比上升。
- 一块都没省——每张卡都有完整的参数、梯度、优化器状态。省显存要用 ZeRO。
- 某张卡的数据比别人少一个 batch 提前退出循环,其他卡在等一个永远不来的 AllReduce。排查:
TORCH_NCCL_BLOCKING_WAIT=1让它超时报错而不是永久挂起。
🛑 可以停在这里
⚡ 走神救援
先记住这几件事
- 数据并行让每卡保留模型副本、处理不同数据,再同步梯度;它本身不分摊模型状态。
- 检查采样器、梯度累积和各卡迭代次数,避免重复通信或集合操作相互等待。
- 通过梯度分桶与计算通信重叠提高效率,再实测增加卡数后的扩展收益。
下一节 👉 11-ZeRO与FSDP.md ⭐