WangYu::Space

cat /dev/mind

大语言模型并行策略(三):张量并行

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

当模型的参数量较大时,模型参数、梯度、优化器状态以及激活值的内存占用量都可能超过单个 GPU 的显存容量,导致模型无法在单个 GPU 上进行训练和推理。为了支持更大参数量的模型,需要使用多张 GPU 进行并行计算。数据并行可以提升训练的吞吐量,但并没有减少单个 GPU 的显存占用量,而本文将要介绍的张量并行(Tensor Parallelism)可以将模型的参数切分到不同的 GPU 上,从而降低单个 GPU 的显存占用量,使得更大参数量的模型可以在多张 GPU 上进行训练和推理。

张量并行的原理

深度学习模型中,大部分计算都是矩阵乘法操作,而矩阵乘法的计算量和参与运算的矩阵大小有关。当某个矩阵的维度较大时,一方面会占用较多的内存,另外一方面也会占用较多的计算资源。张量并行的核心思想是将大矩阵切分为多个小矩阵,利用线性代数中分块矩阵乘法的思想,将这些小矩阵分布到不同的 GPU 上并行地完成计算。这一方面降低单个 GPU 中显存的占用,另外一方面能够利用多个 GPU 并行计算,提升计算速度。

考虑一个全连接层中的矩阵乘法 Y=X×WY = X \times W,其中 XX 是输入矩阵,WW 是权重矩阵,YY 是输出矩阵。

现在假设权重矩阵 WW 过大无法完全放入单个 GPU 的显存中,此时可以将 WW 按列切分为两个小矩阵,然后将切分后的权重矩阵分布到两个 GPU 上。此时全连接层的计算流程如下:

全连接层按列切分示意图

将权重矩阵按列维度切分并分布到不同的 GPU 上后,矩阵乘法可以拆分为两个小矩阵的乘法操作,每个 GPU 只需要独立计算自己负责的部分,计算完成后,对多个 GPU 上的输出矩阵进行汇总,得到最终的输出矩阵 YY。下面是张量并行的计算流程示意图:

张量并行原理示意图

两种切分方式

在张量并行中,对权重矩阵的切分有两种方式,按列切分和按行切分。下面我们以矩阵乘法 Y=X×WY = X \times W 为例,分别介绍按行切分和按列切分的计算流程。

假设 XX 的维度为 B×MB \times MWW 的维度为 M×NM \times NYY 的维度为 B×NB \times N

按列切分

按列切分的方式是将权重矩阵 WW 从列的维度上切分为多个小矩阵,切分后的分块为 W1,W2,...,WnW_1, W_2, ..., W_n,每个小矩阵的维度为 M×(N/n)M \times (N/n),其中 nn 是切分的份数。矩阵乘法的计算流程如下:

因为矩阵乘法可以想象为使用输入矩阵 XX 的每一行对权重矩阵 WW 的所有行做线性组合,因此权重矩阵的列之间是相互独立的,矩阵乘法可以在 GPU 之间并行计算。

使用这种切分方式时,输入矩阵 XX 需要全部发送给每个 GPU,因为每个 GPU 都需要使用完整的输入矩阵 XX 来计算自己的输出矩阵 Yi=X×WiY_i = X \times W_i。计算完成后,每个 GPU 的输出矩阵 YiY_i 是最终输出矩阵 YY 的部分列,来自多个 GPU 的输出矩阵 YiY_i 需要在最后进行拼接,得到最终的输出矩阵 YY

按行切分

按行切分将权重矩阵 WW 从行的维度上切分为多个小矩阵,切分后的分块为 W1,W2,...,WmW_1, W_2, ..., W_m,每个小矩阵的维度为 (M/m)×N(M/m) \times N,其中 mm 是切分的份数。为了完成分块矩阵乘法,需要将输入矩阵 XX 按列切分为多个小矩阵,切分后的分块为 X1,X2,...,XmX_1, X_2, ..., X_m,每个小矩阵的维度为 B×(M/m)B \times (M/m)。矩阵乘法的计算流程如下:

