写在前面(接上篇)
上篇《一文吃透分布式通信原语(上):点对点基础与全局图谱》把点对点原语(Send/Recv、ISend/IRecv、SendRecv、Barrier)和全局对比表讲清了,并反复强调一句话:集合通信底层都是一串 Send/Recv,只是组织方式(环 / 树 / 蝶形)不同。
本篇就把这些积木拼起来,逐个讲透集合通信,并在最后用 Send/Recv 从零实现 Ring-AllReduce、与官方实现对拍验证。沿用上篇的约定:S = 单份数据字节数,N = rank 数;代码用 torch.distributed(有 GPU 用 NCCL,无 GPU 退化到 gloo),每段附实跑输出。
若还没看过上篇,建议先读。小节编号延续上篇(上篇为 0、1,本篇为 2、3),方便两篇对照。
2. 集合通信原语:由点对点拼装的集体操作
集合通信是一组 rank 共同完成的操作。记住一句话:它们底层都是一串 Send/Recv,只是组织方式(环 / 树 / 蝶形)不同,带来不同的延迟-带宽权衡。
两类主流算法先打个底:
- Ring(环状):N 个 rank 串成环,数据分成 N 块,每步每个 rank 发一块给右邻居、收一块。带宽最优(充分利用每条链路),延迟随 N 线性增长,适合大数据。
- Tree(树状)/ Halving-Doubling / Recursive Doubling:用树或蝶形结构,延迟最优(
~logN步),适合小数据或延迟敏感场景。
NCCL 会根据数据量、拓扑自动在 Ring 和 Tree 之间选择。
2.1 Broadcast:广播
1. 定义与语义
broadcast(tensor, src):把 root(src)的数据复制到所有 rank。
操作前 操作后
rank0(root): [1,2,3] ──► rank0: [1,2,3]
rank1: [0,0,0] rank1: [1,2,3]
rank2: [0,0,0] rank2: [1,2,3]
rank3: [0,0,0] rank3: [1,2,3]
2. 底层原理与算法
- 朴素实现:root 依次
Send给每个 rank → N-1 次串行发送,延迟O(N)。 - 树状广播:root 发给 2 个,这 2 个再各发 2 个……
logN步完成,延迟O(logN)。 - 通信量:每个非 root rank 收
S;树状下延迟~S·logN量级。
3. 组合关系
树状 Broadcast 本质是一棵由 Send/Recv 构成的转发树。Scatter 可视为「广播的不同分片版本」(每个 rank 拿不同的一块,而非同一份)。
4. 应用场景
- 参数初始化 / 权重同步:训练开始时,把 rank0 随机初始化(或从 checkpoint 加载)的权重广播给所有 rank,保证各 DP 副本初始一致。
- 为什么选它:目标是「所有人拿到同一份数据」,正是 Broadcast 的语义。
5. 实例代码
# broadcast.py
# 启动:torchrun --nproc_per_node=4 broadcast.py
import torch
import torch.distributed as dist
def main():
dist.init_process_group(backend="nccl" if torch.cuda.is_available() else "gloo")
rank = dist.get_rank()
device = torch.device(f"cuda:{rank}") if torch.cuda.is_available() else torch.device("cpu")
if rank == 0:
tensor = torch.tensor([1.0, 2.0, 3.0], device=device) # root 有数据
else:
tensor = torch.zeros(3, device=device) # 其它 rank 待填充
print(f"[rank{rank}] 操作前 = {tensor.tolist()}")
dist.broadcast(tensor, src=0) # 从 rank0 广播
print(f"[rank{rank}] 操作后 = {tensor.tolist()}")
dist.destroy_process_group()
if __name__ == "__main__":
main()
实测执行(本机无 GPU,走 gloo 后端):
torchrun --nproc_per_node=4 --rdzv-endpoint=127.0.0.1:29506 broadcast.py
真实输出(按 rank 整理,打印顺序可能交错):
[rank0] 操作前 = [1.0, 2.0, 3.0]
[rank1] 操作前 = [0.0, 0.0, 0.0]
[rank2] 操作前 = [0.0, 0.0, 0.0]
[rank3] 操作前 = [0.0, 0.0, 0.0]
[rank0] 操作后 = [1.0, 2.0, 3.0]
[rank1] 操作后 = [1.0, 2.0, 3.0]
[rank2] 操作后 = [1.0, 2.0, 3.0]
[rank3] 操作后 = [1.0, 2.0, 3.0]
rank0(root)的 [1,2,3] 被广播到所有 rank,其余 rank 从全 0 变成 [1,2,3]。
2.2 Scatter / Gather:分发 / 收集
1. 定义与语义
Scatter:root 持有 N 份数据,把第 i 份发给 rank i。(分发,各 rank 拿到不同片)Gather:每个 rank 持有 1 份,root 把 N 份按 rank 顺序收齐。(收集,Scatter 的逆操作)
Scatter(root=rank0 持 [A,B,C,D]):
操作前 rank0:[A,B,C,D] rank1:[] rank2:[] rank3:[]
操作后 rank0:[A] rank1:[B] rank2:[C] rank3:[D]
Gather(收集到 root=rank0):
操作前 rank0:[A] rank1:[B] rank2:[C] rank3:[D]
操作后 rank0:[A,B,C,D] rank1:[B] rank2:[C] rank3:[D]
2. 底层原理与算法
- 通信量:root 侧发出 / 收齐约
(N-1)/N·S(自己那份不走网络),其余每 rank 收/发S/N。 - 朴素实现是 root 对每个 rank 一次
Send/Recv,延迟O(N)。 - ⚠️ 后端限制:
scatter/gather在 NCCL 后端上支持有限(老版本不支持),用 gloo 后端最稳。GPU 上通常用AllGather+ 切片,或all_to_all替代。
3. 组合关系
Gather是Scatter的逆。AllGather=Gather到所有 rank(而非只到 root)=Gather+Broadcast的效果。
4. 应用场景
Scatter:把一批样本 / 一个大张量的不同分片下发给各 rank 处理。Gather:推理/评测时把各 rank 的局部结果汇总到主 rank 统一输出或落盘。
5. 实例代码
# scatter_gather.py
# 启动:torchrun --nproc_per_node=4 scatter_gather.py (建议 gloo 后端)
import torch
import torch.distributed as dist
def main():
# scatter/gather 用 gloo 最稳
dist.init_process_group(backend="gloo")
rank = dist.get_rank()
world = dist.get_world_size()
# ---------- Scatter ----------
recv = torch.zeros(1)
if rank == 0:
# root 准备 N 份数据:[10],[11],[12],[13]
scatter_list = [torch.tensor([10.0 + i]) for i in range(world)]
else:
scatter_list = None
dist.scatter(recv, scatter_list=scatter_list, src=0)
print(f"[rank{rank}] scatter 后收到 = {recv.tolist()}")
# ---------- Gather ----------
send = torch.tensor([float(rank)]) # 每 rank 发自己的编号
if rank == 0:
gather_list = [torch.zeros(1) for _ in range(world)]
else:
gather_list = None
dist.gather(send, gather_list=gather_list, dst=0)
if rank == 0:
print(f"[rank0] gather 收齐 = {[t.tolist() for t in gather_list]}")
dist.destroy_process_group()
if __name__ == "__main__":
main()
实测执行(scatter/gather 用 gloo 后端最稳):
torchrun --nproc_per_node=4 --rdzv-endpoint=127.0.0.1:29507 scatter_gather.py
真实输出(按 rank 整理,打印顺序可能交错):
[rank0] scatter 后收到 = [10.0]
[rank1] scatter 后收到 = [11.0]
[rank2] scatter 后收到 = [12.0]
[rank3] scatter 后收到 = [13.0]
[rank0] gather 收齐 = [[0.0], [1.0], [2.0], [3.0]]
Scatter:root 把 [10,11,12,13] 第 i 份分发给 rank i;Gather:root 把各 rank 的编号 0/1/2/3 按序收齐。
6. 对比补充:Scatter / Gather 与 AllGather 真的没区别吗
有区别,而且关键。三者数据流向完全不同(N=4,root=rank0):
Scatter(一对多·分发,各拿不同片)
操作前 rank0:[A,B,C,D] rank1:[·] rank2:[·] rank3:[·]
操作后 rank0:[A] rank1:[B] rank2:[C] rank3:[D]
Gather(多对一·收集,只有 root 拿到全部)
操作前 rank0:[A] rank1:[B] rank2:[C] rank3:[D]
操作后 rank0:[A,B,C,D] rank1:[B] rank2:[C] rank3:[D]
AllGather(多对多·全收集,人手一份全部)
操作前 rank0:[A] rank1:[B] rank2:[C] rank3:[D]
操作后 rank0:[A,B,C,D] rank1:[A,B,C,D] rank2:[A,B,C,D] rank3:[A,B,C,D]
- Scatter:root 把 N 份不同数据,一人分一份(分发)。
- Gather:各 rank 的 1 份汇总,只到 root(收集);非 root 拿不到完整结果。
- AllGather:各 rank 的 1 份汇总,人手一份完整的(多对多)。
组合关系与方向:
AllGather = Gather(先收集到 root) + Broadcast(再广播给所有人)
Scatter = Gather 的逆方向(分发 vs 收集)
本文两段实跑输出正是最直接的佐证——Gather 只有 rank0 打印完整结果,而 AllGather 每个 rank 都打印 [0,1,2,3]:
# Gather:只有 rank0 齐
[rank0] gather 收齐 = [[0.0], [1.0], [2.0], [3.0]]
# AllGather:人手一份
[rank0] 操作后 收齐 = [0.0, 1.0, 2.0, 3.0]
[rank1] 操作后 收齐 = [0.0, 1.0, 2.0, 3.0]
[rank2] 操作后 收齐 = [0.0, 1.0, 2.0, 3.0]
[rank3] 操作后 收齐 = [0.0, 1.0, 2.0, 3.0]
一句话结论:Gather「只汇到 root」,AllGather「人人都拿到」,Scatter 则是反过来「分发不同片」——三者并不等价。也正因为 AllGather 要「人人都有」,通信比 Gather 重,常用 Ring 算法做到带宽最优。
2.3 AllGather:全收集
1. 定义与语义
all_gather:每个 rank 持有 1 份,操作后每个 rank 都拿到全部 N 份(按 rank 顺序拼接)。
操作前 rank0:[A] rank1:[B] rank2:[C] rank3:[D]
操作后 rank0:[A,B,C,D] rank1:[A,B,C,D] rank2:[A,B,C,D] rank3:[A,B,C,D]
2. 底层原理与算法
- Ring-AllGather:N-1 步,每步每个 rank 把「当前手里最新的一块」发给右邻居、从左邻居收一块。N-1 步后所有块转遍整个环。
- 通信量:每个 rank 发送 / 接收约
(N-1)/N·S(S为聚合后总大小)。带宽最优,几乎与朴素广播无关。
3. 组合关系
AllReduce = ReduceScatter + AllGather——这是最重要的组合等式,后面 AllReduce 会详述。AllGather 负责「把各自算好的分片拼回完整结果并人手一份」。
4. 应用场景
- ZeRO / FSDP 的参数聚合:FSDP 把模型参数按 rank 切分存储(每 rank 只存 1/N)。前向 / 反向用到某层时,临时
AllGather把该层完整参数拼回来算,算完立刻释放。用显存换通信,让单卡能放下超大模型。 - 为什么选它:需要「每个 rank 都临时拿到完整参数」,正是 AllGather 语义。
5. 实例代码
# all_gather.py
# 启动:torchrun --nproc_per_node=4 all_gather.py
import torch
import torch.distributed as dist
def main():
dist.init_process_group(backend="nccl" if torch.cuda.is_available() else "gloo")
rank = dist.get_rank()
world = dist.get_world_size()
device = torch.device(f"cuda:{rank}") if torch.cuda.is_available() else torch.device("cpu")
local = torch.tensor([float(rank)], device=device) # 每 rank 1 份
gather_list = [torch.zeros(1, device=device) for _ in range(world)]
print(f"[rank{rank}] 操作前 local = {local.tolist()}")
dist.all_gather(gather_list, local) # 全收集
result = [t.item() for t in gather_list]
print(f"[rank{rank}] 操作后 收齐 = {result}")
dist.destroy_process_group()
if __name__ == "__main__":
main()
每个 rank 操作后都打印 [0.0, 1.0, 2.0, 3.0]。
实测执行(本机无 GPU,走 gloo 后端):
torchrun --nproc_per_node=4 --rdzv-endpoint=127.0.0.1:29508 all_gather.py
真实输出(按 rank 整理,打印顺序可能交错):
[rank0] 操作前 local = [0.0]
[rank1] 操作前 local = [1.0]
[rank2] 操作前 local = [2.0]
[rank3] 操作前 local = [3.0]
[rank0] 操作后 收齐 = [0.0, 1.0, 2.0, 3.0]
[rank1] 操作后 收齐 = [0.0, 1.0, 2.0, 3.0]
[rank2] 操作后 收齐 = [0.0, 1.0, 2.0, 3.0]
[rank3] 操作后 收齐 = [0.0, 1.0, 2.0, 3.0]
每个 rank 各持 1 份编号,AllGather 后人手一份完整的 [0,1,2,3]。
2.4 Reduce / AllReduce:归约 / 全归约
1. 定义与语义
Reduce(op):把所有 rank 的数据按op(SUM/MAX/…)逐元素归约,结果只放到 root。AllReduce(op):同样归约,但结果人手一份(所有 rank 都拿到)。
归约操作符覆盖:SUM(求和)/ AVG(求平均)/ MAX / MIN / PROD(求积)。
AllReduce + SUM,N=4,每 rank 持一个标量:
操作前 rank0:[1] rank1:[2] rank2:[3] rank3:[4]
操作后 rank0:[10] rank1:[10] rank2:[10] rank3:[10] (1+2+3+4=10,人手一份)
AllReduce + AVG:
操作后 每 rank 都是 10/4 = 2.5
2. 底层原理与算法:为什么 Ring-AllReduce 通信量与 N 几乎无关
Ring-AllReduce = ReduceScatter + AllGather 两阶段,这是带宽最优算法:
把每个 rank 的数据切成 N 块,环上 N 个 rank:
- 阶段一 ReduceScatter(N-1 步):每步每个 rank 把一块「加」到右邻居对应块上并传递。N-1 步后,每个 rank 手里有「某一块的全局归约和」(第 i 块的总和落在某个 rank 上)。
- 阶段二 AllGather(N-1 步):把这些「已归约好的块」沿环转一圈,让每个 rank 都集齐所有块。
通信量推导:每个 rank 在每一步收发一块(大小 S/N)。两阶段各 N-1 步:
单 rank 通信量 = 2 × (N-1) × (S/N) = 2(N-1)/N · S ──(N→∞)──► 2S
结论:当 N 很大时,每个 rank 的通信量趋近于常数 2S,几乎与 N 无关。这就是 Ring-AllReduce 能扩展到上千卡仍高效的根本原因——加再多卡,单卡要搬的字节数不变(代价是延迟随 N 线性增长,所以超大 N 时 NCCL 会切到 Tree)。
- Tree-AllReduce:延迟
~2·logN步,小数据更快。 - 拓扑敏感度:Ring 对带宽敏感、吃满 NVLink;跨机走 IB 时延迟放大,NCCL 会分层(机内 Ring + 机间 Tree)。
3. 组合关系
- AllReduce = ReduceScatter + AllGather(本文 2.5 的 ReduceScatter + 本文 2.3 的 AllGather)。
Reduce= AllReduce 后只保留 root 的那份(或归约树只汇到 root,不再广播回去)。
4. 应用场景
- 数据并行(DP)的梯度同步,用的正是
AllReduce + AVG:- 每个 DP 副本在不同数据上算出各自的梯度,要把它们平均后让所有副本用同一个平均梯度更新 → 保证参数始终一致。
- 为什么是 AVG 而不是 SUM:loss 通常是 batch 内样本的平均,梯度要对「全局 batch」求平均才无偏;SUM 会让梯度放大 N 倍,等价于学习率被偷偷乘了 N。(实现上也可先 SUM 再除以 N,效果等同 AVG。)
- 张量并行(TP)的层内同步,用的是
AllReduce + SUM:一层(如 Megatron 把一个 MLP/Attention 的权重按列或行切到多卡)被切开后,每张卡只算出部分和,前向要AllReduce把各卡的部分结果加起来得到完整输出,反向同理对输入梯度再做一次AllReduce。所以 TP 每层都有 1~2 次 AllReduce,通信极频繁——这也是 TP 通常只在机内高带宽(NVLink)范围内用、不跨机的原因。 Reduce:把各 rank 的 loss / 统计量汇总到主 rank 打印或记录(不需要人手一份时用 Reduce 比 AllReduce 省一半通信)。
5. 实例代码
# all_reduce.py
# 启动:torchrun --nproc_per_node=4 all_reduce.py
import torch
import torch.distributed as dist
def main():
dist.init_process_group(backend="nccl" if torch.cuda.is_available() else "gloo")
rank = dist.get_rank()
world = dist.get_world_size()
device = torch.device(f"cuda:{rank}") if torch.cuda.is_available() else torch.device("cpu")
# 每个 rank 持有自己的「梯度」(这里用 rank+1 模拟:1,2,3,4)
grad = torch.tensor([float(rank + 1)], device=device)
print(f"[rank{rank}] 操作前 grad = {grad.tolist()}")
# ---- AllReduce + SUM ----
s = grad.clone()
dist.all_reduce(s, op=dist.ReduceOp.SUM)
# ---- AllReduce + AVG:梯度同步真正用的(无 AVG 时先 SUM 再除以 world)----
avg = grad.clone()
dist.all_reduce(avg, op=dist.ReduceOp.SUM)
avg /= world
print(f"[rank{rank}] SUM = {s.tolist()} AVG(梯度同步) = {avg.tolist()}")
dist.destroy_process_group()
if __name__ == "__main__":
main()
预期:每个 rank 都打印 SUM = [10.0] AVG = [2.5]。
实测执行(本机无 GPU,走 gloo 后端):
torchrun --nproc_per_node=4 --rdzv-endpoint=127.0.0.1:29509 all_reduce.py
真实输出(按 rank 整理,打印顺序可能交错):
[rank0] 操作前 grad = [1.0]
[rank1] 操作前 grad = [2.0]
[rank2] 操作前 grad = [3.0]
[rank3] 操作前 grad = [4.0]
[rank0] SUM = [10.0] AVG(梯度同步) = [2.5]
[rank1] SUM = [10.0] AVG(梯度同步) = [2.5]
[rank2] SUM = [10.0] AVG(梯度同步) = [2.5]
[rank3] SUM = [10.0] AVG(梯度同步) = [2.5]
四个 rank 的梯度 1/2/3/4,SUM 后人手一份 10,除以 world=4 得平均 2.5——这正是数据并行梯度同步的做法。
其它归约符:把
op换成dist.ReduceOp.MAX / MIN / PRODUCT即可。注意 NCCL 对AVG的原生支持见版本,稳妥做法是「SUM 后除以 world」。
2.5 ReduceScatter:归约分散
1. 定义与语义
reduce_scatter:先把所有 rank 的数据逐元素归约,再把结果按块切开分散——每个 rank 只拿到归约结果的「自己那一块」。
每 rank 持 N=4 块数据,ReduceScatter + SUM:
操作前 rank0: [a0, a1, a2, a3]
rank1: [b0, b1, b2, b3]
rank2: [c0, c1, c2, c3]
rank3: [d0, d1, d2, d3]
操作后 rank0: [a0+b0+c0+d0] ← 只拿第 0 块的全局和
rank1: [a1+b1+c1+d1] ← 只拿第 1 块
rank2: [a2+b2+c2+d2] ← 只拿第 2 块
rank3: [a3+b3+c3+d3] ← 只拿第 3 块
2. 底层原理与算法
- Ring-ReduceScatter(N-1 步):就是 Ring-AllReduce 的前半段。
- 通信量:单 rank
(N-1)/N·S(恰为 AllReduce 的一半)。
3. 组合关系
AllReduce = ReduceScatter + AllGather:ReduceScatter 让每人拿到「一块的全局和」,AllGather 再把这些块拼齐并人手一份。
4. 应用场景
- ZeRO / FSDP 的梯度切分:反向算出梯度后,不做完整 AllReduce,而是
ReduceScatter—— 每个 rank 只保留并更新自己负责的那 1/N 参数对应的梯度。配合参数也切分(AllGather 临时取回),显存占用降到 1/N 量级。 - 为什么选它:FSDP 下每个 rank 只需要「自己那片参数的归约梯度」,不需要完整梯度,用 ReduceScatter 比 AllReduce 省一半通信。
5. 实例代码
# reduce_scatter.py
# 启动:torchrun --nproc_per_node=4 reduce_scatter.py
import torch
import torch.distributed as dist
def main():
dist.init_process_group(backend="nccl" if torch.cuda.is_available() else "gloo")
rank = dist.get_rank()
world = dist.get_world_size()
device = torch.device(f"cuda:{rank}") if torch.cuda.is_available() else torch.device("cpu")
# 每个 rank 持 world 块;第 j 块的值 = rank*10 + j
input_list = [torch.tensor([float(rank * 10 + j)], device=device) for j in range(world)]
output = torch.zeros(1, device=device)
print(f"[rank{rank}] 操作前 = {[t.item() for t in input_list]}")
dist.reduce_scatter(output, input_list, op=dist.ReduceOp.SUM)
# rank i 应拿到「所有 rank 的第 i 块」之和 = Σ_r (r*10 + i)
print(f"[rank{rank}] 操作后 只拿第{rank}块的全局和 = {output.tolist()}")
dist.destroy_process_group()
if __name__ == "__main__":
main()
N=4 时,rank i 的结果 = Σ(10r+i), r=0..3 = 60 + 4i → rank0:60, rank1:64, rank2:68, rank3:72。
实测执行(本机无 GPU,走 gloo 后端):
torchrun --nproc_per_node=4 --rdzv-endpoint=127.0.0.1:29510 reduce_scatter.py
真实输出(按 rank 整理,打印顺序可能交错):
[rank0] 操作前 = [0.0, 1.0, 2.0, 3.0]
[rank1] 操作前 = [10.0, 11.0, 12.0, 13.0]
[rank2] 操作前 = [20.0, 21.0, 22.0, 23.0]
[rank3] 操作前 = [30.0, 31.0, 32.0, 33.0]
[rank0] 操作后 只拿第0块的全局和 = [60.0]
[rank1] 操作后 只拿第1块的全局和 = [64.0]
[rank2] 操作后 只拿第2块的全局和 = [68.0]
[rank3] 操作后 只拿第3块的全局和 = [72.0]
rank i 只拿到「所有 rank 第 i 块之和」:60/64/68/72,与推导 60+4i 完全一致。
2.6 All2All / AllToAllv:全交换
1. 定义与语义
All2All 与 AllGather 的本质区别:AllGather 中每个 rank 发给所有人的是「同一份」数据;All2All 中每个 rank 发给不同人的是「不同」数据。 相当于一次分布式的「矩阵转置」。
- 等长 All2All:每个 rank 发给每个 rank 的块大小相同。
- 变长 AllToAllv:各块大小可以不同(variable)。
All2All,N=3,用 r_sX 表示「rank r 要发给 rank X 的数据」:
操作前 rank0:[0_s0, 0_s1, 0_s2] # 第 j 项准备发给 rank j
rank1:[1_s0, 1_s1, 1_s2]
rank2:[2_s0, 2_s1, 2_s2]
操作后 rank0:[0_s0, 1_s0, 2_s0] # 第 i 项收自 rank i,内容是「发给我 rank0 的」
rank1:[0_s1, 1_s1, 2_s1]
rank2:[0_s2, 1_s2, 2_s2]
即:发送矩阵 M[i][j] = rank i 发给 rank j 的数据,操作后 rank j 收到的是 M[:, j] 这一列——转置。
2. 底层原理与算法
- 通信量:每个 rank 发出
(N-1)/N·S(给别人的部分),收入同量。是所有集合通信里最「重」的——因为每对 rank 之间都有独立数据流。 - 拓扑敏感度极高:N×N 的全连接流量,对互联带宽、对分带宽(bisection bandwidth)要求极高。这也是为什么 MoE 对超节点的高带宽全互联如此依赖。
- 变长(v)版本:需要先交换各段长度(通常先一次小 All2All 交换 count),再按长度做数据交换。
3. 组合关系
All2All可由 N 轮SendRecv(或Scatter)拼成:第 k 轮,每个 rank i 把「发给 rank (i+k)%N 的块」发出去。- 和
AllGather的区别已在上面点明:是否每个 rank 发送不同数据。
4. 应用场景
- MoE 专家并行(EP)的 token 路由(dispatch / combine),用的是变长 AllToAllv:
- dispatch:每个 rank 上的 token 经过 gating 要被送到它选中的专家所在的 rank。由于每个 token 的选择不同,每个 rank 发给各专家 rank 的 token 数量不等 → 必须用变长 AllToAllv。
- combine:专家算完后,把结果按原路由送回 token 原来的 rank,又是一次(反向的)AllToAllv。
- 为什么是变长:路由是数据相关的,负载天然不均衡——专家 A 可能收到 100 个 token,专家 B 只收到 10 个。等长 All2All 无法表达这种不均衡。
5. 实例代码
# all_to_all.py
# 启动:torchrun --nproc_per_node=4 all_to_all.py
import torch
import torch.distributed as dist
def main():
dist.init_process_group(backend="nccl" if torch.cuda.is_available() else "gloo")
rank = dist.get_rank()
world = dist.get_world_size()
device = torch.device(f"cuda:{rank}") if torch.cuda.is_available() else torch.device("cpu")
# ---------- 等长 All2All ----------
# rank i 的第 j 个元素 = i*10 + j,准备发给 rank j
send = torch.arange(world, device=device, dtype=torch.float32) + rank * 10
recv = torch.zeros(world, device=device)
print(f"[rank{rank}] 等长 操作前 send = {send.tolist()}")
dist.all_to_all_single(recv, send) # 等长全交换(本质是转置)
print(f"[rank{rank}] 等长 操作后 recv = {recv.tolist()}")
# ---------- 变长 AllToAllv(MoE 真实场景)----------
# 构造各 rank 发给各 rank 的「不等长」token 数:发给 rank j 的个数 = (rank + j) % 3 + 1
out_splits = [((rank + j) % 3 + 1) for j in range(world)]
# 对端要收多少,需要先交换 split(这里直接按规则算出 in_splits)
in_splits = [((j + rank) % 3 + 1) for j in range(world)] # rank j 发给我 rank 的个数
send_v = torch.full((sum(out_splits),), float(rank), device=device)
recv_v = torch.zeros(sum(in_splits), device=device)
print(f"[rank{rank}] 变长 发送量={out_splits} 接收量={in_splits}")
dist.all_to_all_single(recv_v, send_v,
output_split_sizes=in_splits,
input_split_sizes=out_splits)
print(f"[rank{rank}] 变长 操作后 recv_v = {recv_v.tolist()} (值=来源rank)")
dist.destroy_process_group()
if __name__ == "__main__":
main()
实测执行(本机无 GPU,走 gloo 后端):
torchrun --nproc_per_node=4 --rdzv-endpoint=127.0.0.1:29511 all_to_all.py
真实输出(按 rank 整理,打印顺序可能交错):
# 等长 All2All(本质是转置)
[rank0] 等长 操作前 send = [0.0, 1.0, 2.0, 3.0] 操作后 recv = [0.0, 10.0, 20.0, 30.0]
[rank1] 等长 操作前 send = [10.0, 11.0, 12.0, 13.0] 操作后 recv = [1.0, 11.0, 21.0, 31.0]
[rank2] 等长 操作前 send = [20.0, 21.0, 22.0, 23.0] 操作后 recv = [2.0, 12.0, 22.0, 32.0]
[rank3] 等长 操作前 send = [30.0, 31.0, 32.0, 33.0] 操作后 recv = [3.0, 13.0, 23.0, 33.0]
# 变长 AllToAllv(recv_v 的值 = 来源 rank,个数由 in_splits 决定)
[rank0] 发送量=[1,2,3,1] 接收量=[1,2,3,1] recv_v = [0, 1,1, 2,2,2, 3]
[rank1] 发送量=[2,3,1,2] 接收量=[2,3,1,2] recv_v = [0,0, 1,1,1, 2, 3,3]
[rank2] 发送量=[3,1,2,3] 接收量=[3,1,2,3] recv_v = [0,0,0, 1, 2,2, 3,3,3]
[rank3] 发送量=[1,2,3,1] 接收量=[1,2,3,1] recv_v = [0, 1,1, 2,2,2, 3]
等长版把发送矩阵「转置」:rank j 收到的第 i 项是「rank i 发给它的」。变长版里 rank0 收到 1 个来自 r0、2 个来自 r1、3 个来自 r2、1 个来自 r3 → [0,1,1,2,2,2,3],个数与 in_splits 完全对应,正是 MoE token 路由的负载不均衡形态。
变长版本的关键:
input_split_sizes/output_split_sizes分别告诉库「我发给每个 rank 多少、从每个 rank 收多少」。真实 MoE 里,这两个 split 本身也要先用一次小 All2All 交换 token 计数才能得到。
3. 加餐:用 Send/Recv 从零实现 Ring-AllReduce
前面反复说「集合通信底层是点对点拼的」「AllReduce = ReduceScatter + AllGather」。下面这段教学代码只用 batch_isend_irecv(即 SendRecv) 从零实现 Ring-AllReduce,并和 dist.all_reduce 对拍验证。读懂它,就真正打通了点对点与集合通信的任督二脉。
# ring_allreduce_from_scratch.py
# 启动:torchrun --nproc_per_node=4 ring_allreduce_from_scratch.py
import torch
import torch.distributed as dist
def ring_all_reduce(tensor: torch.Tensor):
"""只用 SendRecv 实现 Ring-AllReduce(SUM)。
分两阶段:ReduceScatter(N-1 步) + AllGather(N-1 步)。
要求 tensor 长度能被 world 整除。"""
rank = dist.get_rank()
world = dist.get_world_size()
left = (rank - 1 + world) % world # 从左邻居收
right = (rank + 1) % world # 发给右邻居
chunks = list(tensor.chunk(world)) # 切成 world 块(view,原地可改)
def sendrecv(send_chunk, recv_like):
"""发一块给右邻居,同时从左邻居收一块(SendRecv,不死锁)。"""
recv_buf = torch.empty_like(recv_like)
ops = [dist.P2POp(dist.isend, send_chunk.contiguous(), right),
dist.P2POp(dist.irecv, recv_buf, left)]
for w in dist.batch_isend_irecv(ops):
w.wait()
return recv_buf
# ---------- 阶段一:ReduceScatter ----------
# 第 step 步,发出块索引 (rank-step),把收到的累加到块索引 (rank-step-1)
for step in range(world - 1):
send_idx = (rank - step + world) % world
recv_idx = (rank - step - 1 + world) % world
recv_buf = sendrecv(chunks[send_idx], chunks[recv_idx])
chunks[recv_idx] += recv_buf # 归约:累加
# 此时 chunks[(rank+1)%world] 持有「该块的全局和」
# ---------- 阶段二:AllGather ----------
# 把已归约好的块沿环转一圈,覆盖式填满每个 rank
for step in range(world - 1):
send_idx = (rank + 1 - step + world) % world
recv_idx = (rank - step + world) % world
recv_buf = sendrecv(chunks[send_idx], chunks[recv_idx])
chunks[recv_idx].copy_(recv_buf) # 覆盖:拷贝(不是累加)
def main():
dist.init_process_group(backend="nccl" if torch.cuda.is_available() else "gloo")
rank = dist.get_rank()
world = dist.get_world_size()
device = torch.device(f"cuda:{rank}") if torch.cuda.is_available() else torch.device("cpu")
n = world * 2 # 长度能被 world 整除
x = torch.arange(n, dtype=torch.float32, device=device) + rank * 100
ref = x.clone()
dist.all_reduce(ref, op=dist.ReduceOp.SUM) # 官方实现作为参照
mine = x.clone()
ring_all_reduce(mine) # 我们从零实现的
ok = torch.allclose(mine, ref)
print(f"[rank{rank}] 自实现 == 官方AllReduce ? {ok}")
if rank == 0:
print(f"[rank0] 官方结果 = {ref.tolist()}")
print(f"[rank0] 自实现结果 = {mine.tolist()}")
dist.destroy_process_group()
if __name__ == "__main__":
main()
跑通后每个 rank 都应打印 == 官方AllReduce ? True。这段代码把「2(N-1) 轮 SendRecv」的内部过程完全摊开:前 N-1 轮边传边加(ReduceScatter),后 N-1 轮边传边盖(AllGather)。
实测执行(本机无 GPU,走 gloo 后端):
torchrun --nproc_per_node=4 --rdzv-endpoint=127.0.0.1:29512 ring_allreduce_from_scratch.py
真实输出(按 rank 整理):
[rank0] 自实现 == 官方AllReduce ? True
[rank1] 自实现 == 官方AllReduce ? True
[rank2] 自实现 == 官方AllReduce ? True
[rank3] 自实现 == 官方AllReduce ? True
# 官方结果与自实现结果完全一致:
[600.0, 604.0, 608.0, 612.0, 616.0, 620.0, 624.0, 628.0]
四个 rank 全部 True——只用 batch_isend_irecv(SendRecv)拼出的 Ring-AllReduce,结果与官方 dist.all_reduce 逐元素一致(每个元素 = 4i + 600,正是 4 个 rank 对应位置之和)。至此「集合通信底层由点对点拼装」得到实证。
小结
一条自底向上的逻辑链:点对点是积木 → 拼成集合通信 → 服务于并行策略。
几个要记住的结论:
- 集合通信都是点对点拼的。Ring 类算法的每一步就是一次
SendRecv;AllReduce = ReduceScatter + AllGather。 - Ring-AllReduce 带宽最优:单卡通信量
2(N-1)/N·S → 2S,几乎与卡数无关,这是它能扩展到千卡的根本;小数据时 Tree 延迟更优,NCCL 自动择优。 - 原语各司其职:DP 用
AllReduce+AVG;TP 用AllReduce+SUM(层内,仅机内);FSDP 用AllGather+ReduceScatter;PP 用Send/Recv;MoE 用变长AllToAllv;初始化用Broadcast;对齐/调试用Barrier。 - All2All 是最重的通信,N×N 全连接流量,对带宽/对分带宽要求极高——这也正是 MoE 场景强烈依赖超节点高带宽全互联的原因。
把这组原语的语义、复杂度、组合关系吃透,再看任何分布式训练框架(Megatron / DeepSpeed / FSDP / vLLM)的通信代码,都只是这几块积木的不同拼法而已。