并行化Transformer训练(39分钟阅读)
摘要
一个交互式、可探索的解释,介绍各种用于训练Transformer的并行化方案,改编自关于模型扩展的学术内容。
本交互式指南探讨了用于训练Transformer的数据并行、FSDP、张量并行、流水线并行和专家并行。它重点关注不同的硬件和通信模式如何决定每种策略何时成为瓶颈。
查看缓存全文
缓存时间: 2026/08/21 15:40
# 如何为训练并行化Transformer——可交互的解释说明
来源:https://ezyang.github.io/interactive-parallelize-transformer/
本文改编自Jacob Austin、Sholto Douglas、Roy Frostig、Anselm Levskaya、Charlie Chen、Sharad Vikram、Federico Lebron、Peter Choy、Vinay Ramasesh、Albert Webson & Reiner Pope (Google DeepMind)所著《How to Scale Your Model》第五章(https://jax-ml.github.io/scaling-book/training/)的一个*可交互*版本。
✦ 我们从最初的密集型TPU方案开始——数据并行、FSDP、张量并行、它们的混合形式以及管道并行——然后融入GPU集群模型和针对混合专家模型的专家并行。对于每种方案,我们都会探讨通信何时成为瓶颈。
(本摘要属于改编部分;该章节原有的摘要描述了其四种密集型方案。)
**本页面是一个可工作的模型,而非对模型的描述。**
所有*绿色数字*都可以左右拖动,或双击后输入精确值。
所有*蓝色数字*都是根据绿色数字实时计算的——可以在这里试试:拖动*批次*并观察*每芯片批次*如何跟随变化(将鼠标悬停在任何蓝色数字上可查看其公式)。
它们共享一个模型与硬件状态,因此任何地方的更改都会传播到各处。
并行度保持方案局部性:密集混合组使用 N = DP·TP,而专家并行和管道并行在各自部分建模;复合示例会明确声明其完整的乘积。
放心拖动吧——它会恢复所有被修改的数字至默认值,同时保留你的模型、硬件以及规格/测量选择(它与顶部栏中的按钮是同一个按钮,当有任何数字被修改时,该按钮会亮起橙色);任何单个数字在双击并确认留空时会单独恢复;浏览器的“后退”按钮可以浏览你之前的配置。
**你正在阅读谁的话?**
原始段落来自TPU和GPU章节(© 2022 Maruan Al-Shedivat, © 2025 Google LLC, MIT许可协议 (https://ezyang.github.io/interactive-parallelize-transformer/LICENSE-scaling-book.txt));AI创作的偏离内容有明确标签,遵循以下约定:凡是章节中印刷的固定数字,本页面都会实时计算(这些原位替换未单独标记);交互式图表及其说明文字替换了原始的静态图表;
✦ 边注以及明确标记为改编的段落是AI撰写的编辑性文字——初始版本由Fable (Anthropic)构建,本次对抗性审查及其更正由OpenAI Codex执行——包括说明、旁注以及新的屋顶线入门 (https://ezyang.github.io/interactive-parallelize-transformer/#roofline);
专家并行 (https://ezyang.github.io/interactive-parallelize-transformer/#expert-parallelism) 和GPU网络 (https://ezyang.github.io/interactive-parallelize-transformer/#gpus) 部分则整合了第12章的原始段落,其AI撰写的连接与改编文字用相同约定标记;
章节中单字母网格轴名称在全文中渲染为命名的并行度——其X轴是DP,其Y轴是TP,管道并行部分的Z轴是PP,第12章的专家轴Z轴是EP(全局替换;每个都是自己的可修改变量,在其部分使用的文本中调整);
在GPU预设下,硬件词汇相应变更——TPU→GPU,ICI→NVLink,DCN→InfiniBand,pod→node,MXU→tensor core——因此文章读起来像一台一致的机器,而任何TPPU预设都会恢复章节的确切措辞(那些刻意*比较*两者的句子从不替换);
由本版编入章节文本的内容带有虚线下划线(就像这样);其旁边的拼接编辑(标准引用实践允许的括号或省略号)未做标记;
为了容纳交互元素而必须修改的句子,会有一个Δ边注引用原文并说明修改;
当本版的附加内容使得章节中的陈述不准确时,会用斜体*(编注:...)*进行更正。
## 我们所说的“规模化”是什么意思?
“模型规模化”的目标是能够增加用于训练或推理的芯片数量,同时实现吞吐量成比例的线性增长(我们称之为*强扩展*)。虽然单芯片性能取决于内存带宽与浮点运算量之间的权衡,但集群级别的性能取决于通过将芯片间通信与有用的浮点运算重叠来隐藏它。这并非易事,因为增加芯片数量会增加通信负载,同时减少可用于隐藏通信的每设备计算量。正如我们在第3节(https://jax-ml.github.io/scaling-book/sharding/)所看到的,分片矩阵乘法通常需要昂贵的AllGather或ReduceScatter操作,这可能阻塞TPU进行有用的工作。本节的目的是找出这些操作何时变得*过于昂贵*。
在本节中,我们将讨论五种常见的并行方案:(纯)**数据并行**、**全分片数据并行**(FSDP / ZeRO 分片)、**张量并行**(也称为模型并行)、**专家并行**(用于混合专家模型),以及(简要提及)**管道并行**。对于每种方案,我们将展示产生的通信成本,以及该成本何时开始成为计算成本的瓶颈。我们将专注于通信界限——因为虽然内存容量约束很重要,但在预训练中使用重计算(激活检查点)和大量芯片时,它们通常不是限制因素。
(编注:本版扩展讨论了专家并行 (https://ezyang.github.io/interactive-parallelize-transformer/#expert-parallelism),这与原版不同。)
对于本节,你可以只关注芯片间通信成本,因为只要我们有足够大的单芯片批次大小,从HBM到MXU的数据传输已经与计算重叠。
我们将使用以下符号来简化本节中的计算。
实时数值显示:(与顶部栏对应)
| 符号 | 含义 | 实时值 |
| :--- | :--- | :--- |
| **模型参数** | | |
| *D* | **d**<sub>model</sub>(隐藏维度/残差流维度) | |
| *F* | **d**<sub>ff</sub>(前馈维度) | |
| | **F 约定(通用):** *一个专家*的宽度(密集时为 d<sub>ff</sub>);数学计算使用 k·*F*,权重存储 E·*F*,章节中的方程是 E = k = 1 的情况(第12章的解析)。一个实在的局限:混合了密集和MoE块的模型具有*两个*真正不同的F——DeepSeek-V3的前三层以更宽的宽度密集运行——而本页面将这类模型近似为统一的MoE。将鼠标悬停在任何*F*上查看实时宽度。 | |
| *B* | 批次维度(批次中的token数量;总数,非每设备) | |
| *T* | 序列长度 | |
| *L* | 模型中的层数 | |
| **硬件特性** | | |
| *C* | 每芯片FLOPS/s | |
| *W* | 网络带宽(TPU网格轴每轴双向/单向GPU或节点出口,常带下标如 *W<sub>ici</sub>* 或 *W<sub>dcn</sub>*) | |
| *ici·dcn* | | |
| *DP* | 沿数据并行网格轴的芯片数(章节的X轴) | |
| *TP* | 沿替代张量并行网格轴的芯片数(章节的Y轴) | |
| *Z* | 沿第三个网格轴的芯片数,标记为Z | |
| *PP* | 管道级数(管道部分的Z轴) | |
| *EP* | 专家并行度(第12章的Z轴;见专家并行部分) | |
✦ 改编——本符号被章节中的密集模型和当今的前沿开源模型所采用(支持的行可点击)
章节中的示例是密集的LLaMA时代模型;此后前沿已转向混合专家模型。
形状来自每个模型在Hugging Face上发布的`config.json`;参数总量来自其safetensors元数据。检索于2026年8月。
E和k计入共享专家,因此k·*F*是所表示架构的激活宽度;列标题解释了每个字段。
顶部栏下拉菜单中的密集模型为了对比列在表前,加载的模型行显示为实时绿色——在此处修改它。
| 模型 | 参数量 | D | F | 激活 k·F | L | E | k |
| :--- | :--- | :--- | :--- | :--- | :--- | :--- | :--- |
| (章节默认值) | 70.6B | 8,192 | 28,672 | 28,672 | 80 | 1 | 1 |
| | 13.0B | 5,120 | 13,824 | 13,824 | 40 | 1 | 1 |
| | 8.54B | 3,072 | 24,576 | 24,576 | 28 | 1 | 1 |
| DeepSeek-V3* | 685B | 7,168 | 2,048 | 18,432 | 61 | 257 | 8+1 |
| Kimi K3 (仅参考) | 2.78T | 7,168 | 3,072 | 55,296 | 93 | 896 | +2 16+2 |
| | 753B | 6,144 | 2,048 | 18,432 | 78 | 257 | 8+1 |
| | 1.60T | 7,168 | 3,072 | 21,504 | 61 | 385 | 6+1 |
| | 2.45T | 8,192 | 2,048 | 22,528 | 92 | 513 | 10+1 |
| | 952B | 6,144 | 3,072 | 24,576 | 66 | 258 | 6+2 |
| | 427B | 6,144 | 3,072 | 15,360 | 60 | 129 | 4+1 |
*计数示例:256个路由 + 1个共享专家 → E 257;Top-8 + 共享 → k 9。其前三层实际上是密集的(见上方F约定说明)。
K3不是实时预设,因为其路由专家在从残差D = 7,168投影到3,584宽度的潜在空间后运行。其路由专家中间宽度为F = 3,072。页面的单个D×F专家模型无法忠实表示两个维度。
点击一个支持的模型,将其形状(D, F, L, E, k)加载到页面的共享状态中(顶部栏随之变化);点击列标题进行排序。
F = 每专家宽度;激活k·F = 每token激活宽度;E / k = 总专家数 / 激活专家数,计入共享。在所有支持的实时MoE预设中,每专家F仅为2,048或3,072,尽管总参数量跨越数千亿到万亿,激活宽度k·F聚类在15k到25k之间。由于本章后面的张量并行界限随激活宽度k·F缩放,这种聚类就是为什么TP限制在所有支持的前沿预设中看起来如此相似。K3作为参考行保留,但其潜在MoE形状故意未加载到这些公式中。
✦ 改编——硬件,附有凭证(点击一行加载)
本页面计算使用的所有硬件数字,包括规格和持续性能及其来源。
完整引文见`SOURCES.md` (https://ezyang.github.io/interactive-parallelize-transformer/SOURCES.md) 与本页并列——每个值都可追溯到供应商规格表、已发布的测量结果或本书自身的基准测试;检索于2026-08-17;点击任何单元格可固定其引文并跟随来源链接。
合成数字的方法论:NVIDIA数据表标头*稀疏*FLOP/s,此处减半为密集;“双向”带宽减半为单方向;每GPU扩展是节点NIC总数除以其GPU。
≈标记的是*估计*因子(已说明依据,无直接公开测量——例如Blackwell集合通信继承H100的已测量NCCL比率,直到存在独立的nccl-tests)而非测量值。
将鼠标悬停在任何单元格上查看其引文。*持续*和*达到*列是规格的实测比例:切换顶部栏的**规格/测量**控件,页面上的每个方程都会按此比例降低(计算 × 持续,带宽 × 达到——加载的硬件行将它们显示为实时绿色修改)。已经假设MFU的墙钟估计继续使用规格峰值,因此没有重复计算。
注意引文引出的结论:TPU的持续性能远比受功耗限制的NVIDIA部件更接近其标称数字。
| 硬件 | C(密集bf16) | × 持续 | W 链路 | × 达到 | W 扩展 | HBM |
| :--- | :--- | :--- | :--- | :--- | :--- | :--- |
| | 459 TF | ≈0.72 | 180 GB/s | ≈0.95 | 6.25 GB/s | 96 GB |
| | 197 TF | ≈0.67 | 90 GB/s | ≈0.95 | 3.13 GB/s | 16 GB |
| | 989 TF | 0.73 | 450 GB/s | 0.82 | 50 GB/s | 80 GB |
| | 2.25 PF | 0.69 | 900 GB/s | ≈0.82 | 50 GB/s | 180 GB |
| | 2.5 PF | ≈0.70 | 900 GB/s | ≈0.82 | ≈50 GB/s | 186 GB |
| | 2.5 PF | ≈0.70 | 900 GB/s | ≈0.82 | 100 GB/s | 288 GB |
| | 989 TF | ≈0.73 | 200 GB/s | 0.80 | 50 GB/s | 80 GB |
为简化起见,**我们将Transformer近似为MLP块的堆栈**——如第4节(https://jax-ml.github.io/scaling-book/transformers/)所见,对于较大的模型,注意力只占FLOPS中相对较小的部分。我们还将忽略门控矩阵乘法,为每层留下以下简单结构:
改编:通过这种简化,每层包含 2·*D*·E·*F* 个权重(密集模型时E=1,即仅为 2·D·F),整个堆栈在当前设置下有 2·*D*·E·*F*·*L*= 参数——这是本页面通信算术中的“P”。内存问题不同:一个真正的检查点还包含门控MLP的第三个矩阵和注意力堆栈,因此内存计量器将权重定价为 P<sub>w</sub>≈ 3·D·E·F·L + 2.5·D<sup>2</sup>·L=,这与模型表发布的总数相差在百分之几以内(词汇嵌入和MHA时代注意力除外)。
一个简化的Transformer层。我们将每个FFW块视为两个矩阵的堆栈**W<sub>in</sub>**: bf16[D, F](上投影)和**W<sub>out</sub>**: bf16[F, D](下投影),输入为**In**: bf16[B, D]。框边按实时维度绘制(对数刻度)——拖动*F*=并观察矩阵变宽。将鼠标悬停在一个边上,可在页面各处突出显示该维度。
这是我们没有并行化的小Transformer的完整算法。
**前向传播:** 需要计算 Loss[B]
1. Tmp[B, F] = In[B, D] ·<sub>D</sub> Win[D, F]
2. Out[B, D] = Tmp[B, F] ·<sub>F</sub> Wout[F, D]
3. Loss[B] = ...
**反向传播:** 需要计算 dWout[F, D], dWin[D, F]
1. dOut[B, D] = ...
2. dWout[F, D] = Tmp[B, F] ·<sub>B</sub> dOut[B, D]
3. dTmp[B, F] = dOut[B, D] ·<sub>D</sub> Wout[F, D]
4. dWin[D, F] = In[B, D] ·<sub>B</sub> dTmp[B, F]
5. dIn[B, D] = dTmp[B, F] ·<sub>F</sub> Win[D, F] (*用于前面的层*)
我们提供此内容是为了与添加通信后的算法进行比较。
以下是我们将讨论的4种并行方案。每种方案都可以通过图中**In**、**Win、Wout和Out**的分片方式来唯一定义。
改编:快速回顾本书的符号:数组维度上的下标表示其被分割的网格轴——In[<sub>BDP</sub>, D] 表示批次维度被分成*DP*份,沿轴*DP*每芯片一份——而·上的下标表示被缩并的维度。
下方的探索器允许你点击切换方案:对于每种方案,它会绘制所有四个数组,其分片按芯片着色,显示在实时*DP*和*TP*下的每芯片局部形状,以及方案在前向和反向传播中支付的集合通信。
**1. 数据并行:***激活沿批次分片,参数和优化器状态在每个设备上复制。通信仅发生在反向传播期间。*
In[<sub>BDP</sub>, D] ·<sub>D</sub> Win[D, F] ·<sub>F</sub> Wout[F, D] → Out[<sub>BDP</sub>, D]
**2. 全分片数据并行(FSDP 或 ZeRO-3):***激活沿批次分片(类似纯数据并行),参数沿同一网格轴分片,并在前向传播中使用前即时进行AllGather。优化器状态也沿批次分片。减少重复内存。*
In[<sub>BDP</sub>, D] ·<sub>D</sub> Win[<sub>DDP</sub>, F] ·<sub>F</sub> Wout[F, <sub>DDP</sub>] → Out[<sub>BDP</sub>, D]
**3. 张量并行(也称为Megatron分片或模型并行):***激活沿D(d<sub>model</sub>)分片,参数沿F(d<sub>ff</sub>)分片。在每个块之前和之后进行AllGather和ReduceScatter激活。与FSDP兼容。*
In[B, <sub>DTP</sub>] ·<sub>D</sub> Win[D, <sub>FTP</sub>] ·<sub>F</sub> Wout[<sub>FTP</sub>, D] → Out[B, <sub>DTP</sub>]
**4. 管道并行:***权重沿层维度分片,激活沿层维度微批次化和滚动。管道级之间的通信最小(仅单跳移动激活)。滥用符号表示:*
改编:注意所有四种方案的共同点:每种都运行*相同*的矩阵乘法——FLOPS从不改变,只有数组的位置和乘法之间必须运行的集合通信不同。因此对于每种方案,问题始终在于这些集合通信能否隐藏在矩阵乘法之后。
在章节逐一剖析方案之前,本改编插入了一个简短的入门——首先,感受屋顶线 (https://ezyang.github.io/interactive-parallelize-transformer/#roofline)——构建一个能回答所有四种方案问题的图像。
In[<sub>LPP</sub>, B, D][i] ·<sub>D</sub> Win[<sub>LPP</sub>, D, F][i] ·<sub>F</sub> Wout[<sub>LPP</sub>, F, D][i] → Out[<sub>LPP</sub>, B, D][i]
四种方案,每种一个标签页。对于当前方案:其分片语法行、绘制的四个数组(分片按芯片着色)、实时*DP*和*TP*下的每芯片局部形状。
相似文章
预训练并行化与失败训练运行笔记(12分钟阅读)
一篇技术深度文章,探讨大型语言模型中预训练运行失败的常见原因,包括专家路由中的因果破坏问题和数值精度错误,并附有Llama 4、Gemini 2 Pro和GPT-4的示例。
使用奇偶瓶颈层扩展可解释的Transformer
介绍了ParityTransformer,这是一种GPT-2规模的架构,具有深度奇偶瓶颈层,使中间表示在设计上即可解释,在无需传统宽瓶颈内存成本的情况下高效强制稀疏性,并在稀疏探测任务上展示了具有竞争力的性能。
Transformer 可扩展性危机:现代语言模型中性能墙的首次全面实证分析
本文对 118 个 Transformer 模型进行了首次大规模实证分析,揭示了关键的性能墙,其中成功率从 512 token 时的 88.1% 下降到 2048 token 时的 0%,挑战了主流的缩放假设。
Transformer Explainer:交互式学习文本生成模型
Transformer Explainer 是一个交互式可视化工具,让非专业人士能够通过浏览器中的实时实验和可视化,理解 GPT-2 模型的内部工作机制。
什么是 Looped Transformers?清晰解释(8分钟阅读)
Looped transformers 在多次传递中重复使用相同的层,以参数数量换取计算量,用更少的权重实现更好的推理。文章追溯了这一想法到 Universal Transformer (2018),并解释了其最初因扩展定律和时机而失败的原因。