观文听傑

返回

上一篇把同一层的矩阵切到多卡,让超宽 MLP 不必在单卡聚齐。另一类模型不是“某层太宽”,而是层数太多、跨节点链路又不适合每层 collective。Pipeline Parallel(流水线并行,PP)把连续层分成 stages(阶段),只在边界传激活与梯度。

但把层各放一张卡并不会自动并行:若一个完整 batch 从头走到尾,后面的 GPU 先等,前面的 GPU 后等。本篇聚焦解决这一等待的核心:把 batch 切成 micro-batches(微批次),并用 GPipe 或 One-Forward-One-Backward(一次前向一次反向,1F1B)安排它们。

01 先区分三种“批次”#

设一次 optimizer update 的全局 batch 为 BglobalB_{global},数据并行大小为 DpD_p,每个 pipeline replica 的 batch 为 B=Bglobal/DpB=B_{global}/D_p。再切成 mm 个 micro-batches,每个大小为

bμ=B/m.b_\mu=B/m.

PP 中的 micro-batch 不是额外的 optimizer step。mm 个 micro-batch 的梯度共同构成一次更新;若中途 optimizer.step(),你改变了目标,也让后面的 micro-batch 使用更新后的权重,形成 weight staleness(权重陈旧/版本不一致)。

02 不切 micro-batch 时哪里在空转?#

三段模型 S0,S1,S2S_0,S_1,S_2 的单个 batch 前向:

时间 ->  t0  t1  t2
S0       F0  --  --
S1       --  F0  --
S2       --  --  F0
text

每个 stage 只有三分之一时间在算。切出多个 micro-batches 后,S0S_0S1S_1 处理 μ0\mu_0 时可处理 μ1\mu_1,这才形成流水线。

边界张量也必须明确。若 S0S_0 输出 A0Rbμ×L×DA_0\in\mathbb R^{b_\mu\times L\times D},forward 发送激活;backward 则从 S1S_1 收到同 shape 的 L/A0\partial\mathcal L/\partial A_0。参数只属于本 stage,不沿边界搬运。

03 GPipe:先全部前向,再全部反向#

以 3 stages、4 micro-batches 为例,Fj/Bj 表示 μj\mu_j 的前向/反向。简化到每格等时:

时间 ->  0   1   2   3   4   5   6   7   8   9   10  11
S0       F0  F1  F2  F3  --  --  --  --  B3  B2  B1  B0
S1       --  F0  F1  F2  F3  --  --  B3  B2  B1  B0  --
S2       --  --  F0  F1  F2  F3  B3  B2  B1  B0  --  --
text

这叫 fill-drain(填充—排空):先让所有 micro-batches 完成 forward,再启动 backward。优点是顺序直观;缺点是前段 stage 必须长期保留多个尚未反传的激活。

只看 forward,pp 个等速 stages、mm 个 micro-batches 需要 m+p1m+p-1 个时间格,理想气泡比例为

ρbubble=p1m+p1.\rho_{bubble}=\frac{p-1}{m+p-1}.

p=3,m=4p=3,m=4 时是 2/6=33.3%2/6=33.3\%;增大 mm 可降低比例,但 micro-batch 太小会让 GEMM 利用率下降,并增加通信启动次数。

04 1F1B:稳定期交替前向与反向#

1F1B 先 warmup(预热)填入若干 micro-batches;一旦某 stage 能对较早 micro-batch 反传,就交替做一次 forward 和一次 backward;最后 cooldown(冷却)排空。

flowchart LR
  W[warmup<br/>只做必要 forward] --> S[steady state<br/>交替 1F1B]
  S --> C[cooldown<br/>排空 backward]
  C --> U[所有 micro-batch 完成<br/>optimizer.step]
mermaid

与 GPipe 相比,经典 1F1B 并不神奇地消除所有气泡;主要收益是让 backward 更早开始,降低同时驻留的 activation 数量。对 stage rr,warmup 长度和 pipeline 深度有关,越靠前通常越需先送入更多 micro-batches。

05 四个 micro-batch 的梯度怎样仍等价?#

若每个 micro-batch loss 是元素平均 j\ell_j,且每块有效 token 数都相同,完整 batch 平均 loss 为

L=1mj=1mj.\mathcal L=\frac1m\sum_{j=1}^{m}\ell_j.

于是每次 backward 的梯度贡献应缩放 1/m1/m,或先累加 sum loss,最后除以完整有效 token 数。PyTorch 2.14 的 pipeline schedule 参数 scale_grads=True 默认按 micro-batch 数缩放梯度,应该与返回平均 loss 的 loss_fn 匹配;若 loss_fn 返回 sum,应设置 scale_grads=False 并自己按全局分母归一化。

