FFN 张量并行:从列切分到行切分

为什么两张 GPU 分工后,结果仍然和完整矩阵乘法相同

Posted by Liu Mengxuan on September 21, 2026

当 FFN 的权重无法放进一张 GPU 时,可以沿中间维度把它切到多张 GPU 上。关键不是把矩阵随意切开,而是让升维矩阵生成的那部分中间特征,与降维矩阵对应的那部分行保持一致。

这篇文章用两张 GPU 和行向量记法说明:W_up 按列切,W_down 按行切,最后对部分输出做 AllReduce 求和。

1. 先看完整的 FFN

忽略偏置,标准 FFN 是:

$h=\phi(xW_{up}),\qquad y=hW_{down}$

设一个 token 的形状为 [1, d_model],中间维度为 d_ff

1
2
3
4
5
x:[1, d_model]
       × W_up:[d_model, d_ff]
    h:[1, d_ff]
       × W_down:[d_ff, d_model]
    y:[1, d_model]

如果一次处理 N 个 token,只需把第一维换成 N

1
2
[N, d_model] × [d_model, d_ff] → [N, d_ff]
[N, d_ff]    × [d_ff, d_model] → [N, d_model]

2. W_up 为什么按列切?

假设两张 GPU 平分中间维度,把升维矩阵按列切成:

$W_{up}=[W_{up,0}\;W_{up,1}]$

1
2
GPU 0:W_up,0 = W_up[:, :d_ff//2]   形状 [d_model, d_ff/2]
GPU 1:W_up,1 = W_up[:, d_ff/2:]   形状 [d_model, d_ff/2]

矩阵的每一列负责生成一个中间特征,因此按列切,正好把中间特征分成两组:

$h_0=\phi(xW_{up,0}),\qquad h_1=\phi(xW_{up,1})$

两张 GPU 都需要一份完整的 x,但各自只生成一半中间特征。因为 φ 是逐元素操作,所以激活可以留在本地完成:

$h=[h_0\;h_1]$

不需要先把 h_0h_1 通信后再激活。

3. W_down 为什么按行切?

完整的降维矩阵形状是 [d_ff, d_model]。它的每一行对应中间向量的一个特征,所以要按照刚才的中间特征分界切成上下两块:

$W_{down}=\begin{bmatrix}W_{down,0}\\W_{down,1}\end{bmatrix}$

1
2
GPU 0:W_down,0 = W_down[:d_ff//2, :]   形状 [d_ff/2, d_model]
GPU 1:W_down,1 = W_down[d_ff//2:, :]   形状 [d_ff/2, d_model]

这里使用 NumPy/PyTorch 的二维切片:W_down[:d_ff//2, :] 中,逗号前取前一半行,逗号后取所有列。// 是整数除法;假设 d_ff 能被 2 整除,切片结束下标不包含在结果中。

于是每张 GPU 可以用自己的中间特征和对应的矩阵块计算部分贡献:

$y_0=h_0W_{down,0},\qquad y_1=h_1W_{down,1}$

两边的形状都是 [N, d_model]

两张 GPU 对 FFN 的 W_up 列切分与 W_down 行切分,最后 AllReduce 求和

4. 为什么相加后与完整结果相同?

把分块结果代回完整矩阵乘:

$\begin{aligned}y&=[h_0\;h_1]\begin{bmatrix}W_{down,0}\\W_{down,1}\end{bmatrix}\\&=h_0W_{down,0}+h_1W_{down,1}=y_0+y_1\end{aligned}$

它只是把原本的一次大求和拆成两组部分和。GPU 0 负责前半部分,GPU 1 负责后半部分,最后逐元素相加即可恢复完整输出:

1
2
3
4
GPU 0:y₀ = h₀ @ W_down,0
GPU 1:y₁ = h₁ @ W_down,1

AllReduce:y = y₀ + y₁

如果把 W_up 按列切,却把 W_down 按列切,中间特征和降维矩阵就无法这样一一对应,也不能用一次求和恢复结果。

5. SwiGLU 如何切分?

SwiGLU 有两条升维分支:

$h=\operatorname{SiLU}(xW_{gate})\odot(xW_{up})$

W_gateW_up 都按列切,W_down 按行切。GPU 0 和 GPU 1 各自完成本地的门控、内容投影和逐元素乘法,再对降维结果做一次 AllReduce:

1
2
3
4
5
6
7
8
9
10
11
12
13
GPU 0:
  gate₀ = SiLU(x @ W_gate,0)
  up₀   = x @ W_up,0
  h₀    = gate₀ * up₀
  y₀    = h₀ @ W_down,0

GPU 1:
  gate₁ = SiLU(x @ W_gate,1)
  up₁   = x @ W_up,1
  h₁    = gate₁ * up₁
  y₁    = h₁ @ W_down,1

最终:y = AllReduce(y₀, y₁)

* 是逐元素乘法,@ 是矩阵乘法。门控计算不会破坏这种切分,因为每个中间通道只需要同一 GPU 上对应的 gate 和 up 值。

6. 通信发生在哪里?

在这个 FFN 前向过程中,输入 x 已经在各 GPU 上有一份完整副本。各卡本地完成投影和激活,只有最后的部分输出需要求和,因此需要一次 AllReduce。

这并不意味着整个 Transformer 只有一次通信:Attention 子层通常也有自己的张量并行通信,反向传播还会有对应的梯度通信。这里的结论只针对这段 FFN 前向数据流。

7. 形状与切分的检查方法

每次分析张量并行切分,可以按三步检查:

  1. W_up 的列切分后,输出维度是否正好是每张 GPU 负责的中间特征数?
  2. W_down 的行切分后,输入维度是否与本地中间特征数相等?
  3. 每张 GPU 的部分输出是否都是 [N, d_model],从而可以逐元素 AllReduce?

切分边界与特征顺序也必须一致。在这些条件下,两种计算数学等价;实际浮点计算可能因求和顺序不同而有微小舍入差异。

参考资料