在大模型分布式训练中,不同的并行策略依赖于不同的集合通信原语来同步数据和模型。All-Reduce 和 All-Gather 是其中最重要的两种。
下面我们来详细拆解这些通信方式。
核心概念:为什么需要通信?
无论是哪种并行方式,其根本目的都是将工作和数据分布到多个设备(GPU/NPU)上。为了确保最终结果的正确性,这些设备之间必须进行通信,以同步梯度、权重或激活值。
主要的通信原语
1. All-Reduce(全规约)
这是数据并行中最为核心的通信操作。
- 目标:将所有设备上的张量进行某种操作(如求和、求平均),并将最终结果同步到所有设备上。
-
过程:
- 每个设备都有一个初始张量(例如,梯度)。
- 通过通信,对所有设备的张量进行规约操作(通常是求和)。
- 求和后的最终结果被发送到每一个参与的设备上。
- 比喻:一个小组开会,每个人都有自己的意见(本地梯度)。经过讨论后,大家最终达成了一致的结论(全局平均梯度),并且每个人都知道了这个一致的结论。
-
在数据并行中的应用:
- 每个GPU在本地计算完梯度后。
- 使用
All-Reduce对所有GPU的梯度进行求和或求平均。 - 每个GPU都得到了全局一致的梯度,然后各自独立地更新本地模型权重。
2. All-Gather(全收集)
这在模型并行、张量并行和 ZeRO 优化中非常常见。
- 目标:每个设备都拥有一部分数据,通过通信后,每个设备都拥有全部数据的完整集合。
-
过程:
- 假设有4个设备,每个设备有一个不同的张量
A,B,C,D。 - 执行
All-Gather后,每个设备都拥有[A, B, C, D]这个完整的列表。
- 假设有4个设备,每个设备有一个不同的张量
- 比喻:拼图游戏。开始时每个人手上有几块拼图,通过交换信息后,每个人都拥有了一副完整的拼图照片。
-
在模型并行中的应用:
- 在反向传播中,计算某一层的梯度可能需要下一层的梯度作为输入。
- 如果下一层被切分到了多个设备上,就需要使用
All-Gather来在某个设备上重构出完整的梯度。
3. Reduce-Scatter(规约散射)
这在 ZeRO 优化的第1阶段(优化器状态分区)和模型并行中会用到。
-
目标:
Reduce-Scatter可以看作是All-Reduce的“另一半”。它先将所有设备上的张量进行规约(如求和),然后将求和后的完整结果按照设备数量进行切分,每个设备只保留其中一块。 -
过程:
- 每个设备有一个张量。
- 对所有张量进行规约操作(求和),得到一个完整的和。
- 将这个“和”切分成N份,每个设备只得到其中一份。
- 比喻:一个小组要完成一份报告,每个人写一部分。大家先把各自写的部分汇总起来(Reduce),然后每个人负责校对和修改其中一部分(Scatter)。
-
在ZeRO-1中的应用:
- 在梯度计算完成后,使用
Reduce-Scatter操作,让每个设备只负责更新一部分参数的梯度。 - 这样,每个设备也只需要保存一部分优化器状态,极大地节省了内存。
- 在梯度计算完成后,使用
4. Broadcast(广播)
这是一个基础操作。
- 目标:将一个设备(通常是 rank 0)上的张量复制到所有其他设备上。
- 过程:一个源设备发送数据,所有目标设备接收相同的数据。
- 应用:初始化模型权重时,将随机种子在 rank 0 生成的初始权重广播到所有设备;或在推理时,将输入提示词广播到所有设备。
通信方式与并行策略的对应关系
| 并行策略 | 主要通信原语 | 通信发生时机与内容 |
|---|---|---|
| 数据并行 | All-Reduce | 反向传播结束后,同步所有设备的梯度。 |
| 张量并行 | All-Reduce 或 All-Gather | 在前向传播和反向传播中都需要通信,用于同步部分计算结果(如激活值)和梯度。具体操作取决于模型层的类型。 |
| 流水线并行 | 点对点通信 | 在流水线的不同阶段(设备)之间,前向传播时传递激活值,反向传播时传递梯度。这更像是一个传送带,而不是全局同步。 |
| ZeRO 优化 | Reduce-Scatter + All-Gather | <ul><li>梯度计算后:使用 Reduce-Scatter 对梯度进行分区。</li><li>参数更新前:使用 All-Gather 在设备间重建完整的参数。</li></ul> |
高级理解:All-Reduce 的实现
你甚至可以认为,一个高效的 All-Reduce 是由 Reduce-Scatter 和 All-Gather 两个阶段组合而成的:
- Reduce-Scatter 阶段:将所有设备的张量求和,但结果被切分,每个设备保留一部分。
- All-Gather 阶段:每个设备将自己持有的那一部分结果分享给所有其他设备。
这样组合起来,最终每个设备都拥有了所有部分的求和结果,即完成了 All-Reduce。
总结
-
All-Reduce: 大家都有数据,合并成一个结果,然后人手一份。 (数据并行的核心) -
All-Gather: 每人有一块碎片,最后人手一份完整的拼图。 (模型并行、ZeRO的核心) -
Reduce-Scatter: 大家都有数据,合并成一个结果,然后每人只拿一块。 (ZeRO的核心) -
Broadcast: 我有一份数据,复制给你们所有人。
理解这些集合通信原语,是深入理解大模型如何被分布式训练和优化的钥匙。它们共同协作,确保了分布在成千上万个计算单元上的模型部件能够高效、一致地工作。