一句话语义
通信组中有 个 rank,每个 rank 持有同 shape 张量 。以求和为例,All-Reduce 结束后,每个 rank 都得到:
它同时完成两件事:先 reduce 出一个结果,再让所有参与者都得到这个结果。因此它在数据并行中用于同步梯度,也在张量并行中用于合并行并行线性层产生的部分和。集合通信总览用同一组输入对比 All-Gather、Reduce-Scatter 与 All-Reduce 的 shape 和所有权。
从最直观的实现到 Ring All-Reduce
最容易想到的是先把所有数据 Reduce 到一个中心 rank,再 Broadcast 回去。语义正确,但中心节点会成为带宽瓶颈。Ring All-Reduce 则把 rank 首尾相连,让每条链路在大部分时间里同时工作。
它把操作拆成两个阶段:
- Reduce-Scatter:边传递边累加。结束时,每个 rank 拥有最终结果的一个不同分块。
- All-Gather:传播这些已经归约完成的分块。结束时,每个 rank 都拥有完整结果。
下面以 为例。颜色表示向量中的块位置,下标表示数据来自哪个 rank。
INTERACTIVE EXPLAINER
Ring All-Reduce:从四份局部向量到四份全局结果
动画把 3 轮 Reduce-Scatter 与 3 轮 All-Gather 压缩为关键状态;重点观察每个阶段拥有的数据和通信语义。
四个 rank 各自持有长度相同但数值不同的向量;目标是逐元素求和,并把完整结果交还给所有 rank。
为什么需要 轮
每个向量被切成 块。在单向环上,一个块每轮只移动到相邻 rank:
- Reduce-Scatter 需要 轮,才能让每个目标块累计来自全部 rank 的贡献。
- All-Gather 再需要 轮,才能让每个完整归约块传播到全部 rank。
因此总轮数为:
如果每个 rank 的输入大小为 字节,那么每轮发送约 字节,每个 rank 总发送量约为:
在常见的 - 模型中, 表示一次通信的固定延迟, 表示传输每字节的时间,Ring All-Reduce 的近似成本为:
这解释了它的特点:对大张量,带宽利用率很好;对很小的张量, 次启动延迟可能占主导,树形算法往往更合适。
代码中的 All-Reduce
PyTorch 的 all_reduce 默认是原地操作:
import torch
import torch.distributed as dist
# 每个 rank 上 tensor 的 shape 和 dtype 必须兼容
tensor = torch.tensor([rank + 1.0], device="cuda")
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
# world_size=4 时,每个 rank 都得到 tensor([10.])
print(tensor)
在异步模式下,返回的 work handle 表示通信已经入队,不代表结果立刻可以安全使用:
work = dist.all_reduce(tensor, async_op=True)
do_independent_compute()
work.wait() # 第一次读取归约结果前必须建立同步关系
算法不是永远只有 Ring
| 算法 | 优势 | 更适合 |
|---|---|---|
| Ring | 大消息带宽利用率高,负载均匀 | 大梯度、大激活、稳定高速链路 |
| Tree | 通信轮数约为 | 小消息或延迟敏感场景 |
| Hierarchical | 先节点内、再节点间,匹配多级拓扑 | 多机多卡集群 |
| 机内专用算法 | 利用 NVLink/NVSwitch 等拓扑 | 高带宽单节点通信组 |
NCCL 等通信库会根据消息大小、拓扑和运行环境选择或调优算法。应用通常声明的是集合通信语义,而不是手写每一轮点对点传输。
工程中最容易踩的坑
- shape 或 dtype 不一致:所有 rank 必须以兼容的张量参与同一次 collective。
- collective 顺序不一致:某些 rank 进入 A,另一些进入 B,会造成死锁或数据错误。
- 把同步时间算错位置:异步通信的等待可能推迟到后续算子,看起来像是后续算子突然变慢。
- 跨节点带宽不足:算法本身带宽最优,不等于物理网络足够快。
- 小张量过多:大量小 All-Reduce 会被启动延迟主导,通常需要 bucket 或 fusion。
与张量并行串起来看
在行并行线性层中,每个 TP rank 得到完整输出 shape 的部分和 。All-Reduce 对这些部分和逐元素求和,并把 返回给全部 TP rank:
回到张量并行的 MLP 动画,最后一步正是这个过程。理解“本地矩阵乘产生部分和,集合通信恢复逻辑张量”,就理解了大量模型并行实现的基本节奏。