MeshSlice: Efficient 2D Tensor Parallelism for Distributed DNN Training

发表时间: 2025-06

速读

一句话结论 本文提出了用于大规模深度神经网络分布式训练的高效 2D 张量并行算法 MeshSlice 及其自动调优器,通过将集体通信切片并在行列双维度上实现通信与计算的完全重叠,在 256 芯片集群上将 GPT-3 和 Megatron-NLG 的训练速度比现有最优基线提升了 12.0% 和 23.4%。

要解决什么问题 大规模语言模型训练严重依赖张量并行(TP),但传统的 1D TP 通信量随芯片数线性增长,需要通信带宽呈二次方扩展,因此通常被限制在 8 路。将矩阵分片到 2D 芯片网格的 2D TP 具有更好的扩展性,但其底层的 2D 通用矩阵乘法(GeMM)算法存在致命卡点。具体而言:Cannon 算法需要方形网格且伴随额外的倾斜移位操作,通信流量极大;SUMMA 算法采用细粒度的广播和归约,流水线气泡多,同步开销随迭代次数呈 $O(K^2)$ 增长,在大规模网格中效率极低;Collective 2D GeMM 算法改用全收集(AllGather)和归约散布(ReduceScatter)等集体通信操作,带宽利用率极高,但由于真实数据依赖,完全无法将通信与计算重叠;Wang 算法则只能在单一维度上拆分集体通信,另一个维度的通信依然无法重叠。此外,2D TP 涉及数据流、网格形状、分片方式等众多参数,搜索空间庞大,极度依赖专家手动试错调优。

怎么做的 MeshSlice 的核心思路是将高效的集体通信划分为多个部分的集体操作,通过软件流水线,在行列两个维度上同时实现通信与计算的重叠。关键设计由两部分构成。第一部分是 MeshSlice 2D GeMM 算法。它不执行单次庞大的集体通信,而是将本地矩阵分片沿特定维度切分为 $K$ 个子分片。在一个包含 $K$ 次迭代的循环中,每次迭代仅对一个子分片执行部分的 AllGather 或 ReduceScatter,并计算部分的 GeMM。这样,当前迭代的部分 GeMM 计算就可以与后续迭代的通信及切片操作完美重叠。为了避免切片导致的非连续内存访问,算法引入了分块切片机制,将内存访问对齐到硬件缓存行大小。第二部分是 MeshSlice LLM 自动调优器,用于替代人工配置。它分为两个阶段:首先,通过启发式规则为前向和反向传播选择能让最大矩阵保持固定的数据流,从而最小化通信流量,这同时也自动确定了张量的分片方式;其次,利用分析成本模型协同优化网格形状和切片数 $K$。集体通信的成本模型定义为: $$T_{comm} = T_{launch} + (P - 1) \times L_{sync} + \frac{size(shard)}{BW}$$ 其中 $T_{launch}$ 为启动开销,$P$ 为行或列的芯片数,$L_{sync}$ 为同步延迟,$BW$ 为测得的链路带宽。自动调优器将总执行时间拆解为前导时间、稳态时间与收尾时间,通过穷举搜索极小的参数空间,在几秒内即可输出最优配置。

效果如何 实验在模拟的 Google TPUv4 集群(最高 256 芯片)上评估了 175B 参数的 GPT-3 和 530B 参数的 Megatron-NLG 模型,序列长度设为 2048。对比基线包括 2D 路线的 Cannon、SUMMA、Collective 2D GeMM 以及代表当前 SOTA 的 Wang 算法,同时对比了 1D 路线的 1D TP 和完全分片数据并行(FSDP)。在 256 芯片的弱扩展实验中,MeshSlice 表现出最高的计算利用率,其端到端训练速度比 Wang 算法在 GPT-3 和 Megatron-NLG 上分别快 12.0% 和 23.4%。从 16 路扩展到 256 路时,MeshSlice 的效率仅分别下降 16.8% 和 5.8%,而 SUMMA 和 1D 方法在大规模下效率发生崩塌。自动调优器也证明了其价值,仅数据流优化一项就为 GPT-3 带来了 21.2% 的性能提升,且成本模型预测的最优网格形状和切片数与模拟器穷举结果完全一致。该方法的局限性在于对底层硬件的异步通信支持有要求。在真实的 4x4 TPUv4 集群上测试时,由于当前 TPUv4 尚不支持 AllGather 和 ReduceScatter 与计算的底层异步重叠,MeshSlice 因为更细粒度的操作开销,比 Collective 算法慢了 4.5%(切片操作本身开销仅占 1.3%)。但作者预估,一旦硬件支持重叠,MeshSlice 将在真实集群上比 Collective 算法快 32.8% 以上。

A1 主要贡献

本文旨在解决大规模DNN模型分布式训练中张量并行(TP)的通信瓶颈问题。现有的1D TP因通信成本高而可扩展性有限,而2D TP虽然能通过将矩阵分片到2D加速器网格中来减少通信,但其核心的通用矩阵乘法(GeMM)算法存在效率问题。具体来说,Cannon算法通信流量大;SUMMA算法同步开销高;而使用集体通信操作的2D GeMM无法将通信与计算重叠。此外,优化2D TP的众多参数(如数据流、网格形状、分片方式)非常困难,通常需要专家手动配置。

