大模型训练中的集合通信

在大模型分布式训练中,不同的并行策略依赖于不同的集合通信原语来同步数据和模型。All-ReduceAll-Gather 是其中最重要的两种。

下面我们来详细拆解这些通信方式。


核心概念:为什么需要通信?

无论是哪种并行方式,其根本目的都是将工作和数据分布到多个设备(GPU/NPU)上。为了确保最终结果的正确性,这些设备之间必须进行通信,以同步梯度、权重或激活值。

主要的通信原语

1. All-Reduce(全规约)

这是数据并行中最为核心的通信操作。

  • 目标:将所有设备上的张量进行某种操作(如求和、求平均),并将最终结果同步到所有设备上。
  • 过程
    1. 每个设备都有一个初始张量(例如,梯度)。
    2. 通过通信,对所有设备的张量进行规约操作(通常是求和)。
    3. 求和后的最终结果被发送到每一个参与的设备上。
  • 比喻:一个小组开会,每个人都有自己的意见(本地梯度)。经过讨论后,大家最终达成了一致的结论(全局平均梯度),并且每个人都知道了这个一致的结论。
  • 在数据并行中的应用
    • 每个GPU在本地计算完梯度后。
    • 使用 All-Reduce 对所有GPU的梯度进行求和或求平均
    • 每个GPU都得到了全局一致的梯度,然后各自独立地更新本地模型权重。

2. All-Gather(全收集)

这在模型并行张量并行ZeRO 优化中非常常见。

  • 目标:每个设备都拥有一部分数据,通过通信后,每个设备都拥有全部数据的完整集合
  • 过程
    1. 假设有4个设备,每个设备有一个不同的张量 A, B, C, D
    2. 执行 All-Gather 后,每个设备都拥有 [A, B, C, D] 这个完整的列表。
  • 比喻:拼图游戏。开始时每个人手上有几块拼图,通过交换信息后,每个人都拥有了一副完整的拼图照片。
  • 在模型并行中的应用
    • 在反向传播中,计算某一层的梯度可能需要下一层的梯度作为输入。
    • 如果下一层被切分到了多个设备上,就需要使用 All-Gather 来在某个设备上重构出完整的梯度。

3. Reduce-Scatter(规约散射)

这在 ZeRO 优化的第1阶段(优化器状态分区)和模型并行中会用到。

  • 目标Reduce-Scatter 可以看作是 All-Reduce 的“另一半”。它先将所有设备上的张量进行规约(如求和),然后将求和后的完整结果按照设备数量进行切分,每个设备只保留其中一块
  • 过程
    1. 每个设备有一个张量。
    2. 对所有张量进行规约操作(求和),得到一个完整的和。
    3. 将这个“和”切分成N份,每个设备只得到其中一份。
  • 比喻:一个小组要完成一份报告,每个人写一部分。大家先把各自写的部分汇总起来(Reduce),然后每个人负责校对和修改其中一部分(Scatter)。
  • 在ZeRO-1中的应用
    • 在梯度计算完成后,使用 Reduce-Scatter 操作,让每个设备只负责更新一部分参数的梯度。
    • 这样,每个设备也只需要保存一部分优化器状态,极大地节省了内存。

4. Broadcast(广播)

这是一个基础操作。

  • 目标:将一个设备(通常是 rank 0)上的张量复制到所有其他设备上。
  • 过程:一个源设备发送数据,所有目标设备接收相同的数据。
  • 应用:初始化模型权重时,将随机种子在 rank 0 生成的初始权重广播到所有设备;或在推理时,将输入提示词广播到所有设备。

通信方式与并行策略的对应关系

并行策略 主要通信原语 通信发生时机与内容
数据并行 All-Reduce 反向传播结束后,同步所有设备的梯度
张量并行 All-ReduceAll-Gather 前向传播反向传播中都需要通信,用于同步部分计算结果(如激活值)和梯度。具体操作取决于模型层的类型。
流水线并行 点对点通信 在流水线的不同阶段(设备)之间,前向传播时传递激活值,反向传播时传递梯度。这更像是一个传送带,而不是全局同步。
ZeRO 优化 Reduce-Scatter + All-Gather <ul><li>梯度计算后:使用 Reduce-Scatter 对梯度进行分区。</li><li>参数更新前:使用 All-Gather 在设备间重建完整的参数。</li></ul>

高级理解:All-Reduce 的实现

你甚至可以认为,一个高效的 All-Reduce 是由 Reduce-ScatterAll-Gather 两个阶段组合而成的:

  1. Reduce-Scatter 阶段:将所有设备的张量求和,但结果被切分,每个设备保留一部分。
  2. All-Gather 阶段:每个设备将自己持有的那一部分结果分享给所有其他设备。

这样组合起来,最终每个设备都拥有了所有部分的求和结果,即完成了 All-Reduce

总结

  • All-Reduce大家都有数据,合并成一个结果,然后人手一份。 (数据并行的核心)
  • All-Gather每人有一块碎片,最后人手一份完整的拼图。 (模型并行、ZeRO的核心)
  • Reduce-Scatter大家都有数据,合并成一个结果,然后每人只拿一块。 (ZeRO的核心)
  • Broadcast我有一份数据,复制给你们所有人。

理解这些集合通信原语,是深入理解大模型如何被分布式训练和优化的钥匙。它们共同协作,确保了分布在成千上万个计算单元上的模型部件能够高效、一致地工作。

©著作权归作者所有,转载或内容合作请联系作者
【社区内容提示】社区部分内容疑似由AI辅助生成,浏览时请结合常识与多方信息审慎甄别。
平台声明:文章内容(如有图片或视频亦包括在内)由作者上传并发布,文章内容仅代表作者本人观点,简书系信息发布平台,仅提供信息存储服务。

相关阅读更多精彩内容

友情链接更多精彩内容