跳到主要内容

一台机器装不下 — 分布式训练(多卡分工)

这一章讲三件事: 为什么必须分布式(三堵墙的三笔硬账);怎么切(数据并行、流水线并行、按层内的数表切——三种各切一块的分工);切完之后显存和网络这两笔账怎么算(ZeRO 与集群拓扑)。 主走查预告: 把一张大数表切成两半放到两块加速卡上,怎么保证算出来的结果和单卡一字不差。

1. 三堵墙:三笔硬账

模型需求与硬件供给的差距在拉大:从 2013 年的 AlexNet 到 2024 年的 DeepSeek-V2,模型参数量大约每 18 个月增长 56 倍,而硬件算力每 18 个月翻一番——需求增速是供给的 56 倍对 2 倍1。书里把矛盾总结为三堵墙,每堵墙都有具体数字2:

账目数字
计算墙单卡算力 vs 总计算量H100 单卡每秒算 2000 万亿次(浮点,即带小数的数的运算),GPT-3 总共要算约 3×10^23 次,差 8 个数量级
显存墙权重 vs 显存GPT-3 的 1750 亿参数用 FP32 存要 700GB,H100 只有 80GB
通信墙梯度同步 vs 网络搬运速度128 个模型副本每轮更新至少传 89.6TB 梯度,单条专用高速网络线只有 800Gbps

顺带记一笔训练成本的真实感:LLaMA-65B 一次预训练要 1022362 加速卡工时——一块卡不吃不喝跑一百多年3

三堵墙的因果关系值得想清楚:计算墙和显存墙逼你把任务拆到多卡(分布式),而拆完立刻撞上通信墙——分布式不是解决方案,是用一个新问题换两个旧问题。分布式训练调度的是整个集群的全部资源(算力、显存、网络),任何一块短板都会拖慢全局。后面所有技术都在管理这个新问题。

2. 三种切法

数据并行:最容易想到的切法

每块加速卡放一份完整的模型副本,把一批(凑在一起送训的一小堆样本)数据平均分给它们,各自算完梯度后广播、取平均、更新参数。它加速比最高、实现最简单(不影响计算逻辑),缺点同样直白:每块卡都要装下整个模型——在千亿参数时代,这条路线单独用已经不成立,只能当外层包装4

流水线并行:按层接力

把模型按层切成几段,每段放一块卡,数据像流水线一样依次流过。朴素版本有个致命浪费:前一段在算的时候,后一段干等;反向传播时又反过来——大量时间设备空转,这叫流水线气泡5

两个补救方案。GPipe:把数据批(每次凑出来送算的一小堆)再切成更小的微批次,第一个微批次流到第二段时,第一段马上开始算第二个——气泡被填掉一部分。1F1B(一个 Forward 配一个 Backward,即 F1——一次前向——配一次后向;Megatron-LM 提出):调度粒度细到「做一次前向、立刻做一次后向」交替,下游设备等待时就能干活,省内存效果更好;还有交错式变体,让每块卡负责多个不连续的层段,进一步填气泡6

张量,即把一层里的计算切开:张量并行

流水线是「你算完我再算」,张量并行是「我们同时算一块大矩阵的各部分」。它必须回答一个问题:切开的计算怎么保证和单卡在数学上完全等价。这正是主走查要走的路。

主走查:矩阵按列切开,两块卡拼出同一个答案

要算 Y = X·A。X 是输入矩阵(每块卡都有完整副本),A 是参数矩阵(太大,单卡装不下)。书里给的切法之一是按列切:A = [A1, A2],A1、A2 分别放到两块卡上7

拿小矩阵走一遍(数值是为演示编的):

┌ 1 2 ┐ ┌ 5 6 ┐
X = │ 3 4 │ A = │ 7 8 │
└ ┘ └ ┘

按列切:A1 = [5 7]ᵀ(第 1 列),A2 = [6 8]ᵀ(第 2 列)

