🏠 总目录📚 本教程 张量并行 ← →
📑 本页目录(点开跳转)

12 · 张量并行

⏱ 26 分钟 | ⭐ 把一层切开放到多张卡上


🎯 一句话

前面的并行都是"每卡一份完整的层"。 张量并行把【一层内部的矩阵】横着或竖着切开,分给多张卡。 它是唯一能救"单层都装不下"的办法 —— 代价是每层都要通信。


✂️ 一、两种切法

设一层是 $Y = \text{GeLU}(XA)$,权重 $A$ 要切开。

① 列并行(按列切 A)

对照

A = [A₁ | A₂] ← 竖着切

GPU0 算 Y₁ = GeLU(X A₁)

GPU1 算 Y₂ = GeLU(X A₂)

⭐ 关键:GeLU 是逐元素的,可以【各算各的】,不需要通信!

结果拼起来就是完整的 Y = [Y₁ | Y₂]

② 行并行(按行切 A)

信息关系

A = ⎡A₁⎤ ← 横着切
⎣A₂⎦
输入也要按列切:X = [X₁ | X₂]
GPU0 算 X₁A₁
GPU1 算 X₂A₂
⚠️ Y = X₁A₁ + X₂A₂→必须 AllReduce 求和

🧩 二、Megatron 的经典组合(⭐ 这一节是核心)

MLP 块:$Y = \text{GeLU}(XA) \cdot B$

⭐ 精妙之处:第一层用【列并行】,第二层用【行并行】 第一层 列并行 第二层 行并行 GPU0 X XA₁ GeLU Y₁ Y₁B₁ GPU1 X XA₂ GeLU Y₂ Y₂B₂ AllReduce 输出 → 中间【完全不需要通信】 → 整个 MLP 块只需要【一次】AllReduce ⭐
两条支路从左跑到右一直没有交叉 —— 中间四步全在各自卡内算完,要看的就是箭头合拢的位置有多靠右:整个 MLP 只在那里跨卡一次。

🔑 为什么这样组合是最优的: 列并行的输出天然是"按列切开的", 而行并行需要的输入恰好就是"按列切开的" —— 严丝合缝。 如果顺序反过来(先行后列),中间就要多一次通信。

Attention 块:更自然

结果对照

多头注意力本来就是【多个独立的头】
直接把头分给不同的卡:GPU0 拿头 1-8,GPU1 拿头 9-16
各自算完整的 attention(Q、K、V 投影都按列并行)
最后的输出投影用行并行→一次 AllReduce ⭐

关键信息

⭐ 一个 Transformer 层的总通信:
前向:MLP 一次 AllReduce + Attention 一次 AllReduce = 2 次
反向:同样 2 次
每层 4 次 AllReduce

📉 三、通信代价:为什么它必须在机内

算一算

每次 AllReduce 的数据量 = batch × seq_len × hidden × 2 字节

例:batch=8, seq=2048, hidden=8192, BF16

→ 8 × 2048 × 8192 × 2 = 268 MB

→ 每层 4 次 = 1.07 GB

→ 80 层 = 86 GB【每一步】⭐

在 NVLink 900GB/s:约 0.1 秒 ✅ 可接受

在 PCIe 64GB/s: 约 1.3 秒 ⚠️ 慢 13 倍

跨机 IB 25GB/s: 约 3.4 秒 💀 完全不可用

🔑 这个数字直接给出了铁律: 张量并行度不要超过单机的 GPU 数(通常是 8)。 这不是建议,是带宽算出来的硬约束。

⭐ 另外注意:通信量和层数、hidden、序列长度都成正比,但和张量并行度无关 —— 也就是说,TP=2 和 TP=8 的通信量一样,但 TP=8 分摊的计算更少 → 并行度越高,通信占比越大,效率越低。


🔧 四、代码长什么样

