先建立一个准确的心智模型
张量并行不是“多张 GPU 各跑一份模型”,而是多张 GPU 合作完成同一个算子。每个 tensor-parallel rank 只保存权重矩阵的一个分片、计算结果的一部分,并在算子边界通过集合通信重新建立下一步需要的数据布局。
假设线性层为:
当 太大或单卡矩阵乘吞吐不够时,可以沿 或 切分。真正决定实现是否高效的,并不是“能不能切”,而是切完后下一个算子需要什么 shape,以及何时必须通信。
两种基本切法
列并行:切分输出特征
把权重按列切成 份:
每个 rank 都接收完整输入 ,但只产生一部分输出特征:
如果下一层也接受按最后一维切分的输入,就不必立即 All-Gather。把通信推迟到真正需要完整张量的位置,是张量并行设计的关键技巧。
行并行:切分输入特征
当输入和权重沿收缩维做相同切分时,每个 rank 计算的是最终输出的一项部分和:
前向传播中的这次求和正是 All-Reduce 的工作。反向传播经过列并行层时,还会出现一次对输入梯度部分和的 All-Reduce。
一个完整的 MLP 切分过程
Transformer 的 MLP 很适合把两种切法配对:第一层 做列并行,第二层 做行并行。以下示例采用经典 Megatron 风格的数据布局,假设没有启用 sequence parallel,并同时展示前向和反向的通信位置。
INTERACTIVE EXPLAINER
两张 GPU 如何完成一次 MLP 前向与反向
前向在行并行输出处归约 Y;反向在列并行输入处归约 ∂X。中间分片布局相容,因此不需要额外 All-Gather。
单设备执行 X → W₁ → GeLU → W₂ → Y。两层权重和中间激活都由同一设备持有。
设 MLP 的中间维度为 。在两个 rank 上:
| 张量 | 未切分 shape | 每个 rank 的 shape | 数据布局 |
|---|---|---|---|
| 输入 | 每个 rank 一份完整副本 | ||
| 沿列切分 | |||
| 隐藏状态 | 沿最后一维切分 | ||
| 沿行切分 | |||
| 部分输出 | 尚未跨 rank 求和 | ||
| 最终输出 | All-Reduce 后每个 rank 相同 |
通信契约,而不只是前向路径
在 gather_output=False、没有 sequence parallel 的经典实现中:
| 并行线性层 | 前向传播 | 反向传播 |
|---|---|---|
| 列并行 | 输出 保持分片,不通信 | 各 rank 产生 部分和,对 做 All-Reduce |
| 行并行 | 各 rank 产生 部分和,对 做 All-Reduce | 产生分片的 ,可直接交给 的反向 |
| 配对后的 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 与通信计算重叠。
成本到底转移到了哪里
若张量并行度为 ,理想情况下每个 rank 只保存约 的对应权重,并承担约 的矩阵乘计算。然而它并不是免费的:
- 输入或输出可能需要复制、All-Gather、Reduce-Scatter 或 All-Reduce。
- 一次请求会同时占用整个 TP 通信组,调度粒度变粗。
- 跨节点链路通常显著慢于节点内 NVLink/NVSwitch,过大的 TP degree 可能让通信吞掉计算收益。
- 小矩阵或小 batch 无法让 GPU 吃满时,切得更碎反而降低效率。
用非常粗略的形式,可以把一层延迟写成:
只有被减少的矩阵乘时间大于新增的通信和调度开销,张量并行才会带来实际加速。
什么时候优先考虑张量并行
- 单层权重或 KV/激活相关工作集无法舒适地放进单卡。
- 模型层数不适合继续做流水线切分,或者流水线气泡过大。
- 节点内有高带宽互联,集合通信成本可以接受。
- 推理服务需要降低单请求延迟,而不仅是提高总吞吐。
与其他并行轴的边界
Sequence/Context/Pipeline/Expert Parallel 总览按 token activation、attention context、层深度与 expert 集合分别定义所有权。TP 切单个矩阵特征轴;这些并行模式即使与 TP 组合,也不能用同一条 All-Reduce 叙述替代。