6.2 张量并行推理:切矩阵、AllReduce 与 NVLink
Megatron 式行/列并行切分、前向内部的两次 AllReduce、TP 通信量与带宽账、为什么 TP 通常不跨节点,以及 vLLM 中的 TP 配置与验证
6.1 说 TP 是推理的第一选择:单请求延迟最优、显存线性分摊。这一节把它的原理讲透——一个矩阵乘法怎么切到多卡、切完怎么通信、通信花多少钱。TP 的实现沿袭 Megatron-LM 的经典方案,理解它之后,你就能自己推算”TP=8 的 70B 模型每层要通信多少数据”这类面试必考题。
📑 目录
- 1. 直觉:把一个大矩阵乘法拆成多个小矩阵乘法
- 2. 列并行:切权重矩阵的列
- 3. 行并行:切权重矩阵的行
- 4. 一次前向里的两次 AllReduce
- 5. 通信账:TP 到底花多少钱
- 6. 为什么 TP 通常不跨节点
- 7. vLLM 中的 TP:配置、显存分布与验证
- 📝 总结
- 🎯 自我检验清单
- 📚 参考资料
1. 直觉:把一个大矩阵乘法拆成多个小矩阵乘法
Transformer 前向里最重的运算是线性层:
其中 是激活(形状 , 为 batch 的 token 数, 为 hidden size), 是权重(形状 ,如 MLP 的升维层 )。TP 的思想一句话:把 切成多块,每卡拿一块,各算各的部分积,最后拼起来。
矩阵乘法有个漂亮的代数性质——列切和行切分别对应两种”拼法”:
- 按列切 : —— 结果横向拼接(每卡算输出的不同列),不需要通信就能拼出完整
- 按行切 :(其中 )—— 结果逐元素相加(每卡算同一输出的部分和),需要一次 AllReduce 归约
💡 提示:列切是”输出并行”,行切是”输入并行”——前者每卡持有完整输入、产出部分输出;后者每卡持有部分输入、产出部分输出再相加。理解这个区分,下面的通信分析才有抓手。
2. 列并行:切权重矩阵的列
2.1 前向计算
设 TP=2,权重 按列切成 :
- 每卡持有完整的输入 (激活需要广播/复制到所有卡)
- 卡 0 计算 ,卡 1 计算
- 输出 直接拼接——不需要通信
2.2 典型位置
Transformer 里的 QKV 投影和 MLP 升维层(Gate/Up) 用列并行:
QKV: W ∈ R^H × 3H → 切 3H 维
MLP Up: W ∈ R^H × 4H → 切 4H 维
这些层的共同点:输出维度比输入大,列切把”大的那一维”摊到各卡,每卡只需算小矩阵乘法。
📌 关键点:列并行”输出拼接不需要通信”是表象——真正的通信发生在下一层。因为每卡只持有 的一部分列,而下一层(如 MLP 的降维层)需要完整的输入。所以 Transformer 里的列并行和行并行总是成对出现,通信在两者之间发生(见第 4 节)。
3. 行并行:切权重矩阵的行
3.1 前向计算
MLP 的降维层 按行切成 :
- 每卡持有部分输入 (即上一层列并行的部分输出)
- 卡 0 计算 ,卡 1 计算
- 输出 —— 需要一次 AllReduce
3.2 典型位置
- MLP 降维层(Down):输入是升维层的部分输出,行切天然对齐
- 输出投影层(Output)、注意力输出投影(O 投影):同理行切
3.3 为什么 AllReduce 而不是别的
是”每卡都持有部分和、需要每卡都拿到完整和”——这正是 **AllReduce(全归约)**的语义:所有卡的结果归约求和后广播给所有卡。它的通信量:
系数 2 来自”归约(收)和广播(发)“两个方向。每层 Transformer Block 恰好有 2 次 AllReduce(自注意力的 O 投影 1 次 + MLP 的 Down 投影 1 次)。
4. 一次前向里的两次 AllReduce
把 QKV/Attention-O/MLP-Up/Down 串起来看一张 TP 前向的通信图:
输入 X(完整复制到每卡)
│
▼
QKV 投影(列并行)──→ 每卡持 Q/K/V 的部分列
│ (Attention 的 q·k^T 在部分列上做,无需通信)
▼
Attention 输出投影 O(行并行)──→ AllReduce #1(每卡拼回完整 Y_attn)
│
▼
MLP Gate/Up(列并行)──→ 每卡持部分升维结果
│
▼
MLP Down(行并行)──→ AllReduce #2(每卡拼回完整 Y_mlp)
│
▼
残差 + LayerNorm(每卡独立,无通信)
│
▼
下一个 Block(重复)
📌 关键点:每层 Transformer Block 恰好 2 次 AllReduce,每次通信的数据量是 的激活(约 字节,FP16)。这是 TP 通信账的核心数字——层数 × 2 就是一次前向的总 AllReduce 次数。Self-Attention 内部的 QK^T、Softmax、AV 都在”部分列”上本地完成,这也是 TP 不需要额外通信的原因。
4.1 隐藏的代价:激活也要复制
列并行的输入 需要完整复制到所有卡—— 在每卡都有一份。这带来两个后果:
- 激活显存 ×TP:每卡存完整激活,只有权重显存 ÷TP
- Prefill 阶段(序列长)激活是显存大头,所以 TP 省权重的效果在长序列场景会被激活膨胀抵消一部分
💡 提示:这也是”序列并行(Sequence Parallelism)“的动机——把 LayerNorm/Dropout 的激活也按序列切分,省掉列并行输入复制的冗余。vLLM 的 PCP(Prefill Context Parallelism)和新版本里对长序列的处理与此相关,理解 TP 的激活复制问题后,再看那些高级特性就顺了。
5. 通信账:TP 到底花多少钱
用 70B 模型、TP=8 算一笔具体的账(hidden = 8192,batch = 1 个请求、序列 1 个 token 的 Decode 场景):
每层每次 AllReduce 的数据量:
每层 2 次 AllReduce → 64KB;70B 模型约 80 层 → 一次 Decode 前向的 TP 通信总量约 5MB。
通信耗时(NVLink 4.0,单对带宽约 100-450GB/s,这里按聚合后单次 AllReduce 等效带宽估算):
对比:Decode 单步本身约 1-10ms 量级(受带宽墙限制)。通信占 1% 左右——这就是 TP 在 NVLink 内”便宜”的定量证明。
如果跨节点(假设 400Gb/s IB ≈ 50GB/s):
加上网络延迟(微秒级 × 多次往返),通信占比从 1% 涨到 10%+,而且 每层都卡在网络往返上(80 层 × 2 次往返的延迟累加,乐观估计每层加 5-10μs 延迟 → 每 token 多 0.4-0.8ms,直接把 TPOT 拖垮)。
📌 关键点:TP 跨节点的两个杀手:带宽降 10 倍(通信时间 ×10)和延迟累加(每层都要等网络往返,80 层 × 2 次 = 160 次串行网络等待)。所以工程上的铁律是——TP 必须锁在 NVLink 节点内。跨节点的事交给 PP(6.3),它每段边界才通信一次。
6. 为什么 TP 通常不跨节点
把上面的账总结成三条理由:
- 带宽鸿沟:NVLink(900GB/s 聚合级)vs 节点间网络(25-400Gb/s 即 3-50GB/s),差 1-2 个数量级。TP 每层都通信,对带宽极渴求
- 延迟串行化:TP 的 AllReduce 是前向路径上的同步点——每一层都要等通信完成才能继续。跨节点时每层多一次网络往返,80 层累加成肉眼可见的 TPOT 劣化
- 可靠性:TP 一个 rank 掉线,整个副本不可用(通信组断裂);跨节点时网络故障概率更高
💡 提示:NVLink 也有代际差异——NVLink 3.0(A100)600GB/s、NVLink 4.0(H100)900GB/s、NVLink 5.0(B200)1.8TB/s 级。TP 支持的最大规模跟随 NVLink 带宽演进:A100 时代 TP=8 常见,H100 时代 TP=8 依然合理,而 Blackwell 的 NVLink 域进一步扩大(NVL72 机柜级互联),TP 的天花板还在长。
7. vLLM 中的 TP:配置、显存分布与验证
7.1 配置
# 单节点 4 卡 TP
vllm serve meta-llama/Llama-3.1-70B-Instruct \
--tensor-parallel-size 4 \
--gpu-memory-utilization 0.9
等价 Python:LLM(model=..., tensor_parallel_size=4)。
7.2 显存分布验证
TP=4 启动后,nvidia-smi 观察每卡显存:权重显存约为单卡的 1/4。70B FP16 权重 140GB ÷ 4 = 35GB/卡,加上每卡约 2GB 的激活与 KV Cache 预留——如果看到某卡显存明显高于其他卡,通常是激活复制不均或调度不均(请求都落在一张卡上)。
7.3 启动日志验证
启动日志里会打印世界大小与 KV cache 配置:
INFO 07-23 13:56:04 [kv_cache_utils.py:775] GPU KV cache size: 643,232 tokens
INFO 07-23 13:56:04 [kv_cache_utils.py:779] Maximum concurrency for 40,960 tokens per request: 15.70x
GPU KV cache size:全部 GPU 累计可存 Token 数(TP 下各卡 KV Cache 加起来)Maximum concurrency:按max_model_len折算的并发上限估计
7.4 TP 与投机解码的叠加注意
第 5 章讲过投机解码——TP 下草稿模型用 draft_tensor_parallel_size 独立配置(只能 1 或与 Target 相同)。TP 增大后,验证前向的通信也增大(每层 2 次 AllReduce),投机解码的验证成本跟着涨,接受率收益被通信摊薄——高 TP 下投机收益需要重新实测(第 5.5 节的流程)。
📌 关键点:TP 是”显存线性降、延迟亚线性降”的策略——显存确确实实 ÷TP,但延迟受通信拖累只能接近 ÷TP 而非等于。实测中 TP 从 1 到 4 通常延迟降 3x 左右,从 4 到 8 只能再降 1.5-1.8x(通信占比上升)。“加卡不加倍”是 TP 的常态,算 ROI 时别按线性期望。
📝 总结
- TP 的本质:把矩阵乘法切到多卡——列切(输出拼接,无需通信)配行切(部分和相加,需 AllReduce)
- 通信拓扑:每层 Transformer Block 恰好 2 次 AllReduce(O 投影 + MLP Down),数据量
- 激活复制:列并行输入完整复制到每卡,权重 ÷TP、激活 ×TP
- 通信账:70B/TP8 单步通信约 5MB——NVLink 内占 ~1%,跨节点涨到 10%+ 且每层串行等待
- 铁律:TP 锁在 NVLink 节点内,跨节点交给 PP
- vLLM:
--tensor-parallel-size,验证看 nvidia-smi 显存分布与启动日志 KV cache 报告
🎯 自我检验清单
- 列并行和行并行分别对应什么”拼法”?为什么输出拼接不需要通信?
- 一次 Transformer 前向里哪两处做 AllReduce?数据量公式是什么?
- 为什么 TP 下激活显存反而 ×TP?什么场景这个代价最明显?
- 能否口算 70B/TP8 的单步 TP 通信总量?
- 用带宽和延迟两个角度解释”TP 不跨节点”
- TP 从 1→4 和 4→8 的延迟收益为什么递减?
📚 参考资料
- Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism(https://arxiv.org/abs/1909.08053)
- Narayanan et al., Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM(https://arxiv.org/abs/2104.04473)
- vLLM 官方文档:Parallelism and Scaling(https://docs.vllm.ai/en/latest/serving/parallelism_scaling/)
- NVIDIA:NVLink 技术文档(https://www.nvidia.com/en-us/data-center/nvlink/)