跳到主要内容
推理优化

6.2 张量并行推理:切矩阵、AllReduce 与 NVLink

Megatron 式行/列并行切分、前向内部的两次 AllReduce、TP 通信量与带宽账、为什么 TP 通常不跨节点,以及 vLLM 中的 TP 配置与验证

张量并行AllReduceNVLinkMegatron列并行行并行

6.1 说 TP 是推理的第一选择:单请求延迟最优、显存线性分摊。这一节把它的原理讲透——一个矩阵乘法怎么切到多卡、切完怎么通信、通信花多少钱。TP 的实现沿袭 Megatron-LM 的经典方案,理解它之后,你就能自己推算”TP=8 的 70B 模型每层要通信多少数据”这类面试必考题。

📑 目录


1. 直觉:把一个大矩阵乘法拆成多个小矩阵乘法

Transformer 前向里最重的运算是线性层:

Y=XWY = X \cdot W

其中 XX 是激活(形状 [B,H][B, H]BB 为 batch 的 token 数,HH 为 hidden size),WW 是权重(形状 [H,H][H, H'],如 MLP 的升维层 H=4HH' = 4H)。TP 的思想一句话:WW 切成多块,每卡拿一块,各算各的部分积,最后拼起来

矩阵乘法有个漂亮的代数性质——列切和行切分别对应两种”拼法”

  • 按列切 W=[W1,W2]W = [W_1, W_2]Y=X[W1,W2]=[XW1,XW2]Y = X \cdot [W_1, W_2] = [X W_1, X W_2] —— 结果横向拼接(每卡算输出的不同列),不需要通信就能拼出完整 YY
  • 按行切 W=[W1W2]W = \begin{bmatrix} W_1 \\ W_2 \end{bmatrix}Y=X1W1+X2W2Y = X_1 W_1 + X_2 W_2(其中 X=[X1,X2]X = [X_1, X_2])—— 结果逐元素相加(每卡算同一输出的部分和),需要一次 AllReduce 归约

💡 提示:列切是”输出并行”,行切是”输入并行”——前者每卡持有完整输入、产出部分输出;后者每卡持有部分输入、产出部分输出再相加。理解这个区分,下面的通信分析才有抓手。


2. 列并行:切权重矩阵的列

2.1 前向计算

设 TP=2,权重 WRH×HW \in \mathbb{R}^{H \times H'} 按列切成 W1,W2RH×H/2W_1, W_2 \in \mathbb{R}^{H \times H'/2}

  • 每卡持有完整的输入 XX(激活需要广播/复制到所有卡)
  • 卡 0 计算 Y1=XW1Y_1 = X W_1,卡 1 计算 Y2=XW2Y_2 = X W_2
  • 输出 Y=[Y1,Y2]Y = [Y_1, Y_2] 直接拼接——不需要通信

2.2 典型位置

Transformer 里的 QKV 投影MLP 升维层(Gate/Up) 用列并行:

QKV: W ∈ R^H × 3H    → 切 3H 维
MLP Up: W ∈ R^H × 4H → 切 4H 维

这些层的共同点:输出维度比输入大,列切把”大的那一维”摊到各卡,每卡只需算小矩阵乘法。

📌 关键点:列并行”输出拼接不需要通信”是表象——真正的通信发生在下一层。因为每卡只持有 YY 的一部分列,而下一层(如 MLP 的降维层)需要完整的输入。所以 Transformer 里的列并行和行并行总是成对出现,通信在两者之间发生(见第 4 节)。


3. 行并行:切权重矩阵的行

3.1 前向计算

MLP 的降维层 WdownR4H×HW_{down} \in \mathbb{R}^{4H \times H} 按行切成 Wdown,1,Wdown,2R2H×HW_{down,1}, W_{down,2} \in \mathbb{R}^{2H \times H}

  • 每卡持有部分输入 X1,X2X_1, X_2(即上一层列并行的部分输出)
  • 卡 0 计算 Y1=X1Wdown,1Y_1 = X_1 W_{down,1},卡 1 计算 Y2=X2Wdown,2Y_2 = X_2 W_{down,2}
  • 输出 Y=Y1+Y2Y = Y_1 + Y_2 —— 需要一次 AllReduce

3.2 典型位置

  • MLP 降维层(Down):输入是升维层的部分输出,行切天然对齐
  • 输出投影层(Output)注意力输出投影(O 投影):同理行切

3.3 为什么 AllReduce 而不是别的

Y1+Y2Y_1 + Y_2 是”每卡都持有部分和、需要每卡都拿到完整和”——这正是 **AllReduce(全归约)**的语义:所有卡的结果归约求和后广播给所有卡。它的通信量:

AllReduce 通信量=2×(TP1)/TP×数据量2×数据量(TP 较大时)\text{AllReduce 通信量} = 2 \times (\text{TP}-1)/\text{TP} \times \text{数据量} \approx 2 \times \text{数据量} \quad (\text{TP 较大时})

系数 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,每次通信的数据量是 [B,H][B, H] 的激活(约 2×B×H2 \times B \times H 字节,FP16)。这是 TP 通信账的核心数字——层数 × 2 就是一次前向的总 AllReduce 次数。Self-Attention 内部的 QK^T、Softmax、AV 都在”部分列”上本地完成,这也是 TP 不需要额外通信的原因。

4.1 隐藏的代价:激活也要复制

列并行的输入 XX 需要完整复制到所有卡——XX 在每卡都有一份。这带来两个后果:

  • 激活显存 ×TP:每卡存完整激活,只有权重显存 ÷TP
  • Prefill 阶段(序列长)激活是显存大头,所以 TP 省权重的效果在长序列场景会被激活膨胀抵消一部分

💡 提示:这也是”序列并行(Sequence Parallelism)“的动机——把 LayerNorm/Dropout 的激活也按序列切分,省掉列并行输入复制的冗余。vLLM 的 PCP(Prefill Context Parallelism)和新版本里对长序列的处理与此相关,理解 TP 的激活复制问题后,再看那些高级特性就顺了。


5. 通信账:TP 到底花多少钱

用 70B 模型、TP=8 算一笔具体的账(hidden HH = 8192,batch = 1 个请求、序列 1 个 token 的 Decode 场景):

每层每次 AllReduce 的数据量

2×B×H×2 bytes=2×1×8192×2=32 KB2 \times B \times H \times 2\text{ bytes} = 2 \times 1 \times 8192 \times 2 = 32\text{ KB}

每层 2 次 AllReduce → 64KB;70B 模型约 80 层 → 一次 Decode 前向的 TP 通信总量约 5MB

通信耗时(NVLink 4.0,单对带宽约 100-450GB/s,这里按聚合后单次 AllReduce 等效带宽估算):

tcomm5 MB450 GB/s11μst_{\text{comm}} \approx \frac{5\text{ MB}}{450\text{ GB/s}} \approx 11\mu s

对比:Decode 单步本身约 1-10ms 量级(受带宽墙限制)。通信占 1% 左右——这就是 TP 在 NVLink 内”便宜”的定量证明。

如果跨节点(假设 400Gb/s IB ≈ 50GB/s)

tcomm5 MB50 GB/s100μst_{\text{comm}} \approx \frac{5\text{ MB}}{50\text{ GB/s}} \approx 100\mu 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 通常不跨节点

把上面的账总结成三条理由:

  1. 带宽鸿沟:NVLink(900GB/s 聚合级)vs 节点间网络(25-400Gb/s 即 3-50GB/s),差 1-2 个数量级。TP 每层都通信,对带宽极渴求
  2. 延迟串行化:TP 的 AllReduce 是前向路径上的同步点——每一层都要等通信完成才能继续。跨节点时每层多一次网络往返,80 层累加成肉眼可见的 TPOT 劣化
  3. 可靠性: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),数据量 2×B×H2 \times B \times H
  • 激活复制:列并行输入完整复制到每卡,权重 ÷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 的延迟收益为什么递减?

📚 参考资料