长度不同、padding 数不同的语言模型不能简单平均 j\ell_j。正确目标仍是

L=jSjjnj,\mathcal L=\frac{\sum_j S_j}{\sum_j n_j},

其中 SjS_j 是有效 token loss sum,njn_j 是有效 token 数。还叠加 DP 时,分母必须跨 DP group 汇总,但不能在 PP stages 上把同一批 token 重复计数。

06 Stage 边界不是“在层列表中间切一刀”#

每个 stage 必须拥有自己使用的参数与 buffer,并定义完整 forward。边界还要传递后续真正需要的值:hidden states、attention mask、position ids 或 auxiliary loss,不能只传主 Tensor 就假设语义完整。

切分点要同时平衡:

维度不能只看什么应实际测什么
计算层数相同每 stage 前向/反向 wall time
参数参数量相同参数、optimizer、临时 workspace 峰值
激活一个 shape边界 bytes、同时驻留 micro-batch 数
通信网络标称带宽send/recv 延迟、等待与计算重叠

Embedding、首层输入、末端 LM head 和词表 loss 往往很不均匀。四等份层数不等于四个等速 stages,最慢 stage 会给所有其他 stages 制造周期性等待。

07 用一个手算调度判断激活峰值#

假设每个 micro-batch 在 S0S_0 产生 2 GB 必须保留到 backward 的激活,忽略重算。

  • GPipe 连续做完 4 个 forward 后,S0S_0 最多保留约 4×2=84\times2=8 GB。
  • 若 1F1B 在完成必要 warmup 后让最早 micro-batch 立即 backward,稳定期保留数被 pipeline 深度限制,而非随 mm 线性增长。

但若开启 Activation Checkpointing(激活检查点),保存量下降、backward 计算变长,原本均衡的 stages 可能失衡。调度、重算和切分点必须一起 profile,不能独立调优。

08 当前 PyTorch 的最小调度骨架#

PyTorch 2.14 的 torch.distributed.pipelining 提供 PipelineStageSchedule1F1B。Stage 负责通信 buffer、send/recv 与本段 backward;schedule 负责 micro-batch 顺序。

这段骨架假设 input_idslabels 可均匀按 batch 轴分成 4 份。当前 API 也允许用 arg_mbskwarg_mbstarget_mbs 传入已经切好的 micro-batches,适合 token 数不等或复杂输入结构。

PipelineStage 需要正确的边界 shape/dtype 来分配通信 buffers。当前文档支持首次 micro-batch 动态推断,也可传 example tensors 做静态约束;混合精度和 TP 会改变实际 dtype/layout,必须按运行时契约构造,否则会触发 PipeliningShapeError 或更隐蔽的通信错误。

09 训练与推理的调度不是一回事#

训练需要 backward,1F1B 的目标是安排前后向并限制激活驻留。自回归推理则有 token 间依赖:下一 token 要等上一 token 采样完成,单请求很难形成训练式 micro-batch pipeline。通常依靠多个并发请求、prefill/decode 分离或模型副本提高利用率。

因此“PP 训练吞吐高”不能推出“单请求生成延迟低”。推理还要传 KV cache 或在各 stage 保存自己层的 cache,测量指标应分为 time-to-first-token 与 inter-token latency。

10 与 TP、FSDP2 组合时谁切什么?#

3D parallelism(3D 并行)常写成 [dp,pp,tp]

  • TP group 内的 ranks 同算一个 stage 内的矩阵;
  • PP group 沿深度传同一 micro-batch 的激活/梯度;
  • DP group 的 pipeline replicas 处理不同样本并同步对应参数 shards。

一个样本不能既被错误地分给 TP ranks,又在 PP ranks 上重复计数。调试时打印每个 rank 的三维坐标、样本 ID、stage ID 和所有 process group 成员,比只打印全局 rank 更有用。

组合 FSDP2 时还要决定 all-gather 与 stage 执行的重叠;组合 TP 时 stage 边界可能是 DTensor。每增加一个维度,先用更小 world size 验证数值,再增加调度复杂度。

11 Checkpoint 必须保存流水线进度吗?#

通常只在完整 optimizer update 后声明 checkpoint 成功,此时所有 micro-batches 已排空,无需保存“正在管道中的激活”。需要保存每个 stage 的参数、optimizer state、scaler、学习率日程、数据游标、随机状态与 mesh/stage 映射。

若系统试图在任意 micro-batch 中途容错恢复,就必须记录权重版本、已完成的 forward/backward 集合与通信状态,复杂度陡增。工程上更常见的边界是让本轮失败并从上一个完整 update 重放,同时保证数据迭代器可重现。

