PyTorch 2.x 深度学习专题【左扬精讲】—— 张量并行:列并行与行并行的前向/反向传播、MHA 子层切分与通信开销最小化
PyTorch 2.x 深度学习专题【左扬精讲】—— 张量并行:列并行与行并行的前向/反向传播、MHA 子层切分与通信开销最小化
本篇深入剖析大模型训练中最重要的层内并行技术—— 张量并行(Tensor Parallelism, TP)。
当模型的单层参数(Transformer 单一层内部全部可学习权重,包含自注意力模块的 QKV 投影权重 $W_q$ / $W_k$ / $W_v$、输出投影 $W_O$,以及 FFN 前馈网络的升维权重 $W_{up}$、降维权重 $W_{down}$。
单层参数量随 hidden 维度 $h$ 的平方增长,例如 $h=4096$ 时单层参数量约 200M,单层权重总大小可达数十 GB)大到一张 GPU 装不下时,需要先理解三种并行策略各自的切分逻辑:
AllReduce 求平均 $\overline{\nabla W} = \frac{1}{N}\sum_i \nabla W_i$ → 各卡用同一份平均梯度同步更新自己的权重副本。send 给 GPU 1 → GPU 1 接着算第 $K_1\!+\!1$ 层……像工厂流水线一样串行接力,层间用 P2P Send/Recv 或 all-to-all 传激活。AllReduce 聚合 → 输出 $Y$ 完整。反向同理,需要
AllReduce 同步梯度。
三者关键区分(一眼对比):
| 并行方式 | 切分对象 | 最小切分单位 | 能否切开单层内部矩阵 |
|---|---|---|---|
| 数据并行(DP) | 训练样本(batch) | 完整模型 | ❌ 不行 每张卡完整复制全部权重 |
| 流水线并行(PP) | 网络层 | 完整 Transformer 层 | ❌ 不行 一层必须完整待在一张卡 |
| 张量并行(TP) | 权重 / 激活张量的特征维度 | 矩阵内部子块 | ✅ 可以 把一层里面矩阵劈开分到多卡 |
张量并行的精髓只有两招:列并行(Column Parallel)和行并行(Row Parallel)。它们以 GEMM(General Matrix Multiplication)为基本单元,对权重矩阵 W 在不同维度做切分,再配合 AllReduce 通信,让数学结果与单卡完全一致。
本篇将一步步讲清楚列/行并行的前向传播、反向传播、MHA 子层切分、连续线性层如何交替使用、以及对残差 Add 与 LayerNorm 的影响。
torch.distributed <- 分布式通信原语(AllReduce/AllGather/ReduceScatter)
torch.distributed.fsdp <- 全分片数据并行(与 TP 正交,本篇不展开)
torch.nn.functional.linear <- 一切列/行并行的数学起点 Y = XA + b
torch.cuda.amp.autocast <- 混合精度上下文(与 TP 通信可融合)
megatron.core.tensor_parallel.ColumnParallelLinear <- Megatron-LM 的列并行实现
megatron.core.tensor_parallel.RowParallelLinear <- Megatron-LM 的行并行实现
megatron.core.tensor_parallel.mappings <- 集合通信算子与 autograd 包装
megatron.core.tensor_parallel.CopyToModelParallelRegion <- 前向 copy、反向 AllReduce
megatron.core.tensor_parallel.ReduceFromModelParallelRegion <- 前向 AllReduce、反向 copy
torch.distributed.tensor.parallel.DistributedTensorParallel <- PyTorch 2.x 原生 TP
torch.nn.parallel.DistributedDataParallel <- 与 TP 正交,常与 TP 组合
# 核心 API 列表:集合通信原语 + TP 实现入口 + autograd 包装 + PyTorch 2.x 原生 TP
张量并行列并行行并行AllReduceMegatron-LMMHALayerNorm通信开销最小化PyTorch 2.x
学习重点提示
本篇深入探讨张量并行的列/行切分策略、前反向传播、MHA 子层切分、对残差和 LayerNorm 的影响、以及通信开销最小化设计。以下是每个主题你需要掌握的深度说明:
张量并行的诞生与演进(What & How & Why):
- What:张量并行是层内(intra-layer)模型并行的一种,把单个权重矩阵沿特征维度切分到多卡协同计算;与之相对的是层间(inter-layer)的流水线并行。
- How:能说出张量并行诞生的三大历史推力(单卡内存墙、DP 失效、PP 颗粒度太粗);能列出演进过程中的关键节点:朴素手工切分 → PipeDream 流水线 → Megatron-LM 列/行切分 → Megatron-LM 序列并行 → Transformer Engine 通信融合。
- Why:理解"为什么不能简单把 W 一刀两半让两卡各算一半"——根本原因是矩阵乘法的耦合性,必须配合集合通信;理解"为什么要在列/行之间交替"——为了把通信次数压到 2 次/层。
列并行 Column Parallel(What & How & Why):
- What:把权重矩阵 A 沿输出维度(列方向)切分,每张卡持有 A[:, k*out/p : (k+1)*out/p],输入 X 在所有卡上完整复制;前向输出 Y_k = X @ A_k,每张卡持有不同的输出特征切片。
- How:能写出 Y = [Y_0, Y_1, ..., Y_{p-1}] 的拼接语义;能画图描述"输入完整复制,权重列切分,输出分片"的物理布局;能解释反向传播时为何变成对 X 的梯度做 AllReduce。
- Why:没有列并行会怎样?→ 单卡必须持有完整 A,显存峰值随 hidden_size^2 爆炸;列并行的设计意图是让每张卡独立计算不同的输出特征,从而避免在前向做通信。
行并行 Row Parallel(What & How & Why):
- What:把权重矩阵 A 沿输入维度(行方向)切分,每张卡持有 A[k*in/p : (k+1)*in/p, :];输入 X 也按对应维度切分;前向每张卡算出部分和,最后通过 AllReduce 求和得到完整 Y。
- How:能写出 Y = sum_k (X_k @ A_k) 的求和语义;能解释为什么行并行的 AllReduce 放在权重之后、激活之前;能讲清反向传播时 A 的梯度是局部的、X 的梯度是 AllReduce 之后还原到分片。
- Why:理解行并行的设计意图是让求和操作只发生一次(避免在多个列并行后重复 AllReduce);没有行并行,连续的列并行会引入 2 次 AllReduce/层,行并行把通信量减半。
列并行的前向 / 反向传播(What & How & Why):
- What:前向 Y_k = X @ A_k + b_k,无通信;反向输入梯度 dX = AllReduce(dY_k @ A_k^T),权重梯度 dA_k = X^T @ dY_k 局部累加。
- How:能写出 autograd 视角的前向是 copy、反向是 AllReduce(对应 CopyToModelParallelRegion);能解释为何 dX 必须 AllReduce 而 dA 不需要。
- Why:理解反向为何不能在本地各自算——因为前向时每张卡只看到了 Y 的切片 Y_k,但 X 对所有切片的输出都有贡献,反向必须把所有切片的梯度求和才能还原 dX。
行并行的前向 / 反向传播(What & How & Why):
- What:前向 Y = AllReduce_k(X_k @ A_k) + b;反向 dA_k = X_k^T @ dY 局部,dX_k = dY @ A_k^T 无通信(已经分片)。
- How:能写出 autograd 视角的前向是 AllReduce、反向是 copy(identity)(对应 ReduceFromModelParallelRegion);能讲清"前向通信反向无通信"与列并行"前向无通信反向通信"的对称性。
- Why:理解列并行/行并行的反向是彼此对偶的——列并行的前向省通信,行并行的反向省通信;两者交替使用可以让总通信量最小化。
列/行并行交替的通信开销最小化(What & How & Why):
- What:在 MLP 的两个 GEMM 之间采用"先列并行后行并行"模式,在 MHA 中 Q/K/V 用列并行、输出投影用行并行,使每个 Transformer 块只有 2 次 AllReduce(1 次 fwd + 1 次 bwd 中)。
- How:能画出"列并行 → GeLU → 行并行"的通信流;能计算总通信量为 2 * (p-1)/p * (b*s*h) bytes(以 AllReduce 等价字节计);能解释为什么不能用两次列并行或两次行并行——会引入额外通信。
- Why:理解这是数学等价性与通信最少性的 trade-off;理解 Megatron-LM 选择这个模式的核心动机是"通信次数固定为 2,与切分度 p 无关"——这对大模型至关重要。
MHA 子层的张量并行(What & How & Why):
- What:把 num_heads 维度的注意力头切分到多卡上——Q/K/V 投影用列并行(沿输出维切出每个 head 的子矩阵),输出投影用行并行。
- How:能写出 [sq, b, h] --> [sq, b, (np * 3 * hn)] 的 reshape 语义;能解释为何 MHA 子层天然适合沿 head 切分——每个 head 的注意力计算互相独立;能讲清 num_attention_heads_per_partition 的含义。
- Why:没有 MHA 子层切分会怎样?→ 单个 head 维度小(典型 128),但 num_heads * 128 合计极大(如 4096);列并行利用了 head 维度的天然可分性,使得 QKV 三组矩阵沿 head 维独立切分时无需通信。
MHA 反向传播的通信(What & How & Why):
- What:MHA 子层反向需要 1 次 AllReduce(来自输出投影的行并行反向 → 还原 dX)。
- How:能写出"列并行的反向 + 行并行的反向"组合的具体通信位置;能解释 Softmax 反向在 TP 视角下是 no-op(因为各 head 独立)。
- Why:理解 MHA 整体通信模式与 MLP 块完全对称——前后两个列/行交替,1 次 fwd 通信 + 1 次 bwd 通信。
对残差 Add 和 LayerNorm 的影响(What & How & Why):
- What:残差 Add Y = X + Sublayer(X) 在 TP 视角下通常在列并行之前完成(输入完整)或在行并行之后完成(输出已汇总),因此无额外通信;LayerNorm 在 [b, s, h] 形式下天然按特征维归一化,对张量并行透明。
- How:能解释为何残差 Add 几乎"免费"——因为它在 GEMM 之外;能讲清 LayerNorm 必须在完整 h 维上做,因此出现在 TP 区域之外;能说出 PyTorch 2.x 序列并行(SP)把 LayerNorm 也沿 sequence 维切分,需要额外 AllGather / ReduceScatter。
- Why:理解这些"看似无关"的算子其实是 TP 区域边界的天然锚点——它们不需要切分,是把多个 TP 子层粘合起来的胶水。
阅读前提 & 建议:
- 前置知识:建议了解 Transformer 块结构(Self-Attention + MLP + 残差 + LayerNorm),了解 PyTorch 基础 torch.nn.Linear、torch.distributed 集合通信;本篇在 "PyTorch 2.x 深度学习专题【左扬精讲】—— 分布式训练基础" 之上展开。
- 不涉及的内容:流水线并行(PP)、专家并行(EP)、ZeRO/FSDP、序列并行(SP)的完整数学推导、Transformer Engine FP8 通信融合内核细节。
- 深度预期:学完本篇后,你能讲清楚张量并行为什么诞生、列/行并行的切分几何、前反向传播的通信模式、MHA 子层如何沿 head 切分、连续线性层为何要交替使用列/行并行、对残差和 LayerNorm 的影响;能区分列并行与行并行的对偶性;能回答 "为什么不能两次都用列并行"。
- 后续延伸:本篇 → 序列并行(SP)→ Pipeline Parallelism → FSDP → ZeRO → 3D 并行(TP+PP+DP)综合实践。
本文目录
一、张量并行的诞生背景与演进史
Why — 大模型为什么必须"切"?
在大模型(数十亿到数千亿参数)时代,单个 GEMM 的权重矩阵就可能达到 GB 级别。
例如 GPT-3 175B 的 MLP 升维矩阵形状是 [hidden_size=12288, ffn_hidden_size=49152],在 fp16 下占 2 * 12288 * 49152 ≈ 1.2 GB,再加上反向所需的梯度、优化器状态(Adam 对每一个模型参数,都要维护同样大小的一阶动量 m、二阶动量 v,所以优化器状态总元素数 = 2 × 模型参数量;这部分是额外开销,DP 每张卡都存全套,所以单卡显存爆炸,必须模型并行 / ZeRO),单卡根本装不下。
传统的 数据并行(Data Parallel, DP) 只能复制整个模型到每张卡上、用更大的 batch 训练,对单卡内存没有任何缓解作用。
这时候只有两条路:层间并行(流水线并行,Pipeline Parallelism)和层内并行(张量并行)。流水线并行颗粒度太粗——一整层只能放在一张卡上,对于单层就 1.2 GB 的矩阵仍然无能为力。因此张量并行应运而生:把单个权重矩阵 W 沿某个维度切成 p 份,让 p 张卡各持一份、协同完成一次 GEMM。
1.1 张量并行的诞生:朴素手工切分
最早的 "层内切分" 思路非常朴素:把 W 切成两半,让 rank 0 算左半、rank 1 算右半,最后通信汇总结果。
这种做法数学上是对的,但通信模式很糟糕:每一个 GEMM 都要做一次 AllReduce,而一次 AllReduce 通信量是 2 * (p-1) / p * msg_size(近线性增长),且通信是同步点会拖慢整体吞吐。
同时期还有一种更朴素的 "naive 模型并行":直接把模型按层切到不同卡上,rank 0 跑 layer 0、rank 1 跑 layer 1…… 这本质上就是流水线并行的雏形,但完全串行,GPU 利用率极低。
1.2 张量并行的关键演进节点
下图给出张量并行从诞生到当前主流方案的演进时间线(从左到右,箭头方向表示演进关系)。
+---------------------------------------------------------------------------+
| 张量并行(Tensor Parallelism)演进时间线 |
+---------------------------------------------------------------------------+
[2014-2017] 朴素手工切分
|
| 痛点:每个 GEMM 都要通信、通信量随切分度线性增长、无成熟实现
v
[2018] PipeDream / GPipe(流水线并行 PP)
|
| 解决了"按层切"的串行问题,引入 micro-batch + bubble
| 但单层过大时仍然束手无策
v
[2019.09] Megatron-LM (Shoeybi et al.) 首次提出列/行切分
| "Training Multi-Billion-Parameter Language Models" (arXiv:1909.08053)
|
| 核心洞察:Transformer 块内 GEMM 沿特定维度切分可让通信量最少
| 创新点:列/行交替 → 2 次 AllReduce / 层,与切分度 p 无关
v
[2021] Megatron-LM v2 (Narayanan et al.) 流水线 + 张量 + 数据 三维并行
|
v
[2022] Megatron Core + Sequence Parallelism (Korthikanti et al.)
| "Reducing Activation Recomputation in Large Transformer Models" (arXiv:2205.05198)
|
| 把 LayerNorm / Dropout 也沿 sequence 维切分
| 通信总量不变(AllReduce 拆成 AllGather + ReduceScatter)
v
[2022-2023] Transformer Engine + FP8 通信融合
| 把 AllReduce 算子融合进 GEMM 内核,隐藏通信延迟
v
[2023-2024] PyTorch 2.x 原生 TP (DTensor) + torchtitan 落地
|
v
[2024-2026] 主流 LLM 训练(Llama-3、Qwen-3、DeepSeek)默认开 TP+PP+DP+EP
# 说明:从朴素切分到 Megatron 列/行交替是质变,从 TP 到 SP 是量变(用通信换内存)
1.3 为什么是 Megatron 找到了最优解?
关键洞察在于 Transformer 块的结构是高度规则的 "对称 GEMM 链":MLP 是 X -> FC1 -> GeLU -> FC2 -> Y 两个 GEMM,Self-Attention 是 X -> QKV -> Softmax(QK^T)V -> O -> Y。在这样的对称结构里,可以精确地选 切哪个维度 才能让通信次数最少。
Megatron 团队发现:对 FC1 列切、对 FC2 行切,可以把通信次数压到 2 次 / 层,与切分度 p 无关。这就是张量并行的灵魂。
注意:"通信次数与 p 无关"指的是每张卡上 AllReduce 通信的次数(每次 1 次),而不是通信量。每张卡的单次 AllReduce 通信量仍然是 2 * (p-1) / p * b * s * h,随 p 增长趋于上限,但不会爆炸到 p 倍。
本节小结
-
- 诞生原因:单层参数过大突破单卡内存墙,DP 无效、PP 颗粒度太粗,必须做"层内切分"。
- 关键节点:朴素手工切分 → Megatron 列/行交替(2 次 AllReduce/层)→ 序列并行(AllGather + ReduceScatter 替代 AllReduce)→ 通信融合。
- 设计精髓:让通信次数与切分度 p 无关,这是张量并行能 scale 到 p=8/16/64 的根本原因。
二、专业名词定义:列/行并行、AllReduce、集合通信
What — 这一节给出后续章节涉及的所有专业名词的精确定义。
2.1 张量并行(Tensor Parallelism, TP)
张量并行 是一种 层内模型并行(intra-layer model parallelism)策略:把单个神经网络层(通常是 torch.nn.Linear)的 权重张量 沿某个维度切分到多张 GPU 上,让多卡协同完成同一次前向/反向计算。
张量并行的 "张量" 二字,指的就是这些被切分的权重和激活张量。它与 "流水线并行(按层切)" 和 "数据并行(按 batch 切)" 正交,三者可组合成 3D 并行。
2.2 切分度(Tensor Parallel Size, TP Size)
记为 p(或 tp_size),表示把单层权重切到几张卡上。
例如 tp_size=8 表示把一个 Linear 切到 8 张卡上协同计算。在 Megatron-LM 里这个值由 --tensor-model-parallel-size 指定。切分度 p 越大,每张卡上的权重越少,但通信次数不变、单次通信量增加(趋于 2 * (p-1)/p * msg_size)。
2.3 列并行(Column Parallel)
列并行把权重矩阵 A ∈ R^{in × out} 沿 输出维度(列方向)切成 p 份,每张卡持有 A_k = A[:, k*out/p : (k+1)*out/p]。
输入 X 在所有卡上完整复制。前向时每张卡独立计算 Y_k = X @ A_k,输出按列拼接成完整 Y = [Y_0, Y_1, ..., Y_{p-1}],但每张卡上只有自己的 Y_k 切片。列并行的核心特征是前向无通信。
2.4 行并行(Row Parallel)
行并行把权重矩阵 A ∈ R^{in × out} 沿 输入维度(行方向)切成 p 份,每张卡持有 A_k = A[k*in/p : (k+1)*in/p, :]。
输入 X 也按对应维度切分为 X_k = X[:, k*in/p : (k+1)*in/p]。前向时每张卡独立计算部分积 Y_k = X_k @ A_k,然后通过 AllReduce 把所有 Y_k 求和得到完整 Y = Σ_k Y_k。行并行的核心特征是 前向做 1 次 AllReduce。
2.5 AllReduce
AllReduce 是一种集合通信原语:参与通信的 p 张卡各自持有一段数据(shape 完全相同),通信结束后 每张卡都获得所有卡数据的逐元素求和(或平均/最大值等)结果。
在张量并行中通常用 sum 或 mean。AllReduce 的通信量为 2 * (p-1) / p * msg_size(band-optimal 算法,如 ring-allreduce),当 p=8 时约为 1.75 倍 msg_size。在 Megatron-LM 里 AllReduce 由 NCCL 的 ncclAllReduce 实现。
2.6 AllGather
AllGather:p 张卡各自持有一段切片(shape 互不重叠),通信结束后 每张卡都获得完整拼接后的张量。与 AllReduce 不同的是 AllGather 不做归约,只做拼接。通信量为 (p-1) / p * msg_size。
2.7 ReduceScatter
ReduceScatter:p 张卡各自持有完整数据,通信结束后每张卡获得 该数据归约后的一段切片。
等价于 "先 Reduce 再 Scatter",通信量与 AllGather 相同((p-1)/p * msg_size),但数据语义相反。AllGather + ReduceScatter 在通信量上等价于一次 AllReduce(2 * (p-1)/p * msg_size),这是序列并行的关键。
2.8 MHA(Multi-Head Attention)
Transformer 块中负责自注意力的子层。把 hidden 维度 h 拆成 num_heads 个独立头,每个头维度 head_dim = h / num_heads。
Q/K/V 各自由一个 GEMM 投影得到;注意力分数由 softmax(Q @ K^T / sqrt(head_dim)) @ V 计算;最后拼接所有头并过一个输出投影 W_O。MHA 的天然结构是 每个 head 独立计算,这使得它非常适合沿 head 维度做张量并行。
2.9 集合通信(Collective Communication)
多个进程协同完成的通信操作总称。
在张量并行场景下通常发生在同一个 ProcessGroup(即 TP group)内部,由 NCCL / Gloo(跨平台集合通信后端,CPU/GPU 都能用,常用于没有 NVLink 的环境或小规模分布式) 等后端实现。常见集合通信原语:AllReduce、AllGather、ReduceScatter、Broadcast、AllToAll。
2.10 ProcessGroup
一个逻辑通信域,由一组 rank 组成。
同一张卡可以加入多个 ProcessGroup(如 TP group、PP group、DP group),从而在不同维度做不同的集合通信。Megatron-LM 在初始化时为每种并行维度建一个独立的 ProcessGroup。
本节小结
-
- 列并行:切输出维、前向无通信、反向通信;权重 [in, out/p],输入 [b*s, in] 完整复制。
- 行并行:切输入维、前向通信、反向无通信;权重 [in/p, out],输入 [b*s, in/p] 同样切分。
- AllReduce:所有卡得到归约结果,通信量 2 * (p-1)/p * msg。
- AllGather / ReduceScatter:AllReduce 可拆解为这对操作,通信量相同但内存布局不同(SP 利用这一点)。
三、列并行(Column Parallel)详解
What — 列并行是什么?
列并行(ColumnParallelLinear,Megatron 命名)是张量并行的第一种切分方式:把权重矩阵沿 输出维度(列方向)切,输入在所有卡上完整复制。
每张卡独立计算不同的输出特征切片,前向完全无通信。这是张量并行里 "看起来最朴素" 但实际最重要的一招。
Why — 为什么需要列并行?
问题一:单卡装不下完整的 A。当 A 形状为 [12288, 49152](GPT-3 175B 风格的 MLP 升维矩阵)时,fp16 下 1.2 GB,加上 Adam 优化器状态(2x 参数量)共约 3.6 GB,单卡放不下。
问题二:数据并行完全无法缓解。DP 把整个模型复制到每张卡上,权重总量和单卡一样,DP 只能扩展 batch size,对单卡内存无帮助。
没有列并行会发生什么?
-
- 后果 1:单卡必须持有完整 A,权重峰值随 hidden × ffn_hidden 平方增长;175B 级别模型训练在单卡完全不可能。
- 后果 2:即使把 8 张卡拼成 "大显存卡"(如 8 * 80GB = 640GB),单次 GEMM 也无法把 1.2 GB 的矩阵放进任何一张卡,必须做切分。
- 后果 3:无法 scale 到 175B / 530B / 1T 参数级别;大模型训练变成"不能做"。
设输入 X ∈ R^{b × s × in}、权重 A ∈ R^{in × out},则 Linear 计算 Y = X @ A ∈ R^{b × s × out}。列并行把 A 沿输出维切成 p 份:
A = [A_0 | A_1 | ... | A_{p-1}],其中 A_k ∈ R^{in × out/p}
每张卡 k 持有 A_k,输入 X 在所有卡完整复制。前向:
Y_k = X @ A_k ∈ R^{b × s × out/p},所有卡独立、无通信。
反向时,损失对 Y 的梯度 dY 也会按列切成 dY = [dY_0 | dY_1 | ... | dY_{p-1}],每张卡只持有 dY_k。权重梯度:
dA_k = X^T @ dY_k ∈ R^{in × out/p},局部计算、无通信。
输入梯度:
dX = AllReduce_k(dY_k @ A_k^T) ∈ R^{b × s × in},必须 AllReduce。
下图是列并行切分几何(以 tp_size=2、A ∈ R^{in × out} 为例):
列并行(Column Parallel, tp_size=2)
=================================
输入 X(所有卡完整复制) 权重 A(沿输出维切分)
+----------------------+ +-----------------+-----------------+
| | | | |
| X ∈ R^{b×s×in} | | A_0 ∈ R^{in×out/2} | A_1 ∈ R^{in×out/2}|
| | | | |
| rank 0: X | | rank 0: A_0 | |
| rank 1: X | | | rank 1: A_1 |
| ... | | | |
+----------+-----------+ +-----------------+-----------------+
| | |
v v v
+----------------------+ +-----------------+-----------------+
| X @ A_0 | | Y_0 = X@A_0 | |
| rank 0 计算 | | ∈ R^{b×s×out/2} | |
+----------------------+ +-----------------+ |
| X @ A_1 | | | Y_1 = X@A_1 |
| rank 1 计算 | | | ∈ R^{b×s×out/2}|
+----------------------+ +-----------------+-----------------+
完整结果(逻辑上) Y = X @ A = [Y_0 | Y_1] 每张卡只有自己的 Y_k 切片
# 说明:输入 X 完整复制到每张卡;A 沿输出维(列)切;每张卡独立算 Y_k;前向无通信
关键观察:Y = X @ A = X @ [A_0 | A_1 | ... | A_{p-1}] = [X@A_0 | X@A_1 | ... | X@A_{p-1}]。因为矩阵乘法对列分块 天然可分,每张卡算自己的那块互不干扰。物理上每张卡只持有 Y_k ∈ R^{b×s×out/p},完整 Y 是所有 Y_k 的"逻辑拼接",但物理上没有任何卡持有完整 Y。
这种 "逻辑拼接但物理分片" 的特性正是列并行的精髓:前向时 GEMM 输出天然落在不同卡上,无需任何通信。代价是 反向时必须 AllReduce 才能还原 dX(见第六章)。
Megatron-LM 中列并行的实现是 ColumnParallelLinear(位于 megatron/core/tensor_parallel/layers.py)。其前向核心逻辑:
class ColumnParallelLinear(torch.nn.Module): # Megatron 列并行实现
def __init__(self, input_size, output_size, *, config, ...): # 构造:input_size, output_size, TP config
tp_size = get_tensor_model_parallel_world_size() # 取 TP 切分度 p
self.output_size_per_partition = output_size // tp_size # 每张卡分到的输出维度
self.weight = Parameter(torch.empty(...)) # 权重形状 [input_size, output_size_per_partition]
def forward(self, input_): # 前向:Y_k = X @ A_k
input_parallel = copy_to_tensor_model_parallel_region(input_) # 输入"广播"到所有卡(autograd hack)
output = F.linear(input_parallel, self.weight, self.bias) # 本地 GEMM,无通信
return output # 每张卡持有 Y_k,物理不完整
关键点:copy_to_tensor_model_parallel_region 是个 autograd Function,前向是 identity(无通信),反向是 AllReduce。这正是列并行"前向无通信、反向通信"的核心机制。
本节小结
- 切分几何:权重 [in, out/p]、输入 [b, s, in] 完整复制、输出 [b, s, out/p] 物理分片。
- 前向无通信:GEMM 沿列天然可分,每张卡独立算。
- 反向必通信:dX = AllReduce(dY_k @ A_k^T),因为 dX 对所有切片的输出梯度都有贡献。
- 典型应用:MLP 的 FC1、Self-Attention 的 Q/K/V 投影。
四、行并行(Row Parallel)详解
What — 行并行是什么?
行并行(RowParallelLinear,Megatron 命名)是张量并行的第二种切分方式:把权重矩阵沿 输入维度(行方向)切,输入按相同维度切分。
每张卡独立计算部分积 Y_k = X_k @ A_k,然后通过 AllReduce 求和得到完整 Y = Σ_k Y_k。行并行的核心特征是 前向做 1 次 AllReduce。
Why — 为什么需要行并行?
问题一:列并行后输出是分片的,下一层怎么用?。
MLP (Multi‑Layer Perceptron,多层感知机 泛指由多个全连接层 + 激活堆叠而成的网络,只要是 FC -> ACT -> FC 就属于 MLP。 FFN Feed-Forward Network 前馈网络的本质上是一个两层的 MLP ) 的 FC1(\(FC1: d_{hidden} \rightarrow 4d_{hidden}\) 把特征维度放大 升维,通常4倍维度) 用列并行后,每张卡上的 Y_k ∈ R^{b×s×out/p} 是分片的。如果 FC2 想要接收完整形状的输入,要么做 AllGather(多一次通信),要么让 FC2 也用列并行——但两个列并行连续使用会引入额外通信。
问题二:GEMM 的求和本质。注意 Y = X @ A 中,Y_{ij} = Σ_r X_{ir} * A_{rj}。对 A 沿行切分后,A_k = A[k*in/p : (k+1)*in/p, :],那么 Y_{ij} = Σ_k Σ_{r ∈ slice_k} X_{ir} * A_{rj} = Σ_k (X_k @ A_k)_{ij}。这意味着不同卡的输出可以通过简单的 sum 拼起来——这就是 AllReduce 的来源。
没有行并行会发生什么?
-
- 后果 1:连续两个线性层都用列并行时,第一层列切输出 Y_k ∈ R^{b×s×out/p},第二层想要对每张卡独立 GEMM,必须在第二层前做 AllGather 把 Y 拼回完整——多一次通信,总通信次数 4 次/层。
- 后果 2:如果连续两个都用行并行,第一层行并行做 AllReduce 还原完整 Y,第二层再行切+再 AllReduce——又是 4 次通信。
- 后果 3:只有列-行交替才能让总通信次数压到 2 次/层(一次 fwd AllReduce + 一次 bwd AllReduce),这是张量并行能 scale 的关键。
行并行把 A ∈ R^{in × out} 沿输入维度(行)切成 p 份:
A = [A_0; A_1; ...; A_{p-1}],其中 A_k ∈ R^{in/p × out}
输入 X ∈ R^{b × s × in} 也按相同维度切:
X = [X_0 | X_1 | ... | X_{p-1}],其中 X_k ∈ R^{b × s × in/p}
前向每张卡独立算 Y_k = X_k @ A_k ∈ R^{b × s × out},再 AllReduce 求和:
Y = AllReduce_k(Y_k) = Σ_k X_k @ A_k ∈ R^{b × s × out},每张卡持有完整 Y。
反向时,dY ∈ R^{b × s × out} 是完整的(每张卡都有)。权重梯度:
dA_k = X_k^T @ dY ∈ R^{in/p × out},局部计算、无通信。
输入梯度:
dX_k = dY @ A_k^T ∈ R^{b × s × in/p},无通信(每张卡算自己的切片)。
下图是行并行切分几何(以 tp_size=2 为例):
行并行(Row Parallel, tp_size=2)
=================================
输入 X(沿输入维切分) 权重 A(沿输入维/行切分)
+-----------+-----------+ +---------+---------+
| | | | | |
| X_0 | X_1 | | A_0 | A_1 |
| ∈ R^{b×s | ∈ R^{b×s | | ∈ R^{in/2×out} |
| ×in/2} | ×in/2} | | | |
| | | | | |
| rank 0 | rank 1 | | rank 0 | rank 1 |
+-----------+-----------+ +---------+---------+
| | | |
v v v v
+-----------+-----------+
| X_0 @ A_0 | X_1 @ A_1 | 每张卡独立算部分积
| ∈ R^{b×s×out} |
+-----------+-----------+
| |
+------+------+
|
v
+--------------+
| AllReduce | sum 归约
| (sum) |
+--------------+
|
v
+--------------+
| Y = Y_0+Y_1 | 完整 Y,每张卡持有
| ∈ R^{b×s×out}|
+--------------+
# 说明:X 和 A 都沿输入维切;每张卡算 Y_k;前向做 AllReduce 求和;每张卡得到完整 Y
关键观察:Y = X @ A = [X_0 | X_1 | ... | X_{p-1}] @ [A_0; A_1; ...; A_{p-1}] = Σ_k X_k @ A_k。因为矩阵乘法对行分块天然是求和结构,每张卡算出的 Y_k 是不完整的"部分积",必须通过 AllReduce 求和才能得到完整 Y。物理上每张卡在 AllReduce 之后持有完整 Y ∈ R^{b×s×out}。
这种"输入输出都不同形状"的特性是行并行的精髓:前向付出 1 次 AllReduce 通信的代价,换来反向时 dX 天然分片、无需通信。
Megatron-LM 中行并行的实现是 RowParallelLinear。其前向核心逻辑:
class RowParallelLinear(torch.nn.Module): # Megatron 行并行实现
def __init__(self, input_size, output_size, *, config, input_is_parallel=True, ...): # 构造:标记输入已分片
tp_size = get_tensor_model_parallel_world_size() # 取 TP 切分度 p
self.input_size_per_partition = input_size // tp_size # 每张卡分到的输入维度
self.weight = Parameter(torch.empty(self.input_size_per_partition, output_size)) # 权重 [in/p, out]
def forward(self, input_): # 前向:Y = AllReduce_k(X_k @ A_k)
output = F.linear(input_, self.weight) # 本地 GEMM,得到部分积 Y_k
output_parallel = reduce_from_tensor_model_parallel_region(output) # AllReduce 求和
if self.bias is not None: # 偏置只在 rank 0 加(避免重复)
output_parallel = output_parallel + self.bias
return output_parallel # 每张卡持有完整 Y
关键点:reduce_from_tensor_model_parallel_region 也是个 autograd Function,前向是 AllReduce,反向是 identity(无通信)。这与列并行的 copy_to_tensor_model_parallel_region 正好对偶。
本节小结
-
- 切分几何:权重 [in/p, out]、输入 [b, s, in/p] 同样切分、输出 [b, s, out] 完整(AllReduce 后)。
- 前向必通信:Y = AllReduce(X_k @ A_k),因为矩阵乘法对行分块是求和结构。
- 反向无通信:dX_k = dY @ A_k^T 天然分片,无需通信。
- 典型应用:MLP 的 FC2、Self-Attention 的输出投影 W_O。
五、列/行并行的前向传播对比
What — 前向传播的核心对比
列并行和行并行的前向传播在通信模式、权重形状、输入形状、输出形状上完全对偶。本节用对比表 + 时序图的方式把差异讲清楚,让读者一眼看出"为什么列-行交替能让通信最少"。
5.1 核心差异对比表
| 维度 | 列并行(Column Parallel) | 行并行(Row Parallel) |
|---|---|---|
| 切分维度 | 权重沿输出维(列)切 | 权重沿输入维(行)切 |
| 权重形状 | [in, out/p] | [in/p, out] |
| 输入形状 | [b*s, in] 完整复制到每张卡 | [b*s, in/p] 沿输入维切分 |
| 每卡输出 | [b*s, out/p] 物理分片 | [b*s, out] 完整(AllReduce 后) |
| 前向通信 | 无(每张卡独立 GEMM) | 1 次 AllReduce(sum 归约) |
| 数学本质 | Y = [X@A_0 | X@A_1 | ... | X@A_{p-1}],列分块拼接 | Y = Σ_k X_k@A_k,行分块求和 |
| 典型应用 | MLP 的 FC1、Self-Attention 的 Q/K/V 投影 | MLP 的 FC2、Self-Attention 的输出投影 W_O |
5.2 前向传播时序对比图
下图把列并行和行并行的前向步骤横向并排画出,方便对比(以 tp_size=2 为例):
列并行前向(Column Parallel Forward) 行并行前向(Row Parallel Forward)
================================== ===================================
X 完整复制到两张卡 X 沿输入维切到两张卡
A_0 在 rank0, A_1 在 rank1 A_0 在 rank0, A_1 在 rank1
| |
v v
rank0: Y_0 = X @ A_0 rank0: y_0 = X_0 @ A_0
rank1: Y_1 = X @ A_1 rank1: y_1 = X_1 @ A_1
| |
v v
(无通信) (AllReduce sum)
| |
v v
rank0 持有 Y_0 (分片) rank0 持有 Y = y_0+y_1 (完整)
rank1 持有 Y_1 (分片) rank1 持有 Y = y_0+y_1 (完整)
| |
v v
逻辑拼接:[Y_0 | Y_1] = Y 直接使用:Y
物理上每卡只有 Y_k 每卡都有完整 Y
通信次数: 0 通信次数: 1 (AllReduce)
通信量: 0 通信量: 2 * (p-1)/p * b*s*out
# 关键对比:列并行"前向无通信但输出分片",行并行"前向通信但输出完整"
5.3 前向传播的"几何直觉"
对权重 A ∈ R^{in × out},考虑一个具体的例子:in=4, out=6, tp_size=2,则:
-
- 列并行:A = [A_0 | A_1],其中 A_0 ∈ R^{4×3}, A_1 ∈ R^{4×3}。
- 行并行:A = [A_0; A_1],其中 A_0 ∈ R^{2×6}, A_1 ∈ R^{2×6}。
列并行的物理意义是"把输出通道分给不同卡算"——rank 0 算出 Y_0 的前 3 个特征,rank 1 算出 Y_1 的后 3 个特征。每张卡看到的输入都是完整的 4 维,只是各自只关心输出维的一部分。
行并行的物理意义是"把输入特征拆给不同卡算"——rank 0 只看到输入的前 2 维 X_0,rank 1 只看到后 2 维 X_1;每张卡独立算一个 2×6 的小矩阵乘,然后通过 AllReduce 把所有 2×6 的结果加起来,得到完整的 1×6 输出。行并行的"行"字就是"沿输入维切"的形象说法。
本节小结
-
- 列并行前向:每张卡独立 GEMM,输出物理分片,无通信。
- 行并行前向:每张卡独立算部分积,AllReduce 求和,输出完整。
- 对偶性:列并行"前向无通信"与行并行"前向通信"互为对偶,输出形状一个分片一个完整。
六、列/行并行的反向传播对比
What — 反向传播的核心对比
反向传播是张量并行最容易绕晕的地方——列并行的反向要通信,行并行的反向反而不要通信。理解这个"对偶性"是掌握张量并行的关键。本节从数学推导、autograd 视角、通信位置三个层面讲清楚。
6.1 反向传播数学推导对比
| 反向项 | 列并行(Column Parallel) | 行并行(Row Parallel) |
|---|---|---|
| 前向:每卡输出 | Y_k = X @ A_k(分片) | Y_k = X_k @ A_k,再 AllReduce |
| 反向:dY 的形状 | dY_k,每卡只有自己那块分片 | dY,每卡持有完整 |
| 反向:权重梯度 dA | dA_k = X^T @ dY_k,局部、无通信 | dA_k = X_k^T @ dY,局部、无通信 |
| 反向:输入梯度 dX | dX = AllReduce(dY_k @ A_k^T),必通信 | dX_k = dY @ A_k^T,无通信 |
| 反向通信次数 | 1 次 AllReduce | 0 次 |
| 反向通信量 | 2 * (p-1)/p * b*s*in | 0 |
6.2 反向传播时序对比图
下图把列并行和行并行的反向步骤横向并排画出:
列并行反向(Column Parallel Backward) 行并行反向(Row Parallel Backward)
==================================== ====================================
每卡持有 dY_k(上游梯度,分片) 每卡持有完整 dY(上游梯度)
| |
v v
本地算 dA_k = X^T @ dY_k 本地算 dA_k = X_k^T @ dY
(无通信) (无通信)
| |
v v
本地算 dY_k @ A_k^T 本地算 dY @ A_k^T
(每张卡得到部分 dX) (每张卡直接得到 dX_k)
| |
v v
(AllReduce sum) (无通信)
把所有切片的 dX 拼回完整 每张卡已经持有自己的 dX_k
| |
v v
rank0 持有完整 dX rank0 持有 dX_0(分片)
rank1 持有完整 dX rank1 持有 dX_1(分片)
通信次数: 1 (AllReduce) 通信次数: 0
通信量: 2 * (p-1)/p * b*s*in 通信量: 0
# 关键对比:列并行"反向必通信还原 dX",行并行"反向无通信天然分片 dX_k"
6.3 autograd 视角的对偶性
Megatron-LM 把列/行并行的通信封装成两个对称的 torch.autograd.Function:
class _CopyToModelParallelRegion(torch.autograd.Function): # 列并行的"输入广播"包装
@staticmethod
def forward(ctx, input_): # 前向:identity(输入已被 broadcast)
return input_
@staticmethod
def backward(ctx, grad_output): # 反向:AllReduce(每张卡的 dX 必须拼起来)
return _all_reduce(grad_output, group=tp_group)
class _ReduceFromModelParallelRegion(torch.autograd.Function): # 行并行的"输出归约"包装
@staticmethod
def forward(ctx, input_): # 前向:AllReduce(部分积求和成完整 Y)
return _all_reduce(input_, group=tp_group)
@staticmethod
def backward(ctx, grad_output): # 反向:identity(每张卡持有自己切片的 dX)
return grad_output
对偶性:
-
- _CopyToModelParallelRegion:前向 identity / 反向 AllReduce。
- _ReduceFromModelParallelRegion:前向 AllReduce / 反向 identity。
这两个函数正好是张量并行的"阴阳两面":列并行把通信放在反向,行并行把通信放在前向。两者交替使用可以让"前向的 AllReduce"和"反向的 AllReduce"分别位于不同的子层,从而让总通信次数压到 2 次/层。
6.4 为什么列并行反向要通信,而行并行反向不要?
Why — 反向通信的"几何"原因
列并行反向为何必须 AllReduce dX?前向时每张卡只算了 Y_k = X @ A_k,但损失 L = loss(Y_0, Y_1, ..., Y_{p-1}) 对 X 的梯度实际上是 dX = Σ_k dY_k @ A_k^T——因为 X 同时参与了所有切片的输出,反向必须把所有切片的贡献加回来。这是"前向无通信 → 反向必通信"的根本原因。
行并行反向为何不需要 AllReduce?前向时每张卡算 Y_k = X_k @ A_k,AllReduce 后 Y 完整。反向时 dX_k = dY @ A_k^T,因为 X_k 只参与了第 k 块输出,反向时各卡的 dX_k 互不干扰,已经是天然分片。这是"前向必通信 → 反向无通信"的根本原因。
没有对偶会怎样?
-
- 如果列并行/行并行的反向都需要通信,那每个子层都会引入 2 次 AllReduce,总通信次数翻倍。
- 如果列并行/行并行的反向都不需要通信,那数学上根本对不上(损失对 X 的梯度算不出来)。
- 对偶设计是张量并行能够"用最少的通信次数还原数学等价性"的根本保证。
6.5 前向 + 反向的完整通信时刻表
下图给出"列并行 + 行并行"组合的完整前反向通信时刻表(一个 MLP 块内):
时间轴 → 列并行(FC1) GeLU 行并行(FC2) 残差+LayerNorm
| | | |
前向 fwd: 独立GEMM(无通信) 逐元素(无通信) AllReduce(通信1) 逐元素(无通信)
| | | |
反向 bwd: AllReduce(通信2) 逐元素(无通信) 独立GEMM(无通信) 逐元素(无通信)
| | | |
合计: 1 次 AllReduce 0 1 次 AllReduce 0
整个 MLP 块: 2 次 AllReduce(1 前向 + 1 反向),与切分度 p 无关!
# 关键洞察:列-行交替让"前向 1 次"和"反向 1 次"分别落在不同子层,总通信 = 2 次 / 层
本节小结
-
- 列并行反向:dX = AllReduce(dY_k @ A_k^T),1 次 AllReduce。
- 行并行反向:dX_k = dY @ A_k^T,无通信。
- autograd 对偶:_CopyToModelParallelRegion(前向 copy / 反向 AllReduce)vs _ReduceFromModelParallelRegion(前向 AllReduce / 反向 copy)。
- 设计意图:列并行的反向通信 + 行并行的前向通信 = 整个 MLP 块只有 2 次 AllReduce / 层,与切分度 p 无关。
七、连续线性层交替使用列/行并行的通信开销最小化
What — 通信开销最小化的标准模式
在 Megatron-LM 风格的张量并行里,连续的线性层(GEMM)之间采用"列并行 → 元素级算子 → 行并行"的固定模式。这个模式让每个 Transformer 块只有 2 次 AllReduce(1 前向 + 1 反向),与切分度 p 无关。这是张量并行能 scale 到 TP=8/16/64 而不爆炸的根本原因。
7.1 MLP 块的标准列-行交替模式
一个标准的 Transformer MLP 块由 FC1 -> GeLU -> FC2 组成。Megatron 风格的切分是:
-
- FC1:hidden_size -> ffn_hidden_size,用 列并行(切输出维)。
- GeLU:逐元素算子,作用在分片后的激活上,无通信。
- FC2:ffn_hidden_size -> hidden_size,用 行并行(切输入维)。
下图给出 MLP 块列-行交替的完整前向流:
MLP 块列-行交替(Column-Row Alternation)
==========================================
输入 X ∈ R^{b×s×h} (完整复制到每张卡)
|
v
[FC1: Column Parallel] 权重 A1 ∈ R^{h × 4h/p},切输出维
| 每张卡独立算: Y1_k = X @ A1_k
| 输出形状: Y1_k ∈ R^{b×s×4h/p} (分片)
v
[GeLU] 逐元素: Y1_k = GeLU(Y1_k)
| 形状不变: Y1_k ∈ R^{b×s×4h/p}
v
[FC2: Row Parallel] 权重 A2 ∈ R^{4h/p × h},切输入维
| 每张卡独立算: Y2_k = Y1_k @ A2_k
| 输出形状: Y2_k ∈ R^{b×s×h} (每张卡是部分积)
v
[AllReduce sum] 把所有 Y2_k 求和
| 通信 1: AllReduce
v
输出 Y ∈ R^{b×s×h} (每张卡持有完整 Y)
通信统计: 前向 1 次 AllReduce(位于 FC2 之后)
反向 1 次 AllReduce(位于 FC1 之后,对应列并行的反向)
总计: 2 次 / MLP 块 / 层
# 列-行交替的精髓:列并行的"前向无通信" + 行并行的"前向通信"分布在不同子层,通信次数最少
7.2 通信量计算
设 batch 大小 b、序列长度 s、hidden 大小 h、切分度 p,则:
单次 AllReduce 通信量(以 fp16/bf16 计,每个元素 2 字节):
comm = 2 * (p-1) / p * b * s * h * 2 bytes
当 p=8 时,(p-1)/p = 0.875,所以单次 AllReduce 通信量约为 1.75 * b * s * h * 2 bytes。一个 MLP 块共 2 次 AllReduce,总通信量约 3.5 * b * s * h * 2 bytes。
对比:朴素列-列模式(FC1 列并行、FC2 也列并行):
-
- FC1 列并行:前向无通信、反向 1 次 AllReduce。
- FC2 列并行:但 FC2 的输入是 FC1 的分片输出 Y1_k,需要 AllGather 把 Y1 拼回完整 → 多 1 次通信。
- 总计 3 次通信 / 层(1 fwd AllGather + 1 bwd AllGather + 1 bwd AllReduce)。
对比:朴素行-行模式(FC1 行并行、FC2 也行并行):
-
- FC1 行并行:前向 1 次 AllReduce 还原完整 Y1。
- FC2 行并行:FC2 又要切输入维,再次前向 1 次 AllReduce。
- 总计 2 次前向 + 2 次反向 = 4 次 / 层。
只有列-行交替能让总通信次数 = 2 次 / 层,这是 Megatron 选择这个模式的核心原因。
7.3 通信开销最小化的数学证明
Why — 为什么列-行交替通信最少?
设一个 MLP 块的两个 GEMM 分别有 M1([h, 4h])和 M2([4h, h])。考虑四种切分组合:
| 组合 | 前向通信 | 反向通信 | 总通信次数 | 总通信量(次数 × bsh) |
|---|---|---|---|---|
| 列-列 | 1 AllGather(中间) | 2 AllReduce | 3 | ~5 bsh |
| 列-行 ★ | 1 AllReduce | 1 AllReduce | 2 | ~3.5 bsh |
| 行-列 | 2 AllReduce | 1 AllGather | 3 | ~5 bsh |
| 行-行 | 2 AllReduce | 2 AllReduce | 4 | ~7 bsh |
结论:列-行交替是四种组合中通信次数最少的(2 次),且这个 2 次与切分度 p 无关!这就是 Megatron 风格的"通信最优切分"。
7.4 工业实践中的常见误区
注意:"通信次数与 p 无关"是一个极强的性质。它意味着即使 p 从 2 增加到 64,每层的总通信次数始终是 2 次。这让张量并行可以在节点内(NVLink 域)放心地使用大 p(如 TP=8),而不用担心通信变成瓶颈。
反过来说:如果选择了列-列或行-行,通信次数会随 p 增加——p=8 时列-列已经是 3 次,行-行 4 次,比列-行多 50%-100%。这就是为什么 Megatron-LM 严格禁止连续同向切分。
本节小结
-
- 标准模式:MLP 块采用"FC1 列并行 + GeLU + FC2 行并行";MHA 块采用"Q/K/V 列并行 + 输出投影行并行"。
- 通信次数:每个 Transformer 块 2 次 AllReduce(1 前向 + 1 反向),与切分度 p 无关。
- 设计精髓:列-行交替让"前向无通信"和"反向无通信"分别落在不同子层,从而把全局通信压到最少。
八、张量并行的模型切分示例(MLP 块完整示例)
What — 用一个具体例子看张量切分
这一节给出一个完整的 MLP 块张量切分示例,从单卡视角到 TP=4 视角,从权重布局到激活形状,从前向到反向,所有数字都列出。这是工业实践中"看 TP 切分图"的样板。
8.1 模型规模设定
-
- 隐藏维度 hidden_size = h = 8192(典型 7B 模型)
- FFN 维度 ffn_hidden_size = 4h = 32768
- 切分度 tp_size = 4
- batch b = 2、序列长度 s = 2048
- 精度 bf16(2 字节)
8.2 单卡视角的 MLP 块
单卡情况下,整个 MLP 块在一张卡上:
单卡 MLP 块(无张量并行)
========================
X ∈ R^{2 × 2048 × 8192} 权重总量:
FC1: W1 ∈ R^{8192 × 32768} W1: 8192*32768*2 = 512 MB
Y1 = X @ W1 + b1 W2: 32768*8192*2 = 512 MB
Y1 ∈ R^{2 × 2048 × 32768} b1 + b2: ~1 MB
GeLU: Y1 = GeLU(Y1)
FC2: W2 ∈ R^{32768 × 8192}
Y2 = Y1 @ W2 + b2
Y2 ∈ R^{2 × 2048 × 8192}
输出 Y = Y2
单卡总权重: 约 1 GB(不含优化器状态)
单卡总激活: 2 * 2048 * (8192+32768+8192) * 2 = 约 200 MB
# 单卡视角:所有权重和激活都在一张卡上
8.3 TP=4 切分后的 MLP 块
下图给出 TP=4 切分后,每张卡上的权重形状、激活形状、通信位置:
TP=4 切分后的 MLP 块(4 张卡协同)
=================================
每张卡持有的权重(4 张卡各 1/4):
+--------------------------------+--------------------------------+
| rank 0: W1_0 ∈ R^{8192×8192} | rank 1: W1_1 ∈ R^{8192×8192} |
| W2_0 ∈ R^{8192×8192} | W2_1 ∈ R^{8192×8192} |
| | |
| rank 2: W1_2 ∈ R^{8192×8192} | rank 3: W1_3 ∈ R^{8192×8192} |
| W2_2 ∈ R^{8192×8192} | W2_3 ∈ R^{8192×8192} |
+--------------------------------+--------------------------------+
W1 = [W1_0 | W1_1 | W1_2 | W1_3] (列切)
W2 = [W2_0; W2_2; ...] (行切,每张卡 8192 行)
每张卡权重大小: 512 MB / 4 = 128 MB(FC1)+ 128 MB(FC2)= 256 MB
前向流(rank 0 视角):
X ∈ R^{2×2048×8192} 所有卡完整复制
|
v
[FC1 列并行] Y1_0 = X @ W1_0 局部
Y1_0 ∈ R^{2×2048×8192} 形状: [b, s, 4h/p] = [2, 2048, 8192]
其他卡算 Y1_1, Y1_2, Y1_3
(无通信,每张卡独立 GEMM)
|
v
[GeLU] Y1_0 = GeLU(Y1_0) 逐元素,形状不变
|
v
[FC2 行并行] Y2_0 = Y1_0 @ W2_0 局部
Y2_0 ∈ R^{2×2048×8192} 部分积
其他卡算 Y2_1, Y2_2, Y2_3
|
v
[AllReduce sum]
comm 1: 2 * (4-1)/4 * 2*2048*8192 * 2 bytes
= 1.5 * 2 * 2048 * 8192 * 2
= 100 MB (单次 AllReduce 通信量)
|
v
Y ∈ R^{2×2048×8192} 每张卡完整持有 Y
反向流(rank 0 视角):
上游 dY ∈ R^{2×2048×8192} 每张卡完整持有
|
v
[FC2 反向]
dA2_0 = Y1_0^T @ dY 局部,无通信
dY1_0_partial = dY @ W2_0^T 每张卡得到部分 dY1
|
v
[AllReduce sum]
dY1 = AllReduce(dY1_0_partial, ...) comm 2: 100 MB
|
v
[FC1 反向]
dA1_0 = X^T @ dY1_0 局部,无通信 (但 dY1_0 = dY1 的切片)
dX = AllReduce(dY1_0 @ W1_0^T, ...) 等等, 这里还有 dX 的 AllReduce
# 注:完整反向流会涉及 2 次 AllReduce(FC1 反向的 dX 还原 + FC2 反向的中间梯度)
8.4 通信量对比
| 方案 | 单卡权重 | 前向 AllReduce | 反向 AllReduce | 总通信/层 |
|---|---|---|---|---|
| 无 TP(单卡) | 1024 MB | 0 | 0 | 0 |
| TP=4(列-行交替) | 256 MB | 100 MB | 100 MB | 200 MB / 步 |
| TP=4(列-列) | 256 MB | 100 MB AllGather | 200 MB AllReduce | 300 MB / 步 |
结论:TP=4 相比单卡,权重降低 4 倍;列-行交替相比列-列,通信降低 33%。这就是工业实践"列-行交替"是默认值的数学根据。
本节小结
-
- 切分几何:TP=4 时,每张卡持有 1/4 权重,FC1 沿输出维切 8192×8192,FC2 沿输入维切 8192×8192。
- 通信位置:前向 AllReduce 在 FC2 后,反向 AllReduce 在 FC1 后。
- 通信量:每张卡每步约 200 MB 通信(100 MB 前向 + 100 MB 反向),与 bsh 线性增长。
九、MHA 子层的张量并行工作机制
What — MHA 子层如何切分?
Multi-Head Attention(MHA)是 Transformer 块的核心子层,由 QKV 投影、注意力计算、输出投影三步组成。MHA 的天然结构是每个 head 独立计算,这使得它可以非常优雅地沿 num_heads 维度切分到多卡上:QKV 投影用列并行,输出投影用行并行。整个 MHA 子层同样只有 2 次 AllReduce / 层,与切分度 p 无关。
Why — 为什么 MHA 适合沿 head 切?
问题:MHA 的"头"为什么天然独立?。注意力计算的核心是 head_i = softmax(Q_i @ K_i^T / sqrt(d_k)) @ V_i,每个 head i 只与自己的 Q_i、K_i、V_i 交互,head 之间不共享任何激活。这意味着:
-
- 把 Q 矩阵沿 head 切分([b, s, num_heads * d_k] → [b, s, num_heads/p, d_k]),不同卡上的 Q 完全独立。
- 注意力分数、Softmax、注意力加权三个算子都在 head 内部完成,无需任何跨卡通信。
- 输出拼接时各 head 互不干扰,拼接后用一个 W_O 投影回 h。
没有沿 head 切会发生什么?
-
- 后果 1:单个 head 维度小(d_k = 128),但 num_heads * d_k = h 合计极大(如 4096),单卡 QKV 投影矩阵 [h, 3h] 放不下。
- 后果 2:只能把整个 QKV 矩阵沿 h 切到多卡,但这会让不同卡上的 Q/K/V 头完全割裂、Softmax 无法计算——必须先 AllGather,引入额外通信。
- 后果 3:沿 head 切是数学上最干净、通信最少的方案,是张量并行 MHA 的"天生选项"。
MHA 子层张量并行的核心是"把 num_heads 维度切到多卡"。具体步骤(以 TP=4、num_heads=32, d_k=128 为例):
- QKV 投影(列并行):用一个 [h, 3h] 的大矩阵一次性产出 Q/K/V,沿输出维切到 4 张卡。每张卡持有 num_heads/4 = 8 个 head 的 Q/K/V 切片,形状 [b, s, 3 * 8 * 128] = [b, s, 3072]。
- reshape & split:把 [b, s, 3072] 切成 Q、K、V 三个 [b, s, 8, 128] 张量(按最后一维 3 等分)。
- 注意力计算:Q @ K^T / sqrt(d_k) -> Softmax -> @ V,全部在每张卡本地完成([b, 8, s, 128] 形状),无通信。
- 输出拼接:每张卡本地把 8 个 head 拼成 [b, s, 1024],但每张卡只有 8 个 head 的输出。
- 输出投影 W_O(行并行):W_O ∈ R^{h × h} 沿输入维切到 4 张卡([1024, 8192]),每张卡算部分积 [b, s, 8192]。
- AllReduce sum:把 4 张卡的部分积求和成完整输出 [b, s, 8192],每张卡持有完整结果。
9.1 MHA 子层切分几何图
MHA 子层张量并行切分图(TP=4, num_heads=32, head_dim=128)
=========================================================
输入 X ∈ R^{b×s×h} = R^{b×s×8192} 完整复制到 4 张卡
|
v
+---------------------------------------------------------------+
| QKV 投影(Column Parallel) |
| W_qkv ∈ R^{8192 × 3×8192} = R^{8192 × 24576} |
| 沿输出维切 4 份,每张卡持有 W_qkv_k ∈ R^{8192 × 6144} |
| 每张卡持有 8 个 head(4 切分),每个 head_dim=128 |
+---------------------------------------------------------------+
| 每张卡独立 GEMM(无通信)
v
+---------------------------------------------------------------+
| 每张卡得到 mixed ∈ R^{b×s×6144} |
| reshape 为 [b, s, 8, 3, 128] |
| split 为 Q ∈ R^{b×s×8×128}, K ∈ ..., V ∈ ... |
+---------------------------------------------------------------+
| 8 个 head 是 rank 0 的子集(head 0,4,8,12,16,20,24,28)
v
+---------------------------------------------------------------+
| 注意力计算(每张卡本地,无通信) |
| scores = Q @ K^T / sqrt(128) |
| ∈ R^{b×8×s×s} |
| attn = softmax(scores) |
| out = attn @ V |
| ∈ R^{b×8×s×128} |
+---------------------------------------------------------------+
| 本地 reshape: [b, s, 1024] (8 * 128 = 1024)
v
+---------------------------------------------------------------+
| 输出投影 W_O(Row Parallel) |
| W_O ∈ R^{8192 × 8192} |
| 沿输入维切 4 份,每张卡持有 W_O_k ∈ R^{1024 × 8192} |
| 每张卡独立 GEMM: y_k = out @ W_O_k ∈ R^{b×s×8192} |
+---------------------------------------------------------------+
|
v
+---------------------------------------------------------------+
| AllReduce sum (1 次) |
| Y = y_0 + y_1 + y_2 + y_3 |
| 每张卡持有完整 Y ∈ R^{b×s×8192} |
+---------------------------------------------------------------+
通信统计: 前向 1 次 AllReduce(位于 W_O 之后)
反向 1 次 AllReduce(位于 QKV 之前)
总计: 2 次 / MHA 子层 / 层
# MHA 切分精髓:QKV 列并行让"按 head 切"成为可能,W_O 行并行让"跨 head 求和"成为可能
9.2 MHA 子层中"列-行交替"的位置
| 子层 | 类型 | 切分维度 | 前向通信 | 反向通信 |
|---|---|---|---|---|
| QKV 投影 W_qkv | 列并行 | 输出维(num_heads) | 无 | 1 次 AllReduce(dX 还原) |
| Softmax + Attn 加权 | 逐元素(head 内) | 无 | 无 | 无 |
| 输出投影 W_O | 行并行 | 输入维(head_dim * num_heads/p) | 1 次 AllReduce | 无 |
| 合计 | 列-行交替 | — | 1 次 | 1 次 |
本节小结
-
- 切分维度:沿 num_heads 切,每张卡持有 num_heads/p 个 head 的 QKV。
- 列-行交替:QKV 投影用列并行、W_O 用行并行,与 MLP 块完全对称。
- 通信次数:每个 MHA 子层 2 次 AllReduce(1 前向 + 1 反向),与 MLP 块合计 4 次 AllReduce / Transformer 块。
- 设计精髓:head 维度的天然独立性 + 列-行交替 = 通信最少。
十、MHA 子层张量并行的反向传播通信
What — MHA 反向通信的精确定位
MHA 子层反向传播共 1 次 AllReduce(来自 QKV 列并行的反向 dX 还原)。这一节讲清楚这次通信出现在哪里、为什么必须有、以及它和 MLP 块反向通信的关系。
10.1 MHA 反向传播的 4 个梯度
MHA 子层反向需要计算 4 类梯度:
- dQ, dK, dV:QKV 投影的输入梯度,分片,每张卡只有自己 head 的梯度。
- dX:进入 MHA 子层的输入梯度,需要 AllReduce 还原(来自 QKV 列并行的反向)。
- dW_qkv:QKV 投影的权重梯度,局部,每张卡只算自己 head 的部分。
- dW_O:输出投影的权重梯度,局部(行并行的反向权重梯度天然分片)。
10.2 MHA 反向传播的具体通信位置
MHA 子层反向传播时序图(TP=4, 视角 rank 0)
==============================================
时间轴 ← 反向传播
上游 dY ∈ R^{b×s×h} 每张卡完整持有(来自 MHA 输出 AllReduce 的反向)
|
v
[W_O 反向 - Row Parallel]
dW_O_0 = out^T @ dY 每张卡局部,无通信
dout_0 = dY @ W_O_0^T 每张卡得到自己 head 的 dout 切片
(无 AllReduce: 因为 W_O 行并行的反向是 identity)
|
v
[注意力反向 - head 内]
dV_0, dK_0, dQ_0 = ... 每张卡本地,无通信
dattn_0 形状 [b, 8, s, s] 分片
dQ, dK, dV 形状 [b, s, 8, 128] 分片
(无通信: head 之间独立)
|
v
[QKV 反向 - Column Parallel]
dW_qkv_0 = X^T @ d_qkv_0 每张卡局部
dX_partial_0 = d_qkv_0 @ W_qkv_0^T 每张卡得到自己切片的 dX
|
v
[AllReduce sum] 通信!
dX = AllReduce(dX_partial_0, dX_partial_1, dX_partial_2, dX_partial_3)
每张卡持有完整 dX ∈ R^{b×s×h}
|
v
传递给前一个子层(LayerNorm 或残差 Add)
通信统计: 反向 1 次 AllReduce(位于 QKV 反向之后)
# 与前向的 AllReduce(W_O 之后)合计 2 次 / MHA 子层
10.3 MHA 反向通信的"对偶性"
MHA 反向通信的位置与前向通信完全对偶:
-
- 前向:W_O 之后 AllReduce(行并行的前向通信)。
- 反向:QKV 之后 AllReduce(列并行的反向通信)。
这正是第六章"列并行/行并行的反向对偶"在 MHA 子层的具体体现。前向 1 次 + 反向 1 次 = 2 次 / MHA 子层 / 步,与 MLP 块完全一致。
10.4 Softmax 反向为什么不需要通信?
Softmax 算子 softmax(x_i) = exp(x_i) / Σ_j exp(x_j) 及其反向 dL/dx_i = softmax(x_i) * (dL/dy_i - Σ_j dL/dy_j * softmax(x_j)) 都是逐元素的(每行内归一化)。在张量并行 MHA 中,Softmax 作用在 [b, num_heads/p, s, s] 形状的注意力分数上,每张卡独立做 Softmax,反向也是独立的。没有任何跨卡依赖。
10.5 完整 Transformer 块的通信总结
把 MHA + MLP 合起来看,每个 Transformer 块共 4 次 AllReduce:
| 子层 | 前向 AllReduce | 反向 AllReduce |
|---|---|---|
| MHA 块(W_O 之后 / QKV 反向) | 1 | 1 |
| MLP 块(FC2 之后 / FC1 反向) | 1 | 1 |
| 合计 | 2 | 2 |
本节小结
-
- MHA 反向 1 次 AllReduce:位于 QKV 列并行反向之后,还原 dX。
- Softmax 反向无通信:因为 head 内部独立。
- 完整 Transformer 块:前向 2 次 + 反向 2 次 = 4 次 AllReduce / 块 / 步。
- 设计对称性:MHA 反向通信位置与前向完全对偶(一个在 W_O、一个在 QKV)。
十一、对残差 Add 和 LayerNorm 的影响
What — 残差 Add 和 LayerNorm 在 TP 视角下"几乎免费"
残差 Add Y = X + Sublayer(X) 和 LayerNorm Y = (X - μ) / σ * γ + β 都是逐元素或按特征维归一化的算子,天然适合张量并行。在标准 TP 中它们几乎不引入额外通信,但在序列并行(SP)中会被进一步沿 sequence 维切分。本节讲清这两类算子对张量并行的具体影响。
11.1 残差 Add:天然无通信
Transformer 块的标准结构是 Y = X + Sublayer(LayerNorm(X)),其中 Sublayer 是 MHA 或 MLP。在 TP 视角下,残差 Add 的两个输入 X 和 Sublayer(X) 都已经处在相同的物理布局上:
-
- 在 TP 区域之前:X 是完整复制,Sublayer(X) 是列并行输出(分片)。此时残差 Add 必须在 X 上做分片(每张卡加自己切片的 Sublayer 输出),即残差本身也要按列切分到多卡上——这隐含要求 X 在进入残差时也按列切分。
- 在 TP 区域之后:X 是列并行输出(分片),Sublayer(X) 是行并行输出(AllReduce 后完整)。此时需要在 X 之后做 AllReduce(变完整)再做残差——但 Megatron 的设计是把残差放在 LayerNorm 之后、下一块的 TP 之前,让残差本身沿分片形式做加法。
关键观察:在 Megatron-LM 标准实现中,残差 Add 无 AllReduce 通信。它的两个输入在残差发生时已经处于相同的物理布局(都是分片或都是完整),所以加法是逐元素完成的。残差只是 TP 区域之间的"胶水"。
下图展示残差 Add 在 TP 视角下的位置:
Transformer 块结构与残差 Add 位置(TP 视角)
==============================================
+----------+
| 输入 X | (列切,每张卡持有 X 的分片)
+----+-----+
|
v
+-----------------------+ |
| LayerNorm 之前 | | (X 分片形式进入)
+-----------------------+ |
v
+-----------------------+ |
| MHA 块 | | QKV 列并行 + W_O 行并行
| (2 次 AllReduce) | |
+-----------------------+ |
v
+----+-----+
| MHA 输出 | (行并行 AllReduce 后完整)
+----+-----+
|
v
+-----------------------+ |
11.2 LayerNorm:天然无通信(标准 TP 下)
LayerNorm Y = (X - μ) / σ * γ + β 在 [b, s, h] 形状下,μ 和 σ 是按 hidden 维 h 计算的(不是 batch 或 sequence 维)。在标准 TP 视角下:
-
- LayerNorm 之前的 X 是完整复制(每张卡都有完整 [b, s, h]):LayerNorm 在每张卡上独立算 μ、σ、Y,无通信。
- LayerNorm 之后的 Y 也是完整复制:可以直接送入列并行的 QKV 投影(输入 X 完整)。
所以在标准 TP(无 SP)下,LayerNorm 完全"免费"——每张卡独立做归一化,无通信、无同步。
11.3 序列并行(SP)对 LayerNorm 的影响
但有个问题:标准 TP 下 LayerNorm 之前/之后的激活都是完整复制到每张卡,意味着 [b, s, h] 形状的激活占满每张卡显存。当 batch 和 sequence 很大时,激活内存可能超过权重,成为新的瓶颈。
序列并行(Sequence Parallelism, SP)的解决方案是:把 LayerNorm 和 Dropout 的激活沿 sequence 维切分。具体做法:
-
- 在进入 Transformer 块前,把 [b, s, h] 沿 s 维切成 p 份,每张卡持有 [b, s/p, h]。
- LayerNorm 在 [b, s/p, h] 形状上独立做归一化(无通信,因为归一化在 h 维)。
- 进入 QKV 列并行前需要 AllGather 把 [b, s/p, h] 拼回 [b, s, h](行并行反向时类似地 ReduceScatter)。
| 方案 | LayerNorm 激活 | QKV 输入通信 | 总通信次数 / 块 | LayerNorm 通信 |
|---|---|---|---|---|
| 标准 TP | [b, s, h] 完整复制 | 无 | 4 AllReduce | 0 |
| TP + SP | [b, s/p, h] 切分 | AllGather + ReduceScatter | 4 AllGather/RS | 0 |
小贴士:SP 的核心洞察是"LayerNorm/Dropout 在 sequence 维度独立"——token 0 的 LayerNorm 结果与 token 4096 完全无关,因此可以沿 sequence 切。通信总量不变(AllGather + ReduceScatter ≈ AllReduce),但激活内存降低 p 倍。
11.4 残差 Add 在 SP 下的小调整
SP 下残差 Add 的两个输入变成了 sequence 维切分形式([b, s/p, h]),加法仍然是逐元素的(每张卡独立加自己的 sequence 切片),无通信。残差 Add 仍然是 SP 下的"免费"算子。
本节小结
-
- 残差 Add:在标准 TP 和 SP 下都无通信,是 TP 区域间的天然胶水。
- LayerNorm:在标准 TP 下无通信(按 h 维归一化,独立可分);在 TP+SP 下沿 s 维切分,仍无通信(按 h 维归一化),但需要 AllGather/ReduceScatter 进出 TP 区域。
- 设计意图:这些"看似无关"的算子是 TP 区域的天然锚点,它们不需要切分,把多个 TP 子层粘合起来。
十二、FAQ(20 组)
FAQ 分组说明。以下 20 组 Q&A 按主题分组:Q1-Q4 讲诞生与动机,Q5-Q8 讲列并行/行并行的数学,Q9-Q12 讲前向/反向通信,Q13-Q15 讲连续交替,Q16-Q17 讲 MHA,Q18-Q20 讲残差/LN/工业实践。
Q1. 张量并行和数据并行(DP)到底有什么区别?
一句话结论:DP 复制整个模型、扩展 batch size;TP 切分单个模型、扩展单层规模。展开:DP 把完整模型复制到 N 张卡上,每张卡处理不同的 batch 切片,适合模型能装下单卡但需要更大 batch 的场景。TP 把单个权重矩阵沿特征维切到 N 张卡上协同计算,适合单层参数过大的场景。两者正交,可组合成 3D 并行(TP×PP×DP)。
Q2. 流水线并行(PP)为什么不能替代张量并行?
一句话结论:PP 按"整层"切、颗粒度太粗,无法解决单层过大问题。展开:PP 把模型按"层"切到不同卡上(rank 0 跑 layer 0-7、rank 1 跑 layer 8-15),但一整层仍必须放在一张卡上。当单层参数达到 GB 级(如 GPT-3 175B 的 FC1 矩阵 1.2 GB)时,PP 无法解决"单卡装不下"的问题。TP 才是颗粒度到"权重矩阵"级别的切分方案。
Q3. 张量并行的切分度 p 选多大合适?
一句话结论:p 通常 ≤ 单机 GPU 数(如 8),跨节点 TP 会因互联带宽下降而性能急剧下降。展开:工业实践通常把 p 限制在单个 NVLink 域内(典型 8 卡),因为 AllReduce 在跨节点(InfiniBand)上的延迟是 NVLink 的 10-100 倍。Megatron-LM 论文(arXiv:2104.04473)明确指出:"当 p 大于单机 GPU 数时,跨节点张量并行的开销可能变得不切实际"。
Q4. 张量并行有哪些已知的局限性?
一句话结论:通信次数随 p 不变但单次通信量趋于上限;通信域必须 NVLink。展开:(1) 通信总次数固定为 2 次 / 层,但单次 AllReduce 通信量 2 * (p-1)/p * msg 随 p 趋于上限,对带宽要求高。(2) 当 p 跨节点时,AllReduce 受限于 InfiniBand 带宽,吞吐急剧下降。(3) LayerNorm、Embedding 算子天然对 TP 不友好(需要 sequence 并行补足)。
Q5. 列并行和行并行的本质区别是什么?
一句话结论:列并行切"输出维"、行并行切"输入维"。展开:列并行把权重 A ∈ R^{in × out} 切成 [in, out/p],每张卡算不同的输出特征;行并行把 A 切成 [in/p, out],每张卡算同一输出特征的不同部分积。前者前向无通信、后者前向必通信。
Q6. 列并行前向为什么无通信?
一句话结论:矩阵乘法对列分块天然可分,每张卡独立算自己的 Y_k = X @ A_k。展开:Y = X @ A = X @ [A_0 | A_1 | ... | A_{p-1}] = [X@A_0 | X@A_1 | ... | X@A_{p-1}]。每张卡算的 Y_k 是完整 Y 的不重叠切片,互不干扰,所以不需要通信。
Q7. 行并行前向为什么必须通信?
一句话结论:矩阵乘法对行分块是求和结构,每张卡只算部分积,必须 AllReduce 求和成完整 Y。展开:Y = X @ A = [X_0 | X_1 | ...] @ [A_0; A_1; ...] = Σ_k X_k @ A_k。每张卡的 Y_k = X_k @ A_k 是不完整的"部分积",必须通过 AllReduce 求和得到完整 Y。
Q8. 列并行和行并行的"对偶性"具体指什么?
一句话结论:列并行"前向无通信、反向必通信",行并行"前向必通信、反向无通信",互为对偶。展开:在 autograd 视角下,列并行用 _CopyToModelParallelRegion(前向 copy、反向 AllReduce),行并行用 _ReduceFromModelParallelRegion(前向 AllReduce、反向 copy)。两者的通信正好分布在"前向"和"反向"两个不同时刻,组合使用可以让总通信最少。
Q9. 列并行的反向为什么必须 AllReduce dX?
一句话结论:因为前向时每张卡的 Y_k = X @ A_k 都依赖完整 X,反向时 X 对所有切片的输出都有贡献,必须把所有切片的梯度加回来。展开:dX = Σ_k dY_k @ A_k^T,每张卡只持有 dY_k,必须 AllReduce 求和才能得到完整 dX。
Q10. 行并行的反向为什么不需要通信?
一句话结论:因为前向时每张卡的 Y_k = X_k @ A_k 只用了 X 的一个切片 X_k,反向时各卡的 dX_k = dY @ A_k^T 天然分片。展开:X_k 只参与了第 k 块输出,反向时各卡 dX_k 互不干扰,已经是天然分片,无需通信。
Q11. 一个 MLP 块(FC1 + GeLU + FC2)总通信次数是多少?
一句话结论:2 次 AllReduce / 层(1 前向 + 1 反向),与切分度 p 无关。展开:FC1 列并行(前向无通信、反向 1 次 AllReduce)+ FC2 行并行(前向 1 次 AllReduce、反向无通信)= 2 次 / 层。这个 2 次与 p 无关,是张量并行能 scale 的关键。
Q12. TP 通信量随 p 怎么变化?
一句话结论:单次 AllReduce 通信量 2 * (p-1)/p * msg 随 p 增长趋于 2*msg 的上限。展开:p=2 时 1.0 * msg、p=4 时 1.5 * msg、p=8 时 1.75 * msg、p=16 时 1.875 * msg、p→∞ 时 2.0 * msg。所以切分度 p 增加的边际通信量快速递减,但单次通信量不会无限增长。
Q13. 为什么不能"连续两个都用列并行"?
一句话结论:连续两个列并行会引入额外 AllGather 把中间激活拼回完整,总通信次数增加到 3-4 次 / 层。展开:FC1 列并行输出 Y1_k 是分片,FC2 列并行需要完整输入,必须在中间加 AllGather 拼回完整 Y1(1 次额外通信)。这破坏了"通信次数 2 次 / 层"的最优解。
Q14. 列-行交替为什么能让通信次数固定为 2?
一句话结论:列并行的"前向无通信"和行并行的"前向通信"分布在不同子层,反向类似,总通信 = 2 次 / 层。展开:FC1 列并行前向无通信、FC2 行并行前向 1 次 AllReduce(合计 1 次前向通信);FC1 反向 1 次 AllReduce、FC2 反向无通信(合计 1 次反向通信)。总通信 2 次 / 层,与切分度 p 无关。
Q15. 工业实践中有没有"破例"用列-列或行-行的场景?
一句话结论:极少。当某个 Linear 切分后输入/输出形状天然匹配时(如 Embedding 的 vocab 切分),可以连续使用同向切分。展开:典型例外是 VocabParallelEmbedding(沿 vocab 维切分)和 VocabParallelLMHead(沿 vocab 维切分),它们的切分维度与中间激活维度不耦合,可以连续使用。但 MLP 和 MHA 的内部 Linear 必须列-行交替。
Q16. MHA 为什么要沿 num_heads 切,而不是沿 head_dim 切?
一句话结论:沿 num_heads 切让每个 head 独立,注意力计算无需跨卡通信;沿 head_dim 切会让 Softmax 内部产生跨 head 依赖。展开:Softmax 在每个 head 内部做 softmax(Q_i @ K_i^T / sqrt(d_k)),head 之间不共享任何激活。沿 num_heads 切让每张卡只算自己 head 的 Q/K/V,Softmax、注意力加权都是本地无通信。如果沿 head_dim 切(即把 Q 的 head_dim 维度切到多卡),Softmax 内部需要 AllReduce,引入额外通信。
Q17. MHA 子层的总通信次数?
一句话结论:2 次 / MHA 子层 / 步(1 前向 + 1 反向),与 MLP 块完全对称。展开:QKV 列并行(前向无通信、反向 1 次)+ W_O 行并行(前向 1 次、反向无通信)= 2 次。完整 Transformer 块(MHA + MLP)= 4 次 / 块 / 步。
Q18. 残差 Add 在 TP 视角下需要通信吗?
一句话结论:不需要。残差 Add 是逐元素加法,在两个输入物理布局相同时无通信。展开:Megatron 标准实现把残差 Add 放在"两个输入都完整复制"的位置上(每张卡独立做加法);或者放在"两个输入都分片"的位置上(每张卡加自己切片的 Sublayer 输出)。两种情况下残差 Add 都不引入 AllReduce。
Q19. LayerNorm 在 TP 视角下需要通信吗?
一句话结论:标准 TP 下不需要;TP+SP 下 LayerNorm 仍然不需要,但需要 AllGather/ReduceScatter 进出 TP 区域。展开:LayerNorm 在 hidden 维 h 上做归一化,每张卡独立做,无通信。SP 沿 sequence 维切分激活,LayerNorm 仍在 h 维归一化(独立、无通信),但 QKV 列并行需要完整 [b, s, h] 输入,所以 TP 区域前要做 AllGather。
Q20. 一个完整的 LLM 训练框架中,TP 通信占总通信的比例?
一句话结论:典型 3D 并行(TP×PP×DP)中,TP 通信约占总通信的 30%-50%,是占比最大的一部分。展开:TP 通信每层 4 次 AllReduce(约 2 * bsh 通信量);PP 通信每层 1 次 P2P Send/Recv(约 2 * bsh);DP 通信每步 1 次 AllReduce 梯度同步(约 2 * 总参数量)。在 Llama-3 70B 这种规模下,TP 通信占比通常最大,因为 PP 的 bubble time 和 DP 的梯度同步频率较低。
FAQ 全篇总纲
-
- 基础概念:Q1-Q4 讲张量并行的诞生与定位(与 DP/PP 的关系、p 选多大、局限性)。
- 数学本质:Q5-Q8 讲列/行并行的切分几何和对偶性。
- 前反向通信:Q9-Q12 讲具体通信位置和通信量公式。
- 交替策略:Q13-Q15 讲为什么列-行交替是通信最优解。
- MHA 切分:Q16-Q17 讲 MHA 沿 head 切的工作机制和通信次数。
- 残差/LN/工业:Q18-Q20 讲残差和 LayerNorm 在 TP 视角下的特性,以及工业实践中的通信占比。
十三、Roadmap 预告
下一篇预告:序列并行(Sequence Parallelism, SP)
- 本篇基础:张量并行(TP)的列/行切分、前反向通信、MHA 子层切分。
- 下一篇目标:用 AllGather + ReduceScatter 替代 AllReduce,把 LayerNorm 和 Dropout 的激活也沿 sequence 维切分,让大模型训练的激活内存降低 p 倍。
- 核心问题:序列并行如何与张量并行正交组合?LayerNorm 在 SP 下的通信模式?SP 的总通信量是否与纯 TP 相同?
- 后续延伸:SP → Pipeline Parallelism(流水线并行)→ FSDP(Fully Sharded Data Parallel)→ 3D 并行综合实践。

浙公网安备 33010602011771号