CUDA 张量并行之后,到底有哪几种计算

用一个 X @ W 小例子讲清张量并行里的本地分片计算、Column Parallel、Row Parallel、all-gather、all-reduce、reduce-scatter,以及它们在 Transformer 里的位置。

一句话结论

张量并行以后,GPU 上不是出现了一种神秘的新计算,而是把一个大计算拆成了几类更小、更清楚的工作:

  1. 每张卡先做自己的 本地分片计算
  2. 有些线性层按 输出维度切,这叫 Column Parallel。
  3. 有些线性层按 输入维度切,这叫 Row Parallel。
  4. 算完以后,有些地方要做 跨卡通信聚合,比如 all-reduce、all-gather、reduce-scatter。
  5. activation、dropout、局部 attention head 这类操作,经常可以继续在本地做。

如果只记一句话:张量并行 = 本地 CUDA kernel 计算 + 必要时用 NCCL 把几张卡的结果拼起来或加起来。

这篇只解决一个问题:看到一个 Transformer 层被张量并行切开以后,你能不能立刻判断每张 GPU 在算什么、什么时候需要通信、通信是在拼接还是求和。

先用一个很土的比喻

假设有一本特别厚的账本,要算全公司一年的总收入。

单卡计算就像一个人从第一页算到最后一页。

张量并行像把账本横着或竖着切给几个人:

  • 有时每个人负责几列,最后把列拼起来。
  • 有时每个人负责几行,最后把小计加起来。
  • 有时每个人算出来的东西本来就只给自己下一步用,那就不用立刻汇总。

GPU 也是这样。所谓“并行”,不是每张卡都完整算一遍,而是每张卡拿到权重或激活的一部分,先算局部答案,再看后面需不需要把答案合并。

最小数学地基:矩阵乘法就是很多个乘加

先不用大模型,用一个小例子。

有一个输入 X

X = [10, 20]

有一个权重矩阵 W

W = [
  [1, 2, 3, 4],
  [5, 6, 7, 8]
]

X @ W 会得到 4 个输出数:

第 1 个输出 = 10*1 + 20*5 = 110
第 2 个输出 = 10*2 + 20*6 = 140
第 3 个输出 = 10*3 + 20*7 = 170
第 4 个输出 = 10*4 + 20*8 = 200

X @ W = [110, 140, 170, 200]

把它放回 Transformer,就是一个线性层:

hidden_states @ weight -> next_hidden_states

张量并行最常见的切法,就是围绕这个 X @ W 做文章。

第一类:本地分片计算

这是最基础的一类。

张量并行之后,每张 GPU 手里只有一部分权重,或者一部分输入。它先不用管别人,直接在自己卡上跑 CUDA kernel:

GPU 0: 算自己的 X0 @ W0
GPU 1: 算自己的 X1 @ W1
GPU 2: 算自己的 X2 @ W2
GPU 3: 算自己的 X3 @ W3

这些本地计算包括:

  • GEMM,也就是矩阵乘。
  • attention 某些 head 的计算。
  • activation,比如 GELU、SiLU。
  • dropout。
  • 一些 reshape、transpose、copy。

这类计算的特点是:数据在本卡就够了,先不需要问别的 GPU 要答案。

所以你看张量并行代码时,第一反应应该是:

这一段是不是每张卡只拿自己的 shard 就能算?

如果能,它就是本地分片计算。

第二类:Column Parallel,按输出列切

现在看第一种经典线性层切法:Column Parallel。

权重矩阵 W 原来是:

W = [
  [1, 2, 3, 4],
  [5, 6, 7, 8]
]

它有 4 列,对应 4 个输出。现在用 2 张 GPU,把列切开:

GPU 0 拿前两列:
W0 = [
  [1, 2],
  [5, 6]
]

GPU 1 拿后两列:
W1 = [
  [3, 4],
  [7, 8]
]

输入 X = [10, 20] 通常每张卡都有一份。

GPU 0 算:

X @ W0 = [110, 140]

GPU 1 算:

X @ W1 = [170, 200]

如果后面需要完整输出,就把两张卡的结果拼起来:

[110, 140] + [170, 200]
拼接成
[110, 140, 170, 200]

这里的关键字是 拼接,不是相加。

所以 Column Parallel 的直觉是:

每张卡负责生产一部分输出列。
最后如果需要完整输出,就 all-gather 拼起来。

