人人都会AI编程

9.7 集合通信原语:AllReduce、AllGather、ReduceScatter 的作用

更新时间:2026-07-09

经过 9.2 到 9.6 节的讨论,你已经看到:无论是数据并行(DP)做梯度同步,张量并行(TP)做激活值拼接,还是 ZeRO 做参数分片与梯度聚合,本质上都是把同一份计算任务拆到多张 GPU 上,再在特定的时机把数据拼回来或统一起来。这种“拼”与“统一”不是自动发生的,它依赖底层一组高度标准化的集合通信原语(Collective Communication Primitives)

如果把分布式训练比作一场多人协作的流水线,那么并行策略(DP/TP/PP/ZeRO)是“车间管理制度”,而 AllReduce、AllGather、ReduceScatter 就是车间里真正搬运货物的叉车与传送带。理解这三个原语,是读懂训练日志里 ncclAllReduce 耗时、定位通信瓶颈、以及合理设计混合并行方案的最后一环。


一、什么是集合通信原语?

在分布式训练中,每个 GPU 通常对应一个独立的计算进程(如 PyTorch DDP 中的一个 Rank)。进程之间的通信分两类:

  • 点对点(P2P):如流水线并行中,前一个 Stage 的 GPU 直接把激活值发给下一个 Stage 的 GPU。特点是只涉及两个节点。
  • 集合通信(Collective):涉及一组节点(如一个通信组内的所有 Rank),大家按照约定同时参与、同时完成。特点是同步性全局性

当前业界几乎所有分布式训练框架(PyTorch DDP、DeepSpeed、Megatron-LM)的底层,都是通过 NCCL(NVIDIA Collective Communications Library) 来调用这些原语,而 NCCL 再映射到 NVLink、PCIe 或 InfiniBand 等物理链路上(详见第 16 章)。

接下来我们聚焦三个最核心的原语。理解它们的最佳方式不是死记公式,而是记住一个灵魂问题:操作完成后,每个节点手里的结果是什么?


二、AllReduce:全员统一,人人有份

定义:通信组内每个节点先贡献自己的本地数据,经过全局归约(通常是求和或求平均)后,每个节点都拿到完全一致的结果

直观类比:4 个小组员分别统计了本季度的销售额,AllReduce 就是大家把数字报给班长汇总、求平均,然后班长把最终平均数同时抄写给每一个人。最后,4 个人的笔记本上写着同一个数字。

在分布式训练中的典型场景

  1. 数据并行(DP)的梯度同步:每张卡算完反向传播后,持有本地梯度 g_i。经过 AllReduce(Sum)后,每张卡都得到全局梯度总和 G = g_1 + g_2 + ... + g_N;如果再除以 N,就是平均梯度。至此,所有卡才能同步进行下一步参数更新。
  2. 张量并行(TP)的激活值同步:在某些 TP 实现中,注意力或 MLP 的前向传播需要对不同卡上的中间结果做归约,确保后续计算基于完整信息。

通信特征
Ring-AllReduce(NCCL 默认算法)的通信量约为 2M(M 为单卡数据量),与节点数 N 基本无关。这是 DP 能 Scale 到千卡级的核心原因之一——通信成本不随卡数线性爆炸。


三、AllGather:各献一块,拼成整图

定义:通信组内每个节点原本只持有一块不同的局部数据,AllGather 操作后,每个节点都拿到所有人数据的完整拼接结果

直观类比:4 个人各持有一张拼图的四分之一。AllGather 完成后,每个人都拿到了完整拼好的大图。注意:最终每个人的手里的图是一模一样的,且是“拼接”而非“叠加”。

在分布式训练中的典型场景

  1. 张量并行(TP)的列切分还原

在 Megatron-LM 的 MLP 实现中,第一层线性层按列切分到多张卡,每张卡算出自己那一列的输出;为了进入下一层,必须通过 AllGather 把这些列按顺序拼成完整的激活矩阵,否则下一层无法继续计算。

  1. ZeRO-Inference / ZeRO-3 的参数收集

当模型参数被分片存储在不同卡上时,前向传播需要先用 AllGather 把当前层所需的完整参数从各卡收集到本地,才能做计算。

通信特征
通信量约为 (N-1)/N × M ≈ M(随 N 增大趋近于 M)。因为每个人都要接收来自其他 N-1 个人的数据,所以数据量会膨胀。在 TP 中,AllGather 的通信开销是限制 TP 并行度不能太大的重要因素——通常 TP 只在一个节点内的 NVLink 高速互联域(如 8 卡)中使用,而很少跨节点。


