全部术语

GLOSSARY ENTRY

张量并行

  • Tensor Parallelism
  • TP

将单个算子的权重和计算切分到多个设备,让每个设备只承担同一次前向传播的一部分。

先建立一个准确的心智模型

张量并行不是“多张 GPU 各跑一份模型”,而是多张 GPU 合作完成同一个算子。每个 tensor-parallel rank 只保存权重矩阵的一个分片、计算结果的一部分,并在算子边界通过集合通信重新建立下一步需要的数据布局。

假设线性层为:

Y=XW,XRM×H,WRH×KY = XW, \qquad X \in \mathbb{R}^{M \times H},\quad W \in \mathbb{R}^{H \times K}

WW 太大或单卡矩阵乘吞吐不够时,可以沿 KKHH 切分。真正决定实现是否高效的,并不是“能不能切”,而是切完后下一个算子需要什么 shape,以及何时必须通信

两种基本切法

列并行:切分输出特征

把权重按列切成 pp 份:

W=[W(0),W(1),,W(p1)]W = [W^{(0)}, W^{(1)}, \ldots, W^{(p-1)}]

每个 rank 都接收完整输入 XX,但只产生一部分输出特征:

Y(r)=XW(r),Y=[Y(0),Y(1),,Y(p1)]Y^{(r)} = XW^{(r)}, \qquad Y = [Y^{(0)}, Y^{(1)}, \ldots, Y^{(p-1)}]

如果下一层也接受按最后一维切分的输入,就不必立即 All-Gather。把通信推迟到真正需要完整张量的位置,是张量并行设计的关键技巧。

行并行:切分输入特征

当输入和权重沿收缩维做相同切分时,每个 rank 计算的是最终输出的一项部分和:

X=[X(0),,X(p1)],W=[W(0)W(p1)]X = [X^{(0)}, \ldots, X^{(p-1)}], \qquad W = \begin{bmatrix} W^{(0)} \\ \vdots \\ W^{(p-1)} \end{bmatrix} Z(r)=X(r)W(r),Y=r=0p1Z(r)Z^{(r)} = X^{(r)}W^{(r)}, \qquad Y = \sum_{r=0}^{p-1} Z^{(r)}

前向传播中的这次求和正是 All-Reduce 的工作。反向传播经过列并行层时,还会出现一次对输入梯度部分和的 All-Reduce。

一个完整的 MLP 切分过程

Transformer 的 MLP 很适合把两种切法配对:第一层 W1W_1 做列并行,第二层 W2W_2 做行并行。以下示例采用经典 Megatron 风格的数据布局,假设没有启用 sequence parallel,并同时展示前向和反向的通信位置。

INTERACTIVE EXPLAINER

两张 GPU 如何完成一次 MLP 前向与反向

前向在行并行输出处归约 Y;反向在列并行输入处归约 ∂X。中间分片布局相容,因此不需要额外 All-Gather。

单设备上的完整前馈网络GeLU参数、激活和计算都集中在一个设备replicated inputcolumn parallelrow parallelreplicated outputFORWARDALL-REDUCE · YALL-REDUCE · ∂X
STEP 01 / 06从完整 MLP 开始

单设备执行 X → W₁ → GeLU → W₂ → Y。两层权重和中间激活都由同一设备持有。

设 MLP 的中间维度为 4H4H。在两个 rank 上:

张量未切分 shape每个 rank 的 shape数据布局
输入 XX[M,H][M,H][M,H][M,H]每个 rank 一份完整副本
W1W_1[H,4H][H,4H][H,2H][H,2H]沿列切分
隐藏状态 HrH_r[M,4H][M,4H][M,2H][M,2H]沿最后一维切分
W2W_2[4H,H][4H,H][2H,H][2H,H]沿行切分
部分输出 ZrZ_r[M,H][M,H][M,H][M,H]尚未跨 rank 求和
最终输出 YY[M,H][M,H][M,H][M,H]All-Reduce 后每个 rank 相同

通信契约,而不只是前向路径

gather_output=False、没有 sequence parallel 的经典实现中:

并行线性层前向传播反向传播
列并行 W1W_1输出 HrH_r 保持分片,不通信各 rank 产生 partialXrpartial X_r 部分和,对 partialXpartial X 做 All-Reduce
行并行 W2W_2各 rank 产生 ZrZ_r 部分和,对 YY 做 All-Reduce产生分片的 partialHrpartial H_r,可直接交给 W1W_1 的反向
配对后的 MLP一次 All-Reduce一次 All-Reduce

所以,若只看推理前向,这个 MLP 对确实只有一次 collective;若看完整训练,则前向和反向各有一次。这里还没有计入同一 Transformer 层中注意力子层的张量并行通信,也没有计入数据并行的梯度同步。

