文章摘要
作者在搭建后训练系统时发现,训练框架代码逐渐臃肿冗余的根源是流水线并行。本文梳理了流水线并行的核心工程问题:本身存在运行气泡,调参难度大,后训练场景下动态输入、额外仅前向计算会进一步提升复杂度;同时流水线并行对输入固定形状的要求催生了序列填充,带来无效计算、易出错等问题,作者认为该限制属于实现缺陷而非本质要求。

在近期的工作中,我从推理系统开发转向了后训练系统搭建。对比之下,原本简洁清晰的推理引擎代码,在训练框架中却变得繁杂臃肿,这让代码维护的难度大幅提升。这套训练框架初衷是精简高效的,但却逐渐朝着复杂的方向发展,类似Megatron的架构复杂度不断攀升。团队和我多次尝试重构代码,但总有一些难以清理的冗余阻碍我们让代码变得优雅,渐渐地,所有的线索都指向了同一个根源——流水线并行(Pipeline Parallelism, PP)。

流水线并行的核心痛点

气泡与参数调优难题

众所周知,流水线并行的结构本身会引入运行气泡,整个机器学习系统领域都有无数研究者为解决这个问题投入了大量精力。

  • • GPipe[1]:首个将batch切分为多个microbatch的开创性工作
  • • PipeDream[2]:提出1F1B调度,但存在新旧权重版本混用的离线策略问题
  • • PipeDream-Flush[3]:修复了1F1B的相关问题
  • • Interleaved 1F1B[4]:将stage进一步切分,让每个rank承载多个计算阶段
  • • TeraPipe[5]:在microbatch内部沿序列维度进一步拆分
  • • Zero Bubble[6]:通过拆分dgrad和wgrad来消除气泡
  • • DualPipe[7]:优化跨RDMA的EP通信开销
  • • (还有更多相关优化方案)

每一种这类算法的理解和实现都需要投入大量精力,而针对不同的模型、硬件和数据场景,还需要进行针对性的参数调优。比如确定需要切分的stage数量、microbatch数目,根据模型各层计算量的差异调整每个stage包含的层数,对于多模态输入还需要额外的适配处理。

如果不想通过盲目尝试来确定参数,那么训练框架的调参本身就是一门复杂的学问:如何通过各组件的性能剖析数据估算整体运行时间、如何定义合理的搜索空间、如何将问题建模为算法优化任务、如何求解NP-hard的规划问题、如何构建性能模拟器,这类研究工作层出不穷,包括DAPPLE[8]、Piper[9]、Alpa[10]、Galvatron[11]、nnScaler[12]、Metis[13]、Zorse[14]、SimAI[15]、Charon[16]等。

后训练场景的动态性挑战

以上的讨论都是基于输入长度和数量固定的假设,但如果考虑后训练尤其是强化学习后训练的需求,情况会变得更加复杂:输入长度和数量都是动态变化的,部分场景下数据量甚至非常小。根据某一输入分布调整的流水线参数,在另一分布下的性能可能会大幅下降。

在Torchtitan中,流水线调度是在训练启动时就确定的:一轮训练包含的microbatch数量、各stage的计算和通信时序、需要准备的中间状态份数都由固定的调度约束决定。而后训练的输入规模往往并不规整,某一步生成的microbatch数目可能超过调度的容量,多出的输入无法直接塞入已编排好的流水线,只能单独再跑一轮。

更糟糕的是,溢出的部分可能很小,极端情况下甚至只有一个microbatch,这意味着溢出的这一轮会产生极高的运行气泡。为此,研究者们开始为RL训练系统添加各类修补方案:比如针对流水线气泡的融合调度[17]、将生成和训练异步化的设计[18,19]、对长尾rollout进行重新编排的方案[20]等。

仅前向传播的额外复杂度

在后训练场景中,部分算法需要额外的仅前向传播计算。例如同策略蒸馏和同策略自蒸馏都会引入teacher或reference模型的额外前向计算;重要性采样、IcePop这类算法也会让一次更新包含不同语义和依赖关系的计算stage。