但注意,不是每次都要马上拼。

比如 MLP 里的第一层通常把 hidden size 放大到 4 倍。如果这个放大的中间结果下一步仍然可以分片处理,那就没必要立刻 all-gather。能不通信就不通信,因为跨卡通信很贵。

第三类:Row Parallel,按输入行切

Row Parallel 刚好相反:它按输入维度切。

还是这组计算:

X = [10, 20]

W = [
  [1, 2, 3, 4],
  [5, 6, 7, 8]
]

这次我们把 X 切开,也把 W 的行切开。

GPU 0 拿:

X0 = [10]
W0 = [
  [1, 2, 3, 4]
]

GPU 1 拿:

X1 = [20]
W1 = [
  [5, 6, 7, 8]
]

GPU 0 算出一份局部贡献:

X0 @ W0 = [10, 20, 30, 40]

GPU 1 算出另一份局部贡献:

X1 @ W1 = [100, 120, 140, 160]

真正的输出要把两份贡献加起来:

[10, 20, 30, 40]
+ [100, 120, 140, 160]
= [110, 140, 170, 200]

这里的关键字是 相加,不是拼接。

所以 Row Parallel 的直觉是:

每张卡负责算最终答案的一部分贡献。
最后必须把贡献求和,常见通信是 all-reduce。

这就是 Row Parallel 和 Column Parallel 最容易混的地方:

Column Parallel: 各卡输出不同列,合并方式是拼接。
Row Parallel: 各卡输出同一组列的部分和,合并方式是求和。

第四类:跨卡通信计算

严格说,通信不是神经网络里的数学算子,但在张量并行里它和计算一样重要。很多时候,模型快不快,不只看 GEMM 多快,还看通信插得好不好。

最常见的通信有四种。

1. all-gather:把分片收集成完整结果

Column Parallel 后,如果每张卡只有一部分输出,而下一步需要完整输出,就要 all-gather。

直觉:

GPU 0 有 [110, 140]
GPU 1 有 [170, 200]

all-gather 后:
GPU 0 有 [110, 140, 170, 200]
GPU 1 有 [110, 140, 170, 200]

它像“大家把自己那一页复印给所有人”。

2. all-reduce:把每张卡的部分和加起来

Row Parallel 后,每张卡算的是局部贡献。要得到完整答案,就要 all-reduce。

直觉:

GPU 0 有 [10, 20, 30, 40]
GPU 1 有 [100, 120, 140, 160]

all-reduce 后:
GPU 0 有 [110, 140, 170, 200]
GPU 1 有 [110, 140, 170, 200]

它像“大家把小计加成总计,然后每个人都拿到总计”。

3. reduce-scatter:先求和,再把结果切开

有时候不需要每张卡都拿完整结果。那就可以 reduce-scatter:

先把各卡贡献求和,
再把总结果切成几片,
每张卡只拿下一步需要的那一片。

它比 all-reduce 更省,因为不用让每张卡都保留完整输出。

4. broadcast:从一张卡发给所有卡

broadcast 比较好理解:

GPU 0 有一份数据
发给 GPU 1、GPU 2、GPU 3

在张量并行里,broadcast 不一定是最核心的通信,但在初始化、同步某些小状态、分发输入时会看到。

第五类:局部逐元素计算

不是所有东西都需要通信。

比如 Column Parallel 之后,每张卡手里有自己那部分输出:

GPU 0: [110, 140]
GPU 1: [170, 200]

如果下一步只是做 GELU:

GELU([110, 140])
GELU([170, 200])

那每张卡自己做就行,不需要先把完整 [110, 140, 170, 200] 拼出来。

这类本地逐元素操作包括:

  • GELU、SiLU、ReLU。
  • dropout。
  • 乘一个 gate。
  • 加 bias。
  • 某些只依赖本 shard 的缩放。

它们的特点是:每个元素自己就能算,不需要看别的元素。

不过 LayerNorm、RMSNorm 这种归一化要更小心。它们通常要看 hidden dimension 上的一组数来算均值、方差或范数。如果 hidden dimension 被切开,就可能需要跨卡同步统计量;如果框架选择在复制的 hidden 上做归一化,那就可以本地做。

所以判断归一化要看具体并行布局,不能一句话说永远本地或永远通信。

放进 Transformer:这几种计算怎么串起来

现在把前面几类放回一个 Transformer block。