可运行的 dense-vs-shard 正确性实验

import platform
import torch

torch.manual_seed(97)
DEVICE=torch.device("cuda" if torch.cuda.is_available() else "cpu")
DTYPE=torch.float32

def main():
    m,h,k,p=3,8,12,2
    x=torch.randn(m,h,device=DEVICE,dtype=DTYPE)
    w1=torch.randn(h,k,device=DEVICE,dtype=DTYPE)
    w2=torch.randn(k,h,device=DEVICE,dtype=DTYPE)
    dense_hidden=torch.nn.functional.gelu(x@w1)
    dense_output=dense_hidden@w2

    # Column parallel: W1 沿输出特征切分,X 在各 rank replicated。
    w1_shards=w1.chunk(p,dim=1)
    hidden_shards=[torch.nn.functional.gelu(x@shard) for shard in w1_shards]
    gathered_hidden=torch.cat(hidden_shards,dim=1)
    torch.testing.assert_close(gathered_hidden,dense_hidden,rtol=1e-5,atol=1e-6)

    # Row parallel: W2 沿收缩维切分,每个 rank 产生同 shape partial Y。
    w2_shards=w2.chunk(p,dim=0)
    partial_outputs=[local_h@local_w for local_h,local_w in zip(hidden_shards,w2_shards)]
    reduced_output=torch.stack(partial_outputs).sum(0)
    torch.testing.assert_close(reduced_output,dense_output,rtol=1e-5,atol=1e-6)

    print(f"python={platform.python_version()} torch={torch.__version__}")
    print(f"device={DEVICE} dtype={DTYPE} world_size={p}")
    print(f"X={tuple(x.shape)} W1={tuple(w1.shape)} W2={tuple(w2.shape)}")
    for rank in range(p):
        print(f"rank={rank} W1_shard={tuple(w1_shards[rank].shape)} "
              f"H_shard={tuple(hidden_shards[rank].shape)} "
              f"W2_shard={tuple(w2_shards[rank].shape)} "
              f"partial_Y={tuple(partial_outputs[rank].shape)}")
    print(f"column_concat_max_abs_diff={(gathered_hidden-dense_hidden).abs().max():.3e}")
    print(f"row_sum_max_abs_diff={(reduced_output-dense_output).abs().max():.3e}")

if __name__=="__main__":
    with torch.inference_mode(): main()

完整文件位于 examples/tensor_parallel_reference.py。实际运行得到列分片 concat 最大差 0,行分片 sum 最大差 9.537e-07。代码把两个 rank 的局部 tensor 放在同一进程/device 以隔离数学语义;生产中 torch.cat 对应在需要完整列并行输出时的 All-Gather,stack(...).sum(0) 对应跨 rank All-Reduce 或 Reduce-Scatter 语义。

本机 集合通信 Gloo 实验 单独验证跨进程 collective 的 shape 与所有权;这台机器没有多卡 NCCL 性能证据。真实框架还会用 autograd 通信算子自动插入前向/反向归约,并处理 bias、随机数状态、权重初始化、参数梯度、sequence parallel 与通信计算重叠。

成本到底转移到了哪里

若张量并行度为 pp,理想情况下每个 rank 只保存约 1/p1/p 的对应权重,并承担约 1/p1/p 的矩阵乘计算。然而它并不是免费的:

  • 输入或输出可能需要复制、All-Gather、Reduce-Scatter 或 All-Reduce。
  • 一次请求会同时占用整个 TP 通信组,调度粒度变粗。
  • 跨节点链路通常显著慢于节点内 NVLink/NVSwitch,过大的 TP degree 可能让通信吞掉计算收益。
  • 小矩阵或小 batch 无法让 GPU 吃满时,切得更碎反而降低效率。

用非常粗略的形式,可以把一层延迟写成:

TlayerTmatmulp+Tcollective+ToverheadT_{layer} \approx \frac{T_{matmul}}{p} + T_{collective} + T_{overhead}

只有被减少的矩阵乘时间大于新增的通信和调度开销,张量并行才会带来实际加速。

什么时候优先考虑张量并行

  • 单层权重或 KV/激活相关工作集无法舒适地放进单卡。
  • 模型层数不适合继续做流水线切分,或者流水线气泡过大。
  • 节点内有高带宽互联,集合通信成本可以接受。
  • 推理服务需要降低单请求延迟,而不仅是提高总吞吐。

与其他并行轴的边界

Sequence/Context/Pipeline/Expert Parallel 总览按 token activation、attention context、层深度与 expert 集合分别定义所有权。TP 切单个矩阵特征轴;这些并行模式即使与 TP 组合,也不能用同一条 All-Reduce 叙述替代。

参考资料