12 一条最短的正确性与性能验证路径#

  1. 两 stages、两个 micro-batches、FP32、无 dropout,和未切分模型比较 logits、loss、每层梯度。
  2. 用样本 ID 日志证明每个 micro-batch 恰好经过所有 stages 一次。
  3. m 从 1 改到 2、4,完整 batch 与 loss 分母不变时,更新后参数应近似相同。
  4. 记录每个 stage 的 F/B 起止时间与 send/recv,画真实时间线并测 bubble。
  5. 分开记录参数、激活、通信 buffer 峰值与 tokens/s,再调整切分点和 mm

13 常见错误与最短调试路径#

症状常见原因最短检查
第二步突然 shape error最后一小批或序列长度改变边界 shape固定/补齐 shape,记录每个 micro-batch 元数据
loss 随 micro-batch 数缩小schedule 和 loss 都除以 mm对照 scale_grads 与 loss reduction
loss 随 micro-batch 数放大sum loss 未按全局有效 token 归一化汇总 Sj,njS_j,n_j 手算
GPU 呈周期性长空洞stage 不均衡或 mm 太小画每 stage 的 F/B/通信时间线
显存仍随 mm 线性增长使用 fill-drain 或 backward 启动太晚统计未反传 activation 数
进程永久等待stage 数、边界结构或调用顺序不一致对齐所有 rank 的 schedule 与 send/recv 日志
恢复后只有部分层改变每 stage 独立保存但缺少全局提交清单校验所有 shards 属于同一 update

14 GPipe、1F1B 与其他调度怎样选?#

GPipe 最容易理解和验证,但 activation 驻留高;1F1B 通常以相近气泡换更低内存。Interleaved 1F1B(交错 1F1B)让每个 rank 持有多个虚拟 stages,可改善负载与气泡,却增加通信和排序复杂度。Zero-bubble 调度进一步拆分 backward-input 与 backward-weight 来填空,前提是运行时能安全调度且两类 backward 成本合适。

不要因为调度名称更“先进”就直接采用。先确认瓶颈究竟是 activation、气泡、stage imbalance 还是网络;若 pp 很小、模型能单节点训练,PP 复杂度可能根本不值得。

15 今天真正需要记住什么?#

  1. PP 按层深度切分;micro-batches 让不同 stages 同时处理不同数据,但共同构成一次 optimizer update。
  2. GPipe 先全 forward 再全 backward,简单但保存更多激活;1F1B 更早反传,主要降低激活驻留。
  3. 增大 micro-batch 数可减气泡,却会缩小 GEMM、增加通信启动;切分点必须按真实时间和内存平衡。
  4. Loss 缩放、边界 shape/dtype、rank group 和 checkpoint 提交边界是最容易静默出错的接口。

16 思考题与小练习#

  1. 画出 4 stages、8 micro-batches 的 GPipe forward 时间线,计算理想 forward 气泡比例;若每格 20 ms,估算理想吞吐。
  2. 四个 micro-batches 的有效 token 数为 [8,4,8,2],loss sum 为 [16,12,8,6]。计算正确全局平均,并说明为何平均四个局部 mean 会错。
  3. 某三段模型每个 micro-batch 的 forward 时间为 [8,20,10] ms,边界传输各 3 ms。提出一种重新切层方案,并说明还需测哪些 backward 数据。

相关工作#

  1. Huang et al., GPipe: Efficient Training of Giant Neural Networks using Pipeline Parallelism,提出以 micro-batch 填充层间流水线的经典方案。
  2. Narayanan et al., PipeDream: Generalized Pipeline Parallelism for DNN Training,研究 1F1B 与权重版本问题。
  3. Narayanan et al., Memory-Efficient Pipeline-Parallel DNN Training,分析同步流水线的激活内存与调度。
  4. Narayanan et al., Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM,讨论交错流水线与 3D 并行。
  5. Qi et al., Zero Bubble Pipeline Parallelism,通过拆分 backward 工作填充流水线气泡。

17 下一篇预告#

DP、TP、PP 与 FSDP 的切分和调度已经连成一套训练架构。下一篇将从“算力不变为何训练仍变快”出发,进入 FlashAttention:如何利用 tiling 减少 HBM 读写,同时保持精确 attention 结果。

模型按层切开后 GPU 为何仍在等待?GPipe、1F1B 与 Micro-batch 调度
https://zwjcode.cn/blog/pipeline-parallel-gpipe-1f1b-microbatch-schedule
作者
发布于 2026年9月17日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。