需要注意的是,仅前向传播的气泡比例和1F1B调度的气泡比例相同。如果直接添加额外的仅前向传播计算,相当于进一步增加了气泡开销;而如果要将这些额外计算和模型训练的正反向传播合并,就需要针对每个算法专门设计流水线并行的调度方案。

序列填充的设计争议

在最初接触训练框架代码时,我一直无法理解为何输入需要被填充为[B, S]的形状,而不是直接使用[T]的扁平格式。早在2022年,Orca[25]就提出无需填充,直接将所有序列拼接即可;我在2023年的博客文章[26]也从第一性原理出发推导了这一结论。但到了2026年,序列填充依然是训练框架的标配,而将序列拍扁拼接反而成为一项高级优化,被称为Packed Sequences[27]、Sequence Packing[28]或THD,且通常带有诸多使用限制。

值得注意的是,在推理领域,不会有推理引擎将支持连续批处理作为卖点时,还附带小字说明:不是所有模型都支持此功能、不是所有配置都兼容此功能、开启后会与其他功能冲突、属于试验性功能需谨慎使用[29]。

填充设计的不优雅之处

我不喜欢序列填充的核心原因在于,这并非一个必要的硬限制。在整个模型的计算流程中,除了注意力机制是基于“序列”概念的,其他所有部分都是直接作用于词元的。而注意力算子本身已经支持不规则输入,例如FlashAttention中的varlen模式。按照奥卡姆剃刀原则,将[T]格式的输入强行填充为[B, S]的格式,本身就是一种不必要的复杂设计。

根据我的设计偏好,如果存在某些部分阻碍了扁平格式的输入,应该在局部解决该问题,而不是将填充变为全局的强制要求。

填充带来的无效计算与负载失衡

一旦引入填充,必然会产生无用的计算开销。为了尽量减少这类浪费,开发者们需要在各个环节进行额外的优化。

在MoE架构中,填充的词元可能被分配到相同的专家,导致严重的负载不均衡。Megatron中有一项优化,允许MoE的路由层接受padding_mask,从而跳过对填充词元的处理,这一设计看似精妙,但转念一想,如果最初就不使用填充,这类优化本就没有存在的必要。

Torchtune的序列打包方案

Torchtune曾采用过一种令人眼前一亮的填充处理方式,其PackedDataset类的文档注释清晰展示了这套设计:

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$量级的显存空间
    • <li style="margin: 0.5em 0;会带来$M^2$量级的访存开销</li> <li style="margin: 0.5em 0;自2022年FlashAttention提出后,额外空间开销可以随token数线性增长,无需随序列长度平方增长
  • • 该方案需要编写更多代码,进一步增加了程序的运行时长和显存占用
    • 序列打包的实现本身就较为复杂,还需要正确构造分块对角矩阵的结构
    • 若不使用填充,只需将所有词元直接拼接,记录好cu_seqlens,并在注意力算子中传入is_causal=True即可完成处理

Torchtune的开发者显然也意识到了这些问题,在PR 1193[32]中引入了FlexAttention和torch.compile,直接带来了1.4倍的性能提升。该打包方案也被沿用到了Torchtitan中,演变为FirstFitPackingConfig。

即便性能得到了提升,我依然不认可这套方案:

  1. 使用FlexAttention+torch.compile来实现常规注意力机制,更像是用高射炮打蚊子,并非最优选择
  2. 对性能的保障被转嫁到了FlexAttention和torch.compile的复杂性中
  3. <li style="margin: 0.5em 0;依然需要编写复杂的打包代码,并构造等价的分块对角结构get_efficient_causal_mask_mod_for_packed_document,只是将原本的稠密矩阵替换为了可编译的结构化元数据</li> <li style="margin: 0.5em 0;整体设计依然不够优雅,更像是在补丁之上叠加的又一层补丁

填充带来的易出错问题

