什么是 Tensor Parallel

Transformer 层里的主要计算,基本都落在几次大矩阵乘法上。Tensor Parallel,简称 TP,就是把这些层内矩阵乘法拆到多张 GPU 上,让多张卡一起完成同一次前向计算。

它切的不是 batch,也不是 token。按 batch 切,是让不同样本走不同设备;TP 则是让同一个 token 经过同一组 GPU,每张卡只负责 hidden 相关权重和中间表示的一段。

因此,TP 的问题不只是拆开算。拆开以后,还要让多张卡的局部结果重新组成等价于单卡的结果,并且通信不能太频繁,否则并行计算省下的时间会被通信吃掉。

从计算结构看,Transformer 层里的 Attention 和 MLP 都可以抽象成同一条路径:

shared hidden state 可切分的中间空间 shared hidden state

TP 的核心,就是围绕这条路径安排矩阵切分通信:进入中间空间时拆开算,在中间空间里尽量留在本卡继续算,回到 shared hidden state 时再把各卡的贡献合起来。下面讨论的是最常见的一维 TP,也就是沿着 hidden 相关维度切 Transformer 层里的矩阵。

TP 中的两种权重切分方式

TP 里最常见的两种切法是 column-wiserow-wise。它们都能并行化一个矩阵乘法,但输出的含义完全不同:column-wise 产出“结果的一段”row-wise 产出“结果的一份贡献”。这个差别决定了后面需要拼接还是求和。

切法切哪里每张 GPU 拿到什么输出代表什么如果要完整结果
column-wise权重的输出维度完整输入 + 一段输出权重中间空间的一段坐标拼接 / all-gather
row-wise权重的输入维度一段输入 + 对应权重同一个输出空间里的一份贡献求和 / all-reduce

column-wise:拆输出空间

column-wise 把权重按输出维度切开:

W=[W0W1Wt1]W = [W_0 \mid W_1 \mid \cdots \mid W_{t-1}]

每张 GPU 都看到完整输入 XX,但只负责输出空间的一段:

Yi=XWiY_i = XW_i

完整输出是这些分片的拼接:

Y=[Y0Y1Yt1]Y = [Y_0 \mid Y_1 \mid \cdots \mid Y_{t-1}]

它的直觉是:每张卡从同一个 hidden state 出发,只生成中间空间的一段坐标。只要后续计算仍然能沿着这段坐标继续做,就不需要马上通信。

row-wise:拆输入空间

row-wise 把输入和权重的输入维度一起切开:

X=[X0X1Xt1],W=[W0W1Wt1]X = [X_0 \mid X_1 \mid \cdots \mid X_{t-1}], \quad W = \begin{bmatrix} W_0 \\ W_1 \\ \vdots \\ W_{t-1} \end{bmatrix}

每张 GPU 本地算一份输出贡献:

Yi=XiWiY_i = X_iW_i

这些贡献都落在同一个输出空间里,所以最终要相加

Y=i=0t1YiY = \sum_{i=0}^{t-1} Y_i

工程上,这个求和通常对应一次 all-reduce

到这里可以先记住一个判断:column-wise 的结果适合继续分片计算,row-wise 的结果适合在最后汇总贡献。 Transformer TP 的通信优化,基本就来自这两个性质的配合。

Transformer TP 的核心流程

如果只看单个矩阵乘法,TP 并不复杂:把矩阵切开,分到多张 GPU 上算,再在需要时通过拼接或求和恢复后续计算需要的形状。真正值得关心的是连续矩阵乘法,因为这里可以少做一次中间通信。

例如:

Y=AM1M2Y = AM_1M_2

如果把两个矩阵乘法都当成独立算子,第一步会先合并 AM1AM_1 的完整中间结果,然后再乘 M2M_2,最后还要再次合并最终结果。

更高效的做法,是不急着合并 AM1AM_1。先把 M1M_1column-wise 切开,让每张 GPU 得到中间空间的一段分片;再把 M2M_2row-wise 切开,让这段分片直接成为下一步的输入。这样中间结果不需要先还原成完整矩阵,只在最终输出时做一次求和。

连续矩阵乘法中的通信捷径

下面这张图展示的就是这个捷径。上半部分是不切分时的完整计算;下半部分把第一层权重按列切,把第二层权重按行切。关键不是模块名字,而是中间的红色分片:column-wise 的输出刚好可以作为 row-wise 的输入

column-wise 和 row-wise Tensor Parallel 矩阵切分示意图
column-wise 的输出分片会继续作为 row-wise 的输入分片;row-wise 产生的是同形状的输出贡献,所以最后通过求和合并。Reference: Wenyi Li, “Dive into Tensor Parallelism”.

用一个抽象公式表示:

Y=ϕ(XW1)W2Y = \phi(XW_1)W_2

W1W_1 负责把输入投到一个可切分的中间空间,W2W_2 负责把这个中间空间投回输出空间。TP 的切法是:

W1=[W1,0W1,1W1,t1]W_1 = [W_{1,0} \mid W_{1,1} \mid \cdots \mid W_{1,t-1}] Hi=ϕ(XW1,i)H_i = \phi(XW_{1,i}) W2=[W2,0W2,1W2,t1]W_2 = \begin{bmatrix} W_{2,0} \\ W_{2,1} \\ \vdots \\ W_{2,t-1} \end{bmatrix} Y=i=0t1HiW2,iY = \sum_{i=0}^{t-1} H_iW_{2,i}

