如何不造出另一个 Megatron(一):流水线并行是怎么污染整个训练框架的
最近我的工作重心从推理系统的开发转移到了后训练系统的开发。看着原本推理引擎里面干净利落的代码,落在训练框架中却变成了弯弯绕绕,我的川字纹加深了。它跑起来很慢,用起来很不灵活,埋伏着各种陷阱,维护起来也很困难。这个训练框架本意是一个精简的系统,然而却朝着 Megatron 的复杂度发展了。
同事们和我几番尝试了重构,总是有些脏东西在阻碍我们把代码变得优雅。渐渐地,越来越多的线索指向了同一个方向——流水线并行(Pipeline Parallelism, PP)。
流水线并行带来的问题
气泡和调参
众所周知,流水线并行这个结构本身就引入了气泡。整个机器学习系统领域充满了无数研究员和气泡搏斗的故事。
- GPipe:开山之作,把 batch 切成多个 microbatch。
- PipeDream:1F1B,但是引入了权重新旧版本混用(有点 off-policy)的问题。
- PipeDream-Flush:修好了 1F1B。
- Interleaved 1F1B:stage 再切细一点,每个 rank 多拿几个 stage。
- TeraPipe:在 microbatch 内再沿序列维度切一刀。
- Zero Bubble:将 dgrad 和 wgrad 拆开。
- DualPipe:尽量掩盖跨 RDMA 的 EP 通信。
- (还有很多)
这里面的每一个算法,光是理解和实现,就已经要花不少力气了。
然后根据每个模型、硬件、数据,还得再调参。到底要切成几个 stage、几个 microbatch。模型的头尾跟中间的计算量不一样,还得调每个 stage 是几层。如果有多模态输入,还得再摘出来处理一下。
如果不想暴力或者拍脑袋一个一个去试,怎么调训练框架的参数也是一门很深的学问。怎么从各个组件的性能剖析数据中估计总体运行时间,怎么定义搜索空间,怎么建模成算法问题,怎么解 NP-hard 的规划问题,怎么构建模拟器……各种研究工作层出不穷:DAPPLE、Piper、Alpa、Galvatron、nnScaler、Metis、Zorse、SimAI、Charon……
后训练的动态性
以上还是假定了输入长度和数量固定的情况下。如果我们考虑后训练,尤其是强化学习后训练的需求,就会变得更加复杂,因为输入长度和数量都是动态的,有的情况甚至有可能数据量很小。按照一个输入分布调整的流水线参数,可能在另一个分布中又变差了。
Torchtitan 中流水线调度是 trainer 启动时就确定的:一轮包含多少个 microbatch、各 stage 何时执行计算和通信,以及需要准备多少份中间状态,都由一个固定的调度约束。后训练的输入规模却未必同样规整,某一步最终切出的 microbatch 数可能刚好超过调度的容量。多出来的输入不能直接塞进已经排好的流水线,只能再跑一轮。
更糟糕的是,溢出的部分可能很小,极端情况下甚至只有一个 microbatch,这就意味着溢出的这一轮会有极高的气泡。
于是大家又开始给 RL 训练系统修修补补:给流水线气泡做融合调度(RLHFuse)、把生成和训练异步化(PipelineRL、StreamRL)、对长尾 rollout 重新编排(RollPacker)……
仅前向传播
在后训练场景下,有的算法要求额外的仅前向传播的计算。同策略蒸馏(On-policy Distillation)和同策略自蒸馏(On-policy Self-distillation)都会引入 teacher 或 reference 模型的额外前向;重要性采样(Importance Sampling)、IcePop 这类算法也会让一次更新包含不同数值语义和依赖关系的计算 stage。
注意到仅前向传播的气泡比例和 1F1B 的是一样大的。如果朴素地加上额外的仅前向传播的计算,那么就相当于又额外增加了一些气泡。而如果要把这些额外的仅前向传播和模型训练正常的正反向传播合并在一起,那又需要针对每一个算法专门设计流水线并行的调度。
序列填充
刚开始看训练框架的代码的时候,我最不能理解的是为什么输入要填充成一个 [B, S] 的形状,而不是直接一个 [T]。早在 2022 年,Orca 就已经告诉了大家不用填充,直接把所有序列接在一起就行了;我 2023 年的博客文章也解释了怎么从第一性原理推导出这个结论。然而到了 2026 年,对于训练框架来说,序列填充依然是标配,拍扁接在一起反而变成了一项高级优化,叫做 Packed Sequences、Sequence Packing 或者 THD,而且还有诸多限制。要知道,在推理那边,可不会有一个推理引擎好意思把支持连续批处理(Continuous Batching)作为卖点,底下再标上小字:不是所有模型都支持此功能,不是所有配置都兼容此功能,开启此功能将与许多其他功能冲突;试验性功能,开启需谨慎。
不优雅
为什么我不喜欢序列填充呢?主要是因为这个根本就不是一个硬限制。
在整个模型的计算过程中,除了注意力机制是作用在“序列”的概念上的,其他的所有部分都是直接作用在词元上的。而注意力算子本身也支持不规则输入,比如 FlashAttention 里面的 varlen。本着“如无需要,勿增实体”的奥卡姆剃刀精神,特地把 [T] 填充成 [B, S] 就是一个不优雅的事情。
按照我的偏好,如果有什么部分在阻碍拍扁的输入,那应该把这个问题在局部解决了,而不是让填充变成一个全局的要求。
浪费、减少浪费
一旦加入了填充,势必引入无用的计算。然后为了尽量减少这些浪费,大家又在各个地方缝缝补补。
在 MoE 里面,这些填充词元有可能被丢给同样的专家,导致严重的负载不均衡。Megatron 里面有个优化,MoE 的路由里面可以接受一个 padding_mask,这样就可以省略填充词元。精妙!
——可是转念一想,如果我们一开始就不填充词元,那这个优化也没有存在的必要。
torchtune / torchtitan / FlexAttention
torchtune 曾经对填充采用过一种最让我眼前一亮的处理方法。下面这段写得非常好的文档注释(docstring)准确地展示了那套设计:
class PackedDataset(Dataset):
"""
Performs greedy sample packing on a provided dataset.
(...)
A packed sample is made up of individual smaller sequence length samples jammed together
within ``max_seq_len``. For example, if max_seq_len is 6 and there are varied
length samples::
tokens = [
[S1, S1, S1, S2, S2, pad],
[S3, S3, S4, S4, pad, pad],
...,
]
To prevent cross-contamination, the following mask would be returned for the
first pack in the example::
mask = [
[1, 0, 0, 0, 0, 0],
[1, 1, 0, 0, 0, 0],
[1, 1, 1, 0, 0, 0],
[0, 0, 0, 1, 0, 0],
[0, 0, 0, 1, 1, 0],
[0, 0, 0, 0, 0, 1],
]
The position ids would be::
input_pos = [
[0, 1, 2, 0, 1, 2],
[0, 1, 0, 1, 2, 3],
...,
]
复制完这段文档注释,我又被硬控了十分钟,因为这个类的功能实在是太逆天了。
- 它让算法复杂度增加了。
- 记每个序列的长度 \(s_i\),总共有 \(b\) 个序列,\(M\) 是打包后的长度,\(n\) 是这段代码打包完之后的行数。
- 注意力机制的复杂度本应该是 \(\Theta\!\left(\sum_{i=1}^{b} s_i^2\right)\)。
- 经过这段代码的处理,直接飙升到了 \(\Theta(n M^2)\)。
- 最坏情况:假设 \(s_i = 1\),\(b = M\)。复杂度本应是 \(\Theta\!\left(\sum_{i=1}^{M} 1^2\right) = \Theta(M)\),现在变成了 \(\Theta\!\left(\frac{\sum_{i=1}^{M} 1}{M} \cdot M^2\right)=\Theta(M^2)\),发生了质的变化!
- 退一步讲,假设 \(s_i = s\) 并且 \(M = k s\),那么每一行就是 \(k\) 个序列。复杂度本应是 \(\Theta(bs^2)\),现在变成了 \(\Theta\!\left( \frac{b}{k}(ks)^2 \right) = \Theta(bs^2 k)\),也就是劣化了 \(\Theta(k)\) 倍。
- 它构造出了 \(M^2\) 的注意力掩码
- 这实打实地占据了 \(M^2\) 的显存
- 也实打实地增加了 \(M^2\) 的访存
- 要知道自从 2022 年 FlashAttention 发明之后,额外空间可以随 token 数线性增长,而不必随序列长度平方增长。
- 它写了更多的代码,让程序变得更慢,显存占用更多
- 光是这个打包已经十分麻烦了,然后还得仔细把分块对角矩阵的构造写对。
- 如果不填充,直接就是把所有的词元拼接在一起,顺便记录好
cu_seqlens,然后注意力算子传入一个is_causal=True就解决了。
显然 torchtune 开发者也意识到了这点,于是在 PR 1193 中引入了 FlexAttention 和 torch.compile,直接提速 1.4 倍。torchtune 的这套打包方式也沿用到了 torchtitan 中,变成了 FirstFitPackingConfig。
就算性能能追上来了,我依然不是很喜欢这一套方案:
- 用 FlexAttention +
torch.compile快速实现各种新奇的注意力机制是一个很好的方案,但是用在常规的注意力机制里面就有点像高射炮打蚊子了。 - 对性能的保障实际上转嫁到了 FlexAttention 和
torch.compile的复杂性里面去了。 - 依然要编写复杂的打包代码,并构造等价的分块对角结构
get_efficient_causal_mask_mod_for_packed_document,只不过它现在是可编译的结构化元数据,而不是一张真的稠密矩阵。 - 听起来依然不优雅,像是一个打在补丁上的补丁。
易错
2024 年,Unsloth 在 Transformers 里面发现了一个 loss 计算错误。原因就是梯度累积的时候,每个 microbatch 各自的平均 loss 被等权地平均了,而没有按有效词元数加权。填充让每个 microbatch 的有效词元数参差不齐,这个等权平均就出了错。Megatron 也修过相同的问题。
- loss = F.cross_entropy(logits, labels, ignore_index=-100)
+ loss = F.cross_entropy(logits, labels, ignore_index=-100, reduction="sum")
+ loss = loss / num_tokens
如果不使用填充,而是让输入长度直接反映真实的有效词元数,那么在设计 loss 的时候就很难绕开一个问题:loss 到底应该按词元平均,还是按序列平均?无论选哪个,至少这是一个显式的设计决定。
这个 bug 容易出现,是因为在固定长度预训练里,每条序列、每个 microbatch 的有效词元数通常都一样。这时“先对每个 microbatch 求平均再平均”和“把所有词元放在一起求平均”恰好等价,于是这个区别长期被掩盖了。到了后训练,序列长度不同、填充后每个 microbatch 的有效词元数也不同,再沿用原来的规约方式就错了。
类似地,今年英伟达的文章里面,也提到了一个忘记扣掉填充词元带来的错误:
此前版本的 Megatron-Core 在计算 FLOPs 时并未考虑 THD 布局,而是假设 max_seqlen 就是有效序列长度,导致在变长场景下对 FLOPs 产生系统性的高估。
还是那个道理——如果数据结构从一开始就直接表达真实的有效词元数,这类错误至少会更难被藏起来。
固定形状
序列填充当然不是为了流水线并行发明的,数据加载器、算子也都有责任。但它能长期从一种局部实现选择,膨胀成整个训练栈的全局形状契约,流水线并行一定有责任。
Megatron 的 ModelParallelConfig 和 pipeline_parallel/schedules.py 都指出了,用上流水线并行最好用固定的 [B, S],不然会有性能损失:
# https://github.com/NVIDIA/Megatron-LM/blob/f2f0f7bfd88fcb1243df55275988d6af52daea35/megatron/core/model_parallel_config.py#L385-L389
class ModelParallelConfig:
###################
# Pipeline Parallel
###################
variable_seq_lengths: bool = False
"""Support for variable sequence lengths across microbatches. Setting this communicates the size
of tensors during pipeline parallelism communication, because of this extra overhead it
should only be set if the sequence length varies by microbatch within a global batch.
"""
# https://github.com/NVIDIA/Megatron-LM/blob/f2f0f7bfd88fcb1243df55275988d6af52daea35/megatron/core/pipeline_parallel/schedules.py#L116-L121
def get_forward_backward_func(...):
"""
seq_length (int, required): Sequence length of the current global batch. If this is a dual-stack
transformer, this is the encoder's sequence length. This is ignored if variable_seq_lengths
in the config is True. Otherwise, each microbatch in the current global batch size must use
this sequence length.
micro_batch_size (int, required): The number of sequences in a microbatch.
"""
torch.distributed.pipelining 文档明确说了,输入的形状必须是静态的:
A
PipelineStageneeds to know the input and output shapes for the stage model, so that it can correctly allocate communication buffers. The shapes must be static, e.g. at runtime the shapes can not change from step to step.
另外,PyTorch 会自动把整个 batch 切成 microbatch:
# https://github.com/pytorch/pytorch/blob/v2.13.0/torch/distributed/pipelining/schedules.py#L505-L527
class _PipelineSchedule(ABC):
@abstractmethod
def step(
self,
*args,
target=None,
losses: list | None = None,
return_outputs=True,
loss_kwargs: dict[str, Any] | None = None,
**kwargs,
):
"""
Run one iteration of the pipeline schedule with *whole-batch* input.
Will chunk the input into microbatches automatically, and go through the
microbatches according to the schedule implementation.
args: positional arguments to the model (as in non-pipeline case).
kwargs: keyword arguments to the model (as in non-pipeline case).
target: target for the loss function.
losses: a list to store the losses for each microbatch.
return_outputs: whether to return the outputs from the last stage.
loss_kwargs: extra keyword arguments forwarded to the loss function.
"""
显然如果没有 B 这个维度的话,就没法进行这个切分了。
对于预训练来说,静态形状输入是可以理解的,但是对于后训练来说,静态形状就变成了一个僵硬的限制。
实现上的缺陷还是本质限制?
其实我并不是很理解为什么 Megatron 和 PyTorch 的流水线并行要求输入的形状是固定的。在我看来,这个限制应该不本质,更像是实现上的缺陷,或者是没有设计好。
首先这个自动切 microbatch,确实没有 B 就没法自动切。但这并不代表不能手动切。可以手工构造每个 microbatch 的 cu_seqlens 以及手工截取输入序列。
第二是 PyTorch 文档里提到的通信缓冲区。在我的理解里面,分配缓冲区并不需要在乎是 [B, S] 两维还是拍扁成 [T] 一维;也不在乎 [T] 是不是可变的。只要 T 有一个上限,按照这个上限分配就行了。然而不管是 PyTorch 还是 Megatron,两边都为了迁就 NCCL send/recv API 的限制,要求使用固定大小的缓冲区。
我在之前的博客以及 MLSys 2026 的 fabric-lib 论文里面都提到了,RDMA 无论是双侧的 SEND/RECV 还是单侧的 WRITE 或 READ,都对缓冲区的大小没有一致性的要求,只要缓冲区足够大就行。另一方面,NVLink 上面的通信基于内存语义,同样也没有要求通信双方有一样的缓冲区大小。所以这个限制还是来自于 NCCL API,并不本质。
为什么 Megatron 的 variable_seq_lengths 文档注释说动态长度会有很大的性能损失呢?我快速看了一眼流水线并行的实现,有两个明显的原因:
# https://github.com/NVIDIA/Megatron-LM/blob/f2f0f7bfd88fcb1243df55275988d6af52daea35/megatron/core/pipeline_parallel/p2p_communication.py
# Simplified
class P2PCommunicator:
def _communicate_shapes(self, ...):
if is_sender:
send_prev_shape_tensor = torch.tensor(tensor_send_prev.size(), ...)
send_next_shape_tensor = torch.tensor(tensor_send_next.size(), ...)
send_prev_op = P2POp(isend, send_prev_shape_tensor, self.prev_rank, self.pp_group)
send_next_op = P2POp(isend, send_next_shape_tensor, self.next_rank, self.pp_group)
ops = [send_prev_op, send_next_op]
else:
recv_prev_shape_tensor = torch.empty((3,), ...)
recv_next_shape_tensor = torch.empty((3,), ...)
recv_prev_op = P2POp(irecv, recv_prev_shape_tensor, self.prev_rank, self.pp_group)
recv_next_op = P2POp(irecv, recv_next_shape_tensor, self.next_rank, self.pp_group)
ops = [recv_prev_op, recv_next_op]
work_list = torch.distributed.batch_isend_irecv(ops) # submit to cuda stream
for work in work_list:
work.wait() # cuda stream wait (non blocking)
if is_sender:
return [0, 0, 0], [0, 0, 0]
recv_prev_shape = recv_prev_shape_tensor.tolist() # block on D2H
recv_next_shape = recv_next_shape_tensor.tolist() # block on D2H
return recv_prev_shape, recv_next_shape
def _communicate(self, ...):
# shape
if config.variable_seq_lengths:
recv_prev_shape, recv_next_shape = self._communicate_shapes(...)
else:
recv_prev_shape, recv_next_shape = tensor_shape, tensor_shape
# payload
if is_sender:
send_prev_op = P2POp(isend, tensor_send_prev, self.prev_rank, self.pp_group)
send_next_op = P2POp(isend, tensor_send_next, self.next_rank, self.pp_group)
ops = [send_prev_op, send_next_op]
else:
tensor_recv_prev = torch.empty(recv_prev_shape, ...)
tensor_recv_next = torch.empty(recv_next_shape, ...)
recv_prev_op = P2POp(irecv, tensor_recv_prev, self.prev_rank, self.pp_group)
recv_next_op = P2POp(irecv, tensor_recv_next, self.next_rank, self.pp_group)
ops = [recv_prev_op, recv_next_op]
work_list = torch.distributed.batch_isend_irecv(ops)
if is_sender:
return None, None, work_list
return tensor_recv_prev, tensor_recv_next, work_list
- 每次点对点通信之前,都需要一次额外的通信来获取形状信息。
- 在提交形状信息的通信操作之后,调用了
.tolist()。这又带来了两个问题:一是触发了一个 D2H 通信;二是 CPU 的控制流被阻塞在这个 D2H 通信上了,实际上就是阻塞在了获取形状信息的通信上。
这么一番实现下来,性能损失可不是一点半点。但真的有必要这样实现吗?
对于流水线并行来说,每个 microbatch 在进入 pipeline 之前,它的 token 数和各 stage 之间需要传递的隐状态的形状都已经可以确定,不像专家并行(Expert Parallelism, EP)那样每一层都会变,因此没必要等到了每次需要在 stage 间传递的时候现问。
这么看起来,流水线并行的这些静态形状限制都只是实现上的问题,不是流水线算法的本质要求。
近期上游的修复
有意思的是,我写这篇文章的过程中,上游刚好修了几个上面提到的问题。
- pytorch PR 188500 Allow explicit pre-split pipeline microbatches:给
schedule.step()加上了arg_mbs、kwarg_mbs和target_mbs,允许调用方把 microbatch 预先切好再传进去,不再强迫调度器沿第 0 维替你切。这样我们就可以自己构造每个 microbatch 的cu_seqlens。 - torchtitan PR 3856 Always Pre-Split Microbatches for PP:把切 microbatch 的责任挪到数据加载器,让数据加载器构造 microbatch 的输入,删掉了“PP 与 varlen 不兼容”的限制。
- torchtitan PR 4121 fold batch dim:把整个数据路径从
[B, S]换成了[T],配置也从按序列数量表达,改成按每个 DP rank、每个 microbatch 的 token 预算表达。
这些修复证明了我前面关于静态形状的判断:[B, S]、自动按第 0 维切分 microbatch、PP 与 varlen 互斥,都不是流水线并行的本质限制,而是 API 选择和实现缺陷。
为了修掉这个局部限制,microbatch 的所有权从调度器移到了数据加载器,trainer、validator、TorchFT、Forge 和测试都要跟着改;检查点间隔之类原本按数据加载器步数表达的概念,也要重新解释。arg_mbs 和 kwarg_mbs 的元素仍然是 Any,模型输入类型仍然不能穿过 schedule.step() 得到静态检查。模型仍然被切成 model_parts,不同 stage 仍然收到不同语义的输入,只有最后一个 stage 才产生真正的 loss。
修掉了一部分历史遗留问题当然是一件值得庆祝的事情,但注意到这些修复依然是牵一发而动全身。流水线并行当然可以用正确的方法实现,问题是为了把它做对,整个框架都必须参与,整个框架的组件都需要理解流水线并行的概念。这是用了流水线并行不得不付出的维护代价。
工程负担
抛开流水线并行本身所需要的工程量不谈,流水线并行也会给整个训练框架的其他部分带来非常多的额外工程负担。这里主要以 torchtitan 举例子。
控制流分叉
开启了流水线并行之后,很多控制流都会跟不开流水线并行时完全不同。为了尽量让两者统一,很多时候正常路径也会变得复杂。
就拿最核心的模型构造和单步前向反向传播举例,哪怕写成过度简化的伪代码,都能看到大量的控制流分叉。甚至还能看到为了保持接口一致,返回一个假的 loss。而且类似的模式会在多个部件中重复。
# https://github.com/pytorch/torchtitan/blob/4c6481182c1c3d7815e3a119a271eebead94c71a/torchtitan/trainer.py
# Over-simplified
class Trainer:
def __init__(self, ...):
model = model_spec.model(config.model) # on meta device
# Fork: how model inits
if parallel_dims.pp_enabled:
pp_schedule, model_parts, pp_has_first_stage, pp_has_last_stage = \
model_spec.pipelining_fn(model, ...)
del model
for m in model_parts:
m.to_empty(device)
cast(BaseModel, m).init_weights()
m.train()
ensure_pp_loss_visible(parallel_dims, pp_schedule_name)
else:
pp_schedule, pp_has_first_stage, pp_has_last_stage = None, None, None
model = model_spec.parallelize_fn(model, ...)
model.to_empty(device)
cast(BaseModel, model).init_weights()
model.train()
model_parts = [model]
# Fork: where lm_head lives
if isinstance(loss_fn, ChunkedLossWrapper):
if parallel_dims.pp_enabled:
if pp_has_last_stage: # lm_head in PP last stage
loss_fn.set_lm_head(model_parts[-1].lm_head)
model_parts[-1]._skip_lm_head = True
else:
pass # non-last stage: no lm_head
else:
# lm_head in the only model part
assert len(model_parts) == 1
loss_fn.set_lm_head(model_parts[0].lm_head)
model_parts[0]._skip_lm_head = True
def forward_backward_step(self, ...):
# Fork: model forward + backward vs schedule step
if parallel_dims.pp_enabled:
with train_context():
# Fork: stage-dependent args passed to schedule step
if pp_has_last_stage:
targets, losses = labels, []
else:
targets, losses = None, None
if pp_has_first_stage:
pp_schedule.step(inputs, target=targets, losses=losses, ...)
else:
pp_schedule.step( target=targets, losses=losses, ...)
# Fork: stage-dependent loss computation
if pp_has_last_stage:
assert losses is not None
loss = sum(stack(losses)).to(device)
else:
loss = tensor([-1.0], device=device) # fake value
else:
assert len(model_parts) == 1
with train_context():
pred = model_parts[0](inputs, ...)
loss, _ = loss_fn(pred, labels, ...)
loss.backward()
return loss
class Validator:
def validate(self, model_parts, ...)
# Fork patterns similar to forward_backward_step
除了控制流分叉以外,上面这段 forward_backward_step 本身也是十分精妙。注意到里面多次判断了流水线并行的 stage,显得十分繁杂。看起来这不就是根据当前 rank 是处于前中后哪一个 stage,给定不同的输入输出而已嘛?为什么不换成下面这个更简单的形式呢:
if parallel_dims.pp_enabled:
with train_context():
if pp_has_first_stage:
pp_schedule.step(inputs, target=None, losses=None, ...)
loss = tensor([-1.0], device=device)
elif pp_has_last_stage:
losses = []
pp_schedule.step( target=labels, losses=losses, ...)
loss = sum(stack(losses)).to(device)
else:
pp_schedule.step( target=None, losses=None, ...)
loss = tensor([-1.0], device=device)
else:
...
实际上,如果考虑了 VPP,那么这段代码就是错误的,因为开启了 VPP 之后一个 rank 可能拥有多个 stage。
模型手术
在元设备(meta device)上构建完整个模型之后,torchtitan 会根据流水线 stage 对模型对象进行裁剪。往好了说,这充分利用了 Python 语言的动态特性。但是在我看来,这种动态操作就是非常脆弱和易错的,每次看到 setattr 我都捏一把汗。
# https://github.com/pytorch/torchtitan/blob/4c6481182c1c3d7815e3a119a271eebead94c71a/torchtitan/distributed/pipeline_parallel.py#L428
# simplified
def _split_module(whole_model: nn.Module, modules_to_keep: set[str]) -> nn.Module:
model = copy.deepcopy(whole_model)
for name, m in model.named_children():
if isinstance(m, (nn.ModuleDict, nn.ModuleList)):
layers_to_keep: set[str] = ...
if layers_to_keep:
# Keep only specified layers
if isinstance(m, nn.ModuleDict):
for layer_name in list(m.keys()):
if layer_name not in layers_to_keep:
del m[layer_name]
elif isinstance(m, nn.ModuleList):
indices_to_keep: list[int] = ...
new_layers = nn.ModuleList([
l for i, l in enumerate(m) if i in indices_to_keep])
setattr(model, name, new_layers)
else:
# No layers from this structure needed, set to empty structure
if isinstance(m, nn.ModuleDict):
setattr(model, name, nn.ModuleDict())
elif isinstance(m, ModuleList):
setattr(model, name, nn.ModuleList())
elif name not in modules_to_keep:
# Replace with None
setattr(model, name, None)
return model
经过这一番裁剪之后,某些子模块可能变成 None,模型的 forward 定义也变得复杂和混乱了起来:
# https://github.com/pytorch/torchtitan/blob/4c6481182c1c3d7815e3a119a271eebead94c71a/torchtitan/models/common/decoder.py#L262
# simplified
class Decoder(BaseModel):
"""Base class for autoregressive decoder-only language models."""
def forward(self, tokens: torch.Tensor, ...):
# Note: `tokens` is int token ids in the first stage,
# but becomes hidden states in later stages.
if self.tok_embeddings is not None:
h = self.tok_embeddings(tokens)
else:
h = tokens
# Note: all stages happen to have an iterable `layers`.
for layer in self.layers.values():
h = layer(h, ...)
# Note: only last stage has `norm`.
if self.norm is not None:
h = self.norm(h)
# Note: only last stage has `lm_head`
if self.lm_head is not None:
output = self.lm_head(h)
else:
output = h
# Note: `output` is hidden states in earlier stages,
# but becomes logits in the last stage.
return output
本来一个简简单单的 forward,愣是每一步都要检查一下 is not None。而且光看这段代码并不能知道哪些模块应该在哪些 stage 出现。
输入输出也很混乱。tokens 这个输入参数听起来像是 token ID,但实际上在后面的 stage 它是隐状态(hidden states)。output 在前面的 stage 是隐状态,而在最后 stage 是 logits。
概念外溢
训练框架其他部分需要针对流水线并行进行特殊适配。举几个例子:
- 原先的
model只是一个对象,现在变成了一个model_parts列表。 - 类似的,
CheckpointManager、Optimizer和LRScheduler都从单一对象变成了一系列的对象。 set_determinism原本只需要给所有 rank 设置相同的随机数种子,但在开启流水线并行之后,不同 stage 的 rank 需要使用不同的随机数种子。clip_grad_norm_需要单独处理流水线并行,因为流水线并行在每个 rank 上删掉了一部分参数。MetricsProcessor需要知道流水线并行的调度策略。
类型系统失效
前面的例子里面其实已经反映了很多让静态类型检查失去作用的场景了:
setattr这类动态操作肯定是不被静态检查保护了。- 由于控制流分叉,许多变量都变成可空的了。
- 一个对象变成了一个列表的对象。
- 在经过流水线并行的切分之后,原先一个
Decoder类型或者BaseModel基类变成了list[nn.Module],损失了基类所拥有的接口。于是之后又通过cast(BaseModel, ...)和cast(Decoder, ...)来获取相应的功能。完全没有类型系统的保障。
当然,这里面有些问题纯粹是 torchtitan 自己的。不过我们之前在改进训练框架的类型检查的时候,遇到了一个只要用了 torch.distributed.pipelining 就做不了的事情。
当时我主要是想改进 Model 抽象类的 Model.forward 抽象方法。
- 一个任务(Task)一般会有一个数据加载器,会产生跟这个任务有关的输入,我们不妨把这个类型叫做
InT。 - 这个任务会对应着一个特定的模型包装,一般来说继承自具体的模型架构。我们不妨把这个模型包装的类型叫做
Model。 - 每个
Model期待的 batch 输入张量都不一样,我们不妨把所有的张量打包成一个BatchT类型。 - 每个
Model也懂得如何将一个InT转换成BatchT,不妨把这个函数称为prepare_batch。 - 每个
Model的输出也可能各不相同,不妨称之为OutT。因为需要反向传播,所以我们要求每个OutT里面都有一个loss张量。
显然每个具体的 Model 实现都知道它自己对应的 InT、BatchT 以及 OutT 具体是什么类型。在 Rust 中,利用关联类型(Associated Types),我们可以定义出这样的泛型接口:
pub trait ModelOutput {
fn loss(&self) -> Tensor;
}
pub trait Model {
type In;
type Batch;
type Output: ModelOutput;
fn prepare_batch(&self, inputs: Self::In) -> Self::Batch;
fn forward(&mut self, batch: Self::Batch) -> Self::Output;
}
在 Python 中,虽然没有关联类型,但是可以用泛型(Generics)来模拟这个行为:
class ModelOutput(Protocol):
@property
def loss(self) -> Tensor: ...
class Model[InT, BatchT, OutT: ModelOutput](ABC):
@abstractmethod
def prepare_batch(self, inputs: InT) -> BatchT: ...
@abstractmethod
def forward(self, batch: BatchT) -> OutT: ...
然后我们考虑 Trainer.step() 大概的流程。基本上是读入下一组数据,进行转换,然后进行前向和反向传播。注意到 Trainer 本身其实并不在意这其中每一步的具体类型,只要它们各自能对上就行。在有了上面的抽象之后,我们便可以写下如下可以通过类型检查的代码:
class Trainer[InT, BatchT, OutT: ModelOutput]:
loader: Iterator[InT]
model: Model[InT, BatchT, OutT]
def step(self) -> OutT:
inputs = next(self.loader)
batch = self.model.prepare_batch(inputs)
out = self.model.forward(batch)
out.loss.backward()
return out
这样我们可以保证,就算 Trainer 的实现变得复杂,或者 Model 的输入发生改变,静态类型检查都能查出不一致的地方。
可惜如果用上 PyTorch 的流水线并行实现,这个美好的类型系统就没法使用了。self.model.forward() 要被替换成 pp_schedule.step()。
首先,前面说到,pp_schedule.step() 会自动切分 microbatch,因此它要求传入的参数是一个个单独列出的 Tensor。显然我们这里传入一个 dataclass 就没法适配。
第二,pp_schedule.step() 用 args 和 kwargs 抹去了类型信息(见前面章节贴出来的代码)。这样一来,哪怕是漏传参数、参数名写错这类简单错误,也不得不等到运行时才暴露。
PP=1
你可能会说,假设我不打算用流水线并行的话,那直接设置 PP=1 就好了,有必要把代码删掉吗?
确实,设置了 PP=1,很多性能上的不足和功能上的限制可能就解除了。但我依然觉得如果没有使用流水线并行的打算,那就应该把代码删掉。
- 如果你不再维护流水线并行的代码,那么那一部分的控制流将逐渐老化。等到未来某一天你真的想用上流水线并行的时候,有可能就没法正确地跑起来了。
- 如果你决定依然维护流水线并行的代码,那么你要付出高昂的维护成本和开发成本,因为你在维护的并不只是流水线并行本身的代码。每一个新的功能,每一个补丁,你都要花更多的时间设计,花更多的时间实现,以保证改动是跟流水线并行兼容的。
- 就算有了编程智能体(coding agent)的帮助,新代码的设计和实现,依然比删掉流水线并行相关代码的情况更难。这里并不是因为我们人类不够聪明或者因为智能体能力不够高,而是因为问题空间变大了。一个新功能一旦引入了新的执行路径,就像流水线并行这样,那么所有功能的组合数量就会呈指数增长,而每一个组合都是一条需要测试和维护的代码路径。
什么时候流水线并行仍然是正确答案
我全篇文章都在抨击流水线并行,但是这里也为流水线并行说几句公道话。
如果纵向扩展域(scale-up domain)显存大小不足、横向扩展域(scale-out domain)速度较慢,那么流水线并行依然有性能上的优势。
比方说 Hopper 系列的 NVLink 一般只能有 8 卡互联,配套的 RDMA 一般最多也就每卡 400 Gbps。在这样的情况下,如果又要训练万亿级别大模型,又要支持百万上下文,流水线并行是最直接的办法了。这也解释了为什么目前为止大部分的模型报告都在用流水线并行,因为 NVL72、B300 以及 800 Gbps RDMA 大规模铺开也没多久。
另一方面,预训练使用流水线并行也是合理的。毕竟预训练的输入形状固定、数据量大,就算用了流水线并行也能达到较高的性能。
总结
本文主要从性能调优和工程复杂性两个方面阐述了流水线并行的问题。
在几种并行方式里面,流水线并行是最不一样的一种。其他并行方式虽然也会切参数、激活和数据,引入通信甚至局部控制流,但大体保留了模型原来的执行结构;性能问题主要讨论的是计算能不能被加速、通信能不能被计算掩盖。流水线并行却直接把计算图沿层切开,对整个前向和反向的控制流进行全局修改。结果就是它的概念会一路泄漏到训练框架的各个角落。流水线并行本身也不会加速计算。为了节省显存、减少通信量,它付出的代价是额外的调度复杂度以及流水线气泡这样的结构性浪费。
考虑新硬件以及后训练工况,我决定删掉流水线并行。当然我也不是在本文一味抱怨,只管挖坑不管埋。下一篇博客我将会推导不用流水线并行怎么训练。