一句话结论
张量并行以后,GPU 上不是出现了一种神秘的新计算,而是把一个大计算拆成了几类更小、更清楚的工作:
- 每张卡先做自己的 本地分片计算。
- 有些线性层按 输出维度切,这叫 Column Parallel。
- 有些线性层按 输入维度切,这叫 Row Parallel。
- 算完以后,有些地方要做 跨卡通信聚合,比如 all-reduce、all-gather、reduce-scatter。
- 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
第一层把维度放大,比如从 h 到 4h。这一步很适合 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
最后的自检
你可以用这几个问题检查自己是否真的懂了:
- 如果每张卡算出来的是不同的输出列,最后应该拼接还是相加?
- 如果每张卡算出来的是同一个输出的不同来源贡献,最后应该拼接还是相加?
- MLP 的第一层为什么常用 Column Parallel?
- MLP 的第二层为什么常用 Row Parallel?
- 为什么 activation 经常可以本地做?
- 为什么张量并行不是 GPU 越多越快?
答案压缩成一句:
Column 负责“各出各的列”,Row 负责“各算各的贡献”;
本地 CUDA kernel 负责把 shard 算快,NCCL collective 负责把必须合并的地方合并。