数据截至 (上游 commit e79cb4c1bae1)
02 · 张量并行
这一章讲什么: 五把刀里通信量最大、设计最精巧的一把。核心是两件事:权重矩阵「按列切还是按行切」的数学,以及把全部通信封进四个 autograd 函数的「前向/反向对偶」技巧。读完你会理解 2019 年 Megatron 论文的全部精华,以及为什么 TP 只敢开在节点内。
1. 它要解决的小问题
TP/PP/DP 里,PP 按层切、DP 按数据切,都很好懂。但即便一层都放不下(或想摊薄单层计算)怎么办?
张量并行 = 把单独一层的矩阵乘法切到多张卡上算。 麻烦在于:矩阵乘不是逐元素操作,切开之后各卡算的是「部分结果」,必须在正确的位置把它们加回去——而且加回的位置不同,通信量天差地别。Megatron 的全部巧思就在「在哪儿加回」上。
2. 思路:列切与行切的对偶
线性层是 Y = XA + b。矩阵 A 有两种切法,性质完全不同:
| 切法 | A 怎么分 | 各卡算出什么 | 出口要做什么 |
|---|---|---|---|
| 按列切(column parallel) | A = [A₁, A₂, …, Aₚ],按输出维分 | Yᵢ = X·Aᵢ,是 Y 的一段列 | 要完整 Y 就 all-gather;不要就不用通信 |
| 按行切(row parallel) | A 按输入维分,X 也相应按列分 | Yᵢ = Xᵢ·Aᵢ,是 Y 的部分和 | 必须 all-reduce 求和才是完整 Y |
关键观察:列切的输出切片,恰好是行切需要的输入切片。
- 列切 A=[A₁…Aₚ] 得到 Y 的第 i 段;
- 下一个矩阵 B 若按行切,第 i 行块要的输入正是 Y 的第 i 段。
所以 Column → Row 配对,中间的激活各卡拿着自己的切片直接往后算,一次通信都不需要。整个通信被推迟到 RowParallelLinear 出口的一次 all-reduce。
图示:一个 MLP 块的通信账
hidden (全量, 各卡相同)
│
▼
┌───────────────────┐
│ ColumnParallel fc1 │ 每卡只存 4h/tp 列, 各算 Y 的一段
│ gather_output=❌ │ ← 出口不 gather
└─────────┬─────────┘
│ 激活切片 (每卡 4h/tp, 互不相同, 不通信!)
▼
activation (GeLU 逐元素, 无需通信)
│
▼
┌───────────────────┐
│ RowParallel fc2 │ 每卡只存 4h/tp 行, 算出部分和
│ input_is_parallel=✅│ ← 入口不 scatter
└─────────┬─────────┘
▼
all-reduce ← 整个 MLP 唯一一次通信
│
hidden (全量)
Attention 同理:QKV 投影是列切(每卡算自己那部分 head 的 Q/K/V),输出投影是行切(attention.py:413 的 linear_proj 设 input_is_parallel=True,:419;linear_qkv 设 gather_output=False,attention.py:1724)。注意力头天然沿 head 维瓜分,一行同步都不用加。
3. 原理演示:四个通信原语的对偶
Megatron 把 TP 需要的全部通信封成四个 torch.autograd.Function。妙处是前向和反向互为对偶——反向通信不用人写,autograd 自动配好:
# 示意,非源码
# 前向做什么 <-> 反向就自动做什么
copy: 前向 恒等 反向 all-reduce # 广播输入; 梯度其实各卡都该有一份
reduce: 前向 all-reduce 反向 恒等 # 求和部分和; 梯度直接往后传
scatter: 前向 切最后维 反向 all-gather # 分发输入; 梯度拼回来
gather: 前向 all-gather 反向 切最后维 # 拼输出; 梯度切回去
为什么 copy 的反向是 all-reduce?因为前向时同一份 X 被各卡「复制」使用(逻辑上),反向时每张卡都算出一份 dX,真值是所有副本贡献之和——所以 all-reduce。这就是「对偶」的含义:前向的通信方向反过来,就是反向的通信。
重点看:有了这四个函数,写 ColumnParallelLinear 的人只需在前向里插一个 copy(或什么都不插),反向的 all-reduce 由 autograd 白送。
4. 真实实现
4.1 四个 autograd 函数
在 megatron/core/tensor_parallel/mappings.py,每个都只有几行,对偶关系一眼可见:
_CopyToModelParallelRegion(mappings.py:201):forward原样返回 input,:213;backward调_reduce(grad_output),:218。_ReduceFromModelParallelRegion(:221):forward调_reduce,:232;backward恒等,:237。_ScatterToModelParallelRegion(:240):forward沿最后维_split_along_last_dim,:252;backward沿最后维_gather_along_last_dim,:257。_GatherFromModelParallelRegion(:260):forward gather / backward split,正好和 scatter 相反。
对外暴露的是四个同名包装函数(copy_to_tensor_model_parallel_region 等在 mappings.py:492-510),底层 _reduce/_split_along_last_dim/_gather_along_last_dim 在 :22、:40、:84。
4.2 ColumnParallelLinear
megatron/core/tensor_parallel/layers.py:986。docstring 第一句就是数学定义:「A is parallelized along its second dimension as A = [A_1, ..., A_p]」。
前向(forward,layers.py:1232)的通信只有两处条件分支:
- 入口:默认调
copy_to_tensor_model_parallel_region(input_)(layers.py:1279)——前向恒等、反向 all-reduce,把「输入是各卡复制的」这件事的正确性悄悄处理掉。开了sequence_parallel或allreduce_dgrad时跳过(因为序列并行下输入已经被切过)。 - 出口:
gather_output=True才 all-gather(layers.py:1335-1338调gather_from_tensor_model_parallel_region),否则各卡留着自己的output_parallel切片。
4.3 RowParallelLinear
layers.py:1382。前向(forward,layers.py:1570)镜像对称:
- 入口:
input_is_parallel=True时直接用输入切片(它约定上游是列切层,:1582-1583);否则scatter_to_tensor_model_parallel_region现场切(:1586)。 - 出口:永远求和——普通模式调
reduce_from_tensor_model_parallel_region(layers.py:1620);开序列并行时改用reduce_scatter_to_sequence_parallel_region(:1616-1618),all-reduce 和 scatter 合成一个 reduce-scatter,通信量减半。
4.4 MLP 里的配对证据
MLP.__init__(megatron/core/transformer/mlp.py:171)就是 §2 那张图的代码形态:
linear_fc1:列切,gather_output=False(mlp.py:224);linear_fc2:行切,input_is_parallel=True(mlp.py:248)。
两个开关一关一开,中间激活的同步就没了。整个 Transformer 层(attention + MLP)前向因此只有 2 次 all-reduce(各在 RowParallel 出口),这是 Megatron 论文的招牌结论。
4.5 词表并行
同一套刀法用在两个「词表维」上:
VocabParallelEmbedding(layers.py:286):embedding 表按词表维切,每卡只管vocab_size/tp个 token 的向量,查表后 mask 掉不属于自己段的再 all-reduce。vocab_parallel_cross_entropy(megatron/core/tensor_parallel/cross_entropy.py:213):logits 按词表切开时,交叉熵的 max 和 sum 都先在本地算、再跨 TP 组 all-reduce 出全局归一化项——避免把[s, b, vocab]的完整 logits gather 到一张卡。
5. 关键细节与坑
- GLU/SwiGLU 要
stride=2。 门控线性单元把 fc1 输出扩成 [gate, up] 两份,列切时必须交错切才能让每卡拿到配对的 gate/up 片段。MLP.__init__里为此给fc1_stride = 2(mlp.py:202-207,注释里明说这是为了跨 TP size 正确 reshard)。用 Kitchen Linear 时不支持 stride≠1,会退化为 1 并有注释告警。 - TP 的通信量决定它只能节点内。 每层 2 次 all-reduce、通信量随
batch × seq × hidden线性增长,与 tp 无关(切得越细,单卡算得越少但总量不变)。NVLink 域内(≤8 卡)才扛得住——这就是 01 章 order 里 tp 总是最内层的原因。 - 序列并行(SP)是 TP 的显存补丁。 LayerNorm 和 residual 之间的激活在纯 TP 下是各卡全量复制的;SP 把它按序列维切开,于是 RowParallel 出口的 all-reduce 换成 reduce-scatter、下一个 ColumnParallel 入口的 copy 换成 all-gather(
_GatherFromSequenceParallelRegion/_ReduceScatterToSequenceParallelRegion,mappings.py:300、:355)。通信总量不变,激活显存除以 tp。 - bias 不做 TP 求和。 注意 RowParallelLinear 里 bias 是 all-reduce 之后才加的(
layers.py:1622-1627)——每卡加完整 bias,不是切片的。ColumnParallel 的 bias 则按列切开(sharded_state_dict里{"weight": 0, "bias": 0},layers.py:1357)。细节虽小,搞反了就是静默错结果。 skip_bias_add=True是性能开关。 MLP 里两个线性层都设了它(mlp.py:226、:250),把 bias 加法和后面的操作融合的机会留给调用方——返回值是(output, bias)二元组而不是加好的张量。- 显式专家通信会关掉这些开关。
explicit_expert_comm为真时,ColumnParallel 入口不 copy、RowParallel 出口不 reduce(layers.py:1267-1273、:1613-1615)——MoE 专家的通信由 token dispatcher 显式接管,见 05 章。
6. 代码地图
| 主题 | 文件路径 | 符号名 |
|---|---|---|
| 四个 TP 通信原语 | megatron/core/tensor_parallel/mappings.py | copy_to_tensor_model_parallel_region、reduce_from_tensor_model_parallel_region、scatter_to_tensor_model_parallel_region、gather_from_tensor_model_parallel_region |
| 原语的 autograd 本体 | megatron/core/tensor_parallel/mappings.py | _CopyToModelParallelRegion、_ReduceFromModelParallelRegion、_ScatterToModelParallelRegion、_GatherFromModelParallelRegion |
| SP 版原语 | megatron/core/tensor_parallel/mappings.py | _GatherFromSequenceParallelRegion、_ReduceScatterToSequenceParallelRegion |
| 列切线性层 | megatron/core/tensor_parallel/layers.py | ColumnParallelLinear |
| 行切线性层 | megatron/core/tensor_parallel/layers.py | RowParallelLinear |
| 词表切 embedding | megatron/core/tensor_parallel/layers.py | VocabParallelEmbedding |
| 词表切交叉熵 | megatron/core/tensor_parallel/cross_entropy.py | VocabParallelCrossEntropy、vocab_parallel_cross_entropy |
| MLP 配对证据 | megatron/core/transformer/mlp.py | MLP(linear_fc1/linear_fc2) |
| Attention 配对证据 | megatron/core/transformer/attention.py | SelfAttention(linear_qkv/linear_proj) |
下一章:03 · 流水并行——把层切段之后,怎么排期才能让卡少等。