这个等式也解释了为什么最后是加法。W2W_2 在输入维度上被切成多段,每张 GPU 只算其中一段输入对最终输出的贡献;这些贡献都落在同一个输出空间里,所以要相加才等价于原来的完整矩阵乘法。通信也因此从“两个矩阵乘法各合并一次”,变成 “中间不合并,最后做一次 all-reduce

这不是所有 TP 的固定顺序,而是连续矩阵乘法里的通信优化。对于 AM1M2AM_1M_2 这种结构,column-wiserow-wise 可以让中间分片自然传下去;如果先 row-wise,第一步产生的是待求和的贡献,通常要先合并,才能作为下一步 column-wise 的完整输入。

映射到 Transformer 层

回到 Transformer 层,Attention 和 MLP 都能套进这条路径。先不看具体公式,只看每张 GPU 手里的东西如何变化:

阶段切法 / 通信每张 GPU 的输入每张 GPU 的输出通信特性
进入中间空间column-wise完整 shared hidden state中间表示的一段分片不需要立即通信
中间计算本地计算本卡的中间分片仍然是本卡的中间分片不需要通信
回到 shared hidden staterow-wise本卡的中间分片shared hidden state 的一份贡献尚不是完整结果
合并结果all-reduce 求和各卡的输出贡献每张 GPU 都得到完整 hidden state需要通信

把这条路径落到具体模块上,就是下面这张表。公式不是为了堆细节,而是为了说明同一个模式反复出现:QKV 和 Gate / Up 负责进入中间空间Attention compute 和 activation 留在本卡Output / Down projection 负责回到 shared hidden state,最后通过 all-reduce 求和。

模块环节切法 / 通信原公式每张 GPU 上的计算
AttentionQKV projectioncolumn-wiseQ=HWQQ=HW^Q
K=HWKK=HW^K
V=HWVV=HW^V
Qi=HWiQQ_i=HW_i^Q
Ki=HWiKK_i=HW_i^K
Vi=HWiVV_i=HW_i^V
Attentionattention compute本地计算A=softmax(QK/dh+mask)VA=\operatorname{softmax}(QK^\top/\sqrt{d_h}+\mathrm{mask})VAi=softmax(QiKi/dh+mask)ViA_i=\operatorname{softmax}(Q_iK_i^\top/\sqrt{d_h}+\mathrm{mask})V_i
Attentionoutput projectionrow-wiseYattn=AWOY_{\mathrm{attn}}=AW^OYattn,i=AiWiOY_{\mathrm{attn},i}=A_iW_i^O
Attentionall-reduce求和Yattn=iYattn,iY_{\mathrm{attn}}=\sum_i Y_{\mathrm{attn},i}每张 GPU 得到完整 YattnY_{\mathrm{attn}}
MLPGate / Up projectioncolumn-wiseG=HWGG=HW^G
U=HWUU=HW^U
Gi=HWiGG_i=HW_i^G
Ui=HWiUU_i=HW_i^U
MLPactivation / gating本地计算M=silu(G)UM=\operatorname{silu}(G)\odot UMi=silu(Gi)UiM_i=\operatorname{silu}(G_i)\odot U_i
MLPdown projectionrow-wiseYmlp=MWDY_{\mathrm{mlp}}=MW^DYmlp,i=MiWiDY_{\mathrm{mlp},i}=M_iW_i^D
MLPall-reduce求和Ymlp=iYmlp,iY_{\mathrm{mlp}}=\sum_i Y_{\mathrm{mlp},i}每张 GPU 得到完整 YmlpY_{\mathrm{mlp}}

所以,Transformer TP 的主干逻辑可以概括成三点:

  1. column-wise 负责把完整 hidden state 拆到多个中间分片里。
  2. 中间计算尽量留在本卡完成,省掉中间结果的合并通信。
  3. row-wise 作为这条路径的收束点,把各分片对 hidden space 的贡献求和。

许多常见 Transformer TP 实现中,一个 Transformer 层通常会出现两次这样的求和通信:一次在 Attention 子层回到 hidden state 时,一次在 MLP 子层回到 hidden state 时。这就是本文讨论的 一维 TP 主线:column-wise 产生中间分片,本地计算沿分片继续,row-wise 产生输出贡献,all-reduce 合成完整 hidden state。

小结

理解 Tensor Parallel,不需要一开始就陷进各种并行策略的名字里。更自然的入口是矩阵乘法本身:column-wise 把输出空间切成多段,row-wise 把输入空间切成多份贡献;前者适合继续分片计算,后者适合在最后求和收束。

Transformer 里的 Attention 和 MLP 正好反复出现“进入中间空间,再回到 hidden state”的结构。于是 TP 的主线就变得很清楚:QKV 或 Gate / Up projection 先用 column-wise 产生中间分片,中间计算留在本卡完成,Output 或 Down projection 再用 row-wise 产生输出贡献,最后通过 all-reduce 合成完整 hidden state。所谓 TP 的内部机理,本质上就是安排好:哪里切开算,哪里合起来

参考资料

  1. Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism
  2. Wenyi Li, “Dive into Tensor Parallelism”
  3. Ashraf Bhuiyan, “Part 9: Tensor Parallelism”