跳到主要内容

数据截至 (上游 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:413linear_projinput_is_parallel=True:419linear_qkvgather_output=Falseattention.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,每个都只有几行,对偶关系一眼可见:

  • _CopyToModelParallelRegionmappings.py:201):forward 原样返回 input,:213backward_reduce(grad_output):218
  • _ReduceFromModelParallelRegion:221):forward_reduce:232backward 恒等,:237
  • _ScatterToModelParallelRegion:240):forward 沿最后维 _split_along_last_dim:252backward 沿最后维 _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]」。

前向(forwardlayers.py:1232)的通信只有两处条件分支:

  • 入口:默认调 copy_to_tensor_model_parallel_region(input_)layers.py:1279)——前向恒等、反向 all-reduce,把「输入是各卡复制的」这件事的正确性悄悄处理掉。开了 sequence_parallelallreduce_dgrad 时跳过(因为序列并行下输入已经被切过)。
  • 出口:gather_output=True 才 all-gather(layers.py:1335-1338gather_from_tensor_model_parallel_region),否则各卡留着自己的 output_parallel 切片。

4.3 RowParallelLinear

layers.py:1382。前向(forwardlayers.py:1570)镜像对称:

  • 入口:input_is_parallel=True 时直接用输入切片(它约定上游是列切层,:1582-1583);否则 scatter_to_tensor_model_parallel_region 现场切(:1586)。
  • 出口:永远求和——普通模式调 reduce_from_tensor_model_parallel_regionlayers.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=Falsemlp.py:224);
  • linear_fc2:行切,input_is_parallel=Truemlp.py:248)。

两个开关一关一开,中间激活的同步就没了。整个 Transformer 层(attention + MLP)前向因此只有 2 次 all-reduce(各在 RowParallel 出口),这是 Megatron 论文的招牌结论。

4.5 词表并行

同一套刀法用在两个「词表维」上:

  • VocabParallelEmbeddinglayers.py:286):embedding 表按词表维切,每卡只管 vocab_size/tp 个 token 的向量,查表后 mask 掉不属于自己段的再 all-reduce。
  • vocab_parallel_cross_entropymegatron/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 = 2mlp.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/_ReduceScatterToSequenceParallelRegionmappings.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.pycopy_to_tensor_model_parallel_regionreduce_from_tensor_model_parallel_regionscatter_to_tensor_model_parallel_regiongather_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.pyColumnParallelLinear
行切线性层megatron/core/tensor_parallel/layers.pyRowParallelLinear
词表切 embeddingmegatron/core/tensor_parallel/layers.pyVocabParallelEmbedding
词表切交叉熵megatron/core/tensor_parallel/cross_entropy.pyVocabParallelCrossEntropyvocab_parallel_cross_entropy
MLP 配对证据megatron/core/transformer/mlp.pyMLPlinear_fc1/linear_fc2
Attention 配对证据megatron/core/transformer/attention.pySelfAttentionlinear_qkv/linear_proj

下一章:03 · 流水并行——把层切段之后,怎么排期才能让卡少等。