并行策略
对应课程 Lecture 7-8:Parallelism(Percy, Tatsu)
单卡放不下 7B 模型的训练状态(见资源估算),更放不下百亿千亿参数的模型本身。并行是把一个大训练问题拆到多张卡上的艺术。每一种并行都对应一个拆分维度:
┌─ 数据并行(DP):切数据
训练的四个切分维度 ──┼─ 张量并行(TP):切每一层的权重矩阵
├─ 流水线并行(PP):切层
└─ 序列并行(SP):切序列长度起点:数据并行与梯度累积
数据并行(Data Parallelism,DP) 最简单:每张卡放一份完整模型,各自处理不同的数据批次,反向传播后对梯度做 all-reduce(求和)再同步更新。
- 优点:实现容易,扩展近线性(只要计算够多);
- 缺点:每张卡都要放完整的模型 + 优化器状态 + 梯度,大模型根本放不下——这个瓶颈由 ZeRO 解决(见下)。
如果单卡放得下模型、只是塞不下大批次,用梯度累积(gradient accumulation):把大批次切成小批次,逐个算梯度累加,凑够一个逻辑批次再 step。注意它和 DP 的区别:累积是串行省显存,DP 是并行省时间,两者可以叠加。
ZeRO:分摊优化器状态
ZeRO(Zero Redundancy Optimizer) 系列思想:数据并行下每张卡都存一模一样的优化器状态,纯冗余。按需分片:
| 阶段 | 分片内容 | 每卡显存(bf16 混合精度,近似) |
|---|---|---|
| ZeRO-1 | 优化器状态 | 参数 + 梯度 + 状态/N |
| ZeRO-2 | 状态 + 梯度 | 参数 +(状态+梯度)/N |
| ZeRO-3 | 状态 + 梯度 + 参数 | 一切都 /N |
代价:前向/反向时要 all-gather 临时拼出完整的层参数,通信量增加。ZeRO-3 本质上等于“参数也切分的数据并行”,PyTorch 中的 FSDP 就是它的工业实现。
张量并行:把矩阵乘法切开
张量并行(Tensor Parallelism,TP) 把每一层的权重矩阵切开,让多张卡共同计算同一层。以 Megatron-LM 的经典做法为例,对
- 列并行:把
按列切成 ,每张卡算 ,输出天然是拼接关系,无需通信; - 行并行:把
按行切,输入 也要按列切,各自算完部分和后做 all-reduce 汇总。
把 MLP 和注意力的多个矩阵巧妙地按列/行交错安排,整个 Transformer 块每个前向只需要两次 all-reduce(MLP 一次、注意力一次),反向同样两次。
TP 用在哪一层?
TP 的通信发生在每个子层内部,对带宽要求极高——只适合放在 NVLink 互联的同一节点内(通常 4 或 8 卡一组)。跨节点用它会被通信时间吞掉。
序列并行(sequence parallelism) 是 TP 的好搭档:LayerNorm、Dropout 这类逐 token 的操作,让每张卡只处理序列的一段,进一步省激活值显存,通信与 TP 复用同一套 all-reduce / reduce-scatter。
流水线并行:按层切
流水线并行(Pipeline Parallelism,PP) 把模型的
- 朴素做法的灾难:前半段算的时候后半段闲着,GPU 利用率只有
; - GPipe:把一个批次切成若干微批次(micro-batch) 填满流水线,各卡交替做前向/后向;
- 1F1B(one-forward-one-backward):更省显存的调度——每张卡做完一次前向就做一次后向,让“在途”的激活值数量有界,而不是攒到最后才释放。工业界(如 Megatron)的标准选择。
流水线气泡(bubble) 无法完全消除:设段数为
激活值重计算
显存还是不够时,用激活值检查点(activation checkpointing):前向只保存少数检查点层的激活值,反向时从最近的检查点重新计算中间激活。典型配置多花约 33% 的计算,换取大量显存——和 PP/长序列搭配几乎是必需品。
三维/四维并行:组合拳
真实的百亿~千亿参数训练把所有维度叠起来:
总并行度 = DP × TP × PP (× SP)
例:512 卡 = 64(DP) × 8(TP, 节点内) × 8(PP, 跨节点)搭配原则:
- TP 最贵(每层通信)→ 放最小的组(节点内 NVLink)​;
- PP 便宜(段间通信)→ 跨节点​;
- DP 无处不在(数据随便加)→ 用剩余的卡​;
- 再配合 ZeRO-1(对优化器状态分片)、梯度累积控制全局批次大小。
通信原语速查
| 原语 | 作用 | 用在哪里 |
|---|---|---|
| all-reduce | 各卡求和,人人得到总和 | DP 梯度同步、TP 行并行 |
| all-gather | 各卡把自己的碎片拼给大家 | ZeRO-3 参数重组、序列并行 |
| reduce-scatter | 求和后每人拿一份碎片 | 序列并行的反向、FSDP |
| all-to-all | 按目的地重排数据 | MoE 专家路由、PP 传输 |
诊断与调优的思路
训练慢的时候,按顺序问:
- 时间花在算还是等? 看每步时间 ÷ 理想计算时间(
/卡算力)的利用率; - 通信占比? all-reduce 时间 vs 计算时间,判断 DP/TP 是否超配;
- 气泡占比? PP 调度是否合理,微批次数够不够;
- 显存瓶颈? 是否触发重计算过多、是否需要 ZeRO-2/3;
- 是不是卡在数据加载? 磁盘 IO 和预处理常常是被忽视的瓶颈。