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 装不下时,需要先理解三种并行策略各自的切分逻辑:

数据并行(DP)不可行
📌 干什么 每张 GPU 都保存完整一份模型权重 $\left\{W_q, W_k, W_v, W_O, W_{up}, W_{down}\right\}_{full}$,分到不同 GPU 上的只是训练样本(mini-batch 被均分成 $N$ 份,每张 GPU 一份)。
🔁 计算流程 各卡独立做前向 → 反向得到局部梯度 $\nabla W_i$ → 跨卡 AllReduce 求平均 $\overline{\nabla W} = \frac{1}{N}\sum_i \nabla W_i$ → 各卡用同一份平均梯度同步更新自己的权重副本
📦 切分维度 batch 维(即"样本维度" —— 一次喂给模型的训练样本数量 $B$,如 32 / 128 / 1024。把这个数字拆成 $N$ 份,每张 GPU 一次只算 $B/N$ 条样本,但模型权重本身不动)不切分模型权重,也不切分 hidden 维度 $h$。
✖ 致命弱点:每张 GPU 都必须持有完整的模型权重——单卡装不下单层权重时,DP 完全失效。
流水线并行(PP)不可行
📌 干什么 把整个 Transformer 栈按切成 $N$ 段,第 $1\!\sim\!K_1$ 层放 GPU 0、第 $K_1\!+\!1\!\sim\!K_2$ 层放 GPU 1,以此类推,每张 GPU 负责完整若干层的前向 + 反向。
🔁 计算流程 GPU 0 算完第 $K_1$ 层 → 把激活 $a_{K_1}$ send 给 GPU 1 → GPU 1 接着算第 $K_1\!+\!1$ 层……像工厂流水线一样串行接力,层间用 P2P Send/Recvall-to-all 传激活。
📦 切分维度 网络层维(即"层维度"——Transformer 总共有 $L$ 层(如 $L=48$ 层),把前 $1\!\sim\!K_1$ 层、第 $K_1\!+\!1\!\sim\!K_2$ 层……依次分配给不同 GPU;每张 GPU 拿到的不是"矩阵的一部分",而是"连续的几层")不切分单层内部矩阵,也不切分 hidden $h$ 或 batch。
✖ 颗粒度太粗:最小切分单位是完整 Transformer 层,无法解决"单层本身参数就超单卡容量"的问题。
张量并行(TP)★ 唯一可行
📌 干什么 把单一层内部的权重矩阵沿特征维度(hidden $h$)切分成 $N$ 块,每张 GPU 拿一份 $\frac{1}{N}$ 的子矩阵 + 一份 $\frac{1}{N}$ 的输入子块,协同完成同一层的 GEMM
🔁 计算流程 前向:各卡算 $Y_i = X_i W_i$($X_i$ 是输入按列切,$W_i$ 是权重按列切)→ 列并行 AllReduce 聚合 → 输出 $Y$ 完整。
反向同理,需要 AllReduce 同步梯度。
📦 切分维度 hidden 维(即"特征维 / 通道维"——矩阵的列方向,如 $h=4096$ 时每个 token 的向量长度是 4096 维;沿这个维度切就是把"一个 token 的完整向量"拆成 $N$ 份,每张 GPU 只算 $\frac{h}{N}$ 维)可以切开单层内部矩阵,最小单位是矩阵内部的子块 $\left(\frac{h}{N}\right) \times h$ 或 $h \times \left(\frac{h}{N}\right)$。
✔ 专门解决单层参数过大而单卡显存不足的问题,能够把权重和激活精细切分到多张 GPU 上。

三者关键区分(一眼对比):

并行方式切分对象最小切分单位能否切开单层内部矩阵
数据并行(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 必须 AllReducedA 不需要。
    • 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.Lineartorch.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 的环境或小规模分布式) 等后端实现。常见集合通信原语:AllReduceAllGatherReduceScatterBroadcastAllToAll

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 参数级别;大模型训练变成"不能做"。
How — 列并行的数学定义与几何切分图

设输入 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=2A ∈ 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;前向无通信
Why — 列并行前向为何能"无通信"?

关键观察: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(见第六章)。

How — Megatron 代码视角

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 的关键。
How — 行并行的数学定义与几何切分图

行并行把 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
Why — 行并行前向为何"必通信"?

关键观察: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 天然分片、无需通信。

How — Megatron 代码视角

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 前向传播的"几何直觉"

How — 用矩阵切分视角理解前向

对权重 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 视角的对偶性

How — autograd Function 的对偶设计

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 风格的切分是:

    • FC1hidden_size -> ffn_hidden_size,用 列并行(切输出维)。
    • GeLU:逐元素算子,作用在分片后的激活上,无通信。
    • FC2ffn_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 的"天生选项"
How — MHA 张量并行的完整工作流

MHA 子层张量并行的核心是"把 num_heads 维度切到多卡"。具体步骤(以 TP=4、num_heads=32, d_k=128 为例):

    1. QKV 投影(列并行):用一个 [h, 3h] 的大矩阵一次性产出 Q/K/V,沿输出维切到 4 张卡。每张卡持有 num_heads/4 = 8 个 head 的 Q/K/V 切片,形状 [b, s, 3 * 8 * 128] = [b, s, 3072]
    2. reshape & split:把 [b, s, 3072] 切成 Q、K、V 三个 [b, s, 8, 128] 张量(按最后一维 3 等分)。
    3. 注意力计算Q @ K^T / sqrt(d_k) -> Softmax -> @ V,全部在每张卡本地完成([b, 8, s, 128] 形状),无通信
    4. 输出拼接:每张卡本地把 8 个 head 拼成 [b, s, 1024],但每张卡只有 8 个 head 的输出。
    5. 输出投影 W_O(行并行)W_O ∈ R^{h × h} 沿输入维切到 4 张卡([1024, 8192]),每张卡算部分积 [b, s, 8192]
    6. 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 类梯度:

    1. dQ, dK, dV:QKV 投影的输入梯度,分片,每张卡只有自己 head 的梯度。
    2. dX:进入 MHA 子层的输入梯度,需要 AllReduce 还原(来自 QKV 列并行的反向)。
    3. dW_qkv:QKV 投影的权重梯度,局部,每张卡只算自己 head 的部分。
    4. 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 反向通信的"对偶性"

How — 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:天然无通信

How — 残差 Add 为什么"免费"?

Transformer 块的标准结构是 Y = X + Sublayer(LayerNorm(X)),其中 Sublayer 是 MHA 或 MLP。在 TP 视角下,残差 Add 的两个输入 XSublayer(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 下)

How — LayerNorm 为什么"免费"?

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 并行综合实践。

posted @ 2026-08-08 17:59  左扬  阅读(21)  评论(0)    收藏  举报