训练 Infra 03:并行策略全景——为什么是 TP1 / PP1 / EP8 / DP40
逐个拆开 DP、ZeRO 1/2/3、Megatron distributed optimizer、TP、PP、EP、CP/SP,然后用 40 = 2³×5 这个整除约束和 PP bubble 公式,把本项目的并行选型完整推演一遍——结论是 TP=2 会让显存变得更差。
训练 Infra 03:并行策略全景
本节属于大模型训练 Infra 实战课。第 1 节给了带宽地图,第 2 节算出"不分片需要 560GB、单卡只有 96GB"。这一节讲怎么分。
上一节的结论是:35B 模型用混合精度 Adam 训练需要 560GB,单卡 96GB。差 5.8 倍。所以并行不是调优选项,是可行性前提。
问题是并行有五六种,每种切的东西不同、通信频率不同、约束不同。这一节先把它们逐个拆开,然后用本项目的真实约束推一遍——你会看到一个反直觉的结论:在这套集群上,TP=2 会让显存变得更差。
1. DP:数据并行与梯度累积
原理
最简单的一种:每张卡拿一份完整的模型副本,喂不同的数据,各自算梯度,然后 all-reduce 求平均,各自更新。
每 rank: 前向 → 反向 → all-reduce(梯度) → optimizer.step()
DP 不省任何显存(每张卡都是完整副本),它只提高吞吐。所以纯 DP 的上限是"单卡装得下整个模型的训练态"——对 35B 来说完全不可能。
梯度累积
当 global batch 需要比"DP 数 × micro batch"更大时,就在本地多跑几个 micro batch,把梯度累加起来,累够了才通信+更新:
本项目:
global batch = 80
DP = 40
micro batch = 1
→ 每个 DP rank 需要处理 80/40 = 2 个样本
→ 梯度累积步数 = 2
梯度累积的两个作用:凑大 batch、摊薄通信(累积 k 步才通信一次,通信占比降到 1/k)。
代价是显存里要长期挂一份梯度 buffer(第 2 节那 13.8GB),以及——梯度累积步数会成为 PP 的死穴,见第 4 小节。
2. ZeRO 1/2/3 与 Megatron distributed optimizer
ZeRO 的三级
ZeRO(Zero Redundancy Optimizer)的思路是:DP 下每张卡存的东西是完全重复的,那就把重复的部分切开,用通信换显存。分三级递进:
| 级别 | 切什么 | 每参数常驻(bf16 权重+梯度+Adam) | 通信代价 |
|---|---|---|---|
| ZeRO-1 | 优化器状态(12B) | 2 + 2 + 12/N | 与 DP 基本相同 |
| ZeRO-2 | + 梯度(2B) | 2 + 2/N + 12/N | 与 DP 基本相同 |
| ZeRO-3 | + 权重(2B) | (2 + 2 + 12)/N | 显著增加 |
关键分界线在 ZeRO-3。前两级切的是"只在 step 时才需要全量"的东西,通信可以合并到原本就要做的梯度同步里。而 ZeRO-3 切的是权重——前向每一层都需要完整权重,所以每层都要 all-gather 一次参数,反向再来一次。通信次数从"每步一次"变成"每层两次",量级完全不同。
Megatron distributed optimizer 是哪一级
大致相当于 ZeRO-1 + ZeRO-2,但实现路径不同,这个差别值得讲清楚:
DeepSpeed ZeRO-2 的典型流程:
反向 → 梯度 reduce-scatter → 各 rank 更新自己那份 → all-gather 权重
Megatron distributed optimizer:
反向(边算边填连续 grad buffer,按 bucket 触发)
→ bucket 满就 reduce-scatter(与后续反向计算重叠)
→ 各 rank 更新自己那份 fp32 主参数
→ all-gather bf16 权重(与下一步前向重叠)
差别在连续 buffer + bucket + 重叠这三件事。Megatron 把梯度预先摊在一块连续显存里(第 2 节那 13.8GB 的来源),这样:
- reduce-scatter 可以按 bucket 分批发起,与还没算完的反向计算重叠
- 通信是大块连续内存,带宽利用率高(小张量逐个通信会被延迟吃掉)
- 代价是那块 buffer 常驻,不能省
这就是 v1(transformers+DeepSpeed) 2.5 天、v2(Megatron) 一夜跑完 2 epoch 的一部分原因——不是 ZeRO 算法差,是通信与计算的重叠做得不同(第 5 节会完整对比两个框架)。
项目实测
优化器状态 12 B/参数,沿 DP 切:
非专家 X=2.9B , DP组=40 → 2.9e9 × 12 / 40 = 0.87 GB
专家 每卡4.0B , DP组=5 → 4.0e9 × 12 / 5 = 9.63 GB
合计 10.5 GB
注意专家的 DP 组只有 5——这是 EP 的副作用,第 2 节算过,代价 8.4GB。
3. TP:张量并行
原理
把单个矩阵乘法切开,让多张卡协作算一层。经典做法是列并行 + 行并行配对:
第一个线性层按列切: Y = X · [W1 | W2] → 各卡得到 Y 的一部分,无需通信
第二个线性层按行切: Z = [Y1;Y2] · [W1';W2'] → 各卡算部分和,需要 all-reduce
一个 transformer block 里 attention 和 FFN 各有这样一对,所以前向 2 次 all-reduce、反向 2 次,每层共 4 次。
特点
- 省显存:权重、梯度、优化器状态、激活全都按 TP 切开
- 通信频率极高:每层都要,所以必须在 NVLink 域内
- 约束严格:hidden size、注意力头数都必须能被 TP 整除
项目:TP 的通信到底贵不贵
先把这笔账算出来,因为我在第 1 节说"TP 会引入通信税",这个说法不够准确,需要更正。
假设 seq=49152, hidden=2048, 48 层, TP=2:
单个激活张量 = 49152 × 2048 × 2 B = 201 MB
ring all-reduce 传输量 = 2(N-1)/N × S = 201 MB / 次
每 micro batch: 48 层 × 4 次 = 192 次 × 201 MB = 38.7 GB
每步 2 个 micro batch → 77.3 GB
按 NVLink 有效带宽 300 GB/s → 0.26 秒
占 120 秒的一步 → 约 0.2%
0.2%。基本可以忽略。 所以"TP 通信太贵"在这套配置上并不成立——一步 120 秒实在太长了,什么通信都被摊薄了。
真正的原因是别的两条,见下面第 5 小节。
4. PP:流水线并行
原理
把模型按层切成若干段,每段放一组卡,像流水线一样传递激活。
PP=4: 卡组0(层1-12) → 卡组1(层13-24) → 卡组2(层25-36) → 卡组3(层37-48)
通信量极小(只传层边界的激活,点对点),但有个致命问题:bubble(气泡)。
bubble 公式
流水线要填满才有效率。开头几个 micro batch 在灌,末尾几个在排空,这段时间大部分卡在空转:
bubble 占比 = (p - 1) / (m + p - 1)
p = 流水线级数
m = 每条流水线要过的 micro batch 数
结论很清楚:m 必须远大于 p,PP 才划算。
项目:为什么 PP=1 是唯一选择
m = global_batch / DP = 80 / 40 = 2
PP=2: bubble = 1/(2+2-1) = 33%
PP=4: bubble = 3/(2+4-1) = 60%
PP=5: bubble = 4/(2+5-1) = 67%
micro batch 数只有 2,PP=2 就要浪费三分之一的算力,PP=4 浪费六成。
而 m 为什么只有 2?因为 micro batch=1(49K 序列把显存吃满了,第 2 节),global batch 80 又不能随意加大(会改变训练动力学)。也就是说:
长序列 → micro batch 只能是 1 → 每 rank 的 micro batch 数极少 → PP 的 bubble 无法摊薄 → PP 出局。
这是一条完整的因果链,值得记住。长序列训练和流水线并行天生不合。
5. EP:专家并行(MoE 特有)
原理
MoE 的每一层有很多专家(独立的 FFN),每个 token 只被路由到少数几个。EP 把专家分到不同卡上:
1. router 决定每个 token 去哪些专家
2. all-to-all: 把 token 发到持有对应专家的卡
3. 各卡用本地专家计算
4. all-to-all: 把结果送回原来的卡
每个 MoE 层两次 all-to-all,频率和 TP 同级——所以同样必须在 NVLink 域内。
all-to-all 为什么特别怕慢链路
all-reduce 有 ring/tree 等算法可以优化,且流量是规则的。而 all-to-all 是 N×N 的全交叉,每对 rank 之间都要传数据,而且:
- 流量不均:路由是数据决定的,某个专家可能突然很热
- 同步性强:所有 rank 必须都完成才能继续
- 一慢全慢:最慢的那条链路决定整体耗时
所以 EP 跨节点的后果不是"慢一点",而是"慢一个数量级"。
容量因子与 drop token
路由不均衡时,某个专家可能收到远超平均的 token。实现上通常设一个容量因子(capacity factor):
每个专家最多接收 = capacity_factor × (总token数 / 专家数) × top_k
超出的 token 被 drop(直接跳过这个专家,走残差)。容量因子的取舍:
- 太小 → drop 多 → 训练信号损失,loss 曲线变差
- 太大 → 显存和计算浪费在 padding 上
这是 MoE 训练里一个容易被忽略但很实在的调参点。建议在日志里加一个 drop rate 指标——它异常升高通常意味着路由塌缩(router collapse),比 loss 更早报警。
6. CP / SP:上下文并行与序列并行
为什么需要
TP 切的是"宽度"(hidden 维),但长序列的压力在"长度"维。49K 序列下,即使模型不大,激活也会爆。CP/SP 就是沿序列维切分。
两个主流做法:
| 方案 | 思路 | 通信 |
|---|---|---|
| Ulysses | 按序列切,attention 前 all-to-all 换成按头切,算完再换回 | 每个 attention 2 次 all-to-all |
| Ring Attention | 序列分块环形传递 KV,每张卡逐块累积 attention | 环形 P2P,可与计算重叠 |
对本项目的意义
这是当前配置里最值得尝试的一个方向,因为:
- 瓶颈确实在序列维(第 2 节算出激活+logits 吃掉 47.5GB)
- CP 能直接降低单卡的序列长度 → 激活和 logits 都按比例下降
- 释放出的显存可以换掉 full recompute(省 25% 算力)或加大 micro batch
但有个 MoE 特有的复杂性:CP 和 EP 都要占用 rank 维度。TP × PP × CP × DP = 40 且 EP | DP,加一个 CP=2 就会把 DP 压到 20,而 8 不整除 20——EP=8 又不可行了。这就引出下一小节。
7. 推演:为什么是 TP1 / PP1 / EP8 / DP40
现在把所有约束摆在一起。
约束清单
1. 硬件: 40 卡 = 5 节点 × 8 卡,节点内 NVLink,节点间 RoCE
2. 分解: TP × PP × DP = 40,而 40 = 2³ × 5
3. MoE : EP 必须整除 DP
4. 拓扑: EP ≤ 8 (否则 all-to-all 跨节点)
5. 显存: 单卡 96GB,峰值必须 < 96
6. 批次: micro batch = 1(序列 49K 的后果)→ m = 2
逐个排除
PP:由约束 6,m=2,PP=2 就浪费 33%。PP=1。
TP:这是最有意思的一步。TP 会减少 DP,而 DP 减少会连锁影响 EP 的可行取值:
| TP | DP=40/TP | EP 可行取值(需整除 DP) | 节点内最大 EP | 每卡参数 | bf16 权重 |
|---|---|---|---|---|---|
| 1 | 40 | 1,2,4,5,8,10,20,40 | 8 | 6.91B | 13.8GB |
| 2 | 20 | 1,2,4,5,10,20 | 5 | 7.87B | 15.7GB |
| 4 | 10 | 1,2,5,10 | 5 | 7.14B | 14.3GB |
| 5 | 8 | 1,2,4,8 | 8 | 4.59B | 9.2GB |
| 8 | 5 | 1,5 | 5 | 6.78B | 13.6GB |
TP=2 的致命处:DP 变成 20,而 8 不整除 20 —— EP=8 直接不可行。 最好只能退到 EP=5(或 EP=4)。而专家参数占 92%(第 2 节反推的结论),EP 变小的损失远大于 TP 变大的收益:
TP=1, EP=8: X + E/8 = 2.90 + 4.01 = 6.91B → 13.8 GB
TP=2, EP=5: X/2 + E/5 = 1.45 + 6.42 = 7.87B → 15.7 GB ← 更差
TP=2, EP=4: X/2 + E/4 = 1.45 + 8.03 = 9.47B → 18.9 GB ← 差得多
所以 TP=2 不是"通信更贵",而是"显存更差"。 这个结论完全反直觉——TP 的教科书定位就是省显存,但在"参数 92% 在专家里 + EP 受整除约束"的组合下,它净亏。
(表里 TP=5 看起来最优,9.2GB。但 TP 要求 hidden size 和注意力头数被 5 整除,而这些数几乎总是 2 的幂——工程上不可行。这也是"漂亮的数学解被硬件现实否决"的一个好例子。)
EP:TP=1 下 DP=40,EP 可以取 1,2,4,5,8,10,20,40。约束 4 要求 EP ≤ 8 才不跨节点,所以EP=8 是最大可行值,把每卡参数压到最低的 6.91B。
DP:剩下的全给它,DP=40。
结论
TP=1 ← 会破坏 EP=8 的整除性,反而更吃显存
PP=1 ← micro batch 数只有 2,bubble 无法摊薄
EP=8 ← 节点内最大可行值,恰好等于单节点卡数
DP=40 ← 剩余全部
这套配置不是调出来的,是被约束推出来的——几乎是唯一解。
8. 常见坑
把 EP 当成"越大越好"。EP=10/20/40 在 TP=1 下都整除 40,显存也更省(EP=10 时每卡 6.11B),但一跨出 8 卡,all-to-all 就走 RoCE,性能雪崩。EP 的上限由物理拓扑决定,不由整除性决定。
忘记 EP 会缩小专家的 DP 组。这是双刃剑:EP 省了权重和梯度,但让优化器状态的分片效果差 8 倍(第 2 节那 8.4GB)。做显存预算时两边都要算。
先定 batch 再定并行。正确顺序是反的:先看序列长度决定 micro batch,再看 micro batch 数决定 PP 可行性。很多人先拍一个 PP=4 再发现 bubble 吃掉一半算力。
忽略整除约束就改一个维度。这套配置里任何一个数字都牵连其他三个。想加 CP=2?DP 立刻从 40 掉到 20,EP=8 随之失效。改并行配置前先把 TP × PP × CP × DP = world 和 EP | DP 两个式子写出来。
在 MoE 上照搬稠密模型的并行经验。稠密模型里 TP 是省显存的主力,MoE 里 EP 才是——因为参数绝大多数在专家里。
思考题
1. 现在假设集群从 5 节点扩到 8 节点(64 卡),其他条件不变(micro batch 仍是 1,global batch 允许调整):
- (a)
64 = 2⁶,TP=2 现在会让 EP=8 不可行吗?为什么? - (b) 如果 global batch 改成 128,PP=2 的 bubble 变成多少?PP 值得开了吗?
- © 你会怎么配这 64 卡?
2. 你想上 CP=2 来解决序列维的显存压力(第 2 节的 47.5GB)。在 40 卡上:
- (a)
TP × PP × CP × DP = 40,CP=2 时 DP 是多少? - (b) EP=8 还可行吗?如果不行,退到最大可行的 EP 是多少,每卡参数变成多少?
- © CP=2 省下的激活显存,和 EP 变小多吃的权重显存,哪个更多?(用第 2 节的数字估)
下一节讲通信层:NCCL 的五个集体原语、通信与计算怎么重叠、watchdog 超时机制,以及那次 5 小时静默死锁到底该怎么在 5 分钟内定位出来——py-spy 全集群扫栈、flight recorder、GPU 利用率形态学。