卡 1 算 Y1 = X·A1: 卡 2 算 Y2 = X·A2:
行1: 1×5 + 2×7 = 19 行1: 1×6 + 2×8 = 22
行2: 3×5 + 4×7 = 43 行2: 3×6 + 4×8 = 50

通信:把 Y1、Y2 拼接(Concat)到一起
Y = [19 22] ← 与单卡直接算 X·A 的结果完全一致
[43 50]

图说:按列切分,每块卡拿到完整 X 和一列参数,
各算各的、互不干扰,最后只做一次「拼接」级别的通信。

等价性成立的原因:矩阵乘法按列分配——X·[A1, A2] = [X·A1, X·A2]。另一半切法是按行切:A 按行切成上下两块,此时 X 也要按列切,各卡算出部分结果后相加(而不是拼接)。两种切法各有用途,真正的巧妙在组合:Transformer 的前馈层有两层全连接,第一层按列切、第二层按行切——第一层的输出恰好是按列分开的形式,正是第二层需要的输入形状,于是中间那次汇总通信被整个省掉了8

嵌入层和损失函数(算「错得多离谱」的公式)也有各自的切法:词表太大时按词切,查不到的词记 0,再跨卡求和(词表 64000、表示宽度 5120 时,仅嵌入参数加梯度就要近 2.5GB);交叉熵(衡量答案错得多离谱的常用损失)按类别切,用减最大值的技巧避免数值溢出,只需三次小通信9

3. 显存账:ZeRO 与省内存的数字表示

切法解决「装不下模型」,还没解决「装不下训练状态」。

数字表示的精度(一个数保留多少位来存)是另一条主线。

混合精度(高低位数字混着用)训练的思路就出自这里。

训练用 Adam 优化器(负责照梯度调参数的程序)时,除了参数和梯度,还要存一阶动量和二阶动量(可以理解为「梯度的移动平均」和「梯度平方的移动平均」,用来自适应调整步长)。书里算了精确的账:混合精度下每个参数要占 16 字节——2 字节参数 + 2 字节梯度 + 12 字节优化器状态,Adam 状态占 75%;一个 75 亿参数的模型,FP16 权重本身只要 15GB,训练状态却要 120GB10

**ZeRO(零冗余优化器)**的思路:数据并行里每个卡都存全套优化器状态,而这套东西在各卡之间是冗余的——那就分片,每张卡只存 1/N:

层级切什么每卡内存趋近通信代价
ZeRO-1优化器状态原来的 1/4不变
ZeRO-2+ 梯度1/8不变
ZeRO-3+ 参数本身趋近 01.5 倍11

混合精度是另一半账:前向反向用 FP16 算(快一倍、省一半显存),但 FP16 数值范围小,容易上溢下溢,所以保留一份 FP32 主权重;配套的动态损失缩放技术,在反向传播前把损失放大 2^K 倍防止小梯度「沉底」,算完再缩回来12

4. 网络账:拓扑决定一切

并行策略能不能跑快,取决于通信走的是哪种网络。书里的层次:节点内(一台服务器——即一台装 8 张卡的机器——里)用 NVLink+NVSwitch 全连接,任意两卡双向 900GB/s

节点间走 InfiniBand(专用高速网络线),单条 200-400Gbps,整个集群用胖树拓扑追求搬运速度无收敛13

作为对照:普通 PCIe(主板上的通用插线标准)5.0 总线只有 128GB/s——节点内和节点间差着几倍到几十倍。

这套物理现实直接决定了放置策略(参数服务器架构与去中心化架构之争,大模型训练实际都用后者,靠 AllReduce 等集合通信原语(最基本的标准化通信操作)同步14):通信量最大的张量并行放节点内,通信量最小的流水线并行放节点间,数据并行套在最外层——这就是 DeepSpeed 的 3D 并行15

真实案例:BLOOM 的 384 卡

