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

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

实际使用(别自己写)

# ⭐ 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。 这已经是它第三次出场了 —— 它是分布式训练里最有用的一个等式。

# 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 也在降,但正则效果没了

🛑 可以停在这里

走神救援

张量并行把一层内部的矩阵切开分给多卡——唯一能救"单层都装不下"的办法,代价是每层都要通信。两种切法列并行(竖着切权重,各算各的不用通信,因为 GeLU 是逐元素的)、行并行(横着切,必须 AllReduce 求和)。⭐Megatron 的精妙组合:MLP 先列并行再行并行——因为列并行的输出天然按列切开,而行并行需要的输入恰好就是按列切开的,严丝合缝 → 整个 MLP 只要一次 AllReduce(反过来要多一次)。Attention 更自然:多头本来就独立,直接把头分给不同卡。⭐每层前向反向共 4 次 AllReduce通信代价(80层/hidden8192/batch8/seq2048):每次 268MB × 4 × 80 = 86GB 每步 → NVLink 0.1 秒可接受、PCIe 1.3 秒、跨机 3.4 秒完全不可用 → ⭐铁律:张量并行度不要超过单机 GPU 数(通常 8),这是带宽算出来的硬约束。⭐通信量和 TP 度无关但计算被分摊TP 越大通信占比越高,TP=8 是甜点、TP=16 收益已很小序列并行:TP 的盲区是 LayerNorm/Dropout 没被切,按序列维度切开与 TP 交替,⭐激活显存再降 TP 倍而通信总量不变——又是 AllReduce = ReduceScatter + AllGather这个等式第三次出场,是分布式里最有用的一个)。⚠️四个坑:维度必须整除、⭐TP 组内 dropout 种子必须不同(相同则被切开的神经元同步丢弃,训练能跑 loss 也降但正则效果没了)、词表并行要 mask 掉不属于自己的 token、TP 度太大。别自己实现——用 Megatron-LM / PyTorch 原生 TP / vLLM 的 tensor_parallel_size

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

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