当 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_0、h_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]。
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_gate 和 W_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. 形状与切分的检查方法
每次分析张量并行切分,可以按三步检查:
W_up的列切分后,输出维度是否正好是每张 GPU 负责的中间特征数?W_down的行切分后,输入维度是否与本地中间特征数相等?- 每张 GPU 的部分输出是否都是
[N, d_model],从而可以逐元素 AllReduce?
切分边界与特征顺序也必须一致。在这些条件下,两种计算数学等价;实际浮点计算可能因求和顺序不同而有微小舍入差异。
参考资料
- Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism
- Megatron-LM 官方仓库