按行切分的原理稍微复杂一些,需要回顾一下矩阵外积的计算原理。将矩阵 XX 的每一列与矩阵 WW 的对应行做外积(XiX_iXX 的第 ii 列,WiW_iWW 的第 ii 行),每个外积得到一个 B×NB \times N 的矩阵,然后将这些矩阵求和,得到最终的输出矩阵 YY。假设 XX 的列数为 mm,则可以将矩阵乘法拆分为 mm 个小矩阵的乘法操作,计算流程如下:

Y=i=1mXi×WiY = \sum_{i=1}^{m} X_i \times W_i

这里 Xi×WiX_i \times W_i 是一个 B×NB \times N 的矩阵,多个 Xi×WiX_i \times W_i 的结果求和后得到最终的输出矩阵 YY

下图是矩阵外积的计算流程示意图:

图片来自:https://zhuanlan.zhihu.com/p/441943479

使用按行切分时,每个 GPU 持有权重矩阵 WW 的部分行,在计算时只需要将输入矩阵 XX 的对应列发送给对应的 GPU。计算完成后,每个 GPU 的输出矩阵 Yi=Xi×WiY_i = X_i \times W_i 是一个 B×NB \times N 的矩阵,来自多个 GPU 的输出矩阵 YiY_i 需要在最后进行求和,得到最终的输出矩阵 YY

通信模式

使用张量并行时,GPU 之间需要进行通信操作,将各个 GPU 的计算结果汇总,得到最终的输出矩阵,不同的切分方式对应不同的通信模式。

按列切分

使用按列切分时,每个 GPU 的输入是完整的输入矩阵,这需要使用 broadcast 的通信模式,将输入矩阵发送给每个 GPU。在每个 GPU 上,使用自己持有的权重矩阵计算输出矩阵的部分列。计算完成后,每个 GPU 的输出矩阵是最终输出矩阵的部分列,需要使用 All-Gather 的通信模式,将各个 GPU 的输出矩阵收集到一起,得到完整的输出矩阵。下面是一个简单的示意图,展示了按列切分时前向计算的通信模式:

按列切分前向计算的示意图

在反向传播阶段,为了计算 WiW_i 的梯度,只需要回传输出矩阵 YiY_i 的梯度即可,输入矩阵 XX 的梯度和权重矩阵 WiW_i 的梯度计算方式如下:

LWi=XT×LYi\frac{\partial L}{\partial W_i} = X^T \times \frac{\partial L}{\partial Y_i} LX=LYi×WiT\frac{\partial L}{\partial X} = \frac{\partial L}{\partial Y_i} \times W_i^T

下面是一个简单的示意图,展示了按列切分时反向传播的通信模式:

按列切分反向传播的示意图

输入矩阵 XX 的梯度需要使用 All-Reduce 的通信模式,将各个 GPU 上的 XX 的梯度进行求和,得到最终的 XX 的梯度。

按行切分

使用按行切分时,首先需要将输入矩阵从列的维度上切分为多个小矩阵,然后使用 Scatter 的通信模式,将这些小矩阵发送给对应的 GPU。在每个 GPU 上,使用自己持有的权重矩阵计算输出矩阵。计算完成后,每个 GPU 的输出矩阵是一个完整的输出矩阵,需要使用 All-Reduce 的通信模式,将各个 GPU 的输出矩阵进行求和,得到最终的输出矩阵。下面是一个简单的示意图,展示了按行切分前向计算的通信模式:

按行切分前向计算的示意图

在反向传播阶段,输入为 YY 的梯度 LY\frac{\partial L}{\partial Y},而此前每个 GPU 计算出的结果是 YiY_i,但因为 Y=i=1mYiY = \sum_{i=1}^{m} Y_i,因此 LYi\frac{\partial L}{\partial Y_i} 就等于 LY\frac{\partial L}{\partial Y}

因为 Yi=Xi×WiY_i = X_i \times W_i,所以 WiW_i 的梯度和 XiX_i 的梯度可以通过链式法则计算出来,其结果为:

LWi=XiT×LYi\frac{\partial L}{\partial W_i} = X_i^T \times \frac{\partial L}{\partial Y_i} LXi=LYi×WiT\frac{\partial L}{\partial X_i} = \frac{\partial L}{\partial Y_i} \times W_i^T

下面是一个简单的示意图,展示了按行切分时反向传播的通信模式:

按行切分反向传播的示意图