QKV 投影:通常像 Column Parallel

attention 开始时要从 hidden states 算出 Q、K、V:

hidden -> Q, K, V

很多张量并行实现会把 QKV 投影按输出维度切开。也就是每张卡负责一部分 head。

直觉:

GPU 0 负责一部分 attention heads
GPU 1 负责另一部分 attention heads

这很适合 attention,因为不同 head 本来就可以分开算。

Attention heads:很多计算是本地的

如果某个 head 的 Q/K/V 都在同一张卡上,那这个 head 的 attention 可以本地算:

Q_head @ K_head^T
softmax
乘 V_head

这就是张量并行很漂亮的地方:切 head 以后,每张卡可以独立算自己那部分 head,不必每一步都通信。

Attention 输出投影:通常像 Row Parallel

attention heads 算完后,要经过输出投影,把多头结果映射回 hidden size。

这一步常常用 Row Parallel。因为输入的 head 已经分在不同 GPU 上,每张卡先用自己那部分 head 算一份对最终 hidden 的贡献,然后用 all-reduce 把贡献加起来。

直觉:

GPU 0: 我的 heads 对最终 hidden 的贡献
GPU 1: 我的 heads 对最终 hidden 的贡献

all-reduce:
把贡献加起来,得到完整 hidden

MLP:一扩一收,刚好 Column + Row

Transformer 的 MLP 通常有两层线性层:

hidden -> intermediate -> hidden

第一层把维度放大,比如从 h4h。这一步很适合 Column Parallel:

每张卡负责 intermediate 的一部分列。

中间的 activation 通常可以本地做:

每张卡对自己的 intermediate shard 做 GELU 或 SiLU。

第二层把 4h 收回 h。这一步很适合 Row Parallel:

每张卡用自己的 intermediate shard 算一份 hidden 贡献。
最后 all-reduce 相加。

所以 MLP 的心智模型很简单:

先切输出:Column Parallel
中间本地激活:local activation
再切输入:Row Parallel
最后求和:all-reduce

最容易错的三个点

错误 1:把 Column 和 Row 的合并方式搞反

记住:

Column Parallel -> 拼接
Row Parallel -> 求和

为什么?

Column Parallel 是每张卡负责不同输出列,所以结果互不重叠,要拼起来。

Row Parallel 是每张卡负责同一个输出的不同来源贡献,所以结果重叠在同一组输出位置上,要加起来。

错误 2:以为每个地方都要通信

不是。

张量并行性能好的关键,就是尽量让一串操作都留在 shard 状态下继续算。比如:

Column Parallel Linear
-> local activation
-> Row Parallel Linear
-> all-reduce

中间 activation 不需要通信。能把通信推迟,就不要急着 all-gather。

错误 3:以为通信只是小开销

通信非常关键。

本地 GEMM 在 GPU 上很快,但跨 GPU 要走 NVLink、PCIe 或网络。张量并行的收益来自“把大计算分摊到多张卡”,成本来自“卡和卡之间必须交换结果”。

所以张量并行不是卡越多越快。卡越多,每张卡算得少了,但通信次数和通信压力可能上升。什么时候划算,要看模型大小、batch、seq length、硬件互联和具体 kernel。

一张总表

类型                     每张 GPU 在做什么                  合并方式
本地分片计算              算自己的 shard                     不一定合并
Column Parallel           算不同输出列                       all-gather 拼接,或继续保持分片
Row Parallel              算同一输出的部分贡献                all-reduce 求和,或 reduce-scatter
局部逐元素计算             对自己的元素做 activation/dropout    不需要通信
跨卡通信                  交换、拼接、求和、切分结果            NCCL collective

最后的自检

你可以用这几个问题检查自己是否真的懂了:

  1. 如果每张卡算出来的是不同的输出列,最后应该拼接还是相加?
  2. 如果每张卡算出来的是同一个输出的不同来源贡献,最后应该拼接还是相加?
  3. MLP 的第一层为什么常用 Column Parallel?
  4. MLP 的第二层为什么常用 Row Parallel?
  5. 为什么 activation 经常可以本地做?
  6. 为什么张量并行不是 GPU 越多越快?

答案压缩成一句:

Column 负责“各出各的列”,Row 负责“各算各的贡献”;
本地 CUDA kernel 负责把 shard 算快,NCCL collective 负责把必须合并的地方合并。