
“Megatron”现在几乎成了大模型训练框架的代名词凡是涉及百亿参数以上的训练任务第一反应就是把模型切成几段rank 0 跑前几层rank 1 跑中间几层rank 2 跑最后几层层与层之间用 send/recv 把激活值一个接一个往下传。这是流水线并行Pipeline Parallelism, PP。它确实解决了“一张卡放不下整个模型”的物理问题但如果你照着 Megatron 的架构去搭自己的训练框架用不了多久就会品出一句扎心的结论流水线并行不是一个局部功能它更像一瓶泼出去的水会顺着框架的每一条结构缝往下渗最后所有子系统都得看它的脸色。这篇文章写给两类人一类是自己动手搭过分布式训练栈、被 PP 调度器改到怀疑人生的工程师另一类是正准备从零设计训练框架、但还在犹豫“把并行策略直接写进模型定义里到底行不行”的同学。我会把 PP 是怎么一步步“污染”调度、通信、显存、checkpoint、调试乃至后续扩展性的过程拆开讲清楚然后给出一个能让你少走几年弯路的解耦思路。1. 流水线并行为什么绕不开又为什么让人头疼1.1 模型大到一张卡装不下流水线并行成了“必经之路”先回到最原始的问题为什么一定要上流水线并行因为大模型训练遇到的第一道墙是显存墙。一张 H100 80G就算把参数、梯度、优化器状态全部做混合精度压缩一个 175B 参数的模型光参数和梯度就要吃掉远超过单卡显存的空间更别说中间还会产生巨大的激活值。遇到这种规模数据并行DP不够用因为每张卡都要放一份完整模型张量并行TP能拆但 TP 的通信量是 O(模型维度) 级别的每层都要做 all-reduce跨节点带宽立刻成为瓶颈。剩下的就是流水线并行把模型按层切成多个 stage每个 GPU 只负责其中一段。它的通信量远小于 TP只在 stage 边界传一次激活值和梯度理论上这是个“显存友好、通信友好”的方案。但注意这只是理论上。真正动手实现时你会发现PP 有一个绕不开的痛点bubble。如果整批数据串行地过一个又一个 stage前面 stage 算的时候后面在干等GPU 利用率惨不忍睹。所以就有了 GPipe 的 micro-batch 切分、PipeDream 的 1F1B 调度、Megatron-LM 的 interleaved 1F1B 等一堆补丁。到这里问题已经从“怎么让模型装下”变成了“怎么调度这么多 micro-batch 才能让流水线不空转”。1.2 Megatron 的成功反而不幸它把 PP 变成了唯一默认答案Megatron-LM 在 3D 并行DP TP PP上的工程实践非常成功成功到后来者几乎都把它的代码结构当作标准答案。你去看很多开源训练框架里的 pipeline 实现基本逻辑几乎一样手写if rank ...判断自己属于哪个 stage在 forward 里直接调用send/recv在训练循环里手动维护 warmup、steady、cooldown 三个阶段。这不是说 Megatron 写得不好。它的目标非常明确在 NVIDIA 自家集群上用固定结构的 Transformer把训练速度压榨到极致。在这种前提下把 PP 深度耦合进框架是合理取舍。但问题在于很多人都从“参考”变成了“复制”把这种为特定模型和特定硬件打磨的耦合结构搬进了自己的框架而后者的需求可能是多模型、多并行策略、动态结构。于是你造出的不是“另一个加速训练系统”而是“另一个 Megatron”——继承的不只是它的性能还有它的所有设计债。这篇文章标题里的“污染”指的就是这个一旦你在框架早期把 PP 以某种很具体的方式写死后面几乎每一个子系统都会被它牵着走。1.3 “污染”的准确定义耦合不是 bug是成本的转移这里要澄清一下我并不是说 PP 本身是脏东西也不是说框架不该支持 PP。我想说的是“污染”在软件工程意义上到底是什么当一个设计决策从它原本该待的位置执行策略扩散到了不该去的位置模型定义、数据流、调试接口时框架就失去了对等性uniformity和可组合性composability。举个例子本来一个模型的 forward 函数应该是“输入 tensor输出 tensor”这种干净语义。但在被 PP 污染的框架里forward 函数里会多出if dist.get_rank() 0、if is_pipeline_stage_last()、send_forward_recv_backward(...)这类分支。这等于把并行策略烙进了业务代码。以后你想复用这个模型做单卡测试不行。想换一种 schedule要改模型代码。想加一个没有 PP 的训练任务整个模型类都被迫背着一个用不上的包袱。这种耦合不是 bug因为它在特定条件下能跑得又快又稳。它是一笔成本转移把“框架实现并行的复杂度”转移给了“每个使用框架的人”。短时间看省事长时间看所有后来者都要为这笔债买单。2. 实际盘点流水线并行污染了框架里的哪些层2.1 调度器从“顺序循环”变成“状态机迷宫”一个没被 PP 污染的朴素训练循环长这样for batch in dataloader: loss model(batch) loss.backward() optimizer.step()简单、清晰、完全顺序执行。但一旦引入 PP训练循环就变成了一个三阶段状态机warmup 阶段只做 forward是为了填满流水线steady 阶段一个 forward 配一个 backward是 1F1B 的核心cooldown 阶段只做 backward是为了排空流水线。每一个 stage 处在哪个阶段取决于它自己在流水线中的位置和当前 micro-batch 的编号。最常见的实现方式是给每个 rank 写一个巨大无比的while循环里面塞满了“当前应该 send 还是 recv”“是先算前向还是先算后向”“这个 micro-batch 激活值要不要存下来”之类的分支判断。第一次写出来能跑通但一旦想换一种调度策略比如 interleaved 1F1B、zero-bubble就要把整个循环推翻重来。我在实际中见过最典型的污染表现框架的Trainer类里全是和 PP 相关的参数——num_micro_batches、num_pipeline_stages、pipeline_rank、schedule而真正和数据并行、优化器相关的逻辑反而被挤到一边。框架的入口不是“训练一个模型”而是“跑一条流水线”。调度已经从手段变成了目的这就是污染最直接的证据。2.2 通信层p2p 与集合通信杂糅在一起流水线并行在通信上和 DP、TP 有本质区别。DP 和 TP 用的都是集合通信collective communicationall-reduce、all-gather这些操作由通信库统一调度语义是“一组卡共同完成一件事”。而 PP 在 stage 边界用的是点对点通信point-to-point本质是“这张卡把数据传给下一张卡”是一个消息传递过程而且带依赖关系下游 stage 必须等上游算完才能开始。这两种通信混在同一个框架里最直接的后果是通信顺序和计算顺序被绑死了。为了不 deadlock每个 rank 必须严格按照预定义的先后次序调用 send/recv。一旦哪里顺序错了轻则卡死重则在多节点上出现 NCCL 超时、通信组初始化不一致等玄学问题。在我看过的许多被 PP 污染的框架里通信调用是散落在模型 forward 和 backward 各个角落的torch.distributed.send()在某两行之间、recv()在另一些函数里中间还夹着数学计算。调试的时候只能靠打印每个 rank 的执行流一条条对痛苦程度不亚于解一个随机死锁。更麻烦的是当你后续想做异步流水线overlap 通信和计算你会发现在这种代码结构里根本无处下手因为通信和计算在代码层面已经焊死了。2.3 显存规划microbatch 参数全面入侵PP 的另一个深坑是显存管理。模型并行通常要配合 activation checkpoint激活重计算来压显存但具体哪些层需要重算、哪些层不需要在 PP 场景下不能一概而论。首先每个 stage 只需要保存自己这部分的激活值这个天然省显存。但问题在于 micro-batch 数量如果你一个 batch 切成 16 个 micro-batch前向过程中可能会有多个 micro-batch 的激活值同时驻留在显存里具体驻留几个取决于调度策略。1F1B 能把驻留数量压到较低但 warmup 阶段依然会积累。为了把显存控制在稳定范围内很多人开始手动调num_micro_batches、activation_checkpoint_interval、recompute_frequency这些参数。这些参数一旦进入框架就会像病毒一样扩散显存分配器要认识它、调度器要处理它、用户配置文件要暴露它、连日志系统都要打印它。到最后框架的显存优化策略不再是从“计算图”出发做全局规划而是围绕“某一种 PP schedule”做人工打补丁。换一个 schedule显存表现立刻不同所有参数又得重新调一遍。我有一个印象很深的经历曾经为了跑一个 70B 模型团队花了整整两周调num_micro_batches和重计算策略最后发现瓶颈根本不是算力而是 pipeline stage 之间的激活驻留互相挤占显存。这件事本身可以用算法改进但当时框架的设计逼着我们只能靠手动调参绕路。2.4 checkpoint、日志和调试分布式状态不再对等没有 PP 时分布式训练框架的 checkpoint 可以做得“单机很像”每个 rank 保存自己那份数据并行分片最后汇总成一个完整模型。但 PP 会让状态对等性彻底破裂——每个 rank 只拥有模型的某一个纵向切片没有哪个 rank 天然拥有完整模型。于是你会遇到一堆问题保存 checkpoint 时要额外维护一个“stage 到参数映射表”恢复训练时每个 rank 要根据自己的 pipeline rank 加载对应切片日志里的 loss 如果每个 stage 都打印一遍你会发现不同 rank 的数值不一样因为计算中断点不同。为了给出一个“全局 loss”还得从最后一个 stage 做跨 rank 通信汇聚。更坑的是调试。框架不提供“按单卡语义跑一遍”的模式时你想验证模型定义对不对只能整个集群一起跑然后面对一堆分散在不同 rank 上的报错。TensorBoard 里的数值也是碎片化的需要自己手动拼接。这些都是 PP 这瓶水洒出来的“隐形面积”——表面上模型能训了但维护、排障、恢复的成本全在看不见的地方。3. 更隐蔽的代价被 PP 绑架后框架失去了哪些可能性3.1 动态控制流和异构网络结构被“劝退”PP 之所以实现简单是因为它默认了一个前提模型是一个规整的、串行的层列表。GPT 这种 decoder-only 结构正好满足但一旦模型里出现分支、条件判断、动态循环PP 的静态切分就会立刻失灵。举个例子如果你在做多模态模型视觉塔和文本塔可能结构不同、计算量不同合并之后还要过几层融合模块。这种模型怎么切 stage按层均分会导致两个塔的边界正好落在某个 module 中间切不碎按智能切分又需要非常了解每个算子耗时。在 PP 结构中这意味着用户必须手工指定每个 stage 包含哪些层而这些信息一旦写死在模型代码里后续模型结构一变切分方案又要重来。我在自研框架时见过更崩溃的情况模型里有 condition 控制流某个分支在 train 和 eval 时的行为不同。开发同学为了适配已经写死的 PP 切分不得不把控制流改成静态 mask整个代码的可读性和可维护性急剧下降。这个代价是用户用脚投票的新模型根本不敢往这个框架上迁移因为光切分就能耗掉好几周。3.2 MoE、异步调度、自动并行每一个都踩在 PP 的痛点上混合专家模型MoE是当下扩展模型容量的热门方向但它和 PP 的兼容性非常差。MoE 的核心是 token 级别的动态路由专家本身可能是稠密的、也可能被切到不同设备上。如果你把 PP 阶段固定成“前几十层算子”那么当某个 token 要路由到位于另一个 stage 的专家时就产生了跨阶段的小粒度通信。这种通信在 PP 模型里没有现成通道硬要做就只能绕回集合通信结果就是通信图变得极其复杂。异步流水线也是一个典型案例。异步流水线希望把 send/recv 和矩阵计算重叠起来缩短 bubble。但这要求调度器能提前预判哪一部分计算不依赖通信结果可以先把无关的算子发出去。如果你在框架早期把调度器和通信调用写死成同步顺序异步化改造就几乎等于重写一遍执行引擎。最可惜的是自动并行。这几年很多团队都在尝试“给定一个模型自动搜出最优的切分策略”。但如果框架在设计之初就把 PP 的 stage 切分方式焊死在用户代码里自动并行能发挥的空间就非常小——它只能优化已有的切分而不是从全局视角重新规划 DP/TP/PP 的组合。这就像一个房子已经按固定户型盖好了你再怎么重新装修也改不了承重墙的位置。3.3 框架团队的日常所有需求都被“bubble 怎么办”挡回来这是我在团队里听过最多的一句话。无论是想加新算子、做新的显存优化还是支持动态 shape评审时总会被问一句这个改动会让流水线 bubble 变大吗bubble 是 PP 的固有代价关注它本身没错。但当它变成一切架构决策的单一否决项时就说明框架已经失衡了。优化目标不是“训练系统整体吞吐最高”而是“当前这种 PP schedule 的 bubble 最小”。两者在局部可能是等价的但在全局往往互相矛盾。比如为了减少 bubble 把 micro-batch 变大可能导致显存不够、激活值溢出为了全局负载均衡支持异构 stage又可能让某一阶段计算时间差异变大。框架团队被这种问题困住本质上是框架被 PP 的局部最优钳制了。一个健康的训练框架应该让“并行策略”是可替换的优化选项而不是让整个团队围绕某个固定的调度方案打转。4. 解耦思路如何不造出另一个 Megatron4.1 第一原则模型定义只描述顺序计算不表达设备映射想摆脱 PP 污染第一步也是最难的一步立下规矩——用户的模型定义就是一段顺序计算逻辑里面不允许出现任何和“第几个 stage”“在哪张卡”相关的代码。forward 的输入是 tensor输出是 tensor中间调用的都是普通算子。class MyModel(nn.Module): def __init__(self): super().__init__() self.layers nn.ModuleList([TransformerBlock(...) for _ in range(24)]) self.norm nn.LayerNorm(...) def forward(self, x): for layer in self.layers: x layer(x) return self.norm(x)这段代码就是用户需要写的全部。它不知道 DP 存在不知道 TP 存在不知道 PP 存在。它只描述一件事数据从第 0 层流到第 23 层。这会带来一个直接的后果单卡调试成为可能。用户可以在单卡上先验证模型逻辑再交给框架去分布式执行。而框架侧并行策略完全由外部配置驱动模型代码本身不需要因为跑 1B 还是 100B 而改动。这条原则听起来简单但对已有框架是伤筋动骨的改动因为很多框架已经把并行逻辑写得到处都是。4.2 运行时三层分离Schedule、Stage、Transport做完了模型定义和并行策略的分离接着要分离运行时的三个职责。第一层是 Schedule调度器它决定“每个 micro-batch 什么时候 forward、什么时候 backward、和谁通信”。GPipe 的串行调度、1F1B、interleaved 都应该只是 Schedule 的实现类而不是框架的骨架。第二层是 Stage阶段抽象它负责任意一段连续子模型的执行封装“输入从哪来、输出到哪去”。Stage 不应该关心模型内部是什么结构它只需要接受一个 nn.Module 切片和它的输入输出接口。第三层是 Transport通信传输它负责 p2p send/recv 的细节包括通信顺序、异步化、buffer 管理。Schedule 不需要自己手写底层通信只需要告诉 Transport“我需要从 rank 1 拿数据给 rank 2。”三者的关系可以类比成模型定义是剧本Stage 是演员Schedule 是导演Transport 是剧场后勤。导演只负责喊开始结束演员只管演自己的部分后勤负责道具搬运。谁都不需要知道另外两方的完整实现细节。4.3 用户侧只需要三样东西切分注解、schedule 配置、统一执行入口落地到工程接口我认为用户侧需要暴露的东西越少越好但有三样不能省。第一样是切分注解。用户可以用一个极简的 API 标注出“哪些层可以被切到不同 stage”比如stage_boundary4表示第 4 层之后可以断开为一个 stage。框架拿到注解后自动生成每个 stage 的子模块引用而不是逼用户手写model[:4]。第二样是 schedule 配置。用户只需要选择用gpipe还是1f1b还是interleaved-1f1b以及指定micro_batch_size。至于每个 micro-batch 在某个阶段该 forward 还是 backward全部由 Schedule 实现内部处理。第三样是一个统一的执行入口。理想的用户体验是这样scheduler PipelineScheduler( modelmodel, # 但用户只传一台“模型坯子” partition[8, 8, 8], # 假设 24 层切 3 段每段 8 层 scheduleinterleaved-1f1b, microbatch_size4 ) scheduler.run(data_loader)run()内部会去创建 PipelineStage初始化 Transport决定每个 rank 该执行哪个子模块。用户不感知dist.get_rank()不感知send_forward_recv_backward更不需要在一个 nn.Module 里写 if-else。注意这张图里没有画任何“手动切分状态”的代码核心逻辑全部在框架侧完成。4.4 让 PP 成为编译期的一个 pass而不是运行时的到处 if比“提供接口”更进一步的思路是把 PP 当成编译优化 pass。就是说用户前端正常写模型框架内部拿到计算图后根据用户给的硬件拓扑和显存预算自动决定切分方案然后生成带有 send/recv 的分布式执行图。这其实就是 PyTorch 最近几年在推的torch.distributed.pipelining以及很多自动并行框架的方向。它把 PP 从“模型结构里长出来的东西”变成了“编译器后端做的一次变换”。用户写的是普通模型框架编译后往图里插入通信算子。好处是模型定义彻底干净单卡执行、多卡执行用同一份代码换 schedule 只是换一个编译选项自动并行可以在编译阶段做整体搜索。代价是编译器复杂度很高需要框架对计算图有完整掌控。如果你的框架还没到这一步至少可以先做到 4.1–4.3 的层次让外部接口保持干净。真正的目标是一致的PP 是一种“执行策略”它应该藏在框架内部而不是暴露给用户去维护。5. 一次实际解耦的复盘从被污染到可扩展5.1 症状清单什么信号说明你的框架已经被 PP 污染了我自己从一个“重度 Megatron 化”的框架里解耦过一次如果你不确定自己的框架有没有被污染可以对照下面这张清单症状判断标准模型代码里出现 rank 判断在forward里找dist.get_rank()、if is_first_stage等有就是污染训练循环无法独立于 PP 运行去掉 PP 支持后Trainer 类只剩一个空壳新增一个调度策略要改 20 文件说明调度逻辑和框架内部强耦合用户在配置里手动设置 micro-batch 数当“微批大小”成为所有实验必填项时显存和调度已经被焊死checkpoint 恢复依赖当前 rank 编号说明状态与设备绑定而非与逻辑数据绑定模型增加一层都要重新设计切分说明 stage 切分没有自动化这些症状如果中了三条以上基本可以判断框架已经从“训练框架”变成了“Megatron 定制版”。好的方面是问题暴露得早还好改如果等到多套模型都在上面跑起来了再重构成本会成倍增加。5.2 解耦实操我按这个顺序把 PP 从“结构”里捞了出来我当时做的第一件事是抽出所有模型代码里的 rank 判断和通信调用。把model.forward里与设备相关的代码全部删除得到一份“单卡语义”的模型。过程很痛苦但只要模型本身不是特别绕花一周左右就能清理完。第二步是抽象 PipelineStage。我把原框架中“一个 rank 对应一层 nn.Module 切片”的逻辑独立成类这个类只接受子模块、stage 编号、前后邻居编号。通信和计算被拆成两个接口step(input)负责算子计算tick(input)负责 send/recv。这样 schedule 的实现不需要关心模型内部。第三步是实现通用 Schedule。我先把 GPipe 和 1F1B 各写成一个类内部统一维护recv_buf/send_buf等状态。用户只需要在配置里指定 schedule 名称真正切换时只改一行配置。最后一步我加了编译期校验给用户的模型和切分注解做一个静态 pass检查每一层切分是否可行比如 stage 边界不能落在被多次调用的共享模块中间。校验不过就抛出明确的错误信息而不是在运行时报莫名其妙的 p2p 超时。5.3 常见问题速查与避坑记录解耦过程中最容易踩的坑我整理成一张速查表问题原因处理建议单卡开始正常多卡跑起来 loss 对不上有 BN 或某些带统计量的模块在 PP 下每个 stage 统计口径不同统一改用 LayerNorm 等不含跨 batch 统计的算子或在 stage 边界做参数同步出现 NCCL 超时或 hangp2p send/recv 顺序和调度逻辑不一致用确定性调度器先 recv 再 send 的顺序要在 schedule 里显式定义并给 Transport 增加超时自动 dump 现场引入异步通信后显存暴涨异步 send/recv 会占用额外的 buffer 等待数据区分“同步通信”和“异步通信”两种模式异步模式需要显式管理 buffer pool改一个 stage 内层数后后续 stage 训练效果突变stage 内部的参数初始化和输入分布不匹配增加跨 stage 的 warmup小学习率先跑几百步观察 loss 是否稳定下降用户模型里多模块共享参数tie_weight切分后同一参数被复制到多个 stage在编译期禁止跨 stage 的 tie或自动转成 all-gather 同步但要在文档明确指出性能代价最想强调的一条经验是别急着在框架初期就把 PP 的调度器优化到极致。先让整个架构保持单卡语义的可运行状态再逐步引入并行策略。这样即使后面要加新模型也能先在单卡上验证逻辑再交给并行运行时而不是一上来就在一个“充满 if rank 判断”的模型里大海捞针。另外如果你要在训练框架里同时支持多种模型架构稠密 Transformer、多模态、MoE请务必将 PP 设计为“可替换的调度策略”而不是“唯一的执行方式”。框架的核心应该是“怎么让计算图高效执行”而不是“怎么让流水线少冒泡”。我个人在实际动手时最大的体会是一个训练框架的长期竞争力不在于它把某一种并行策略调得多快而在于它能不能让新模型、新硬件、新调度方案快速接入。流水线并行是手段不是目的它应该像数据并行一样成为框架里一个普普通通的选项而不是定义框架形态的那根承重梁。如果你正准备设计自己的训练框架我建议把“如果明天要支持一种新的并行方式这个框架的核心结构会不会被动摇”当成最优先的架构问题来想。想清楚了你就不会在造另一个 Megatron 的路上越走越远了。