# 概念示意:列并行的线性层
import torch
import torch.nn as nn
import torch.nn.functional as F
class ColumnParallelLinear(nn.Module):
    def __init__(self, in_f, out_f, tp_size, tp_rank):
        super().__init__()
        assert out_f % tp_size == 0
        self.weight = nn.Parameter(torch.empty(out_f // tp_size, in_f))
        #                                      ↑ 只存自己那一片

    def forward(self, x):
        return F.linear(x, self.weight)      # ⭐ 不需要通信

# 行并行
class RowParallelLinear(nn.Module):
    def __init__(self, in_f, out_f, tp_size, tp_rank):
        super().__init__()
        assert in_f % tp_size == 0
        self.weight = nn.Parameter(torch.empty(out_f, in_f // tp_size))

    def forward(self, x):
        out = F.linear(x, self.weight)
        dist.all_reduce(out, group=TP_GROUP)  # ⭐ 这里必须通信
        return out

实际使用(别自己写):

# 🧩 骨架:`model` 来自你自己的代码,这一段只看写法
# ⭐ PyTorch 原生(2.x)
from torch.distributed.tensor.parallel import (
    parallelize_module, ColwiseParallel, RowwiseParallel)

parallelize_module(model, tp_mesh, {
    "attn.qkv_proj": ColwiseParallel(),
    "attn.out_proj": RowwiseParallel(),      # ⭐ 列→行的配对
    "mlp.gate_proj": ColwiseParallel(),
    "mlp.up_proj":   ColwiseParallel(),
    "mlp.down_proj": RowwiseParallel(),
})

💡 别自己实现张量并行。用 Megatron-LM、PyTorch 原生 TP、 或推理侧的 vLLM/TensorRT-LLM(tensor_parallel_size=8)。 自己写极容易在权重初始化、随机种子、梯度缩放上出微妙的错。


🧵 五、序列并行:一个重要的补充

因果链

⚠️ 张量并行有个盲区:LayerNorm 和 Dropout 【没有被切】
每张卡都要存完整的激活
长序列时这部分显存很可观
✅ 序列并行(Sequence Parallel):
在 LayerNorm/Dropout 这些地方,按【序列维度】切开
和张量并行【交替使用】,用 AllGather/ReduceScatter 转换
⭐ 效果:激活显存再降 TP 倍,而【通信总量不变】
(因为 AllReduce 被拆成了 ReduceScatter + AllGather)

💡 又是那个恒等式(第 5 章、第 11 章): AllReduce = ReduceScatter + AllGather。 这已经是它第三次出场了 —— 它是分布式训练里最有用的一个等式。

# 🧩 骨架:`sequence` 来自你自己的代码,这一段只看写法
# Megatron 里开启
--sequence-parallel     # ⭐ 几乎总是该开,和 TP 配套

⚠️ 六、四个坑

坑 说明
维度不能整除 hidden、头数必须能被 TP 度整除 ⭐
随机种子 Dropout 在 TP 组内必须用不同的种子(否则等于没 dropout);但数据并行组内要相同
词表并行的边界 Embedding 按词表切时,要 mask 掉不属于自己的 token id
TP 度太大 通信占比上升,TP=8 通常是甜点,TP=16 收益已很小

💥 随机种子那条特别隐蔽:如果 TP 组内所有卡的 dropout mask 相同, 那么被切开的那些神经元会被同步地丢弃,等效的 dropout 率完全不对 —— 训练能跑,loss 也在降,但正则效果没了。


🔗 和站内其他章的关系

相关的地方 这里的位置
第 5 章 NVLink vs PCIe 差 15 倍 TP 必须在机内的原因 ⭐
第 5 章 AllReduce=RS+AG 序列并行的原理
第 11 章 ZeRO-3 对比:ZeRO 能跨机,TP 不能
全景导论主线 2 · 模型怎样看懂一句话 MLP 和多头注意力 本章切的就是它们
第 14 章 TP 在 3D 并行里的位置

✅ 检查点

  1. 列并行和行并行的区别?哪个需要通信?
  2. Megatron 的 MLP 为什么用"先列后行"?反过来会怎样?
  3. Attention 怎么切?为什么它比 MLP 更自然?
  4. 一个 Transformer 层前向反向共几次 AllReduce?
  5. 算一下:80 层、hidden=8192、batch=8、seq=2048 时每步通信多少?在 NVLink 和跨机各要多久?
  6. 张量并行度为什么不该超过 8?
  7. 通信量和 TP 度的关系是什么?这有什么含义?
  8. 序列并行解决什么问题?它靠哪个恒等式?
  9. TP 组内的 dropout 种子为什么必须不同?
👀 答案
  1. 列并行按列切权重,各卡各算各的、不需要通信(因为 GeLU 是逐元素的);行并行按行切,需要 AllReduce 求和。
  2. 因为列并行的输出天然是按列切开的,而行并行需要的输入恰好就是按列切开的——严丝合缝,中间完全不用通信,整个 MLP 只需一次 AllReduce。反过来(先行后列)中间要多一次通信。
  3. 直接把头分给不同的卡(Q/K/V 投影列并行,输出投影行并行)。更自然是因为多头注意力本来就是多个独立的头。
  4. 4 次:前向 MLP 一次 + Attention 一次,反向同样 2 次。
  5. 每次 AllReduce = 8×2048×8192×2 = 268MB,每层 4 次 = 1.07GB,80 层 = 86GB/步。NVLink 900GB/s ≈ 0.1 秒;跨机 IB 25GB/s ≈ 3.4 秒(完全不可用)。
  6. 因为跨机带宽撑不住——TP 度超过单机 GPU 数就要跨机通信,而跨机比 NVLink 慢 35 倍。这是带宽算出来的硬约束,不是建议。
  7. 通信量和 TP 度无关(TP=2 和 TP=8 通信量一样),但 TP 越大每卡分摊的计算越少 → 通信占比上升,效率下降。所以 TP=8 通常是甜点。
  8. 解决LayerNorm 和 Dropout 没被切、每卡都要存完整激活的问题。按序列维度切,和 TP 交替使用。靠 AllReduce = ReduceScatter + AllGather——所以激活显存降 TP 倍而通信总量不变。
  9. 如果 TP 组内 dropout mask 相同,被切开的神经元会被同步丢弃,等效 dropout 率完全不对。隐蔽在于训练能跑、loss 也在降,但正则效果没了。

🛑 可以停在这里

⚡ 走神救援

张量并行把一层内部的矩阵切开分给多卡——⭐ 唯一能救「单层都装不下」的办法,代价是每层都要通信。

两种切法:一种切完各算各的、不用通信;另一种切完必须求和。⭐⭐ Megatron 的精妙之处是把两者按顺序组合:前一种的输出天然就是后一种需要的输入形状,严丝合缝——⭐ 于是整个前馈块只要一次通信,反过来排就要多一次。 注意力那边更自然:多个头本来就独立,直接分给不同卡。

⭐⭐ 通信代价直接推出一条铁律:把每层前向反向的同步量乘上层数,在机内高速互联上是可接受的,走 PCIe 就慢一个量级,跨机则完全不可用——⭐ 所以张量并行度不要超过单机的卡数。这不是经验,是带宽算出来的硬约束。

⭐ 另一条也是算出来的:通信量和并行度无关,而计算被分摊了——⭐ 所以并行度越大,通信占比越高,到某个点之后再加收益已经很小。

序列并行补的是它的盲区(归一化和随机丢弃没被切),⭐ 激活显存再降一档而通信总量不变——又是「一次同步等于两次半同步」那个等式,这已经是它第三次出场,是分布式里最有用的一个。

⚠️ 四个坑里最阴的一个:⭐⭐ 同一组内的随机种子必须不同——相同的话,被切开的神经元会被同步丢弃,训练照跑、loss 也降,但正则效果没了。

⭐ 最后一句:别自己实现,用现成的。

下一节 👉 13-流水线并行.md

打卡记录保存在你的浏览器里,首页能看到总进度