为应对这些挑战,本文做出了以下核心贡献:

  1. 提出新颖的MeshSlice算法:这是一种为分布式DNN训练中的2D TP设计的高效2D GeM M算法。MeshSlice通过将AllGather (AG) / ReduceScatter (RdS) 等集体通信操作切分为多个部分的集体操作,从而实现了通信与计算的重叠。这种方法有效隐藏了大部分通信延迟,解决了现有算法无法在行列两个维度上同时实现重叠的问题。
  2. 开发MeshSlice LLM自动调优器(Autotuner):该工具能够自动为大型语言模型(LLM)的训练找到最优的2D TP配置。它首先选择一个高效的2D GeMM数据流,然后利用分析性成本模型协同优化网格形状和通信粒度,从而替代了繁琐的人工调优过程。
  3. 全面的评估与实现:通过模拟训练GPT-3和Megatron-NLG模型的TPUv4集群,本文验证了MeshSlice的性能。结果显示,MeshSlice在高达256路的2D TP中仍保持高效率。在一个256个TPU的集群中,MeshSlice训练GPT-3和Megatron-NLG模型的速度分别比现有最先进的算法快12.0%和23.4%。此外,本文还在真实的Google TPUv4集群上实现了MeshSlice,验证了其切片操作的开销很小,且自动调优器的成本模型能准确估算通信和计算成本。

A3 背景知识/关键Observation/设计原则

2.1 分布式训练方法

2.2 2D张量并行

2.3 2D GeMM 算法

2.3.1 通用方面

2.3.2 Cannon算法

2.3.3 SUMMA 算法

2.3.4 Collective 2D GeMM

A2 方法细节

本文为2D TP做出了两项贡献。首先,提出了一种新的2D GeMM算法,解决了现有2D GeMM算法的局限性。其次,设计了一个LLM自动调优器,为LLM训练找到一个高效的2D TP配置。该LLM自动调优器优化了数据流、网格形状和通信粒度的配置。

我们提出的2D GeMM算法称为MeshSlice。图4可视化了先前算法的时间线,并与MeshSlice进行了比较。该图显示了计算、行间通信和列间通信的时间进展。Cannon需要进行倾斜操作且只支持方形网格形状,因此其流量高于其他算法,增加了总执行时间。SUMMA使用低效的bcast/reduce通信操作,由于细粒度的数据包而产生流水线气泡和同步开销。Collective算法不将集体通信与计算重叠。Wang的算法只划分了一个方向上的集体通信,因此另一个方向上的通信没有被重叠。最后,MeshSlice能够在两个方向上都将通信与计算重叠,从而实现最快的执行速度。

图4:五种2D GeMM算法的时间线对比:Cannon、SUMMA、Collective、Wang和MeshSlice。
图4:五种2D GeMM算法的时间线对比:Cannon、SUMMA、Collective、Wang和MeshSlice。

3.1 MeshSlice 2D GeMM 算法

3.1.1 MeshSlice算法的数学描述

公式1
公式1

算法1
算法1

3.1.2 MeshSlice算法的详细实现

公式2
公式2

公式3
公式3

公式4
公式4

3.2 MeshSlice LLM 自动调优器

3.2.1 阶段1:数据流和分片

3.2.2 阶段2:网格形状和切片数

A4 实验环境

A5 实验结果

5.1 分布式GeMM算法性能

5.2 LLM自动调优器和成本模型

5.3 在真实硬件上的MeshSlice性能

A6 结论

本文提出了MeshSlice算法,一种为分布式DNN训练设计的高效2D张量并行方法。MeshSlice通过将通信操作切分为多个部分,并利用软件流水线在行列两个维度上都实现了通信与计算的高效重叠,从而解决了现有2D GeMM算法(如Cannon、SUMMA、Collective GeMM)存在的流量大、同步开销高或无法重叠等问题。此外,本文还设计了MeshSlice LLM自动调优器,该工具能够通过选择高效的数据流,并利用精确的成本模型协同优化加速器网格形状和通信粒度,从而自动化了复杂的性能调优过程。

在模拟的256个TPUv4集群上的评估表明,MeshSlice在训练GPT-3和Megatron-NLG模型时,端到端性能分别比当前最先进的算法快12.0%和23.4%。

未来的工作方向包括:
1. 扩展到GPU集群:通过在GPU集群的物理网络上构建逻辑网格,将MeshSlice应用于更广泛的硬件平台,并相应调整自动调优器以考虑网络竞争。
2. 应用于推理场景:调整MeshSlice及其自动调优器以适应推理任务中更可能出现的内存瓶颈。
3. 支持其他DNN层:将MeshSlice应用于可转换为GeMM操作的其他层,如卷积层,或用于优化GNN中的2D分布式稀疏GeMM。
4. 结合专家混合(MoE)模型:将MeshSlice的2D TP与MoE的专家并行(EP)相结合,以支持更大规模模型的训练。