张量并行(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=[15​26​37​48​]2×4​W=​1010​0101​2101​1010​​4×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=[412​614​824​412​]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=[412​614​824​412​]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​=​1010​0101​​4×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​=​2101​1010​​4×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​=[15​26​37​48​]×​1010​0101​​=[412​614​]

验算:第一行 =[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​=[15​26​37​48​]×​2101​1010​​=[824​412​]

验算:第一行 =[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​]=[412​614​824​412​]

和我们之前单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​=[10​01​21​10​]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​=[10​01​01​10​]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​=[15​26​]2×2​

GPU 1 持有 XXX 的后2列:

X2=[3478]2×2 X_2 = \begin{bmatrix} 3 & 4 \\ 7 & 8 \end{bmatrix}_{2 \times 2} X2​=[37​48​]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​=[15​26​]×[10​01​21​10​]=[15​26​416​15​]

验算:第一行 =[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​=[37​48​]×[10​01​01​10​]=[37​48​48​37​]

验算:第一行 =[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+7​2+46+8​4+416+8​1+35+7​]=[412​614​824​412​]

和我们之前的结果完全一致! ✅

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⋅[W1​W2​​]=[X1​,X2​]⋅[W1​W2​​]=X1​W1​+X2​W2​

矩阵乘法按行切分权重,等价于把输入也切分后分别乘,最后 相加。这就是 分块矩阵乘法 的性质!

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=​1010​2101​0110​1011​​4×4​B=​1010​0110​1001​0101​​4×4​

输入:

X=[1111]1×4X = \begin{bmatrix} 1 & 1 & 1 & 1 \end{bmatrix}_{1 \times 4}X=[1​1​1​1​]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​=​1010​2101​​(GPU 0)A2​=​0110​1011​​(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​=[10​01​10​01​](GPU 0,前2行)B2​=[10​10​01​01​](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]×[10​01​10​01​]=[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]×[10​10​01​01​]=[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层:

组件通信操作通信量
MLP1次 All-Reduce2bsd2bsd2bsd(前向)
Self-Attention1次 All-Reduce2bsd2bsd2bsd(前向)
总计(前向+反向)4次 All-Reduce8bsd8bsd8bsd

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

查找过程:

  1. 每个GPU检查输入token是否在自己的范围内
  2. 如果在范围内,查找并返回嵌入向量;否则返回零向量
  3. All-Reduce求和,得到完整的嵌入结果

输出层的并行

输出层是一个 (d,V)(d, V)(d,V) 的大矩阵,同样可以按词汇表维度切分(实际上就是列并行),每个GPU计算部分logits,最后All-Gather拼接。


十一、张量并行的局限与最佳实践

11.1 局限性

  1. 通信带宽要求高:All-Reduce需要GPU之间高速通信,通常要求NVLink(600GB/s),而非PCIe(64GB/s)
  2. GPU数量限制:一般限制在单机内(4~8卡),跨节点通信延迟太大
  3. 切分粒度限制: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集群上的工程实现。理解了矩阵可以怎么"切"、结果怎么"合",就理解了张量并行的全部精髓。


参考文献:

  1. Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism (Shoeybi et al., 2019)
  2. 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辅助下完成。

更多推荐