Skip to content

分布式训练总论与显存账本 ​

标签
AI/infra/显存管理
字数
2545 字
阅读时间
11 分钟

训练的本质是反复「前向算 loss → 反向算梯度 → 优化器更新参数」。模型小时一块卡全程装得下,不需要分布式;但模型一旦膨胀到几百亿参数,单卡会在两个维度同时撞墙。

这两个问题相互独立,解法方向也相反 —— 看清这点是理解所有并行策略的前提:

瓶颈症状解法方向
装不下(空间)训练状态超过单卡显存切分状态 —— 把一份东西拆到多卡
跑不完(速度)算力不足以在可接受时间内训完复制并行 —— 多卡各干一摊活

用一个类比:单卡训大模型,像让一个人用一张小书桌整理一整座图书馆 —— 桌子太小摊不开(显存),一个人搬一辈子也搬不完(算力)。前者要「拆开放到几张桌上」,后者要「多叫几个人一起搬」。

装不下:显存瓶颈有多大 ​

直觉上「多少参数占多少显存」是严重低估。粗略地说,混合精度 + Adam 下每个参数约需 16–18 字节静态显存。以 LLaMA 系列估算(仅静态显存,不含激活值):

模型参数量 Ψ静态显存 ≈ 16Ψ单卡 H100(80 GB)
LLaMA 7B7×109约 112 GB装不下
LLaMA 70B70×109约 1,120 GB远超
LLaMA 405B405×109约 6,480 GB需数十卡

小号模型就已经装不下

7B 模型光静态显存就约 112 GB,超过单张 H100 的 80 GB —— 而这还没算前向堆积的激活值。「7B 只要 14 GB」那个直觉只在推理时成立(参见 01-GPU 硬件架构与存储层次)。

跑不完:算力与时间瓶颈 ​

即便显存无限大,单卡还有「算得太慢」这道坎。大模型训练的计算量有一个广为引用的经验公式:

C≈6ΨD

(Ψ 参数量,D 训练 token 数。)系数 6 来自前向约 2 倍、反向约 4 倍参数量的乘加。

代入 70B 模型训 1T(1012)token:

C=6×(70×109)×1012=4.2×1023 FLOPs

单张 H100 的 BF16 算力峰值约 1015 FLOP/s,乐观假设 50% 有效利用率(MFU):

T=4.2×10231015×0.5≈8.4×108 s≈26.6 年

单卡训 70B 要二十多年。 而 N 卡数据并行理论上把时间压到接近 1/N —— 1024 卡可降到约 9 天量级。这是第二个动因。

两面墙的解法方向相反,而且正交:

  装不下(空间)                       跑不完(速度)
  训练状态 > 单卡显存                   算力不够,训不完
        │                                   │
        ▼ 切分状态                          ▼ 复制并行
    把一份东西拆到多卡                    多卡各干一摊活
        │                                   │
     ZeRO / FSDP / TP / PP                DDP
        │                                   │
        └────── 两者正交,可以叠加 ──────────┘
                        │
                        ▼
              3D 并行 = DP × TP × PP

静态显存账本:手算一遍 ​

以 BF16 混合精度 + Adam 为例,设参数量 Ψ,逐项(单位 Bytes):

组成精度每参数说明
模型参数BF162Ψ前向/反向的工作副本
梯度BF162Ψ反向产出,与参数同形状
优化器 · FP32 主权重FP324ΨAdam 维护的高精度副本
优化器 · 一阶动量 mFP324Ψ梯度的指数滑动平均
优化器 · 二阶动量 vFP324Ψ梯度平方的指数滑动平均
Mstatic=2Ψ⏟参数+2Ψ⏟梯度+12Ψ⏟Adam=16Ψ Bytes

训练显存的大头是优化器状态(12Ψ,占 16 里的 12),而非参数本身。这正是 ZeRO 第一刀就切优化器状态的原因 —— 它是性价比最高的下手处。

7B 模型手算 ​

项大小
参数2×7=14 GB
梯度2×7=14 GB
Adam 状态12×7=84 GB
合计约 112 GB

16Ψ 与 18Ψ 的差别在哪 ​

同一个 7B 模型,不同资料会算出 112 GB 或 126 GB。差别只有一处 —— 梯度用哪种精度存:

配置参数梯度FP32 主权重mv合计7B
BF16 梯度BF16 · 2BF16 · 244416 B112 GB
FP32 梯度FP16 · 2FP32 · 444418 B126 GB

两种都是真实配置。差距的来源是梯度累加用的缓冲区精度不同,与对错无关。估算时先确认自己的框架走的是哪条路径。

与库内另一篇的口径差异

