Skip to content

并行策略

对应课程 Lecture 7-8:Parallelism(Percy, Tatsu)

单卡放不下 7B 模型的训练状态(见资源估算),更放不下百亿千亿参数的模型本身。​并行是把一个大训练问题拆到多张卡上的艺术。每一种并行都对应一个拆分维度:

text
                  ┌─ 数据并行(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 的经典做法为例,对 Y=XW

  • 列并行​:把 W 按列切成 [W1,W2],每张卡算 XWi,输出天然是拼接关系,无需通信;
  • 行并行​:把 W 按行切,输入 X 也要按列切,各自算完部分和后做 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) 把模型的 L 层切成几段(stage),各段放不同的卡。数据像工厂流水线一样流过:

  • 朴素做法的灾难​:前半段算的时候后半段闲着,GPU 利用率只有 1/(段数);
  • GPipe​:把一个批次切成若干微批次(micro-batch) 填满流水线,各卡交替做前向/后向;
  • 1F1B(one-forward-one-backward)​:更省显存的调度——每张卡做完一次前向就做一次后向,让“在途”的激活值数量有界,而不是攒到最后才释放。工业界(如 Megatron)的标准选择。

流水线气泡(bubble) 无法完全消除:设段数为 p、微批次数为 m,气泡占比约为 p1m+p1,所以微批次要足够多。PP 的通信只发生在段边界(传激活值),​量小、对带宽不敏感,适合跨节点​。

激活值重计算

显存还是不够时,用激活值检查点(activation checkpointing)​:前向只保存少数检查点层的激活值,反向时从最近的检查点重新计算中间激活。典型配置多花约 33% 的计算,换取大量显存——和 PP/长序列搭配几乎是必需品。

三维/四维并行:组合拳

真实的百亿~千亿参数训练把所有维度叠起来:

text
总并行度 = 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 传输

诊断与调优的思路

训练慢的时候,按顺序问:

  1. 时间花在算还是等? 看每步时间 ÷ 理想计算时间(6ND/卡算力)的利用率;
  2. 通信占比? all-reduce 时间 vs 计算时间,判断 DP/TP 是否超配;
  3. 气泡占比? PP 调度是否合理,微批次数够不够;
  4. 显存瓶颈? 是否触发重计算过多、是否需要 ZeRO-2/3;
  5. 是不是卡在数据加载? 磁盘 IO 和预处理常常是被忽视的瓶颈。

小结

  • 数据并行 + ZeRO 解决“状态放不下”;张量并行解决“单层算不动”(节点内);流水线并行解决“层数放不下”(可跨节点);序列并行和激活重计算管住激活值。
  • 一切并行策略的取舍都围绕三个量:​算力、显存、通信带宽​,对应资源估算GPU 与 TPU的 roofline 思维。
最近更新