张量并行与序列并行
数据并行(06-数据并行:DP、DDP 与 FSDP)只切训练状态的冗余副本。但如果单层的前向计算本身就超出单卡(比如 FFN 的中间激活),它无能为力 —— 因为每张卡仍要算完整的一层。
张量并行(TP)切的是层内的矩阵运算,序列并行(SP)进一步解决 TP 没覆盖到的激活冗余。
TP 的定位与约束
把
TP 的带宽约束最严
TP 每一层的每个子模块都要通信(一个 80 层模型约 320 次 AllReduce,见 03-多卡互联与集群网络)。对延迟极其敏感 —— 走跨机 IB 会比机内 NVLink 慢 18 倍。
所以 TP 通常限制在 NVLink 互联的单机内(8 卡以内),这不是约定,是带宽算出来的。
两种切法
Column Parallel Linear(按列切)
| 项 | 状态 |
|---|---|
| 权重 | 按列切成 |
| 输入 | 每卡都是完整的 |
| 输出 | 各卡得到 |
| 通信 | 无(若各卡本就有完整输入) |
关键性质:激活函数可以独立施加 —— 因为激活是逐元素操作,不依赖其他维度的值。列切分后不需要任何通信就能继续算非线性。
这是 Column Parallel 被用在「第一个线性层」的原因。
Row Parallel Linear(按行切)
| 项 | 状态 |
|---|---|
| 权重 | 按行切成 |
| 输入 | 也需按列分片 —— 恰好与上一个 Column Parallel 的输出对接 |
| 计算 | |
| 通信 | AllReduce(求和) 得到完整 |
行切分正好接在列切分之后 —— 前一个的输出分片,就是后一个需要的输入分片。这个「接口对接」是 TP 切分方案的精髓:中间不需要任何数据重排。
切分恒等式
推导见 12-FFN 与激活函数 —— 那里用 FFN 演示了完整的「列切 → 激活 → 行切 → AllReduce」流程。
Column Parallel 与 Row Parallel 的「接口对接」(
第一层: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 的完整切分
| 子模块 | 第一个线性层 | 第二个线性层 |
|---|---|---|
| Attention | QKV 投影 按列切(多头天然可切,每卡负责部分 Head) | 输出投影 按行切 |
| FFN | 第一个线性层 按列切(切中间维度) | 第二个线性层 按行切 |
通信插入点:每个 Transformer Block 前向需要 2 次 AllReduce —— 一次在 Attention 的输出投影后,一次在 FFN 的第二个线性层后。
(2 次 AllReduce,每次传
反向对称:前向的 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 切分 —— 它们的输入/输出在每张卡上都是完整的
TP 只切了权重矩阵与 Attention/FFN 的计算,却把非 TP 区域的激活留成了完整的 —— 这一段激活显存没有跟着 TP 度一起降。
解法:沿序列维切
在非 TP 区域(LayerNorm、Dropout),沿序列维度切分激活值。
通信模式的改变:
| 位置 | 通信 |
|---|---|
| 离开 TP 区域(进入 LayerNorm) | AllReduce → ReduceScatter |
| 回到 TP 区域 | AllGather |
总通信量不变 —— 因为
AllReduce = ReduceScatter + AllGather(见 03-集合通信与 NCCL)。但激活显存大幅降低:非 TP 区域的激活从完整
降到 ,总激活显存接近线性缩放。
这是一次「同样通信量、换更多显存」的交换 —— 换个说法:原来那次 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 度不匹配」时的实际问题 —— 显存账本要按「参数切分不均匀」重新算,不能假设
。MQA(KV Head = 1)时更极端:所有 rank 的 KV 权重完全相同。
SP 与 CP 不是一回事
这一条容易混:
| 切什么 | 在哪个区域 | |
|---|---|---|
| SP(序列并行) | 非 Attention 区域的激活(LayerNorm / Dropout) | TP 没覆盖到的部分 |
| CP(上下文并行) | Attention 计算本身( | 序列太长时 |
SP 是 TP 的补丁,CP 是应对长序列的独立维度。 CP 的三种方案见 11-长序列训练与上下文并行。
相关
- 12-FFN 与激活函数 —— 列切/行切的完整推导与切分恒等式
- 11-Self-Attention 机制 —— 多头为什么天然可切
- 03-多卡互联与集群网络 —— 「TP 为什么必须机内」的 18 倍量化估算
- 03-集合通信与 NCCL —— AllReduce = ReduceScatter + AllGather
- 10-KV Cache 与推理优化 —— GQA / MQA 的 KV Head 布局
参考
- https://caomaolufei.github.io/AIInfraGuide/guides/模块三-分布式训练/第6章-张量并行与序列并行
- https://arxiv.org/abs/1909.08053
- https://arxiv.org/abs/2205.05198
- https://github.com/NVIDIA/Megatron-LM
YJ