01-GPU 硬件架构与存储层次 用的是 18 B/参数(7B → 126 GB),依据是来源教程 GPU 篇的逐项表;本篇用 16Ψ(7B → 112 GB),依据是来源教程分布式训练篇。两者的差异来源就是上表这一行。两处说法均保留,估算时按自己的框架配置取对应口径。

别忘激活值 ​

激活值是前向中为反向保留的中间结果,随 batch size b 与序列长度 s 增长:

项公式取决于
参数 / 梯度 / Adam 状态16Ψ(或 18Ψ)只取决于参数量 —— 固定成本
激活值∝b⋅s⋅h⋅Lbatch / 序列长度 —— 可调成本

OOM 时优先压激活值:减小 batch、开梯度累积、开激活重计算(Activation Checkpointing)。这三条不动模型规模,只动「一次前向要保留多少中间量」。

16Ψ 的构成(Ψ 是参数量):

  ┌──────┬──────┬────────────────────────────────────────────┐
  │ 参数 │ 梯度 │              优化器状态(Adam)              │
  │  2Ψ  │  2Ψ  │   FP32 主权重 4Ψ   │  一阶动量 m 4Ψ          │
  │ BF16 │ BF16 │                    │  二阶动量 v 4Ψ          │
  └──────┴──────┴────────────────────────────────────────────┘
     12.5%  12.5%                  75%

  ⇒ 大头是优化器状态($12\Psi$,占 16 里的 12)
  ⇒ 这正是 ZeRO 第一刀就切优化器状态的原因 —— 性价比最高
  ⇒ 另外「16Ψ 还是 18Ψ」只差梯度那一格:梯度存 BF16 是 2Ψ,存 FP32 是 4Ψ

五大并行策略全景 ​

每种策略都回答两个问题:切什么(决定省不省显存)、怎么通信(决定快不快)。

策略切什么主要通信解决的问题
数据并行 DP / DDP切数据,模型不切梯度 AllReduce跑不完(加速)
ZeRO / FSDP切优化器状态 / 梯度 / 参数ReduceScatter + AllGather装不下(省显存)
张量并行 TP切单层内的矩阵运算层内 AllReduce单层太大
流水线并行 PP切网络层(按深度分段)段间传激活值(点对点)层数太多
序列 / 上下文并行 SP / CP沿序列维度切序列维 AllGather / All-to-All序列太长

逐个建立直觉:

  • DDP —— 每卡一份完整模型,各算各的数据分片,反向时 AllReduce 同步梯度。它不省显存(每卡仍装完整模型),主攻「跑不完」。
  • ZeRO / FSDP —— DDP 的显存优化版。既然每卡都存了一份完全相同的优化器状态,那不如分片存、用时再 AllGather 拼回来。
  • TP —— 把一层(如一个大 Linear)的权重矩阵按行/列切到多卡。通信频繁且量大,只适合机内 NVLink。
  • PP —— 按层切成几段,激活值像流水线一样在段间传递。通信量小,可跨机。
  • SP / CP —— 沿序列长度切分,让 128K 这类超长上下文也能放下。

每个策略的通信量与原语选择在 03-集合通信与 NCCL 有对照表;「为什么 TP 只能机内」的量化估算在 03-多卡互联与集群网络。

怎么选:带宽决定作用域 ​

选型只有一条核心判据:通信带宽决定策略的作用域。

互联带宽量级适配策略
NVLink / NVSwitch(机内)数百 GB/s(NVLink 4.0 为 900 GB/s)TP(通信最频繁,必须机内)
InfiniBand / RoCE(机间)数百 Gbps(NDR 400 Gbps ≈ 50 GB/s)PP、DP(通信稀疏,可跨机)
以太网更低仅 DP 这类低频通信勉强可用

判据的推导:TP 每一层的每个子模块都要 AllReduce(一个 80 层模型约 320 次),放到跨机带宽上会被通信拖垮;PP 只在段边界传一次激活值,DP 一步只同步一次梯度。这是 3D 并行布局的根本依据。

实用决策树 ​

单卡装得下完整训练状态?
├─ 是 → DDP(最简单,优先选)
└─ 否 → 分片后单卡能装下?
        ├─ 是 → FSDP / ZeRO
        └─ 否 → 瓶颈在哪?
                ├─ 单层太大   → 叠加 TP(机内)
                ├─ 层数太多   → 叠加 PP(可跨机)
                └─ 序列太长   → 叠加 SP / CP
                        ↓
                  3D 并行:DP × TP × PP

推荐路径:从最简单的够用方案起步 —— 单卡装得下就 DDP;装不下先上 FSDP/ZeRO;还不行再按瓶颈叠 TP/PP/CP;千亿级才需要全套 3D 并行。

不要一上来就堆 3D 并行。每多一个并行维度都会显著增加调试复杂度与通信开销。

相关 ​

参考 ​

贡献者 ​

文件历史 ​