把本章所有概念放进一个真实部署:BLOOM(1750 亿参数)用 48 台 DGX-A100 服务器(共 384 块 A100)训练 3.5 个月。切法是:48 组数据并行 × 每组内部 12 个流水线阶段 × 每阶段 4 卡张量并行,再叠 ZeRO 降显存——4×12×48 正好 2304 个并行位,384 块卡各就各位16

5. 作者的判断与证据

  • 有证据的: 三堵墙的全部数字、BLOOM/LLaMA/OPT 的集群配置、ZeRO 三级的内存公式与通信倍数,书里都给了论文出处。
  • 作者的判断: 「ZeRO-3 通信量 1.5 倍」来自原论文的分析;实践中换来的显存收益通常值得,但书没有讨论对小集群场景的适用性。
  • 书里的坦白: 异步(不等所有卡交齐作业就先干)训练——参数服务器架构里的选项——能提速但「训练效果有所波动」——快和稳的取舍没有被定量回答。

6. 边界与局限

  • 流水线并行的 1F1B 与 GPipe 的时间对比是定性的(省内存、总时间相近),没有给出气泡率的公式推导。
  • 集群硬件以 NVIDIA DGX/HGX 为中心,国产芯片、云上弹性集群的拓扑差异未覆盖。
  • 容错(长时间训练中设备故障怎么办)只在开头点了一句,没有展开——这在数月级训练里是真实痛点(我们按通用知识补充)。
  • DeepSpeed 一节的大量篇幅是工程配置代码,机制层面本章已提炼;真要上手跑,直接看书里 4.4 节的七步流程即可。

7. 可带走的

  1. 三堵墙的三笔账:算力差 8 个数量级、显存 700GB 对 80GB、128 副本一轮 89.6TB 梯度——分布式是被数字逼出来的,不是架构偏好;
  2. 三种切法一句话:数据并行切数据、流水线切层、张量并行切矩阵;大模型训练三者混用;
  3. 张量并行的命门是数学等价:按列切拼结果、按行切加结果;FFN「一列一行」的组合能省掉中间通信;
  4. 训练显存的大头不是参数,是 Adam 状态(占 16 字节/参数中的 12 字节)——ZeRO 分片就是切这块;
  5. ZeRO 三级的阶梯:切优化器状态(省到 1/4)→ 加梯度(1/8)→ 加参数(趋近零,通信 1.5 倍);
  6. 混合精度 = FP16 干活 + FP32 主权重 + 动态损失缩放;省的是计算和显存,险的是数值溢出;
  7. 搬运速度决定放置:按层内数表切(通信最重)必须放节点内 NVLink,流水线(通信最轻)走节点间;
  8. 读任何「XX 卡训了 XX 天」的新闻,先问三件事:数据并行几组、流水线几段、张量并行几卡——BLOOM 的答案是 48×12×4。

8. 原文地图

