张量并行(Tensor Parallelism):从原理到实践的完全指南
张量并行(Tensor Parallelism):从原理到实践的完全指南
一、引言:为什么需要张量并行?
想象一下,你有一个拥有 1750亿参数 的GPT-3模型,单精度(FP32)下需要约 700GB 显存。而目前最强的单块GPU(如A100)也只有 80GB 显存。
这就好比你有一张 超大的桌子,但你家的门太小,根本搬不进去。怎么办?
把桌子拆开,分几块搬进去,再组装起来!
这就是 张量并行 的核心思想——把模型的参数矩阵(张量)切开,分散到多块GPU上,每块GPU只负责一部分计算,最后再把结果合起来。
二、基础回顾:线性层的本质
在深入张量并行之前,我们先回顾一下神经网络中最基本的操作——线性变换:
Y=XW+bY = XW + bY=XW+b
其中:
- XXX:输入矩阵,形状为 (batch_size,din)(batch\_size, d_{in})(batch_size,din)
- WWW:权重矩阵,形状为 (din,dout)(d_{in}, d_{out})(din,dout)
- bbb:偏置向量,形状为 (dout)(d_{out})(dout)
- YYY:输出矩阵,形状为 (batch_size,dout)(batch\_size, d_{out})(batch_size,dout)
一个具体的例子:
假设 batch_size=2batch\_size = 2batch_size=2,din=4d_{in} = 4din=4,dout=4d_{out} = 4dout=4:
X=[12345678]2×4W=[1021011010010110]4×4 X = \begin{bmatrix} 1 & 2 & 3 & 4 \\ 5 & 6 & 7 & 8 \end{bmatrix}_{2 \times 4} \quad W = \begin{bmatrix} 1 & 0 & 2 & 1 \\ 0 & 1 & 1 & 0 \\ 1 & 0 & 0 & 1 \\ 0 & 1 & 1 & 0 \end{bmatrix}_{4 \times 4} X=[15263748]2×4W=10100101210110104×4
Y=XW=[468412142412]2×4 Y = XW = \begin{bmatrix} 4 & 6 & 8 & 4 \\ 12 & 14 & 24 & 12 \end{bmatrix}_{2 \times 4} Y=XW=[412614824412]2×4
验算第一行:[1,2,3,4]×W=[1⋅1+2⋅0+3⋅1+4⋅0, 1⋅0+2⋅1+3⋅0+4⋅1, 1⋅2+2⋅1+3⋅0+4⋅1, 1⋅1+2⋅0+3⋅1+4⋅0]=[4,6,8,4][1,2,3,4] \times W = [1\cdot1+2\cdot0+3\cdot1+4\cdot0,\ 1\cdot0+2\cdot1+3\cdot0+4\cdot1,\ 1\cdot2+2\cdot1+3\cdot0+4\cdot1,\ 1\cdot1+2\cdot0+3\cdot1+4\cdot0] = [4, 6, 8, 4][1,2,3,4]×W=[1⋅1+2⋅0+3⋅1+4⋅0, 1⋅0+2⋅1+3⋅0+4⋅1, 1⋅2+2⋅1+3⋅0+4⋅1, 1⋅1+2⋅0+3⋅1+4⋅0]=[4,6,8,4]
让我重新计算一下,确保数值准确:
Y1,:=[1×1+2×0+3×1+4×0, 1×0+2×1+3×0+4×1, 1×2+2×1+3×0+4×1, 1×1+2×0+3×1+4×0] Y_{1,:} = [1 \times 1 + 2 \times 0 + 3 \times 1 + 4 \times 0,\ 1 \times 0 + 2 \times 1 + 3 \times 0 + 4 \times 1,\ 1 \times 2 + 2 \times 1 + 3 \times 0 + 4 \times 1,\ 1 \times 1 + 2 \times 0 + 3 \times 1 + 4 \times 0] Y1,:=[1×1+2×0+3×1+4×0, 1×0+2×1+3×0+4×1, 1×2+2×1+3×0+4×1, 1×1+2×0+3×1+4×0]
=[4,6,8,4]= [4, 6, 8, 4]=[4,6,8,4]
Y2,:=[5×1+6×0+7×1+8×0, 5×0+6×1+7×0+8×1, 5×2+6×1+7×0+8×1, 5×1+6×0+7×1+8×0] Y_{2,:} = [5 \times 1 + 6 \times 0 + 7 \times 1 + 8 \times 0,\ 5 \times 0 + 6 \times 1 + 7 \times 0 + 8 \times 1,\ 5 \times 2 + 6 \times 1 + 7 \times 0 + 8 \times 1,\ 5 \times 1 + 6 \times 0 + 7 \times 1 + 8 \times 0] Y2,:=[5×1+6×0+7×1+8×0, 5×0+6×1+7×0+8×1, 5×2+6×1+7×0+8×1, 5×1+6×0+7×1+8×0]
=[12,14,24,12]= [12, 14, 24, 12]=[12,14,24,12]
所以:
Y=XW=[468412142412]2×4 Y = XW = \begin{bmatrix} 4 & 6 & 8 & 4 \\ 12 & 14 & 24 & 12 \end{bmatrix}_{2 \times 4} Y=XW=[412614824412]2×4
记住这个结果,我们接下来要用两种并行方式分别得到同样的结果。
三、列并行(Column Parallelism)
3.1 核心思想
列并行 是把权重矩阵 WWW 沿着 列方向(输出维度) 切分。
┌──────────┐
│ W │ (d_in × d_out)
│ │
└──────────┘
↓ 按列切分
┌─────┐ ┌─────┐
│ W₁ │ │ W₂ │
│ │ │ │
└─────┘ └─────┘
GPU 0 GPU 1
每块GPU拿到的是 WWW 的若干列,形状为 (din,dout/p)(d_{in}, d_{out}/p)(din,dout/p),其中 ppp 是GPU数量。
3.2 具体数值例子
假设我们用 2块GPU 来并行,将 WWW 按列均分:
GPU 0 持有 WWW 的前2列:
W1=[10011001]4×2 W_1 = \begin{bmatrix} 1 & 0 \\ 0 & 1 \\ 1 & 0 \\ 0 & 1 \end{bmatrix}_{4 \times 2} W1=101001014×2
GPU 1 持有 WWW 的后2列:
W2=[21100110]4×2 W_2 = \begin{bmatrix} 2 & 1 \\ 1 & 0 \\ 0 & 1 \\ 1 & 0 \end{bmatrix}_{4 \times 2} W2=210110104×2
3.3 计算过程
关键:每块GPU都需要完整的输入 XXX。
步骤1:将完整的 XXX 广播(Broadcast)到所有GPU
┌─────────────────┐
│ X (完整输入) │
└────────┬────────┘
┌─────┴─────┐
GPU 0: X GPU 1: X
步骤2:各GPU独立计算部分输出
GPU 0 计算:
Y1=X⋅W1=[12345678]×[10011001]=[461214] Y_1 = X \cdot W_1 = \begin{bmatrix} 1 & 2 & 3 & 4 \\ 5 & 6 & 7 & 8 \end{bmatrix} \times \begin{bmatrix} 1 & 0 \\ 0 & 1 \\ 1 & 0 \\ 0 & 1 \end{bmatrix} = \begin{bmatrix} 4 & 6 \\ 12 & 14 \end{bmatrix} Y1=X⋅W1=[15263748]×10100101=[412614]
验算:第一行 =[1×1+2×0+3×1+4×0, 1×0+2×1+3×0+4×1]=[4,6]= [1 \times 1+2 \times 0+3 \times 1+4 \times 0,\ 1 \times 0+2 \times 1+3 \times 0+4 \times 1] = [4, 6]=[1×1+2×0+3×1+4×0, 1×0+2×1+3×0+4×1]=[4,6] ✓
GPU 1 计算:
Y2=X⋅W2=[12345678]×[21100110]=[842412] Y_2 = X \cdot W_2 = \begin{bmatrix} 1 & 2 & 3 & 4 \\ 5 & 6 & 7 & 8 \end{bmatrix} \times \begin{bmatrix} 2 & 1 \\ 1 & 0 \\ 0 & 1 \\ 1 & 0 \end{bmatrix} = \begin{bmatrix} 8 & 4 \\ 24 & 12 \end{bmatrix} Y2=X⋅W2=[15263748]×21011010=[824412]
验算:第一行 =[1×2+2×1+3×0+4×1, 1×1+2×0+3×1+4×0]=[8,4]= [1 \times 2+2 \times 1+3 \times 0+4 \times 1,\ 1 \times 1+2 \times 0+3 \times 1+4 \times 0] = [8, 4]=[1×2+2×1+3×0+4×1, 1×1+2×0+3×1+4×0]=[8,4] ✓
步骤3:拼接(All-Gather)得到完整输出
Y=[Y1,Y2]=[468412142412] Y = [Y_1, Y_2] = \begin{bmatrix} 4 & 6 & 8 & 4 \\ 12 & 14 & 24 & 12 \end{bmatrix} Y=[Y1,Y2]=[412614824412]
和我们之前单GPU计算的结果完全一致! ✅
3.4 图解总结
X (完整)
┌────┴────┐
↓ ↓
┌─────────┐ ┌─────────┐
GPU 0: │ Y₁=X·W₁ │ │ Y₂=X·W₂ │ :GPU 1
│ (2×2) │ │ (2×2) │
└────┬────┘ └────┬────┘
↓ ↓
All-Gather (拼接列)
↓
Y = [Y₁, Y₂]
(2×4) 完整输出
3.5 通信模式
| 阶段 | 操作 | 说明 |
|---|---|---|
| 前向-输入 | Broadcast / Identity | 每个GPU需要完整的 XXX |
| 前向-输出 | All-Gather | 将部分输出拼接为完整输出 |
四、行并行(Row Parallelism)
4.1 核心思想
行并行 是把权重矩阵 WWW 沿着 行方向(输入维度) 切分。
┌──────────┐
│ W │ (d_in × d_out)
└──────────┘
↓ 按行切分
┌──────────┐
│ W₁ │ GPU 0 (前半行)
├──────────┤
│ W₂ │ GPU 1 (后半行)
└──────────┘
每块GPU拿到 WWW 的若干行,形状为 (din/p,dout)(d_{in}/p, d_{out})(din/p,dout)。
但这里有个关键区别:输入 XXX 也必须被切分!
4.2 具体数值例子
同样使用 2块GPU,将 WWW 按行均分:
GPU 0 持有 WWW 的前2行:
W1=[10210110]2×4 W_1 = \begin{bmatrix} 1 & 0 & 2 & 1 \\ 0 & 1 & 1 & 0 \end{bmatrix}_{2 \times 4} W1=[10012110]2×4
GPU 1 持有 WWW 的后2行:
W2=[10010110]2×4 W_2 = \begin{bmatrix} 1 & 0 & 0 & 1 \\ 0 & 1 & 1 & 0 \end{bmatrix}_{2 \times 4} W2=[10010110]2×4
相应地,输入 XXX 也按列切分:
GPU 0 持有 XXX 的前2列:
X1=[1256]2×2 X_1 = \begin{bmatrix} 1 & 2 \\ 5 & 6 \end{bmatrix}_{2 \times 2} X1=[1526]2×2
GPU 1 持有 XXX 的后2列:
X2=[3478]2×2 X_2 = \begin{bmatrix} 3 & 4 \\ 7 & 8 \end{bmatrix}_{2 \times 2} X2=[3748]2×2
4.3 计算过程
步骤1:各GPU独立计算部分结果
GPU 0 计算:
Z1=X1⋅W1=[1256]×[10210110]=[124156165] Z_1 = X_1 \cdot W_1 = \begin{bmatrix} 1 & 2 \\ 5 & 6 \end{bmatrix} \times \begin{bmatrix} 1 & 0 & 2 & 1 \\ 0 & 1 & 1 & 0 \end{bmatrix} = \begin{bmatrix} 1 & 2 & 4 & 1 \\ 5 & 6 & 16 & 5 \end{bmatrix} Z1=X1⋅W1=[1526]×[10012110]=[152641615]
验算:第一行 =[1×1+2×0, 1×0+2×1, 1×2+2×1, 1×1+2×0]=[1,2,4,1]= [1 \times 1+2 \times 0,\ 1 \times 0+2 \times 1,\ 1 \times 2+2 \times 1,\ 1 \times 1+2 \times 0] = [1, 2, 4, 1]=[1×1+2×0, 1×0+2×1, 1×2+2×1, 1×1+2×0]=[1,2,4,1] ✓
GPU 1 计算:
Z2=X2⋅W2=[3478]×[10010110]=[34437887] Z_2 = X_2 \cdot W_2 = \begin{bmatrix} 3 & 4 \\ 7 & 8 \end{bmatrix} \times \begin{bmatrix} 1 & 0 & 0 & 1 \\ 0 & 1 & 1 & 0 \end{bmatrix} = \begin{bmatrix} 3 & 4 & 4 & 3 \\ 7 & 8 & 8 & 7 \end{bmatrix} Z2=X2⋅W2=[3748]×[10010110]=[37484837]
验算:第一行 =[3×1+4×0, 3×0+4×1, 3×0+4×1, 3×1+4×0]=[3,4,4,3]= [3 \times 1+4 \times 0,\ 3 \times 0+4 \times 1,\ 3 \times 0+4 \times 1,\ 3 \times 1+4 \times 0] = [3, 4, 4, 3]=[3×1+4×0, 3×0+4×1, 3×0+4×1, 3×1+4×0]=[3,4,4,3] ✓
步骤2:求和(All-Reduce)得到完整输出
Y=Z1+Z2=[1+32+44+41+35+76+816+85+7]=[468412142412] Y = Z_1 + Z_2 = \begin{bmatrix} 1+3 & 2+4 & 4+4 & 1+3 \\ 5+7 & 6+8 & 16+8 & 5+7 \end{bmatrix} = \begin{bmatrix} 4 & 6 & 8 & 4 \\ 12 & 14 & 24 & 12 \end{bmatrix} Y=Z1+Z2=[1+35+72+46+84+416+81+35+7]=[412614824412]
和我们之前的结果完全一致! ✅
4.4 为什么是"求和"而不是"拼接"?
这是理解行并行的关键。让我们从数学上看:
Y=XW=X⋅[W1W2]=[X1,X2]⋅[W1W2]=X1W1+X2W2 Y = XW = X \cdot \begin{bmatrix} W_1 \\ W_2 \end{bmatrix} = [X_1, X_2] \cdot \begin{bmatrix} W_1 \\ W_2 \end{bmatrix} = X_1 W_1 + X_2 W_2 Y=XW=X⋅[W1W2]=[X1,X2]⋅[W1W2]=X1W1+X2W2
矩阵乘法按行切分权重,等价于把输入也切分后分别乘,最后 相加。这就是 分块矩阵乘法 的性质!
4.5 图解总结
X₁ (前半列) X₂ (后半列)
↓ ↓
┌─────────┐ ┌─────────┐
GPU 0: │ Z₁=X₁·W₁│ │ Z₂=X₂·W₂│ :GPU 1
│ (2×4) │ │ (2×4) │
└────┬────┘ └────┬────┘
↓ ↓
All-Reduce (逐元素求和)
↓
Y = Z₁ + Z₂
(2×4) 完整输出
4.6 通信模式
| 阶段 | 操作 | 说明 |
|---|---|---|
| 前向-输入 | Scatter / Split | 将输入 XXX 按列切分到各GPU |
| 前向-输出 | All-Reduce | 将部分结果求和 |
五、列并行 vs 行并行 对比
| 维度 | 列并行 (Column) | 行并行 (Row) |
|---|---|---|
| 切分维度 | WWW 按列切 (din,dout/p)(d_{in}, d_{out}/p)(din,dout/p) | WWW 按行切 (din/p,dout)(d_{in}/p, d_{out})(din/p,dout) |
| 输入要求 | 每个GPU需要 完整 XXX | 每个GPU只需 部分 XXX |
| 输出形式 | 每个GPU得到 部分 YYY | 每个GPU得到 完整 YYY 的一个分量 |
| 合并操作 | All-Gather(拼接) | All-Reduce(求和) |
| 通信数据量 | 输出大小 × (p−1)/p(p-1)/p(p−1)/p | 输出大小 × 2(p−1)/p2(p-1)/p2(p−1)/p |
六、实际应用:Transformer中的张量并行
在实际的Transformer模型中(如Megatron-LM),列并行和行并行往往 配合使用,形成精妙的组合。
6.1 MLP层的并行
Transformer的MLP(前馈网络)通常由两个线性层组成:
h=GeLU(XA)⋅Bh = \text{GeLU}(XA) \cdot Bh=GeLU(XA)⋅B
Megatron-LM的做法是:
输入 X (完整)
│
↓ ──── Identity (不通信) ────
│ │
GPU 0 GPU 1
│ │
↓ ↓
X·A₁ (列并行) X·A₂ (列并行)
│ │
↓ ↓
GeLU GeLU
│ │
↓ ↓
·B₁ (行并行) ·B₂ (行并行)
│ │
↓ ──── All-Reduce ─────── ↓
│
Y (完整输出)
巧妙之处:
- 第一层用 列并行:输入需要完整 XXX(直接广播即可),输出是部分结果
- 第二层用 行并行:恰好可以接收上一层的部分输出作为输入!
- 整个MLP只需要一次All-Reduce通信,大大降低了通信开销
6.2 自注意力层的并行
多头注意力天然适合张量并行——不同的注意力头可以分配到不同的GPU上:
假设 8 个注意力头,2 块 GPU:
GPU 0: Head 0, 1, 2, 3 → 计算各自的 Q, K, V, Attention
GPU 1: Head 4, 5, 6, 7 → 计算各自的 Q, K, V, Attention
最后通过 All-Reduce 合并输出
具体来说:
- WQ,WK,WVW_Q, W_K, W_VWQ,WK,WV 使用 列并行(按注意力头切分)
- WOW_OWO(输出投影)使用 行并行
七、一个完整的数值走通示例
让我们用一个更贴近实际的例子,走通MLP的完整张量并行过程。
设置:
- 输入维度:dmodel=4d_{model} = 4dmodel=4
- 隐藏层维度:dff=4d_{ff} = 4dff=4
- 2块GPU
权重矩阵:
A=[1201011010110101]4×4B=[1010010111000011]4×4 A = \begin{bmatrix} 1 & 2 & 0 & 1 \\ 0 & 1 & 1 & 0 \\ 1 & 0 & 1 & 1 \\ 0 & 1 & 0 & 1 \end{bmatrix}_{4 \times 4} \quad B = \begin{bmatrix} 1 & 0 & 1 & 0 \\ 0 & 1 & 0 & 1 \\ 1 & 1 & 0 & 0 \\ 0 & 0 & 1 & 1 \end{bmatrix}_{4 \times 4} A=10102101011010114×4B=10100110100101014×4
输入:
X=[1111]1×4X = \begin{bmatrix} 1 & 1 & 1 & 1 \end{bmatrix}_{1 \times 4}X=[1111]1×4
步骤1:列并行切分 AAA
A1=[12011001](GPU 0)A2=[01101101](GPU 1) A_1 = \begin{bmatrix} 1 & 2 \\ 0 & 1 \\ 1 & 0 \\ 0 & 1 \end{bmatrix} \quad (\text{GPU 0}) \qquad A_2 = \begin{bmatrix} 0 & 1 \\ 1 & 0 \\ 1 & 1 \\ 0 & 1 \end{bmatrix} \quad (\text{GPU 1}) A1=10102101(GPU 0)A2=01101011(GPU 1)
步骤2:各GPU计算第一层
GPU 0: H1=X⋅A1=[1,1,1,1]×A1=[1+0+1+0,2+1+0+1]=[2,4]H_1 = X \cdot A_1 = [1,1,1,1] \times A_1 = [1+0+1+0, 2+1+0+1] = [2, 4]H1=X⋅A1=[1,1,1,1]×A1=[1+0+1+0,2+1+0+1]=[2,4]
GPU 1: H2=X⋅A2=[1,1,1,1]×A2=[0+1+1+0,1+0+1+1]=[2,3]H_2 = X \cdot A_2 = [1,1,1,1] \times A_2 = [0+1+1+0, 1+0+1+1] = [2, 3]H2=X⋅A2=[1,1,1,1]×A2=[0+1+1+0,1+0+1+1]=[2,3]
步骤3:应用激活函数(假设用ReLU简化)
GPU 0: H1′=ReLU([2,4])=[2,4]H_1' = \text{ReLU}([2, 4]) = [2, 4]H1′=ReLU([2,4])=[2,4]
GPU 1: H2′=ReLU([2,3])=[2,3]H_2' = \text{ReLU}([2, 3]) = [2, 3]H2′=ReLU([2,3])=[2,3]
步骤4:行并行切分 BBB
注意,列并行的输出 H1′H_1'H1′ 和 H2′H_2'H2′ 恰好对应 BBB 行并行的输入切分:
B1=[10100101](GPU 0,前2行)B2=[11000011](GPU 1,后2行) B_1 = \begin{bmatrix} 1 & 0 & 1 & 0 \\ 0 & 1 & 0 & 1 \end{bmatrix} \quad (\text{GPU 0,前2行}) \qquad B_2 = \begin{bmatrix} 1 & 1 & 0 & 0 \\ 0 & 0 & 1 & 1 \end{bmatrix} \quad (\text{GPU 1,后2行}) B1=[10011001](GPU 0,前2行)B2=[10100101](GPU 1,后2行)
步骤5:各GPU计算第二层
GPU 0:
Z1=H1′⋅B1=[2,4]×[10100101]=[2,4,2,4]Z_1 = H_1' \cdot B_1 = [2, 4] \times \begin{bmatrix} 1 & 0 & 1 & 0 \\ 0 & 1 & 0 & 1 \end{bmatrix} = [2, 4, 2, 4]Z1=H1′⋅B1=[2,4]×[10011001]=[2,4,2,4]
GPU 1:
Z2=H2′⋅B2=[2,3]×[11000011]=[2,2,3,3]Z_2 = H_2' \cdot B_2 = [2, 3] \times \begin{bmatrix} 1 & 1 & 0 & 0 \\ 0 & 0 & 1 & 1 \end{bmatrix} = [2, 2, 3, 3]Z2=H2′⋅B2=[2,3]×[10100101]=[2,2,3,3]
步骤6:All-Reduce求和
Y=Z1+Z2=[2+2,4+2,2+3,4+3]=[4,6,5,7]Y = Z_1 + Z_2 = [2+2, 4+2, 2+3, 4+3] = [4, 6, 5, 7]Y=Z1+Z2=[2+2,4+2,2+3,4+3]=[4,6,5,7]
验证:单GPU完整计算
H=X⋅A=[1,1,1,1]×A=[2,4,2,3]H = X \cdot A = [1,1,1,1] \times A = [2, 4, 2, 3]H=X⋅A=[1,1,1,1]×A=[2,4,2,3]
H′=ReLU(H)=[2,4,2,3]H' = \text{ReLU}(H) = [2, 4, 2, 3]H′=ReLU(H)=[2,4,2,3]
Y=H′⋅B=[2,4,2,3]×B=[2×1+4×0+2×1+3×0, 2×0+4×1+2×1+3×0, 2×1+4×0+2×0+3×1, 2×0+4×1+2×0+3×1]Y = H' \cdot B = [2,4,2,3] \times B = [2 \times 1+4 \times 0+2 \times 1+3 \times 0,\ 2 \times 0+4 \times 1+2 \times 1+3 \times 0,\ 2 \times 1+4 \times 0+2 \times 0+3 \times 1,\ 2 \times 0+4 \times 1+2 \times 0+3 \times 1]Y=H′⋅B=[2,4,2,3]×B=[2×1+4×0+2×1+3×0, 2×0+4×1+2×1+3×0, 2×1+4×0+2×0+3×1, 2×0+4×1+2×0+3×1]
=[4,6,5,7]= [4, 6, 5, 7]=[4,6,5,7]
完全一致! ✅ 整个过程只在最后做了 一次All-Reduce,中间无需通信。
八、通信开销分析
8.1 通信量计算
假设模型隐藏维度为 ddd,序列长度为 sss,批次大小为 bbb,GPU数量为 ppp。
对于一个Transformer层:
| 组件 | 通信操作 | 通信量 |
|---|---|---|
| MLP | 1次 All-Reduce | 2bsd2bsd2bsd(前向) |
| Self-Attention | 1次 All-Reduce | 2bsd2bsd2bsd(前向) |
| 总计(前向+反向) | 4次 All-Reduce | 8bsd8bsd8bsd |
8.2 计算与通信的权衡
计算量 ∝ d² (矩阵乘法)
通信量 ∝ d (All-Reduce)
当 d 足够大时,计算量远大于通信量
→ 通信可以被计算"掩盖"
→ 并行效率高
这就是为什么张量并行在 大模型 上效果更好。
九、代码示意
下面用PyTorch伪代码展示列并行线性层的实现:
import torch
import torch.distributed as dist
class ColumnParallelLinear(torch.nn.Module):
"""列并行线性层"""
def __init__(self, in_features, out_features, world_size, rank):
super().__init__()
self.rank = rank
self.world_size = world_size
# 每个GPU只存储 out_features / world_size 列
self.local_out_features = out_features // world_size
# 本地权重:完整输入维度 × 部分输出维度
self.weight = torch.nn.Parameter(
torch.randn(in_features, self.local_out_features)
)
self.bias = torch.nn.Parameter(
torch.randn(self.local_out_features)
)
def forward(self, x):
# x: (batch_size, in_features) — 完整输入
# 本地矩阵乘法:得到部分输出
local_output = torch.matmul(x, self.weight) + self.bias
# local_output: (batch_size, local_out_features)
# All-Gather: 收集所有GPU的部分输出
output_list = [torch.empty_like(local_output) for _ in range(self.world_size)]
dist.all_gather(output_list, local_output)
# 拼接得到完整输出
output = torch.cat(output_list, dim=-1)
# output: (batch_size, out_features)
return output
class RowParallelLinear(torch.nn.Module):
"""行并行线性层"""
def __init__(self, in_features, out_features, world_size, rank):
super().__init__()
self.rank = rank
self.world_size = world_size
# 每个GPU只存储 in_features / world_size 行
self.local_in_features = in_features // world_size
# 本地权重:部分输入维度 × 完整输出维度
self.weight = torch.nn.Parameter(
torch.randn(self.local_in_features, out_features)
)
self.bias = torch.nn.Parameter(
torch.randn(out_features)
) if rank == 0 else None # 偏置只需一份
def forward(self, x):
# x: (batch_size, local_in_features) — 部分输入
# 本地矩阵乘法
local_output = torch.matmul(x, self.weight)
# local_output: (batch_size, out_features)
# All-Reduce: 对所有GPU的结果求和
dist.all_reduce(local_output, op=dist.ReduceOp.SUM)
# 加偏置(只加一次)
if self.bias is not None:
local_output += self.bias
return local_output
MLP的组合使用:
class ParallelMLP(torch.nn.Module):
"""张量并行的MLP"""
def __init__(self, d_model, d_ff, world_size, rank):
super().__init__()
# 第一层:列并行(不需要在中间通信)
self.fc1 = ColumnParallelLinear(d_model, d_ff, world_size, rank)
# 第二层:行并行(最后做All-Reduce)
self.fc2 = RowParallelLinear(d_ff, d_model, world_size, rank)
self.activation = torch.nn.GELU()
def forward(self, x):
# x: (batch, seq_len, d_model)
h = self.fc1(x) # 列并行,此处不做All-Gather!
h = self.activation(h) # 各GPU独立激活
y = self.fc2(h) # 行并行 + All-Reduce
return y
注意:在实际的Megatron-LM实现中,列并行层输出后不做All-Gather,而是直接传给行并行层,从而将两次通信合并为一次。这就是上面代码中
fc1输出不做gather的原因。
十、词汇表并行(Vocabulary Parallelism)
大语言模型的词汇表通常很大(如GPT-3有50257个token),对应的嵌入层和输出层参数量巨大。
嵌入层的并行
Embedding 矩阵E∈RV×d \text{Embedding 矩阵} E \in \mathbb{R}^{V \times d} Embedding 矩阵E∈RV×d
按词汇表维度切分:
GPU 0: E[0:V/2, :] — 负责 token 0 ~ V/2-1
GPU 1: E[V/2:V, :] — 负责 token V/2 ~ V-1
查找过程:
- 每个GPU检查输入token是否在自己的范围内
- 如果在范围内,查找并返回嵌入向量;否则返回零向量
- All-Reduce求和,得到完整的嵌入结果
输出层的并行
输出层是一个 (d,V)(d, V)(d,V) 的大矩阵,同样可以按词汇表维度切分(实际上就是列并行),每个GPU计算部分logits,最后All-Gather拼接。
十一、张量并行的局限与最佳实践
11.1 局限性
- 通信带宽要求高:All-Reduce需要GPU之间高速通信,通常要求NVLink(600GB/s),而非PCIe(64GB/s)
- GPU数量限制:一般限制在单机内(4~8卡),跨节点通信延迟太大
- 切分粒度限制:dmodeld_{model}dmodel 和注意力头数必须能被GPU数整除
11.2 最佳实践
┌───────────────────────────────────────────┐
│ 混合并行策略 │
│ │
│ 节点内(NVLink):张量并行 (TP=4或8) │
│ 节点间(InfiniBand):流水线并行 (PP) │
│ 全局:数据并行 (DP) │
│ │
│ 总GPU数 = TP × PP × DP │
└───────────────────────────────────────────┘
例如: 用256块GPU训练一个大模型
- 每台机器8卡,TP=8(机器内张量并行)
- 4台机器串联,PP=4(机器间流水线并行)
- 8组并行副本,DP=8(数据并行)
- 总计:8×4×8=2568 \times 4 \times 8 = 2568×4×8=256 块GPU
十二、总结
| 要点 | 内容 |
|---|---|
| 核心思想 | 把大矩阵切小,分给多个GPU并行计算 |
| 列并行 | 按输出维度切分,需要完整输入,输出拼接(All-Gather) |
| 行并行 | 按输入维度切分,需要切分输入,输出求和(All-Reduce) |
| 最佳搭配 | 列并行 → 激活 → 行并行,只需一次通信 |
| 适用场景 | 单机多卡、NVLink互连、超大模型 |
张量并行的本质就是 分块矩阵乘法 在GPU集群上的工程实现。理解了矩阵可以怎么"切"、结果怎么"合",就理解了张量并行的全部精髓。
参考文献:
- Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism (Shoeybi et al., 2019)
- An Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM (Narayanan et al., 2021)
后记
2026年4月17日12点40分于上海,在opus 4.6辅助下完成。
更多推荐



所有评论(0)