跳转至

Tensor Parallel

张量并行 (Tensor Parallel, TP) 把单层矩阵计算切到多张 GPU,使单层权重和计算跨卡执行。它解决「单层权重或激活超过单卡显存」的问题,通信频率在各类并行中最高。

张量并行的由来是:当模型单层权重都放不进一张卡时,数据并行与流水线并行都无能为力,只能把「一个矩阵乘法本身」切开。这一做法由 Megatron-LM 系统化提出,把 Transformer 的 MLP 与 self-attention 都映射到列切分、行切分两种线性层上,并给出了对应的通信量分析,见 原论文 。

快速开始

把通信密集的 TP group 放在 NVLink 或同机高速互联内,先验证单层输出与单卡一致。建议步骤:

  1. 确定 TP 度(如 2、4、8),把 TP group 的 rank 严格映射到同一台机器或 NVLink 域内。
  2. 对注意力层和 MLP 层分别做列切分与行切分,确保前后层之间的切分方式匹配。
  3. 先用单卡跑一个 batch 记录输出,再在多卡 TP 下跑同一 batch,比较输出是否一致(误差在数值精度范围内)。
  4. 记录每层通信(all-reduce 或 all-gather)的耗时,确认其占总时间比例。

验证成功:TP 输出与单卡一致,单层可正常加载,通信耗时在预期范围内。

机制

张量并行的基本做法是:对线性层权重矩阵按列切分(各卡各算一部分输出)或按行切分(各卡处理输入的一部分再汇总)。以 MLP 的两次矩阵乘为例,第一次按列切分后各卡得到部分激活,经非线性后再按行切分第二次矩阵乘,各卡算出部分结果,最后通过 all-reduce 求和得到完整输出;注意力层中的多组注意力头也可分别放到不同卡上计算。

记输入为 \(X\),权重为 \(W\),输出为 \(Y=XW\)。列并行把 \(W\) 沿列切为 \([W_1, W_2]\),各卡计算 \(Y_i = X W_i\),得到的 \([Y_1, Y_2]\) 按列拼接即完整输出,中间无需通信;行并行则把 \(W\) 沿行切为 \([W_1; W_2]\),各卡计算 \(Y_i = X_i W_i\),最终输出是各卡部分和的累加:

\[ Y = Y_1 + Y_2 = X_1 W_1 + X_2 W_2 \]

这个求和通过一次 all-reduce 完成。Megatron-LM 的 MLP 与 self-attention 都用「列并行 -> 非线性/注意力 -> 行并行」的组合:列并行的输出天然分散在各卡,被逐元素运算消耗后,行并行只做部分和,再由 all-reduce 汇总,从而避免在设备间搬运完整激活。

层与层之间的切分方式需要匹配,否则需要在中间插入额外的通信来对齐。由于每个 token 的前向都要经过这些层,张量并行几乎每一步都会产生通信,通信量与 token 数和层数成正比,因此对互联带宽和延迟极其敏感。

通信量上,一个 Transformer 层在 MLP 的「行并行 g」与 self-attention 的「输出投影」各做一次 all-reduce,前向共 2 次、反向对称 2 次,合计 4 次;每次 all-reduce 的规模为 \(b \times s \times h\)(batch × 序列长度 × 隐藏维度),因此每层通信量正比于 \(b s h\)。关键点是这个量与 TP 度基本无关:切分后每个 rank 只传输 \(1/t\) 的数据,但参与通信的 rank 数也是 \(t\),两者相抵。这意味着 TP 不能靠加卡降低通信总量,只能换来「每卡显存与计算量下降」,这正是它必须留在高速互联域内的原因。

它解决的是「单层放不下」的问题:权重和激活被切分后,单卡只需保存和计算一部分。但当 TP 度超过单机规模、通信跨越较慢网络时,通信开销会迅速吞噬并行收益,此时应缩小 TP 域并配合其他并行维度。

案例

以 70B 模型单层超过单卡显存的场景为例,使用 8 卡 TP 让每卡只保存约 1/8 的层权重。先确认单层能正常加载且输出正确,再观察每步耗时:若跨节点的 all-reduce 延迟主导了总时间,说明 TP 域跨过了高速互联边界,应缩小 TP 域(如降到 4 卡),把剩余的并行需求交给流水线并行或数据并行。

最终目标是让 TP 通信留在低延迟域内,把跨节点的开销转移给通信频率更低的维度。

相关主题