主题原书章原文位置
需求增速 56 倍4 分布式训练text/04-ch04.txt:26(搜「DeepSeek-V2 发布」)
三堵墙4 分布式训练text/04-ch04.txt:64(搜「计算墙」) · text/04-ch04.txt:69(搜「显存墙」) · text/04-ch04.txt:72(搜「通信墙」)
LLaMA-65B GPU 小时4 分布式训练text/04-ch04.txt:59(搜「1022362」)
数据并行加速比4 分布式训练text/04-ch04.txt:133(搜「加速比最高」)
流水线气泡4 分布式训练text/04-ch04.txt:274(搜「流水线气泡」)
GPipe 微批次4 分布式训练text/04-ch04.txt:279(搜「GPipe」)
1F1B4 分布式训练text/04-ch04.txt:291(搜「1F1B」)
按列/按行切分4 分布式训练text/04-ch04.txt:361(搜「按列切块」) · text/04-ch04.txt:365(搜「按行切块」)
FFN 一列一行省通信4 分布式训练text/04-ch04.txt:392(搜「按列切块」)
Embedding 1250MB4 分布式训练text/04-ch04.txt:343(搜「1250MB」)
交叉熵三次通信4 分布式训练text/04-ch04.txt:415(搜「所有类别的损失」)
Adam 16Φ 与 120GB4 分布式训练text/04-ch04.txt:512(搜「16Φ」) · text/04-ch04.txt:516(搜「120GB」)
动态损失缩放4 分布式训练text/04-ch04.txt:506(搜「动态损失缩放」)
ZeRO 三级4 分布式训练text/04-ch04.txt:527(搜「优化器状态进行分区」)
Zero-3 通信 1.5 倍4 分布式训练text/04-ch04.txt:544(搜「1.5 倍」)
集群拓扑与胖树4 分布式训练text/04-ch04.txt:641(搜「胖树」)
NVLink 900GB/s4 分布式训练text/04-ch04.txt:652(搜「900GB/s」)
集合通信与 NCCL4 分布式训练text/04-ch04.txt:359(搜「集合通信」) · text/04-ch04.txt:741(搜「NCCL」)
3D 并行放置逻辑4 分布式训练text/04-ch04.txt:852(搜「张量并行计算组放置在节点内」)
BLOOM 384 卡4 分布式训练text/04-ch04.txt:478(搜「48 个 NVIDIA DGX-A100」)

Footnotes

  1. 出处:「4 分布式训练」第 26 段(text/04-ch04.txt:26,搜「DeepSeek-V2 发布」)。

  2. 出处:「4 分布式训练」第 64-75 段(text/04-ch04.txt:64,搜「计算墙」;text/04-ch04.txt:69,搜「显存墙」;text/04-ch04.txt:72,搜「通信墙」)。

  3. 出处:「4 分布式训练」第 59 段(text/04-ch04.txt:59,搜「1022362」)。

  4. 出处:「4 分布式训练」第 133 段(text/04-ch04.txt:133,搜「加速比最高」)。

  5. 出处:「4 分布式训练」第 274 段(text/04-ch04.txt:274,搜「流水线气泡」)。

  6. 出处:「4 分布式训练」第 279 段(text/04-ch04.txt:279,搜「GPipe」)与第 291 段(text/04-ch04.txt:291,搜「1F1B」)。

  7. 出处:「4 分布式训练」第 361 段(text/04-ch04.txt:361,搜「按列切块」)。走查矩阵数值为演示编造。

  8. 出处:「4 分布式训练」第 392 段(text/04-ch04.txt:392,搜「按列切块」)。原文:对第一个 FC 层的参数矩阵按列切块,对第二个按行切块,可以省去第一个 FC 层后的汇总通信。

  9. 出处:「4 分布式训练」第 343 段(text/04-ch04.txt:343,搜「1250MB」)与第 415 段(text/04-ch04.txt:415,搜「所有类别的损失」)。

  10. 出处:「4 分布式训练」第 512 段(text/04-ch04.txt:512,搜「16Φ」)与第 516 段(text/04-ch04.txt:516,搜「120GB」)。

  11. 出处:「4 分布式训练」第 527-537 段(text/04-ch04.txt:527,搜「优化器状态进行分区」)与第 544 段(text/04-ch04.txt:544,搜「1.5 倍」)。

  12. 出处:「4 分布式训练」第 506 段(text/04-ch04.txt:506,搜「动态损失缩放」)与第 513 段(text/04-ch04.txt:513,搜「2K 倍」)。

  13. 出处:「4 分布式训练」第 641 段(text/04-ch04.txt:641,搜「胖树」)与第 652 段(text/04-ch04.txt:652,搜「900GB/s」)。

  14. 出处:「4 分布式训练」第 359 段(text/04-ch04.txt:359,搜「集合通信」)。八种通信原语见同节列表。

  15. 出处:「4 分布式训练」第 852 段(text/04-ch04.txt:852,搜「张量并行计算组放置在节点内」)。

  16. 出处:「4 分布式训练」第 478 段(text/04-ch04.txt:478,搜「48 个 NVIDIA DGX-A100」)。