什么是 Tensor Parallel
Transformer 层里的主要计算,基本都落在几次大矩阵乘法上。Tensor Parallel,简称 TP,就是把这些层内矩阵乘法拆到多张 GPU 上,让多张卡一起完成同一次前向计算。
它切的不是 batch,也不是 token。按 batch 切,是让不同样本走不同设备;TP 则是让同一个 token 经过同一组 GPU,每张卡只负责 hidden 相关权重和中间表示的一段。
因此,TP 的问题不只是拆开算。拆开以后,还要让多张卡的局部结果重新组成等价于单卡的结果,并且通信不能太频繁,否则并行计算省下的时间会被通信吃掉。
从计算结构看,Transformer 层里的 Attention 和 MLP 都可以抽象成同一条路径:
TP 的核心,就是围绕这条路径安排矩阵切分和通信:进入中间空间时拆开算,在中间空间里尽量留在本卡继续算,回到 shared hidden state 时再把各卡的贡献合起来。下面讨论的是最常见的一维 TP,也就是沿着 hidden 相关维度切 Transformer 层里的矩阵。
TP 中的两种权重切分方式
TP 里最常见的两种切法是 column-wise 和 row-wise。它们都能并行化一个矩阵乘法,但输出的含义完全不同:column-wise 产出“结果的一段”,row-wise 产出“结果的一份贡献”。这个差别决定了后面需要拼接还是求和。
| 切法 | 切哪里 | 每张 GPU 拿到什么 | 输出代表什么 | 如果要完整结果 |
|---|---|---|---|---|
| column-wise | 权重的输出维度 | 完整输入 + 一段输出权重 | 中间空间的一段坐标 | 拼接 / all-gather |
| row-wise | 权重的输入维度 | 一段输入 + 对应权重 | 同一个输出空间里的一份贡献 | 求和 / all-reduce |
column-wise:拆输出空间
column-wise 把权重按输出维度切开:
每张 GPU 都看到完整输入 ,但只负责输出空间的一段:
完整输出是这些分片的拼接:
它的直觉是:每张卡从同一个 hidden state 出发,只生成中间空间的一段坐标。只要后续计算仍然能沿着这段坐标继续做,就不需要马上通信。
row-wise:拆输入空间
row-wise 把输入和权重的输入维度一起切开:
每张 GPU 本地算一份输出贡献:
这些贡献都落在同一个输出空间里,所以最终要相加:
工程上,这个求和通常对应一次 all-reduce。
到这里可以先记住一个判断:column-wise 的结果适合继续分片计算,row-wise 的结果适合在最后汇总贡献。 Transformer TP 的通信优化,基本就来自这两个性质的配合。
Transformer TP 的核心流程
如果只看单个矩阵乘法,TP 并不复杂:把矩阵切开,分到多张 GPU 上算,再在需要时通过拼接或求和恢复后续计算需要的形状。真正值得关心的是连续矩阵乘法,因为这里可以少做一次中间通信。
例如:
如果把两个矩阵乘法都当成独立算子,第一步会先合并 的完整中间结果,然后再乘 ,最后还要再次合并最终结果。
更高效的做法,是不急着合并 。先把 用 column-wise 切开,让每张 GPU 得到中间空间的一段分片;再把 用 row-wise 切开,让这段分片直接成为下一步的输入。这样中间结果不需要先还原成完整矩阵,只在最终输出时做一次求和。
连续矩阵乘法中的通信捷径
下面这张图展示的就是这个捷径。上半部分是不切分时的完整计算;下半部分把第一层权重按列切,把第二层权重按行切。关键不是模块名字,而是中间的红色分片:column-wise 的输出刚好可以作为 row-wise 的输入。
用一个抽象公式表示:
负责把输入投到一个可切分的中间空间, 负责把这个中间空间投回输出空间。TP 的切法是:
这个等式也解释了为什么最后是加法。 在输入维度上被切成多段,每张 GPU 只算其中一段输入对最终输出的贡献;这些贡献都落在同一个输出空间里,所以要相加才等价于原来的完整矩阵乘法。通信也因此从“两个矩阵乘法各合并一次”,变成 “中间不合并,最后做一次 all-reduce”。
这不是所有 TP 的固定顺序,而是连续矩阵乘法里的通信优化。对于 这种结构,column-wise 接 row-wise 可以让中间分片自然传下去;如果先 row-wise,第一步产生的是待求和的贡献,通常要先合并,才能作为下一步 column-wise 的完整输入。
映射到 Transformer 层
回到 Transformer 层,Attention 和 MLP 都能套进这条路径。先不看具体公式,只看每张 GPU 手里的东西如何变化:
| 阶段 | 切法 / 通信 | 每张 GPU 的输入 | 每张 GPU 的输出 | 通信特性 |
|---|---|---|---|---|
| 进入中间空间 | column-wise | 完整 shared hidden state | 中间表示的一段分片 | 不需要立即通信 |
| 中间计算 | 本地计算 | 本卡的中间分片 | 仍然是本卡的中间分片 | 不需要通信 |
| 回到 shared hidden state | row-wise | 本卡的中间分片 | shared hidden state 的一份贡献 | 尚不是完整结果 |
| 合并结果 | all-reduce 求和 | 各卡的输出贡献 | 每张 GPU 都得到完整 hidden state | 需要通信 |
把这条路径落到具体模块上,就是下面这张表。公式不是为了堆细节,而是为了说明同一个模式反复出现:QKV 和 Gate / Up 负责进入中间空间,Attention compute 和 activation 留在本卡,Output / Down projection 负责回到 shared hidden state,最后通过 all-reduce 求和。
| 模块 | 环节 | 切法 / 通信 | 原公式 | 每张 GPU 上的计算 |
|---|---|---|---|---|
| Attention | QKV projection | column-wise | ||
| Attention | attention compute | 本地计算 | ||
| Attention | output projection | row-wise | ||
| Attention | all-reduce | 求和 | 每张 GPU 得到完整 | |
| MLP | Gate / Up projection | column-wise | ||
| MLP | activation / gating | 本地计算 | ||
| MLP | down projection | row-wise | ||
| MLP | all-reduce | 求和 | 每张 GPU 得到完整 |
所以,Transformer TP 的主干逻辑可以概括成三点:
- column-wise 负责把完整 hidden state 拆到多个中间分片里。
- 中间计算尽量留在本卡完成,省掉中间结果的合并通信。
- 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 的内部机理,本质上就是安排好:哪里切开算,哪里合起来。
参考资料