大语言模型并行策略(三):张量并行
当模型的参数量较大时,模型参数、梯度、优化器状态以及激活值的内存占用量都可能超过单个 GPU 的显存容量,导致模型无法在单个 GPU 上进行训练和推理。为了支持更大参数量的模型,需要使用多张 GPU 进行并行计算。数据并行可以提升训练的吞吐量,但并没有减少单个 GPU 的显存占用量,而本文将要介绍的张量并行(Tensor Parallelism)可以将模型的参数切分到不同的 GPU 上,从而降低单个 GPU 的显存占用量,使得更大参数量的模型可以在多张 GPU 上进行训练和推理。
张量并行的原理
深度学习模型中,大部分计算都是矩阵乘法操作,而矩阵乘法的计算量和参与运算的矩阵大小有关。当某个矩阵的维度较大时,一方面会占用较多的内存,另外一方面也会占用较多的计算资源。张量并行的核心思想是将大矩阵切分为多个小矩阵,利用线性代数中分块矩阵乘法的思想,将这些小矩阵分布到不同的 GPU 上并行地完成计算。这一方面降低单个 GPU 中显存的占用,另外一方面能够利用多个 GPU 并行计算,提升计算速度。
考虑一个全连接层中的矩阵乘法 ,其中 是输入矩阵, 是权重矩阵, 是输出矩阵。
现在假设权重矩阵 过大无法完全放入单个 GPU 的显存中,此时可以将 按列切分为两个小矩阵,然后将切分后的权重矩阵分布到两个 GPU 上。此时全连接层的计算流程如下:
将权重矩阵按列维度切分并分布到不同的 GPU 上后,矩阵乘法可以拆分为两个小矩阵的乘法操作,每个 GPU 只需要独立计算自己负责的部分,计算完成后,对多个 GPU 上的输出矩阵进行汇总,得到最终的输出矩阵 。下面是张量并行的计算流程示意图:
两种切分方式
在张量并行中,对权重矩阵的切分有两种方式,按列切分和按行切分。下面我们以矩阵乘法 为例,分别介绍按行切分和按列切分的计算流程。
假设 的维度为 , 的维度为 , 的维度为 。
按列切分
按列切分的方式是将权重矩阵 从列的维度上切分为多个小矩阵,切分后的分块为 ,每个小矩阵的维度为 ,其中 是切分的份数。矩阵乘法的计算流程如下:
因为矩阵乘法可以想象为使用输入矩阵 的每一行对权重矩阵 的所有行做线性组合,因此权重矩阵的列之间是相互独立的,矩阵乘法可以在 GPU 之间并行计算。
使用这种切分方式时,输入矩阵 需要全部发送给每个 GPU,因为每个 GPU 都需要使用完整的输入矩阵 来计算自己的输出矩阵 。计算完成后,每个 GPU 的输出矩阵 是最终输出矩阵 的部分列,来自多个 GPU 的输出矩阵 需要在最后进行拼接,得到最终的输出矩阵 。
按行切分
按行切分将权重矩阵 从行的维度上切分为多个小矩阵,切分后的分块为 ,每个小矩阵的维度为 ,其中 是切分的份数。为了完成分块矩阵乘法,需要将输入矩阵 按列切分为多个小矩阵,切分后的分块为 ,每个小矩阵的维度为 。矩阵乘法的计算流程如下:
按行切分的原理稍微复杂一些,需要回顾一下矩阵外积的计算原理。将矩阵 的每一列与矩阵 的对应行做外积( 是 的第 列, 是 的第 行),每个外积得到一个 的矩阵,然后将这些矩阵求和,得到最终的输出矩阵 。假设 的列数为 ,则可以将矩阵乘法拆分为 个小矩阵的乘法操作,计算流程如下:
这里 是一个 的矩阵,多个 的结果求和后得到最终的输出矩阵 。
下图是矩阵外积的计算流程示意图:

图片来自:https://zhuanlan.zhihu.com/p/441943479
使用按行切分时,每个 GPU 持有权重矩阵 的部分行,在计算时只需要将输入矩阵 的对应列发送给对应的 GPU。计算完成后,每个 GPU 的输出矩阵 是一个 的矩阵,来自多个 GPU 的输出矩阵 需要在最后进行求和,得到最终的输出矩阵 。
通信模式
使用张量并行时,GPU 之间需要进行通信操作,将各个 GPU 的计算结果汇总,得到最终的输出矩阵,不同的切分方式对应不同的通信模式。
按列切分
使用按列切分时,每个 GPU 的输入是完整的输入矩阵,这需要使用 broadcast 的通信模式,将输入矩阵发送给每个 GPU。在每个 GPU 上,使用自己持有的权重矩阵计算输出矩阵的部分列。计算完成后,每个 GPU 的输出矩阵是最终输出矩阵的部分列,需要使用 All-Gather 的通信模式,将各个 GPU 的输出矩阵收集到一起,得到完整的输出矩阵。下面是一个简单的示意图,展示了按列切分时前向计算的通信模式:
在反向传播阶段,为了计算 的梯度,只需要回传输出矩阵 的梯度即可,输入矩阵 的梯度和权重矩阵 的梯度计算方式如下:
下面是一个简单的示意图,展示了按列切分时反向传播的通信模式:
输入矩阵 的梯度需要使用 All-Reduce 的通信模式,将各个 GPU 上的 的梯度进行求和,得到最终的 的梯度。
按行切分
使用按行切分时,首先需要将输入矩阵从列的维度上切分为多个小矩阵,然后使用 Scatter 的通信模式,将这些小矩阵发送给对应的 GPU。在每个 GPU 上,使用自己持有的权重矩阵计算输出矩阵。计算完成后,每个 GPU 的输出矩阵是一个完整的输出矩阵,需要使用 All-Reduce 的通信模式,将各个 GPU 的输出矩阵进行求和,得到最终的输出矩阵。下面是一个简单的示意图,展示了按行切分前向计算的通信模式:
在反向传播阶段,输入为 的梯度 ,而此前每个 GPU 计算出的结果是 ,但因为 ,因此 就等于 。
因为 ,所以 的梯度和 的梯度可以通过链式法则计算出来,其结果为:
下面是一个简单的示意图,展示了按行切分时反向传播的通信模式:
使用 broadcast 的通信模式发送给每个 GPU,然后每个 GPU 使用自己持有的权重矩阵计算出 和 。 可以在 GPU 上本地更新,而不需要与其他 GPU 进行通信。而 需要使用 All-Gather 的通信模式,从各个 GPU 上收集 ,以得到完整的 。
通信模式的对比
使用 TP 时,每个 GPU 首先需要将输入矩阵发送给其他 GPU,然后在每个 GPU 上计算出输出矩阵的一部分,最后将这些结果收集回来并汇总。
按列切分时,前向计算是使用 broadcast + all-gather,反向传播时使用 scatter + all-reduce。按行切分时,前向计算是使用 scatter + all-reduce,反向传播时使用 broadcast + all-gather。
| 切分方式 | 前向计算 | 反向传播 |
|---|---|---|
| 按列切分 | broadcast + all-gather | scatter + all-reduce |
| 按行切分 | scatter + all-reduce | broadcast + all-gather |
张量并行在 Transformer 中的应用
在 Transformer 架构中,张量并行主要应用于注意力机制和全连接层中,下面分别描述在这两种计算中如何使用张量并行。
注意力机制中的张量并行
首先回顾 Attention 机制的计算流程。输入特征 分别乘以三个权重矩阵得到查询(Query)、键(Key)和值(Value):
然后计算输出:
最后将输出 乘以一个投影矩阵 ,得到最终的输出:
目前实际应用中主要使用的是多头注意力机制(Multi-Head Attention, MHA),或者 Grouped Query Attention(GQA),整个计算流程如下:
由于不同注意力头的计算天然独立,因此非常适合从 head 维度切分。
可以将 、、 三个矩阵按 head 维度切分成 N 份(本质上是按列切分),然后将切分后的参数分布到不同的 GPU 上。每个 GPU 独立计算其所负责的注意力头的 、、,然后计算出对应的输出 。
当注意力计算完成后,每个 GPU 上的结果 是部分 head 的输出。而 还需要乘以一个投影矩阵 ,为了减少通信开销,可以将 按行切分成 N 份,每个 GPU 持有 的部分行。在各个 GPU 上将 的分片与对应的 的分片进行矩阵乘法,得到各个 GPU 上的局部结果。最终执行 All-Reduce 操作,将各个 GPU 上的局部结果进行求和,得到最终的输出。下面是 Attention 机制中张量并行的计算流程示意图:
多层感知机(MLP)中的张量并行
MLP 层在不同的 LLM 架构中存在差异,但基本都是可以看作是由两层全连接层组成。第一个全连接层将输入映射到高维空间,第二个全连接层将高维空间映射回输出空间。有些模型中会引入一个 Gate 机制,将第一个全连接层的输出与一个 Gate 向量进行逐元素相乘,然后再输入到第二个全连接层中。比如 Qwen3 中的 MLP 层计算流程如下:
gate = X * W_gate
up = X * W_up
x = F.silu(gate) * up
down = x * W_down
这里可以将 W_gate 和 W_up 从列维度拼接起来,这样 gate 和 up 的计算可以在同一个矩阵乘法中完成。因此无论是否有 Gate 机制,MLP 层的计算都可以看作是两个全连接层的计算。
为了减少在全连接层中的通信开销,可以对第一个矩阵乘法按列切分,这样每个 GPU 的计算结果是输出矩阵的部分列。考虑前面提到的按行切分的方式,输入矩阵 是按列切分的。我们只需要对第二个矩阵乘法进行按行切分,这样前一个矩阵乘法的输出就可以直接作为后一个矩阵乘法的输入。对第一个矩阵按列切分,对第二个矩阵按行切分,这样两个矩阵乘法之间不需要进行额外的通信操作。
下面是一个简单的示意图,展示了 MLP 层中张量并行的计算流程:
总结
张量并行的核心思想是将大矩阵切分为多个小矩阵,利用线性代数中分块矩阵乘法的思想,将这些小矩阵分布到不同的 GPU 上并行地完成计算。这一方面降低单个 GPU 中显存的占用,另外一方面能够利用多个 GPU 并行计算,提升计算速度。在 Transformer 中,Attention 模块和 MLP 模块均可以使用 TP 进行加速。为了降低通信开销,会组合使用列切分和行切分两种切分方式,这样在两个矩阵乘法之间不需要进行额外的通信操作。