四、ReduceScatter:先汇总,再分家

定义:通信组内每个节点先贡献自己的本地数据,全局做归约(如求和),然后将归约后的结果按指定维度切开,每个节点只保留自己对应的那一块

直观类比:4 个会计各自做了一本账。ReduceScatter 先把 4 本账的对应科目加总,然后把加总后的完整账本撕成 4 份,每人只领走自己负责的那几页。最终,4 个人手里的东西加起来才是完整结果,但单个人只有局部

在分布式训练中的典型场景

  1. ZeRO-1/2 的梯度分片

每张卡算完本地梯度后,通过 ReduceScatter 对梯度做全局求和并按卡分片。最终 Rank 0 只保存第 0 块梯度,Rank 1 只保存第 1 块……这样每张卡只需存储 1/N 的优化器状态,从而大幅节省显存。

  1. 与 AllGather 的“黄金搭档”

ReduceScatter 和 AllGather 在数学上互补。注意到一个恒等式:

   AllReduce = ReduceScatter + AllGather
   

ZeRO 正是利用这一点,把传统 DP 中“全员存完整梯度”的 AllReduce,拆成了“先分片汇总,再按需收集”的两阶段操作,把显存压力从 O(N) 降到 O(1/N)。

通信特征
ReduceScatter 的通信量与 AllReduce 同级,约为 2M。它的价值不在于省带宽,而在于改变了数据存放的形态——从“每人全量”变为“每人一份”。


五、实战视角:在你的训练框架里,它们藏在哪儿?

如果你去翻 PyTorch DDP 或 DeepSpeed 的源码,会发现这些原语被封装得很好,但它们留下的痕迹无处不在:

| 并行策略 | 你看到的上层操作 | 底层通信原语 | 目的 |
|---------|-----------------|-------------|------|
| 数据并行(DP) | loss.backward() + optimizer.step() | AllReduce(梯度同步) | 保证各卡参数更新一致 |
| 张量并行(TP) | RowParallelLinear / ColumnParallelLinear | AllReduce(行切分反向)、AllGather(列切分前向拼接) | 拼接激活值、同步梯度 |
| ZeRO-1/2 | deepspeed.backward() | ReduceScatter(梯度分片归约) | 将完整梯度拆成 shard |
| ZeRO-3 | forward() 前参数收集 | AllGather(参数收集) | 按需把分片参数 assemble 到本地 |

性能调优提示
如果你在 NVIDIA Nsight Systems 或 PyTorch Profiler 里看到某个通信原语耗时占比超过 30%,通常意味着:

  • AllReduce 太慢:可能是卡间互联带宽不足(如误把 DP 组跨节点放在 PCIe 上),或梯度太大(考虑梯度累积、混合精度)。
  • AllGather 频繁:可能是 TP 切分过细,或 ZeRO-3 的参数收集粒度太碎(可尝试 param_persistence_threshold 调优)。
  • ReduceScatter 阻塞:往往伴随显存碎片,导致通信与计算无法重叠(overlap)。

六、小结与速查

集合通信原语是分布式训练的“物理定律”。再精妙的并行策略,最终都要落实到这三种基础操作上:

  • AllReduce:人人出数据,人人得完整且相同的归约结果。

定位:梯度同步的基石。

  • AllGather:人人出不同数据块,人人得完整拼接大图。

定位:激活值/参数拼接的利器。

  • ReduceScatter:人人出数据,全局归约后各领一块

定位:显存优化的关键拆分手段。

记住那个恒等式 AllReduce = ReduceScatter + AllGather,你就理解了为什么 ZeRO 能在不增加通信总量的前提下显著节省显存——它只是把一次“全员存全量”的操作,拆解成了“先分片归约、再按需收集”的两个阶段。

在 9.1 到 9.6 节的并行策略与第 16 章的 GPU 互联技术之间,集合通信原语正是承上启下的桥梁。下一步,当你在设计千卡集群的并行方案时,请带着这一节的视角去审视:你的 DP、TP、PP 组合,最终会在物理链路上产生多少 AllReduce 和 AllGather?它们是否能被 NVLink 的带宽消化?这决定了你的 MFU(模型算力利用率)是 50% 还是 80%。