WangYu::Space

cat /dev/mind

大语言模型并行策略(四):序列并行

分类:机器学习标签: LLM创建时间:2026-06-21 22:29:00

在训练模型时,为了能够计算梯度,需要在前向计算时保存每一层的激活值,模型中每一层的输入或输出的激活值会消耗大量的显存。对于大语言模型而言,激活值的显存占用量会随着序列长度线性增长。当输入序列较长时,激活值的显存占用量可能远超过模型参数、梯度和优化器状态的显存占用量,这就导致单个 GPU 很难处理超长序列的推理。

序列并行(Sequence Parallelism)是为了解决超长序列的训练而提出的一种并行策略。它的核心思想是将输入序列切分为多个子序列,然后将子序列分布到不同的 GPU 上进行计算。每个 GPU 只处理自己负责的子序列,这就降低了每个 GPU 的计算量和显存占用。

序列并行最早由 Reducing Activation Recomputation in Large Transformer Models 在 ACL 2023 上正式提出,本文参考了该论文的部分内容。

序列并行的原理

序列并行中,输入序列被切分为多个子序列,然后将子序列分布到不同的 GPU 上进行计算。和数据并行类似,每个 GPU 都持有完整的模型参数和优化器状态,但每个 GPU 只处理自己负责的子序列。由于激活值所用内存随序列长度线性增长,序列并行可以降低每个 GPU 的显存占用量,从而支持超长序列的训练和推理。

使用张量并行时可以将 Transformer 中的 Attention 和 MLP 模块的计算分布到多个 GPU 上进行计算,从而降低每个 GPU 的显存占用量。在 Attention 和 MLP 计算完成之后,每个 GPU 需要将张量并行计算的结果进行汇总,此时每个 GPU 上持有整个序列的完整激活值。

source: https://arxiv.org/pdf/2205.05198

上图中 Dropout 层和 LayerNorm 层的计算中,每个 GPU 的输入是完全相同的,是整个序列的隐状态,其维度为 B×S×HB \times S \times H,其中 BB 是 batch size,SS 是序列长度,HH 是隐状态的维度。虽然这部分的计算量比较小,但它们需要保存输入的激活值用于反向传播计算梯度。

另外我们可以观察到,在 TP 之外的计算中,如 Dropout 层和 LayerNorm 层,各个 token 是相互独立计算的。基于这个观察,我们可以从序列维度进行切分,将输入序列切分成多个子序列,在不同的 GPU 上完成 Dropout 和 LayerNorm 的计算。这样可以降低每个 GPU 在 Dropout 和 LayerNorm 上的计算量和显存占用。

Attention 模块的计算需要完整的序列,因此在 Attention 模块中,会使用张量并行的方式,此时会将输入从隐状态维度切分。其他模块的计算中,序列维度是相互独立的,因此可以使用序列并行的方式,将输入从序列维度切分。因此通常序列并行需要配合张量并行一起使用。

序列并行中的通信模式

不使用序列并行时,Transformer 的一个 block 中的计算流程如下:

其中 Attention 模块和 MLP 模块的计算可以使用张量并行,计算前后涉及到两次通信,分别是使用 broadcast 将输入分发到多个 GPU 上,以及使用 All-Reduce 将输出汇总到每个 GPU 上。

使用序列并行后,当 TP 部分计算完成后,每个 GPU 上持有完整序列的激活值,在接下来的 SP 部分,每个 GPU 只需要部分序列的激活值。因此在 TP 结束时,需要使用 Reduce-Scatter 将不同的子序列的结果分发到不同的 GPU 上。SP 部分计算完成后,每个 GPU 上持有部分序列的激活值,在接下来的 TP 部分,每个 GPU 需要完整序列的激活值。因此在 SP 结束时,需要使用 All-Gather 将不同的子序列的结果汇总到每个 GPU 上。

下图中,演示了 TP 和 SP 之间的通信模式:

上图中,红色和蓝色表示序列的不同子序列,在 SP 部分,每个 GPU 处理一个子序列。进入 TP 区域,需要使用 All-Gather 将不同的子序列的结果汇总到每个 GPU 上。TP 部分计算完成后,不需要使用 All-Reduce 将完整序列的结果汇总到每个 GPU 上,而是使用 Reduce-Scatter 将不同子序列的结果分发到不同的 GPU 上。

使用 SP 和 TP 结合时,Transformer 的一个 block 中的计算流程如下:

在 TP 部分计算完成后,此前需要使用 All-Reduce 将完整序列的结果汇总到每个 GPU 上。而使用了 SP 后,会使用 Reduce-Scatter 将不同子序列的结果分发到不同的 GPU 上。使用 SP 后,通信量并没有增加。

总结

序列并行通过将输入序列切分为多个子序列,并在不同的 GPU 上并行计算,从而降低每个 GPU 的显存占用量,支持超长序列的训练和推理。序列并行通常需要与张量并行结合使用,在计算 Attention 和 MLP 模块时使用张量并行,而在计算 Dropout 和 LayerNorm 时使用序列并行。不使用序列并行时,张量并行计算完成后,每个 GPU 上持有完整序列的激活值,每个 GPU 需要计算整个序列的 Dropout 和 LayerNorm。而使用序列并行后,每个 GPU 只需要计算自己负责的子序列的 Dropout 和 LayerNorm,从而降低每个 GPU 的计算量和显存占用。

评论 评论内容仅博主可见,不会公开显示)