Skip to content

张量并行与序列并行 ​

标签
AI/infra/并行策略
AI/infra/显存管理
字数
1887 字
阅读时间
8 分钟

数据并行(06-数据并行:DP、DDP 与 FSDP)只切训练状态的冗余副本。但如果单层的前向计算本身就超出单卡(比如 FFN 的中间激活),它无能为力 —— 因为每张卡仍要算完整的一层。

张量并行(TP)切的是层内的矩阵运算,序列并行(SP)进一步解决 TP 没覆盖到的激活冗余。

TP 的定位与约束 ​

把 Y=XA 拆分到多卡并行计算,每卡只存 1/N 的权重和对应的部分激活。

TP 的带宽约束最严

TP 每一层的每个子模块都要通信(一个 80 层模型约 320 次 AllReduce,见 03-多卡互联与集群网络)。对延迟极其敏感 —— 走跨机 IB 会比机内 NVLink 慢 18 倍。

所以 TP 通常限制在 NVLink 互联的单机内(8 卡以内),这不是约定,是带宽算出来的。

两种切法 ​

Column Parallel Linear(按列切) ​

A=[A1,A2,…,AN],Yi=X⋅Ai
项状态
权重按列切成 N 份,每卡持 Ai
输入 X每卡都是完整的
输出各卡得到 Y 的不同列分片
通信无(若各卡本就有完整输入)

关键性质:激活函数可以独立施加 —— 因为激活是逐元素操作,不依赖其他维度的值。列切分后不需要任何通信就能继续算非线性。

这是 Column Parallel 被用在「第一个线性层」的原因。

Row Parallel Linear(按行切) ​

项状态
权重按行切成 N 份
输入 X也需按列分片 —— 恰好与上一个 Column Parallel 的输出对接
计算Yi=Xi⋅Ai,各卡算部分结果
通信AllReduce(求和) 得到完整 Y

行切分正好接在列切分之后 —— 前一个的输出分片,就是后一个需要的输入分片。这个「接口对接」是 TP 切分方案的精髓:中间不需要任何数据重排。

切分恒等式 ​

[h0h1][W0W1]=h0W0+h1W1

推导见 12-FFN 与激活函数 —— 那里用 FFN 演示了完整的「列切 → 激活 → 行切 → AllReduce」流程。

Column Parallel 与 Row Parallel 的「接口对接」(N=2 示意):

  第一层:Column Parallel(权重按列切)

    X(每卡都是完整的)          A = [ A₁ │ A₂ ]
         │                        │        │
         ├─▶ GPU0:Y₁ = X·A₁      │        │   ⇒ 各卡持有 Y 的不同列分片
         └─▶ GPU1:Y₂ = X·A₂      └────────┘
    激活函数可以就地施加(逐元素操作,不依赖别的分片)

  第二层:Row Parallel(权重按行切)

                 ┌ W₁ ┐
    Y = [Y₁│Y₂]  │    │ = Y₁·W₁ + Y₂·W₂
                 └ W₂ ┘
         │          │
      GPU0 用 Y₁ 算 Y₁·W₁ ┐
      GPU1 用 Y₂ 算 Y₂·W₂ ┘──▶ AllReduce 求和 ──▶ 完整的输出

  ⇒ 前一层的输出分片正好是后一层需要的输入分片,中间不需要任何数据重排

Transformer Block 的完整切分 ​

子模块第一个线性层第二个线性层
AttentionQKV 投影 按列切(多头天然可切,每卡负责部分 Head)输出投影 按行切
FFN第一个线性层 按列切(切中间维度)第二个线性层 按行切

通信插入点:每个 Transformer Block 前向需要 2 次 AllReduce —— 一次在 Attention 的输出投影后,一次在 FFN 的第二个线性层后。

每层前向通信量=2×2bsh

(2 次 AllReduce,每次传 2bsh;2 是 BF16 的字节数。)

反向对称:前向的 AllReduce 在反向对应同样一次 AllReduce,通信量对称。

这个「每层 2 次」的数字是后面所有并行配置估算的基础 —— 它决定了 TP 能开多大(见 13-3D 并行与混合并行策略)。

一个 Transformer Block 里:TP 切在哪、通信插在哪

  输入 X(完整)
     │
  QKV 投影      ── 按列切 ──▶ 各卡负责部分 Head
     │
  Attention(各卡独立算自己那部分头)
     │
  输出投影      ── 按行切 ──▶  ┌── AllReduce #1 ──┐   ← 通信插入点
     │                         └──────────────────┘
  LayerNorm / Dropout(非 TP 区域:SP 就在这里按序列维切)
     │
  FFN 第一层    ── 按列切(切中间维度)
     │
  FFN 第二层    ── 按行切 ──▶  ┌── AllReduce #2 ──┐
     │                         └──────────────────┘
  输出(完整)

  ⇒ 每个 Block 前向 2 次 AllReduce,每次传 2bsh;反向对称,通信量相同
  ⇒ 这个「每层 2 次」是所有并行配置估算的基础

序列并行(SP) ​

TP 留下的漏洞 ​

LayerNorm 和 Dropout 不参与 TP 切分 —— 它们的输入/输出在每张卡上都是完整的 bsh。这意味着:

TP 只切了权重矩阵与 Attention/FFN 的计算,却把非 TP 区域的激活留成了完整的 —— 这一段激活显存没有跟着 TP 度一起降。

解法:沿序列维切 ​

在非 TP 区域(LayerNorm、Dropout),沿序列维度切分激活值。

通信模式的改变:

位置通信
离开 TP 区域(进入 LayerNorm)AllReduce → ReduceScatter
回到 TP 区域AllGather

总通信量不变 —— 因为 AllReduce = ReduceScatter + AllGather(见 03-集合通信与 NCCL)。

但激活显存大幅降低:非 TP 区域的激活从完整 bsh 降到 bsh/N,总激活显存接近线性缩放。

这是一次「同样通信量、换更多显存」的交换 —— 换个说法:原来那次 AllReduce 的中间结果本来就要在卡间流动,现在把它拆成两半,中间那半恰好落在「按序列切分」的形态上。免费的显存收益。

GQA / MQA 下的一个坑 ​

GQA 的 KV Head 数可能小于 TP 度 —— 例如 KV Head = 4 而 TP = 8。

处理策略做法
复制 KV Head每个 TP rank 复制一份完整 KV Head
分组共享部分 rank 共享同一份 KV Head

影响:KV Head 复制导致 TP 的参数切分不完全均匀(KV 权重没有真正切成 1/8),但计算仍可并行 —— 因为 Q Head 是切开的,每个 rank 算自己那部分 Q 对全部 KV 的注意力。

这是「GQA 与 TP 度不匹配」时的实际问题 —— 显存账本要按「参数切分不均匀」重新算,不能假设 1N。MQA(KV Head = 1)时更极端:所有 rank 的 KV 权重完全相同。

SP 与 CP 不是一回事 ​

这一条容易混:

切什么在哪个区域
SP(序列并行)非 Attention 区域的激活(LayerNorm / Dropout)TP 没覆盖到的部分
CP(上下文并行)Attention 计算本身(QK⊤、PV)序列太长时

SP 是 TP 的补丁,CP 是应对长序列的独立维度。 CP 的三种方案见 11-长序列训练与上下文并行。

相关 ​

参考 ​

贡献者 ​

文件历史 ​