LYi\frac{\partial L}{\partial Y_i} 使用 broadcast 的通信模式发送给每个 GPU,然后每个 GPU 使用自己持有的权重矩阵计算出 LWi\frac{\partial L}{\partial W_i}LXi\frac{\partial L}{\partial X_i}WiW_i 可以在 GPU 上本地更新,而不需要与其他 GPU 进行通信。而 LX\frac{\partial L}{\partial X} 需要使用 All-Gather 的通信模式,从各个 GPU 上收集 LXi\frac{\partial L}{\partial X_i},以得到完整的 LX\frac{\partial L}{\partial X}

通信模式的对比

使用 TP 时,每个 GPU 首先需要将输入矩阵发送给其他 GPU,然后在每个 GPU 上计算出输出矩阵的一部分,最后将这些结果收集回来并汇总。

按列切分时,前向计算是使用 broadcast + all-gather,反向传播时使用 scatter + all-reduce。按行切分时,前向计算是使用 scatter + all-reduce,反向传播时使用 broadcast + all-gather。

切分方式前向计算反向传播
按列切分broadcast + all-gatherscatter + all-reduce
按行切分scatter + all-reducebroadcast + all-gather

张量并行在 Transformer 中的应用

在 Transformer 架构中,张量并行主要应用于注意力机制和全连接层中,下面分别描述在这两种计算中如何使用张量并行。

注意力机制中的张量并行

首先回顾 Attention 机制的计算流程。输入特征 XX 分别乘以三个权重矩阵得到查询(Query)、键(Key)和值(Value):

Q=X×WQ,K=X×WK,V=X×WVQ = X \times W_Q, \quad K = X \times W_K, \quad V = X \times W_V

然后计算输出:

O=softmax(QKTdk)VO = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V

最后将输出 OO 乘以一个投影矩阵 WOW_O,得到最终的输出:

Y=O×WOY = O \times W_O

目前实际应用中主要使用的是多头注意力机制(Multi-Head Attention, MHA),或者 Grouped Query Attention(GQA),整个计算流程如下:

多头注意力机制计算流程

由于不同注意力头的计算天然独立,因此非常适合从 head 维度切分。

可以将 WQW_QWKW_KWVW_V 三个矩阵按 head 维度切分成 N 份(本质上是按列切分),然后将切分后的参数分布到不同的 GPU 上。每个 GPU 独立计算其所负责的注意力头的 QQKKVV,然后计算出对应的输出 OiO_i

当注意力计算完成后,每个 GPU 上的结果 OiO_i 是部分 head 的输出。而 OiO_i 还需要乘以一个投影矩阵 WOW_O,为了减少通信开销,可以将 WOW_O 按行切分成 N 份,每个 GPU 持有 WOW_O 的部分行。在各个 GPU 上将 OiO_i 的分片与对应的 WOW_O 的分片进行矩阵乘法,得到各个 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_gateW_up 从列维度拼接起来,这样 gateup 的计算可以在同一个矩阵乘法中完成。因此无论是否有 Gate 机制,MLP 层的计算都可以看作是两个全连接层的计算。

为了减少在全连接层中的通信开销,可以对第一个矩阵乘法按列切分,这样每个 GPU 的计算结果是输出矩阵的部分列。考虑前面提到的按行切分的方式,输入矩阵 XX 是按列切分的。我们只需要对第二个矩阵乘法进行按行切分,这样前一个矩阵乘法的输出就可以直接作为后一个矩阵乘法的输入。对第一个矩阵按列切分,对第二个矩阵按行切分,这样两个矩阵乘法之间不需要进行额外的通信操作。

下面是一个简单的示意图,展示了 MLP 层中张量并行的计算流程:

MLP 张量并行示意图

总结

张量并行的核心思想是将大矩阵切分为多个小矩阵,利用线性代数中分块矩阵乘法的思想,将这些小矩阵分布到不同的 GPU 上并行地完成计算。这一方面降低单个 GPU 中显存的占用,另外一方面能够利用多个 GPU 并行计算,提升计算速度。在 Transformer 中,Attention 模块和 MLP 模块均可以使用 TP 进行加速。为了降低通信开销,会组合使用列切分和行切分两种切分方式,这样在两个矩阵乘法之间不需要进行额外的通信操作。

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