模型按层切开后 GPU 为何仍在等待?GPipe、1F1B 与 Micro-batch 调度
从流水线气泡出发,手算 GPipe 与 1F1B 的时间线、激活驻留和梯度累积,解释 stage 边界契约、PyTorch PipelineStage/Schedule1F1B 用法与调试方法。
上一篇把同一层的矩阵切到多卡,让超宽 MLP 不必在单卡聚齐。另一类模型不是“某层太宽”,而是层数太多、跨节点链路又不适合每层 collective。Pipeline Parallel(流水线并行,PP)把连续层分成 stages(阶段),只在边界传激活与梯度。
但把层各放一张卡并不会自动并行:若一个完整 batch 从头走到尾,后面的 GPU 先等,前面的 GPU 后等。本篇聚焦解决这一等待的核心:把 batch 切成 micro-batches(微批次),并用 GPipe 或 One-Forward-One-Backward(一次前向一次反向,1F1B)安排它们。
01 先区分三种“批次”#
设一次 optimizer update 的全局 batch 为 ,数据并行大小为 ,每个 pipeline replica 的 batch 为 。再切成 个 micro-batches,每个大小为
PP 中的 micro-batch 不是额外的 optimizer step。 个 micro-batch 的梯度共同构成一次更新;若中途 optimizer.step(),你改变了目标,也让后面的 micro-batch 使用更新后的权重,形成 weight staleness(权重陈旧/版本不一致)。
02 不切 micro-batch 时哪里在空转?#
三段模型 的单个 batch 前向:
时间 -> t0 t1 t2
S0 F0 -- --
S1 -- F0 --
S2 -- -- F0text每个 stage 只有三分之一时间在算。切出多个 micro-batches 后, 在 处理 时可处理 ,这才形成流水线。
边界张量也必须明确。若 输出 ,forward 发送激活;backward 则从 收到同 shape 的 。参数只属于本 stage,不沿边界搬运。
03 GPipe:先全部前向,再全部反向#
以 3 stages、4 micro-batches 为例,Fj/Bj 表示 的前向/反向。简化到每格等时:
时间 -> 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, 个等速 stages、 个 micro-batches 需要 个时间格,理想气泡比例为
时是 ;增大 可降低比例,但 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 ,warmup 长度和 pipeline 深度有关,越靠前通常越需先送入更多 micro-batches。
05 四个 micro-batch 的梯度怎样仍等价?#
若每个 micro-batch loss 是元素平均 ,且每块有效 token 数都相同,完整 batch 平均 loss 为
于是每次 backward 的梯度贡献应缩放 ,或先累加 sum loss,最后除以完整有效 token 数。PyTorch 2.14 的 pipeline schedule 参数 scale_grads=True 默认按 micro-batch 数缩放梯度,应该与返回平均 loss 的 loss_fn 匹配;若 loss_fn 返回 sum,应设置 scale_grads=False 并自己按全局分母归一化。
长度不同、padding 数不同的语言模型不能简单平均 。正确目标仍是
其中 是有效 token loss sum, 是有效 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 在 产生 2 GB 必须保留到 backward 的激活,忽略重算。
- GPipe 连续做完 4 个 forward 后, 最多保留约 GB。
- 若 1F1B 在完成必要 warmup 后让最早 micro-batch 立即 backward,稳定期保留数被 pipeline 深度限制,而非随 线性增长。
但若开启 Activation Checkpointing(激活检查点),保存量下降、backward 计算变长,原本均衡的 stages 可能失衡。调度、重算和切分点必须一起 profile,不能独立调优。
08 当前 PyTorch 的最小调度骨架#
PyTorch 2.14 的 torch.distributed.pipelining 提供 PipelineStage 与 Schedule1F1B。Stage 负责通信 buffer、send/recv 与本段 backward;schedule 负责 micro-batch 顺序。
import os
import torch
from torch import nn
from torch.distributed.pipelining import PipelineStage, Schedule1F1B
torch.distributed.init_process_group("nccl")
rank = torch.distributed.get_rank()
world = torch.distributed.get_world_size()
assert world == 2
device = torch.device("cuda", int(os.environ["LOCAL_RANK"]))
torch.cuda.set_device(device)
# build_stage_module 必须只返回本 rank 拥有的连续层
stage_module = build_stage_module(rank).to(device)
stage = PipelineStage(
stage_module,
stage_index=rank,
num_stages=world,
device=device,
)
def loss_fn(logits, labels):
return nn.functional.cross_entropy(
logits.flatten(0, 1), labels.flatten(), reduction="mean"
)
schedule = Schedule1F1B(
stage,
n_microbatches=4,
loss_fn=loss_fn,
scale_grads=True,
)
optimizer = torch.optim.AdamW(stage_module.parameters(), lr=3e-4)
optimizer.zero_grad(set_to_none=True)
if rank == 0:
# 只有第一 stage 接收完整输入;schedule 自动沿 batch 维切块
schedule.step(input_ids.to(device))
else:
# 只有最后 stage 持有 target 并计算每个 micro-batch loss
losses = []
logits = schedule.step(target=labels.to(device), losses=losses)
optimizer.step() # 所有 micro-batches 完成后,每个 stage 更新自己的参数python这段骨架假设 input_ids、labels 可均匀按 batch 轴分成 4 份。当前 API 也允许用 arg_mbs、kwarg_mbs、target_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 一条最短的正确性与性能验证路径#
- 两 stages、两个 micro-batches、FP32、无 dropout,和未切分模型比较 logits、loss、每层梯度。
- 用样本 ID 日志证明每个 micro-batch 恰好经过所有 stages 一次。
- 将
m从 1 改到 2、4,完整 batch 与 loss 分母不变时,更新后参数应近似相同。 - 记录每个 stage 的 F/B 起止时间与 send/recv,画真实时间线并测 bubble。
- 分开记录参数、激活、通信 buffer 峰值与 tokens/s,再调整切分点和 。
13 常见错误与最短调试路径#
| 症状 | 常见原因 | 最短检查 |
|---|---|---|
| 第二步突然 shape error | 最后一小批或序列长度改变边界 shape | 固定/补齐 shape,记录每个 micro-batch 元数据 |
| loss 随 micro-batch 数缩小 | schedule 和 loss 都除以 | 对照 scale_grads 与 loss reduction |
| loss 随 micro-batch 数放大 | sum loss 未按全局有效 token 归一化 | 汇总 手算 |
| GPU 呈周期性长空洞 | stage 不均衡或 太小 | 画每 stage 的 F/B/通信时间线 |
| 显存仍随 线性增长 | 使用 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 复杂度可能根本不值得。
15 今天真正需要记住什么?#
- PP 按层深度切分;micro-batches 让不同 stages 同时处理不同数据,但共同构成一次 optimizer update。
- GPipe 先全 forward 再全 backward,简单但保存更多激活;1F1B 更早反传,主要降低激活驻留。
- 增大 micro-batch 数可减气泡,却会缩小 GEMM、增加通信启动;切分点必须按真实时间和内存平衡。
- Loss 缩放、边界 shape/dtype、rank group 和 checkpoint 提交边界是最容易静默出错的接口。
16 思考题与小练习#
- 画出 4 stages、8 micro-batches 的 GPipe forward 时间线,计算理想 forward 气泡比例;若每格 20 ms,估算理想吞吐。
- 四个 micro-batches 的有效 token 数为
[8,4,8,2],loss sum 为[16,12,8,6]。计算正确全局平均,并说明为何平均四个局部 mean 会错。 - 某三段模型每个 micro-batch 的 forward 时间为
[8,20,10]ms,边界传输各 3 ms。提出一种重新切层方案,并说明还需测哪些 backward 数据。
相关工作#
- Huang et al., GPipe: Efficient Training of Giant Neural Networks using Pipeline Parallelism ↗,提出以 micro-batch 填充层间流水线的经典方案。
- Narayanan et al., PipeDream: Generalized Pipeline Parallelism for DNN Training ↗,研究 1F1B 与权重版本问题。
- Narayanan et al., Memory-Efficient Pipeline-Parallel DNN Training ↗,分析同步流水线的激活内存与调度。
- Narayanan et al., Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM ↗,讨论交错流水线与 3D 并行。
- Qi et al., Zero Bubble Pipeline Parallelism ↗,通过拆分 backward 工作填充流水线气泡。
17 下一篇预告#
DP、TP、PP 与 FSDP 的切分和调度已经连成一套训练架构。下一篇将从“算力不变为何训练仍变快”出发,进入 FlashAttention:如何利用 tiling 减少 HBM 读写,同时保持精确 attention 结果。