集合通信从零:ring allreduce 与三大原语 本节摘要:撑起分布式训练的四个集合通信原语是 allreduce、broadcast、allgather、reducescatter,训练框架提供的其他原语都是它们的包装。本节在 网格上把它们各实现一次:ring allreduce 拆成两遍(reduce-scatter 再 allgather),并证明每 rank 通信量是 字节每元素;broadcast、allgather、reducescatter 在点到点发送上构建;每个原语都用 gloo 参考实现验证同输入产出一致。读完本节,你能为集群形状(大张量胖管 vs 小张量高延迟)辩护选 ring 还是 tree,并为后续 DDP、ZeRO、流水线并行的全部行为打下地基。
本节摘要:撑起分布式训练的四个集合通信原语是 allreduce、broadcast、allgather、reduce_scatter,训练框架提供的其他原语都是它们的包装。本节在
multiprocessing.Queue网格上把它们各实现一次:ring allreduce 拆成两遍(reduce-scatter 再 allgather),并证明每 rank 通信量是2(N-1)/N字节每元素;broadcast、allgather、reduce_scatter 在点到点发送上构建;每个原语都用torch.distributedgloo 参考实现验证同输入产出一致。读完本节,你能为集群形状(大张量胖管 vs 小张量高延迟)辩护选 ring 还是 tree,并为后续 DDP、ZeRO、流水线并行的全部行为打下地基。
对应原课程:Phase 19 · Lesson 76 ·
collective-ops-from-scratch(原英文phases/19-capstone-projects/76-collective-ops-from-scratch/docs/en.md)。本节属第 20 章「毕业项目」的分布式训练赛道(Track H),是本赛道的地基。
阅读完本节,你应当能够:
2(N-1)/N 字节每元素。multiprocessing.Queue 的点到点发送之上构建 broadcast、allgather、reduce_scatter。torch.distributed gloo 参考验证同输入产出一致。N 个 rank 上的朴素 allreduce 把张量发 N 次到 root、再广播 N 次回来:每 rank 带宽 O(N),root 成瓶颈,墙钟下限是最慢链路乘 N。Ring allreduce 把它压平成 2(N-1) 个大小 T/N 的块,每 rank 字节数降到 2T(N-1)/N,与集群规模无关。Tree allreduce 在小 N 与高延迟链路上赢,因为深度是 log2(N) 跳而非 2(N-1)。给集群形状选错拓扑,最慢的 GPU 就主导步时间。
你在这条赛道读到的每个分布式训练框架都依赖这四个原语:PyTorch DDP 对参数桶做一次 allreduce 同步梯度;ZeRO 用 reduce_scatter 分片优化器状态、用 allgather 广播更新后参数;FSDP 把整个 forward 变成 allgather 加 reduce_scatter;流水线并行用 broadcast 跨阶段组传激活。你实现不了这四个集体通信,就推理不了为什么训练卡住、为什么梯度错配出现在 rank 3、为什么换拓扑后流水线气泡翻倍。
把张量切成 N 个等大块,下标 0..N-1,每 rank 拥有等于自己 rank 的块。第一遍 reduce-scatter 跑 N-1 步:第 s 步,rank r 把块 (r - s) mod N 发给 (r + 1) mod N,从 (r - 1) mod N 收块 (r - s - 1) mod N 并累加进本地副本。N-1 步后,rank r 拥有块 r 的完整和。第二遍 allgather 再跑 N-1 步,把完工的块沿环转直到每 rank 都持有每块的完整和。
| 原语 | 每 rank 字节 | 步数 | 何时用 |
|---|---|---|---|
| Ring allreduce | 2T(N-1)/N | 2(N-1) | 大 T、胖管同构集群 |
| Tree allreduce | T log2(N) | 2 log2(N) | 小 T 或高延迟链路 |
| Broadcast | T | log2(N) 树 | 参数初始化、标量配置 |
| Allgather | T(N-1)/N | N-1 | 分片 forward、ZeRO 解分片 |
| Reduce_scatter | T(N-1)/N | N-1 | ZeRO 梯度分片 |
NCCL 跑在 PCIe 与 NVLink 上,带硬件卸载归约。CPU 上你没有这些。每条环边一个 multiprocessing.Queue 给你单生产者单消费者的有序点到点投递;归约在用户态发生,你付 Python 开销,但线缆模式与 NCCL ring allreduce 完全一致。在队列版上推理正确性,集群行为随之而来。
def ring_allreduce(mesh, rank, world_size, tensor): chunks = split(tensor, world_size) # 切 N 块 # 第一遍:reduce-scatter for s in range(world_size - 1): send_idx = (rank - s) % world_size recv_idx = (rank - s - 1) % world_size mesh.send((rank + 1) % world_size, chunks[send_idx]) chunks[recv_idx] += mesh.recv((rank - 1) % world_size) # 第二遍:allgather for s in range(world_size - 1): send_idx = (rank - s + 1) % world_size recv_idx = (rank - s) % world_size mesh.send((rank + 1) % world_size, chunks[send_idx]) chunks[recv_idx] = mesh.recv((rank - 1) % world_size) return concat(chunks)
每个原语落地都带一个单元测试,把它的输出与 torch.distributed(gloo 后端、同张量、同 world size)比对。你的 ring allreduce 与 gloo 偏差超过 float32 epsilon,测试就失败。对参考实现验证不可商量——没有它,原语看起来对,直到真实训练的第 10000 步才露馅。
业界对比:PyTorch torch.distributed(NCCL/gloo 后端)、Horovod、DeepSpeed、Megatron-LM 都构建在这四个原语之上。NCCL 自带拓扑检测:消息 > ~1MB 用 ring,< ~1MB 用 tree——交叉点是带宽对延迟:大消息带宽项 2T(N-1)/N 主导 ring 赢;小消息 log2(N) 跳数赢。硬编码一种拓扑会在错的消息大小上损吞吐。它们的共识与本节一致:按消息大小而非信仰选 ring 或 tree。
三种模式把原语硬化到可上线。
allreduce 前把梯度装桶。 一个 10 亿参数模型有数万个梯度张量,每张量一次 allreduce 付 N 次延迟下限。DDP 把梯度装进 ~25MB 桶,每桶一次 allreduce,小张量搭大张量的便车。不装桶,延迟开销主导步时间。
通信与计算重叠。 反向逐层逆序算梯度,最后一层梯度一就绪,立刻 kick off 它的 allreduce,同时下一层继续算。PyTorch DDP 用桶就绪钩子接线,网络有空时重叠把可见通信时间砍半。
按消息大小选 ring 或 tree。 NCCL 的拓扑检测对 >1MB 选 ring、<1MB 选 tree。交叉是带宽对延迟。硬编码一种拓扑在错消息大小上损吞吐。
Mesh 类:把 N 个 multiprocessing.Queue 接成环,暴露每 rank 的 send(dst, tensor) 与 recv(src)。是 77~81 节全部原语的共用底座。ring_allreduce、broadcast、allgather、reduce_scatter,77 节 DDP 用 allreduce、78 节 ZeRO 用 reduce_scatter/allgather、79 节流水线用 broadcast、81 节端到端组合全部四个。_gloo_reference:字节级比对的参考实现,任何重写都有回归网。code/main.py 实现:Mesh、ring_allreduce、broadcast(对数树)、allgather(N-1 轮转)、reduce_scatter(allreduce 前半)、_gloo_reference。运行 python3 code/main.py 输出每原语验证表(队列网格 vs gloo)与每 rank 字节计数器,证明 2T(N-1)/N 缩放。
recv_timeout_ms:让停滞的 rank 抛截止错误而非永远挂起。multiprocessing.Queue 换成 TCP socket 实现四个原语,同测试,真实线缆。2T(N-1)/N,与集群规模无关。log2(N),小 T 或高延迟赢;按消息大小选 ring(>1MB)或 tree(<1MB)。下一节,我们将进入「数据并行 DDP」——把本节的 allreduce 包成
DistributedDataParallel:广播初始参数、反向后 allreduce 梯度、除以 world_size 取均值,200 行实现。