2024年,Unsloth在Transformers库中发现了一处loss计算错误[34,35]。问题的根源在于梯度累积时,每个microbatch的平均loss被简单等权平均,而没有按照有效词元数进行加权。填充导致每个microbatch的有效词元数参差不齐,这种等权平均的方式就会产生错误,Megatron也修复过同类问题[36]。

# 错误的实现
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的有效词元数也存在差异,继续沿用原来的规约方式就会出现错误。

类似地,今年的一篇英伟达文章[37]中也提到了一个因未扣除填充词元导致的错误:此前版本的Megatron-Core在计算FLOPs时未考虑THD布局,而是假设max_seqlen即为有效序列长度,导致在变长场景下系统性高估了FLOPs数值。

此前版本的Megatron-Core在计算FLOPs时并未考虑THD布局,而是假设max_seqlen就是有效序列长度,导致在变长场景下对FLOPs产生系统性的高估。

这再次印证了一个道理:如果数据结构从一开始就直接表达真实的有效词元数,这类错误会更难被隐藏。

固定形状的强制限制

序列填充当然不是为流水线并行而发明的,数据加载器和算子本身也有责任,但它能够从一个局部实现选择演变为整个训练栈的全局形状契约,流水线并行难辞其咎。

Megatron的ModelParallelConfig和pipeline_parallel/schedules.py都指出,使用流水线并行时最好使用固定的[B, S]格式输入,否则会带来性能损失:

# 摘自Megatron-LM的model_parallel_config.py
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.
    """

摘自pipeline_parallel/schedules.py

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 PipelineStage needs 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:

# 摘自PyTorch的pipelining/schedules.py
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维度,就无法进行自动的batch切分。对于预训练场景,静态形状输入是可以接受的,但对于后训练来说,静态形状就变成了僵硬的限制。

实现缺陷而非本质限制

我始终无法理解为何Megatron和PyTorch的流水线并行要求输入形状固定。在我看来,这类限制并非流水线算法的本质要求,更像是实现上的缺陷或设计不足。

首先,自动切分microbatch确实依赖B维度,但这并不代表无法手动进行切分。我们可以手工构造每个microbatch的cu_seqlens,并手动截取输入序列。

其次是PyTorch文档中提到的通信缓冲区问题。在我的理解中,分配缓冲区并不需要关心输入是[B, S]还是拍扁为[T]格式,也不需要关心[T]是否可变。只要T存在上限,按照该上限分配缓冲区即可。

然而无论是PyTorch[38]还是Megatron[39],为了适配NCCL的send/recv API限制,都要求使用固定大小的缓冲区。我在之前的博客[40]和MLSys 2026的fabric-lib论文[41]中都提到,RDMA无论是双侧的SEND/RECV还是单侧的WRITE/READ,都对缓冲区大小没有一致性要求,只要缓冲区足够大即可。NVLink上的通信基于内存语义,同样不要求通信双方的缓冲区大小一致,因此这类限制本质上来自NCCL API,并非算法本身的问题。

为何Megatron的variable_seq_lengths文档注释称动态长度会带来显著性能损失?通过查看流水线并行的实现,我发现两个核心原因:

# 摘自Megatron-LM的p2p_communication.py
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
  1. 每次点对点通信前,都需要额外的一次通信来获取形状信息
  2. 这类实现带来的性能损失确实显著,但这真的有必要吗?对于流水线并行来说,每个microbatch在进入流水线前,其token数和各stage间传递的隐状态形状就已经可以确定,不像专家并行那样每层都会变化,因此无需在每次传递时都查询形状信息。

    由此可见,流水线并行的静态形状限制本质上是实现问题,而非算法的本质要求。

    上游社区的近期修复

    有趣的是,在我撰写本文的过程中,上游社区刚好修复了几个相关问题:

    • PyTorch PR 188500[42]:Allow explicit pre-split pipeline microbatches,为schedule.step()添加了arg_mbs、kwarg_mbs和target_mbs参数,允许调用方预先切分microbatch后再传入,不再强制调度器沿第0维进行切分,这样我们就可以自行构造每个microbatch的cu_seqlens
    • Torchtitan PR 3856[43]:Always Pre-Split Microbatches for PP,将切分microbatch的责任转移到数据加载器,由数据加载器构造microbatch输入,移除了“PP与varlen不兼容”的限制
    • Torchtitan PR 4121[44]: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值。这类模式会在框架的多个组件中重复出现:

    # 摘自Torchtitan的trainer.py
    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:
                        # non-last stage: no lm_head
                        pass
                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。

    模型的动态裁剪与复杂性

    在元设备上构建完整模型后,Torchtitan会根据流水线stage对模型对象进行裁剪。从好的方面来说,这充分利用了Python语言的动态特性,但在我看来,这类动态操作非常脆弱且容易出错,每次看到setattr调用都让我感到担忧。

    # 摘自Torchtitan的pipeline_parallel.py
    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方法也变得复杂混乱:

    # 摘自Torchtitan的decoder.py
    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方法,现在需要在每一步都检查模块是否存在。仅通过这段代码,我们无法直观知道哪些模块应该出现在哪些stage中。同时输入输出的语义也变得混乱:tokens参数看似是token ID,但在后续stage中会变为隐状态;output在早期stage是隐状态,在最后stage则变为logits。

    概念向整个框架的外溢

    训练框架的其他部分都需要针对流水线并行进行特殊适配,举例来说:

    • 原本单一的model对象,现在变为model_parts列表
    • 类似CheckpointManager、Optimizer和LRScheduler,都从单一对象变为一系列对象
    • set_determinism原本只需为所有rank设置相同的随机种子,但开启流水线并行后,不同stage的rank需要使用不同的随机数种子
    • clip_grad_norm_需要单独处理流水线并行的情况,因为每个rank上会删除部分参数
    • MetricsProcessor需要了解流水线并行的调度策略

    类型系统的失效

    前面的例子已经反映了多个让静态类型检查失效的场景:

    • setattr这类动态操作无法被静态检查保护
    • <li style="margin: 0.5em 0;控制流分叉导致大量变量变为可空类型</li> <li style="margin: 0.5em 0;对象从单一实例变为列表实例
    • 当然,部分问题仅存在于Torchtitan的具体实现中。不过我们在改进训练框架的类型检查时,曾遇到过只要使用torch.distributed.pipelining就无法解决的问题。

      当时我们尝试改进Model抽象类的forward抽象方法:

      我们可以做如下抽象定义:

      • 一个任务通常会有一个数据加载器,产生与任务相关的输入,我们将其类型称为InT
      • <li style="margin: 0.5em 0;该任务对应特定的模型包装,通常继承自具体的模型架构,我们将其类型称为Model</li> <li style="margin: 0.5em 0;每个Model期待的batch输入张量各不相同,我们将所有张量打包为BatchT类型 <li style="margin: 0.5em 0;每个Model也知道如何将InT转换为BatchT,我们将该函数称为prepare_batch</li> <li style="margin: 0.5em 0;每个Model的输出也各不相同,我们将其称为OutT,且要求每个OutT包含一个loss张量以支持反向传播

      每个具体的Model实现都知道自身对应的InT、BatchT和OutT的具体类型。在Rust中,可以通过关联类型定义如下泛型接口:

      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中,虽然没有原生的关联类型,但可以通过泛型来模拟:

      class ModelOutput(Protocol):
          @property
          def loss(self) -> Tensor: ...
      

      class ModelInT, BatchT, OutT: ModelOutput:
      @abstractmethod
      def prepare_batch(self, inputs: InT) -> BatchT: …
      @abstractmethod
      def forward(self, batch: BatchT) -> OutT: …

      我们可以为Trainer.step()编写如下类型安全的代码:

      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
      

      但只要使用torch.distributed.pipelining,这类类型安全的抽象就无法实现,因为流水线并行会将模型拆分为多个部分,破坏了原有的类型抽象。

以上内容不代表本平台立场,仅供读者参考