并行度汇总

187次阅读
没有评论

Megatron 分布式训练全维度并行技术文档

数值假设(贯穿全文计算示例):

符号 含义 值
BB 每 GPU batch size 2
SS 序列长度 8192
HH hidden_size 8192
LL Transformer 层数 80
NhN_h 注意力头数 64
dhd_h 每头维度 128
EE 专家数 64
kk top-k 2
dtype 数据精度 bf16 (2 bytes)
架构 Llama-style Pre-RMSNorm + SwiGLU + RoPE 现代 Transformer
TP Tensor Parallel 4
SP Sequence Parallel 4 (与 TP 共用通信组)
PP Pipeline Parallel 4
CP Context Parallel 2
EP Expert Parallel 8
DP Data Parallel 4
总 GPU 数 4×4×2×8×4=10244 \times 4 \times 2 \times 8 \times 4 = 1024

目录


1. Tensor Parallelism + Sequence Parallelism (TP+SP)

1.1 总体原理

Tensor Parallelism (TP) 将单个 Transformer 层内的权重矩阵按列或按行切分到多张卡上,每张卡只需持有部分权重并计算部分结果,最后通过通信合并输出。

Sequence Parallelism (SP) 是 TP 的扩展:对 非 TP 参与计算的区域(RMSNorm、残差连接的激活值)按序列维度切分,减少每张卡上的激活值显存占用。SP 与 TP 共用同一个通信组——在 TP 通信组内,激活值按序列切分;在进入 TP 计算区域前通过 All-Gather 恢复完整序列。

1.2 数学推导

1.2.1 标准 TP 数学模型

考虑一个线性层 Y=XWY = XW,其中 X∈ℝB×S×HX \in \mathbb{R}^{B \times S \times H},W∈ℝH×H′W \in \mathbb{R}^{H \times H’}。

列切分 (Column Parallelism):

将 WW 按列切分为 NN 份:W=[W1,W2,…,WN]W = [W_1, W_2, \dots, W_N],其中 Wi∈ℝH×H′/NW_i \in \mathbb{R}^{H \times H’/N}。

Y=XW=X[W1,W2,…,WN]=[XW1,XW2,…,XWN]Y = XW = X[W_1, W_2, \dots, W_N] = [XW_1, XW_2, \dots, XW_N]

每张卡 ii 只需计算 Yi=XWiY_i = XW_i,得到 Yi∈ℝB×S×H′/NY_i \in \mathbb{R}^{B \times S \times H’/N}。输出在特征维度上被切分。

特征维度:每个 token 用一个长度为  hidden_size  的向量表示,向量的每个元素可以理解为一个“特征”。

行切分 (Row Parallelism):

将 WW 按行切分为 NN 份:W=[W1 W2 ⋮ WN]W = \begin{bmatrix} W_1 \ W_2 \ \vdots \ W_N \end{bmatrix},其中 Wi∈ℝH/N×H′W_i \in \mathbb{R}^{H/N \times H’}。

对应地,输入 XX 也按特征维度切分:X=[X1,X2,…,XN]X = [X_1, X_2, \dots, X_N],其中 Xi∈ℝB×S×H/NX_i \in \mathbb{R}^{B \times S \times H/N}。

Y=XW=∑i=1NXiWiY = XW = \sum_{i=1}^{N} X_i W_i

每张卡 ii 计算 Yi=XiWiY_i = X_i W_i,得到完整特征维度但只有部分贡献的 Yi∈ℝB×S×H′Y_i \in \mathbb{R}^{B \times S \times H’},需要 All-Reduce 求和得到最终 YY。

1.2.2 SP 的激活值切分推导

在标准 TP 中,每张卡在非 TP 区域(RMSNorm)仍持有完整的 [B,S,H][B, S, H] 激活值。

SP 切分的是  序列维度  seq_len,而不是特征维度  hidden_size。
而 RMSNorm 是对  每个 token 自己的特征向量  做归一化,不同 token 之间没有依赖。因此可以安全地按 SS 切分。

因此,在非 TP 区域,激活值只需 [B,S/Ntp,H][B, S/N_{tp}, H],而非 [B,S,H][B, S, H]。

SP 的通信优化:

标准 TP 流程:

[SP 激活: B*S/N*H] -> All-Gather -> [B*S*H] -> QKV 列切分 -> Attention -> All-Reduce -> [B*S*H] -> RMSNorm

SP 优化后流程(All-Gather + Reduce-Scatter 融合):

[SP 激活: B*S/N*H] -> All-Gather(S) -> [B*S*H] -> QKV 列切分 -> Attention -> Reduce-Scatter(S) -> [B*S/N*H] -> RMSNorm

这样,TP 的 All-Reduce 被分解为 All-Gather + Reduce-Scatter,正好与 SP 的区域切换对齐,不引入额外通信。

数学等价性证明:

  • All-Reduce = All-Gather + Reduce-Scatter(Ring All-Reduce 的分解)
  • 标准 TP 的 All-Reduce 对 YiY_i 求和得 YY,然后每卡持有完整的 B×S×HB \times S \times H
  • SP 的 Reduce-Scatter 对 YiY_i 求和并按 SS 切分,每卡得到 B×S/Ntp×HB \times S/N_{tp} \times H
  • 计算结果等价,只是存储方式不同(SP 下每卡只存 S/NtpS/N_{tp} 片段)

1.3 切分内容精确描述

1.3.1 参数切分

层 / 组件 权重 切分方式 每卡形状 切分维度
QKV 投影 Wqkv∈ℝH×3HW_{qkv} \in \mathbb{R}^{H \times 3H} 列切分 ℝH×3H/TP\mathbb{R}^{H \times 3H/TP} 输出维度按列
Output 投影 Wo∈ℝH×HW_{o} \in \mathbb{R}^{H \times H} 行切分 ℝH/TP×H\mathbb{R}^{H/TP \times H} 输入维度按行
MLP WgateW_{gate} Wgate∈ℝH×4HW_{gate} \in \mathbb{R}^{H \times 4H} 列切分 ℝH×4H/TP\mathbb{R}^{H \times 4H/TP} 输出维度按列
MLP WupW_{up} Wup∈ℝH×4HW_{up} \in \mathbb{R}^{H \times 4H} 列切分 ℝH×4H/TP\mathbb{R}^{H \times 4H/TP} 输出维度按列
MLP WdownW_{down} Wdown∈ℝ4H×HW_{down} \in \mathbb{R}^{4H \times H} 行切分 ℝ4H/TP×H\mathbb{R}^{4H/TP \times H} 输入维度按行
Embedding E∈ℝV×HE \in \mathbb{R}^{V \times H} 按词表行切分 ℝV/TP×H\mathbb{R}^{V/TP \times H} 词表维度 V

切分原理详解:

  • QKV 列切分 :WqkvW_{qkv} 的输出维度 3H3H 对应 3×Nh3 \times N_h 个 head 的 QKV。每个矩阵按列切分TPTP 份,每卡得到 Nh/TP=64÷4N_h/TP = 64 \div 4,即 16 个 head。每卡只需计算自己负责的 16 个 head 的 attention。
  • Output 行切分:WoW_o 的输入维度 HH 对应所有 head 的拼接结果。每卡只有 16 个 head 的 attn 输出,因此 WoW_o 按行切分,每卡 Wo,i∈ℝH/TP×HW_{o,i} \in \mathbb{R}^{H/TP \times H},计算 Yi=attni⋅Wo,iY_i = \text{attn}i \cdot W{o,i} 得到部分和,All-Reduce 求和。
  • MLP WgateW_{gate}列切分 :WgateW_{gate} 输出维度 4H4H 按列切TPTP 份,每卡 4H/TP=4H/4=H4H/TP=4H/4=H,∈RH×H\in R^{H \times H}。每卡计算部分门控状态。
  • MLP WupW_{up}列切分 :WupW_{up} 输出维度 4H4H 按列切TPTP 份,每卡 4H/TP=4H/4=H4H/TP=4H/4=H,∈RH×H\in R^{H \times H}。每卡计算部分上投影状态。
  • MLP WdownW_{down}行切分 :WdownW_{down} 输入维度 4H4H 按行切TPTP 份,每卡 4H/TP=4H/4=H4H/TP=4H/4=H,∈RH×H\in R^{H \times H}。每卡将 SiLU 门控后的中间状态乘以对应行块得到部分和,All-Reduce 求和。

SwiGLU 结构说明:现代大模型使用 SwiGLU 替代传统 FFN:

SwiGLU(x)=(SiLU(x@Wgate)⊙(x@Wup))@Wdown\text{SwiGLU}(x) = (\text{SiLU}(x @ W_{gate}) \odot (x @ W_{up})) @ W_{down}

其中 SiLU(x)=x⋅σ(x)\text{SiLU}(x) = x \cdot \sigma(x)(也叫 Swish 激活函数),σ(x)\sigma(x) 为 sigmoid 函数。三个权重分别为 Wgate∈ℝH×4HW_{gate} \in \mathbb{R}^{H \times 4H}, Wup∈ℝH×4HW_{up} \in \mathbb{R}^{H \times 4H}, Wdown∈ℝ4H×HW_{down} \in \mathbb{R}^{4H \times H}。TP 切分时 WgateW_{gate} 和 WupW_{up}按列切分,WdownW_{down} 按行切分。

为了保持参数量一致,Llama 实际上把 SwiGLU 的中间维度从 4H4H 调小为 23×4H≈2.67H\frac{2}{3} \times 4H \approx 2.67H​(再向上取整到 256 的倍数)

1.3.2 梯度切分

梯度的切分方式与参数完全一致——因为梯度形状与参数相同。列切分参数的梯度也是列切分,行切分参数的梯度也是行切分。TP 组内每张卡只计算和持有自己那部分参数的梯度。

1.3.3 优化器状态切分

TP 本身不切分优化器状态。优化器状态由 DP(或 ZeRO)负责切分。TP 组内的每张卡独立维护自己持有的参数所对应的优化器状态(如 Adam 的 m 和 v)。

1.3.4 激活值切分(精确到层内位置)

激活值切分总结表:

激活值位置 切分方式 每卡形状 区域
RMSNorm1 输入 / 输出 SP: 序列切分 B×S/TP×HB \times S/TP \times H SP 区域
All-Gather 后, QKV 投影前 完整序列 B×S×HB \times S \times H TP 区域
QKV 投影后 TP: 特征切分 B×S×3H/TPB \times S \times 3H/TP TP 区域
Attention 中间(Q,K,V,attn) TP: head 切分 B×S×H/TPB \times S \times H/TP TP 区域
Output 投影后(RS 前) TP: 完整特征, 部分和 B×S×HB \times S \times H TP 区域
Reduce-Scatter 后(残差前) SP: 序列切分 B×S/TP×HB \times S/TP \times H SP 区域
RMSNorm2 输入 / 输出 SP: 序列切分 B×S/TP×HB \times S/TP \times H SP 区域
MLP WgateW_{gate}/WupW_{up}后 TP: 特征切分 B×S×4H/TPB \times S \times 4H/TP TP 区域
MLP WdownW_{down}后(RS 前) TP: 完整特征, 部分和 B×S×HB \times S \times H TP 区域
层间传递的激活值 SP: 序列切分 B×S/TP×HB \times S/TP \times H SP 区域

1.4 通信组构成

  • TP 通信组: 同一节点内的 TP=4TP=4 张卡
  • SP 与 TP 共用同一通信组: SP 不是独立的并行维度,而是 TP 区域内激活值的存储策略
  • 物理拓扑建议: 同一 NVLink 域内(TP 通信频繁,需要高带宽)

1.5 通信方式

通信点 通信方式 数据流
SP->TP 过渡(Attention 前) All-Gather (沿 S 维度) 4 卡各持 B×S/TP×HB \times S/TP \times H -> 每卡得 B×S×HB \times S \times H
TP 区域结束(Attention 后) Reduce-Scatter (沿 S 维度) 4 卡各持 B×S×HB \times S \times H(部分和) -> 每卡得 B×S/TP×HB \times S/TP \times H
SP->TP 过渡(MLP 前) All-Gather (沿 S 维度) 同上
TP 区域结束(MLP 后) Reduce-Scatter (沿 S 维度) 同上

标准 TP(无 SP)使用 All-Reduce 替代上述 All-Gather + Reduce-Scatter 对。SP 将 All-Reduce 分解为 All-Gather + Reduce-Scatter,并与 SP 区域的序列切分自然对齐。

1.6 通信数据量计算

1.6.1 单次通信量公式

All-Gather (SP->TP):

  • 每卡发送: B×SNtp×H×dtype_bytesB \times \frac{S}{N_{tp}} \times H \times \text{dtype\_bytes}
  • 每卡接收: B×SNtp×H×(Ntp−1)×dtype_bytesB \times \frac{S}{N_{tp}} \times H \times (N_{tp}-1) \times \text{dtype\_bytes}
  • 每卡总通信量(发送 + 接收): B×S×H×dtype_bytesB \times S \times H \times \text{dtype\_bytes}

Reduce-Scatter (TP->SP): 与 All-Gather 通信量相同(对称操作)

1.6.2 具体数值计算

代入 B=2,S=8192,H=8192,Ntp=4,dtype=2 bytes (bf16)B=2, S=8192, H=8192, N_{tp}=4, \text{dtype}=2\text{bytes (bf16)}:

每卡发送量:
Vsend=B×SNtp×H×2=2×2048×8192×2=67,108,864 bytes=64 MiBV_{send} = B \times \frac{S}{N_{tp}} \times H \times 2 = 2 \times 2048 \times 8192 \times 2 = 67,108,864 \text{bytes} = 64 \text{MiB}

每卡接收量:
Vrecv=B×SNtp×H×(Ntp−1)×2=2×2048×8192×3×2=201,326,592 bytes=192 MiBV_{recv} = B \times \frac{S}{N_{tp}} \times H \times (N_{tp}-1) \times 2 = 2 \times 2048 \times 8192 \times 3 \times 2 = 201,326,592 \text{bytes} = 192 \text{MiB}

每卡总通信量: Vtotal=64+192=256 MiBV_{total} = 64 + 192 = 256 \text{MiB}

或等价: Vtotal=B×S×H×2=2×8192×8192×2=268,435,456 bytes=256 MiBV_{total} = B \times S \times H \times 2 = 2 \times 8192 \times 8192 \times 2 = 268,435,456 \text{bytes} = 256 \text{MiB}

1.6.3 单层总通信量

每层: 2 次 All-Gather + 2 次 Reduce-Scatter = 4 次通信

单层每卡总通信量: 4×256=1024MiB=1GiB4 \times 256 = 1024 MiB = 1 GiB

80 层前向 TP 通信总量: 80×1=80GiB80 \times 1 = 80 GiB

反向通信量与前向相同: 80 GiB

前向 + 反向 TP 通信总量: 160 GiB/ 卡

1.7 优化手段

1.7.1 All-Gather + Reduce-Scatter 替代 All-Reduce

标准 TP 使用 2 次 All-Reduce(每层),SP 方案用 All-Gather + Reduce-Scatter 替代,总通信量相近但激活值显存降低 4 倍。

1.7.2 通信 - 计算重叠

  • 在 Attention 计算 QKV 投影的同时,可以启动下一层的 All-Gather
  • Megatron 使用 CUDA stream 实现异步通信

1.7.3 SP 与选择性重计算结合

SP 将激活值降低 NtpN_{tp} 倍。结合选择性重计算(不存储 attention 矩阵,反向时重计算),可进一步降低显存。

1.8 前向与反向流程详解

前向流程(单层)

输入: x∈[B,S/TP,H]x \in [B,S/TP,H] SP 切分
(1) RMSNorm1(x) → [B,S/TP,H][B,S/TP,H] SP
(2) All-Gather(S) → xfull∈[B,S,H]x_{full} \in [B,S,H] Vsend=B×SNtp×H×2=64 MiBV_{send} = B \times \frac{S}{N_{tp}} \times H \times 2 = 64 \text{MiB}
Vrecv=B×SNtp×H×(Ntp−1)×2=192 MiBV_{recv} = B \times \frac{S}{N_{tp}} \times H \times (N_{tp}-1) \times 2 = 192 \text{MiB}
(3) QKV=xfull@WqkviQKV = x_{full} @ W_{qkv_i} → [B,S,(3H/TP)][B,S,(3H/TP)]
reshape → Qi,Ki,Vi∈[B,S,(Nh/TP)∗dh]Q_i, K_i, V_i \in [B,S,(N_h/TP)*d_h]
TP 列切分
(4) Attentioni=softmax(Qi@KiT/sqrt(dh))@ViAttention_i = softmax(Q_i @ {K_i}^T / sqrt(d_h)) @ V_i
→ attni∈[B,S,(H/TP)]attn_i \in [B,S,(H/TP)]
TP
(5) Yi=attni@WoiY_i = attn_i @ W_{o_i} → [B,S,H][B,S,H](部分和) TP 行切分
(6) Reduce-Scatter(S) → attnout∈[B,S/TP,H]attn_{out} \in [B,S/TP,H] Vsend=B×SNtp×H×2=64 MiBV_{send} = B \times \frac{S}{N_{tp}} \times H \times 2 = 64 \text{MiB}
Vrecv=B×SNtp×H×(Ntp−1)×2=192 MiBV_{recv} = B \times \frac{S}{N_{tp}} \times H \times (N_{tp}-1) \times 2 = 192 \text{MiB}
(7) x=xsp+(attnout)x = x_{sp} + (attn_{out}) → [B,S/TP,H] [B,S/TP,H] SP
(8) RMSNorm2(x) → [B,S/TP,H][B,S/TP,H] SP
(9) All-Gather(S) → xfull∈[B,S,H]x_{full} \in [B,S,H] Vsend=B×SNtp×H×2=64 MiBV_{send} = B \times \frac{S}{N_{tp}} \times H \times 2 = 64 \text{MiB}
Vrecv=B×SNtp×H×(Ntp−1)×2=192 MiBV_{recv} = B \times \frac{S}{N_{tp}} \times H \times (N_{tp}-1) \times 2 = 192 \text{MiB}
(10a) gate=SiLU(xfull@Wgatei)gate = SiLU(x_{full} @ W_{gate_i}) → [B,S,4H/TP][B,S,4H/TP] TP 列切分
(10b) up=xfull@Wupiup = x_{full} @ W_{up_i} → [B,S,4H/TP][B,S,4H/TP] TP 列切分
(10c) h=gate∗uph = gate * up → [B,S,4H/TP][B,S,4H/TP] TP, 逐元素
(11)Yi=h@Wdowni Y_i = h @ W_{down_i}→ [B,S,H][B,S,H] (部分和) TP 行切分
(12) Reduce-Scatter(S) → out∈[B,S/TP,H]out \in [B,S/TP,H] Vsend=B×SNtp×H×2=64 MiBV_{send} = B \times \frac{S}{N_{tp}} \times H \times 2 = 64 \text{MiB}
Vrecv=B×SNtp×H×(Ntp−1)×2=192 MiBV_{recv} = B \times \frac{S}{N_{tp}} \times H \times (N_{tp}-1) \times 2 = 192 \text{MiB}

反向流程(单层)

对应关系: 正向 All-Gather 的反向 =Reduce-Scatter, 正向 Reduce-Scatter 的反向 =All-Gather

梯度输入: ∂L∂out=dYi∈[B,S/TP,H]\frac {\partial L} {\partial out} = dY_i \in [B,S/TP,H] SP 切分
(12) All-Gather(S) → dYfull∈[B,S,H]dY_{full} \in [B,S,H] Vsend=B×SNtp×H×2=64 MiBV_{send} = B \times \frac{S}{N_{tp}} \times H \times 2 = 64 \text{MiB}
Vrecv=B×SNtp×H×(Ntp−1)×2=192 MiBV_{recv} = B \times \frac{S}{N_{tp}} \times H \times (N_{tp}-1) \times 2 = 192 \text{MiB}
(11)下投影 Wdown 的梯度与 dHi →
dHi=dYfull@Wdown_iTdH_i =dY_{full} @ W_{down\_i}^T
dWdowni=HiTdYfulldW_{down_i} = {H_i}^T dY_{full}
TP
(10)
计算 dWgatei、dWupi与 dXidW_{gate_i}、dW_{up_i} 与 dX_i
dAi=dHi⊙UidUi=dHi⊙Ai
dGi=dAi⊙[σ(Gi)+Ai⊙(1−σ(Gi))]
dWgatei=XTdGidW_{gate_i} = X^T dG_i
dWupi=XTdUidW_{up_i} = X^T dU_i
dXgate(i)=dGi (Wgate(i))TdXup(i)=dUi (Wup(i))TdX(i)=dXgate(i)+dXup(i)dXi∈[B,S,H]dX_i \in [B,S,H]
TP
(9) Reduce-Scatter(S) → dL/dx in B*S/TP*H


  +--(9) Reduce-Scatter(S) -> dL/dx in B*S/TP*H     [ 通信: 64MiB send, 192MiB recv]
  +--(8) RMSNorm2 反向: 逐位置计算, 无通信
  +--(7) 残差反向: dL/dattn_out = dL/dx; 残差直通
  +--(6) All-Gather(S) -> dL/dY_full in B*S*H      [ 通信: 64MiB send, 192MiB recv]
  +--(5) dL/dattn_i = dL/dY_i @ W_o_i^T; dL/dW_o_i   [TP, 无通信]
  +--(4) Attention 反向: dL/dQ_i, dL/dK_i, dL/dV_i   [TP, 无通信]
  +--(3) QKV 反向: dL/dx_full_i; dL/dW_qkv_i          [TP, 无通信]
  +--(2) Reduce-Scatter(S) -> dL/dx in B*S/TP*H     [ 通信: 64MiB send, 192MiB recv]
  +--(1) RMSNorm1 反向: 逐位置计算, 无通信
  +-- 输出梯度传给上一层

反向通信量 = 正向通信量 = 4 次 x 256 MiB = 1024 MiB/ 层


2. Pipeline Parallelism (PP)

2.1 原理

Pipeline Parallelism 将 Transformer block 按层 (80 层) 切分到 Npp=4N_{pp}=4 个流水线阶段(stage),每个 stage 负责连续的 20 层。数据像流水线上的产品一样,逐 stage 传递。

核心问题 – 气泡 (Bubble): 如果严格串行(等待所有前向完成后再反向),则 GPU 利用率极低。Megatron 采用 1F1B (One Forward, One Backward) 调度策略来减少气泡。

1F1B 调度原理:

  • warmup 阶段:stage i(i 从 0 开始)先执行 Npp−i−1N_{pp} – i – 1 次前向
  • 稳态阶段:交替执行 1 次前向 + 1 次反向
  • cooldown 阶段:执行剩余的反向

气泡数量: Npp−1N_{pp}-1 次前向的时间

1F1B 调度时序图 (以 Npp=4N_{pp}=4, M=8M=8 micro-batches 为例):

其中 F= 前向, B= 反向, 数字是 micro-batch 编号。

2.2 激活值

在 Transformer 架构中,PP 沿着层(Layer)维度进行切分。假设 Stage i 负责 Layer 0~19,Stage i+1 负责 Layer 20~29。

  • 前向传播 (Forward):Stage i 将 Layer 19 的输出 hidden_states 发送给 Stage i+1 作为其输入。
  • 反向传播 (Backward):Stage i+1 计算完 Layer 20 的梯度后,将关于该 hidden_states 的梯度 grad_hidden_states 发送回 Stage i。

在最基础的设定下,这个边界激活值的形状(Shape)为:[B,S,H][B, S, H] , 单次 P2P 传输的数据量为:S×B×H×2S \times B \times H \times 2 Bytes。

2.2.1 结合 TP、SP、CP 时的激活值形态

当 PP 与 Tensor Parallelism (TP)、Sequence Parallelism (SP) 和 Context Parallelism (CP) 组合时,由于输入序列在不同维度上被分片(Sharding),跨 PP 边界的激活值 Shape 和通信拓扑会发生显著变化。

1. PP + TP (无 Sequence Parallelism)

在传统的 Megatron-LM 纯 TP 模式下,每个 Transformer 层的输入和输出在 TP 组内是完全复制(Replicated)的。

  • 边界形态:每个 TP Rank 都在 PP 边界处持有完整的 $[S, B, H]$ 激活值。
  • 通信模式:Stage i 的 TP Rank k 直接将大小为 [B,S,H][B,S , H] 的数据通过 send/recv 发送给 Stage i+1 的 TP Rank k。

2. PP + TP + SP (Sequence Parallelism)

SP 将 AllReduce 拆解为 Reduce-Scatter 和 All-Gather,使得各层之间的激活值在序列维度 SS 上被切分。

  • 边界形态:在 PP 边界处,激活值被 TP 组切分成了 TPTP 份。每个 Rank 持有的 Shape 为:[B,S/TP,H][B,S / TP, H]
  • 通信模式:Stage i 的 TP Rank k 只需发送 [B,S/TP,H][B,S / TP, H] 给 Stage i+1 的 TP Rank k。

3. PP + TP + SP + CP (Context Parallelism)

CP 进一步在序列维度上对数据进行切分(例如 DeepSpeed Ulysses 或 Megatron CP 方案)。

  • 边界形态:序列维度被 SP 和 CP 共同切分。每个 Rank 在 PP 边界处的激活值 Shape 变为:[B,S/(TP×CP),H][B, S / (TP \times CP), H]
  • 通信模式:Stage i 中的 (CP Rank j, TP Rank k) 将这一小块激活值发给 Stage i+1 中对应的 (CP Rank j, TP Rank k)。

2.2.2 不同 PP 调度策略对激活值显存的影响

PP 的调度策略决定了在同一时刻,GPU 显存中必须驻留多少个 Micro-batch 的激活值。这是 AI Infra 系统设计中权衡“显存占用”与“Pipeline Bubble(气泡)”的核心考量。

假设 Micro-batch 总数为 MM。

1. GPipe (Default Batching)

并行度汇总
  • 调度逻辑:前向传播所有的 MM 个 Micro-batch,然后再反向传播 MM 个 Micro-batch。
  • 激活值驻留量:Stage 1 必须将所有 MM 个 Micro-batch 的前向激活值全部保存在显存,直到对应的反向传播到来。
  • 显存占用峰值:O(M)O(M)。通常 M≫pM \gg p,这会导致极大的显存压力,目前主流 LLM 训练已基本弃用。

2. 1F1B (One Forward One Backward)

并行度汇总

也有这种:

并行度汇总
  • 调度逻辑:进入稳态后,交替执行 1 个 Forward 和 1 个 Backward。
  • 激活值驻留量:每个 Stage 需要驻留的 Micro-batch 激活值数量取决于它在 Pipeline 中的位置。Stage i(从 0 开始)在稳态前需要执行 PP−i−1PP – i – 1 个 Warmup Forward。
  • 显存占用峰值:最多只需缓存 PPPP 个 Micro-batch 的激活值。即 O(PP)O(PP)。

3. Interleaved 1F1B (Megatron-LM)

并行度汇总
  • 调度逻辑:将一个物理 GPU 映射到 vv 个虚拟 Stage(Virtual Stages / Model Chunks)。通过让不同阶段的计算交错,缩小最终的 Pipeline Bubble。
  • 激活值驻留量 :由于模型被切得更碎(每块层数变为 L/(PP×v)L / (PP \times v)),但需要处理的 Chunk 数量增加了 vv 倍。驻留的 Micro-batch 数量依然由 PPPP 决定,但由于每个 Chunk 变小,整体激活值 峰值显存占用与基础 1F1B 相当。
  • 网络特征 (关键):Interleaved 1F1B 的跨 Stage 边界数量增加了 vv 倍,导致 P2P 通信频率和通信总量增加了 vv 倍。

4. Zero-Bubble (如 ZB-H1 / V-Schedule)

并行度汇总
  • 调度逻辑:将 Backward 进一步拆解为计算激活值梯度的 BB 和计算权重梯度的 WW, 利用 WW 来填补传统 1F1B 留下的气泡。
  • 激活值驻留量 :为了填补气泡,WW 被大幅延后执行。因为权重梯度的计算需要依赖 前向激活值,延后 WW 意味着前向激活值的生命周期被拉长。
  • 显存占用峰值 : 显著高于 1F1B。为了换取接近零的气泡,系统必须在显存中缓存多于 PPPP 个 Micro-batch 的激活值(具体取决于 WW 的排布位置)。

2.3 数学推导

2.3.1 显存分析

假设 L=80L=80 层均匀切分到 Npp=4N_{pp}=4 个 stage,每 stage L/Npp=20L/N_{pp}=20 层。

1F1B 下,stage i 的最大激活值份数:
Mi=Npp−i(i=0,1,…,Npp−1)M_i = N_{pp} – i \quad (i = 0, 1, \dots, N_{pp}-1)

  • Stage 0: 缓存最多 Npp=4N_{pp} = 4 份 micro-batch 的激活值
  • Stage 1: 3 份
  • Stage 2: 2 份
  • Stage 3: 1 份

2.3.2 气泡占比

设单次前向时间为 tft_f,单次反向时间为 tbt_b,micro-batch 数为 MM。

气泡占比:
Bubble Fraction=Npp−1M+Npp−1\text{Bubble Fraction} = \frac{N_{pp}-1}{M + N_{pp}-1}

当 M=8,Npp=4M=8, N_{pp}=4: Bubble=3/11≈27.3%\text{Bubble} = 3/11 \approx 27.3\%

当 M≫NppM \gg N_{pp} 时,气泡占比趋近于 0。

2.3.3 通信量

PP 的通信发生在相邻 stage 之间,是 P2P 通信。

前向: stage i -> stage i+1 传递激活值

注意:如果同时使用了 SP,则层间传递的激活值是 B×(S/Ntp)×HB \times (S/N_{tp}) \times H,因为 SP 区域下层间传递的是序列切分后的激活值。

实际前向 P2P 通信量(含 SP):
Vforward=B×SNtp×H×2=2×2048×8192×2=67,108,864 bytes=64 MiBV_{forward} = B \times \frac{S}{N_{tp}} \times H \times 2 = 2 \times 2048 \times 8192 \times 2 = 67{,}108{,}864 \text{bytes} = 64 \text{MiB}

反向: stage i+1 -> stage i 传递梯度,通信量相同:
Vbackward=B×SNtp×H×2=64 MiBV_{backward} = B \times \frac{S}{N_{tp}} \times H \times 2 = 64 \text{MiB}

2.4 切分内容精确描述

切分对象 切分方式 说明
参数(权重) 按层切分 每 stage 持有连续的 L/Npp=20L/N_{pp}=20 层的全部权重
梯度 按层切分 与参数一致,每 stage 只持有所在层的梯度
优化器状态 按层切分 与参数一致
激活值(层间) 按层切分 每 stage 只持有自己 20 层的中间激活值
激活值(micro-batch 缓存) 1F1B 调度缓存 stage i 最多缓存 Npp−iN_{pp}-i 份 micro-batch 的激活值

精确说明:

  • PP 切分的是 层间 的划分,不是层内的切分
  • 每 stage 的权重是完整的(在该 stage 内的 20 层中,权重不被 TP 以外的并行度切分)
  • 激活值方面,PP 切的是哪些层的激活值在哪张卡上——stage 0 的卡只存储 layer 1-20 的激活值
  • 层内激活值 不被 PP 切分(被 TP/SP 切分)

2.5 通信组构成

  • PP 通信组: 相邻 stage 之间两两通信
  • 通信是 P2P Send-Recv
  • stage i 只与 stage i-1 和 stage i+1 通信
  • stage 0 只与 stage 1 通信(前向发送)
  • stage Npp−1N_{pp}-1 只与 stage Npp−2N_{pp}-2 通信(反向接收梯度)

2.6 通信方式

通信点 通信方式 方向 数据
前向: stage i -> i+1 P2P Send-Recv 前向 激活值 B×S/Ntp×HB \times S/N_{tp} \times H
反向: stage i+1 -> i P2P Send-Recv 反向 梯度 B×S/Ntp×HB \times S/N_{tp} \times H

2.7 优化手段

2.7.1 1F1B 调度

显存从 MM 份降低到 Npp−iN_{pp}-i 份。

2.7.2 Interleaved 1F1B (交错式调度)

Megatron-LM v2 的优化:将每 stage 的连续层再切分为多个 chunk(如 2 个 chunk),交错执行不同 chunk 的前向 / 反向。气泡从 Npp−1N_{pp}-1 降低到 Npp−1V\frac{N_{pp}-1}{V},其中 VV 是每 stage 的 chunk 数。代价:P2P 通信次数增加 VV 倍,但每次通信量减少。

2.7.3 Pipeline Bubble Fill

在气泡时间内执行其他有用计算(如 DP 的梯度同步),隐藏通信延迟。

2.7.4 micro-batch 数量优化

增大 MM 可以降低气泡比例,但增加显存。需要平衡。

2.8 前向与反向流程

前向流程

[micro-batch m 输入] -> Stage 0 (Layer 1-20)
  | 计算 20 层前向
  | 缓存中间激活值(用于反向)
  | P2P Send: 激活值 -> Stage 1  [64 MiB]
  v
Stage 1 (Layer 21-40)
  | 计算 20 层前向
  | 缓存中间激活值
  | P2P Send: 激活值 -> Stage 2  [64 MiB]
  v
Stage 2 (Layer 41-60)
  | 计算 20 层前向
  | 缓存中间激活值
  | P2P Send: 激活值 -> Stage 3  [64 MiB]
  v
Stage 3 (Layer 61-80)
  | 计算 20 层前向
  | 计算 loss
  | 开始反向

反向流程

Stage 3 (Layer 80->61)
  | 反向计算 20 层梯度
  | 释放对应前向激活值
  | P2P Send: 梯度 -> Stage 2  [64 MiB]
  v
Stage 2 (Layer 60->41)
  | 反向计算 20 层梯度
  | 释放激活值
  | P2P Send: 梯度 -> Stage 1  [64 MiB]
  v
Stage 1 (Layer 40->21)
  | 反向计算 20 层梯度
  | 释放激活值
  | P2P Send: 梯度 -> Stage 0  [64 MiB]
  v
Stage 0 (Layer 20->1)
  | 反向计算 20 层梯度
  | 释放激活值
  | 梯度就绪,等待 DP 同步

每层前向计算流程(Pre-RMSNorm + SwiGLU):

[ 输入 x] (来自上游 stage 或 embedding)
  |
  +-- RMSNorm1(x) -> x_norm           [ 无通信]
  +-- Attention(x_norm) -> attn_out    [TP 通信: AG+RS, 见 Ch1]
  +-- x = x + attn_out                 [ 残差连接, 无 dropout]
  +-- RMSNorm2(x) -> x_norm2          [ 无通信]
  +-- SwiGLU FFN(x_norm2) -> ffn_out  [TP 通信: AG+RS, 见 Ch1]
  +-- x = x + ffn_out                  [ 残差连接, 无 dropout]
  |
[ 输出 x] (传递给下游 stage 或输出层)

每层反向流程:

[ 梯度 dL/dx] (来自下游 stage)
  |
  +-- dL/dx_norm2 = dL/dx (残差直通)
  +-- RMSNorm2 反向: dL/dx += dL/dx_norm2   [ 无通信]
  +-- SwiGLU FFN 反向: dL/dffn_out, dL/dW_gate, dL/dW_up, dL/dW_down  [TP 通信]
  +-- dL/dx += dL/dffn_out (残差直通)
  +-- RMSNorm1 反向: dL/dx_norm = dL/dx     [ 无通信]
  +-- Attention 反向: dL/dattn_out, dL/dW_qkv, dL/dW_o  [TP 通信]
  +-- dL/dx += dL/dattn_out (残差直通)
  |
[ 梯度 dL/dx] (传递给上游 stage)

2.9 一些其他 overlap 相关优化

1. deepseek dualpipe https://zhuanlan.zhihu.com/p/28277195890

并行度汇总
并行度汇总

2. 1f1b 实现类似 dualpipe https://developer.nvidia.cn/blog/1f1b-moe-a2a-computing-overlap/

并行度汇总

3. 首尾不均衡问题

首 stage(Stage 0):

  • 有 Embedding 层——额外的参数和计算(词表 V × H 的 lookup + dropout)
  • 接收原始输入数据,需要做 tokenization / positional embedding
  • 反向传播结束时,梯度从这里开始传给上游(没有上游 stage 传梯度给它)

尾 stage(Stage 3):

  • 有 Output LM Head(Wout∈[H,V]W_{out} \in [H, V]W​out​​∈[H,V])——计算 logits 和 loss
  • 计算 Cross-Entropy Loss
  • 反向传播从这里开始(第一个做反向的 stage)
  • 没有下游 stage 需要它发送激活值

中间 stage(Stage 1, 2):

  • 只有 Transformer 层(Attention + MLP/MoE)
  • 前向收上游激活值,反向收下游梯度

解决:非均匀层切分(Non-uniform Pipeline Split)

不按层数均分,而是按 计算量 或显存 切分:

如果首 stage 有 Embedding(额外显存)→ 给它少几层 Transformer
如果尾 stage 有 LM Head(额外计算)→ 给它少几层 Transformer

例如:
Stage 0: Embedding + 18 层 (减轻以补偿 Embedding 显存)
Stage 1: 20 层
Stage 2: 20 层
Stage 3: 22 层 + LM Head (增加层数但 LM Head 计算量相对小)

Megatron-LM 支持 num_layers_per_virtual_pipeline_stage 和自定义层分配。

3. Context Parallelism (CP)

Context Parallelism(上下文并行)将输入序列沿序列维度 SS 切分到多张卡上,使每张卡只需处理序列的一个子段,从而突破单卡显存对长序列的限制。CP 主要切分的是 Attention 层的激活值(具体为 Q,K,VQ, K, V 及其中间产物),对 MLP 层和 RMSNorm 层的激活值不做跨卡通信(这些层在序列维度上是独立的,各卡只需处理自己那段序列即可)。

CP 有三种主流实现变体,下面分别详述。

3.1 公共基础

序列切分: 将长度为 SS 的序列均匀切分到 NcpN_{cp} 张卡上,每卡处理 Slocal=S/NcpS_{local} = S / N_{cp} 个 token。

Slocal=SNcpS_{local} = \frac{S}{N_{cp}}

数值假设(本节通用):

  • B=2,S=8192,H=8192,Nh=64,dh=128B = 2, S = 8192, H = 8192, N_h = 64, d_h = 128
  • Ncp=2N_{cp} = 2(CP 并行度)
  • Slocal=4096S_{local} = 4096
  • dtype = bf16(2 bytes)

切分内容精确描述:

  • 参数:各项权重矩阵均不切分。Attention 的 Wq,Wk,Wv,WoW_q, W_k, W_v, W_o 和 MLP 的 Wgate,Wup,WdownW_{gate}, W_{up}, W_{down} 在所有 CP 卡上完整复制。
  • 梯度:不切分(梯度在 CP 组内通过 All-Reduce 同步)。
  • 优化器状态:不切分。
  • 激活值切分(精确到层内位置):
    • Attention 层的输入 X∈[B,S,H]X \in [B, S, H]:沿序列切分为 Xi∈[B,Slocal,H]X_i \in [B, S_{local}, H]。 这是层间激活值在进入 Attention 前被切分。
    • Q,K,VQ, K, V:每卡计算自己那段的 Qi,Ki,Vi∈[B,Slocal,H]Q_i, K_i, V_i \in [B, S_{local}, H],这些是 层内中间激活值,标准 Attention 需要全序列 K,VK, V 才能计算 softmax(QKT/dh)V\text{softmax}(QK^T/\sqrt{d_h})V,因此需要跨卡通信。
    • Attention 矩阵 P=QKT∈[B,Nh,Slocal,S]P = QK^T \in [B, N_h, S_{local}, S]:每卡需全序列 KK,这是 CP 通信的核心来源。
    • Attention 输出 O∈[B,Slocal,H]O \in [B, S_{local}, H]:序列维度已切分,后续 WoW_o 投影在各卡独立完成。
    • MLP 层 :MLP(SwiGLU)在序列维度上完全独立,各卡只需处理自己的 SlocalS_{local} 个 token, 无需跨卡通信。
    • RMSNorm:序列维度独立运算,各卡处理自己的子段,无需跨卡通信。

3.2 变体一:Ring-Attention

3.2.1 原理

Ring-Attention 采用 P2P 环形通信,将 K,VK, V 在 CP 组的卡间循环传递。每卡在收到其他卡的 K,VK, V 后,立即计算本地 QQ 与这部分 K,VK, V 的 Attention 部分分数,累加到本地 OO 上。

核心思想:计算与通信重叠——当卡 i 在用卡 j 的 K,VK, V 计算 Attention 时,同时在接收卡 j+1 的 K,VK, V。

3.2.2 数学推导

标准 Attention:Attn(Q,K,V)=softmax(QKTdh)V\text{Attn}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_h}}\right)V

分块计算(卡 i 持有 QiQ_i,需要所有 Kj,VjK_j, V_j):

Pij=QiKjT/dh∈[B,Nh,Slocal,Slocal]P_{ij} = Q_i K_j^T / \sqrt{d_h} \in [B, N_h, S_{local}, S_{local}]

Oi=∑j=0Ncp−1softmaxj(Pij)VjO_i = \sum_{j=0}^{N_{cp}-1} \text{softmax}_{j}(P_{ij}) V_j

其中 softmaxj\text{softmax}_{j} 表示在线 softmax 的第 j 块更新(需维护全局最大值和指数和的运行状态)。

在线 softmax 更新公式:

设当前已处理块的最大值为 mm,指数和为 ℓ\ell,累加输出为 OO。处理新块 jj 时:

mnew=max⁡(m,max⁡(Pij))m_{new} = \max(m, \max(P_{ij}))
ℓnew=ℓ⋅em−mnew+∑tePij,t−mnew\ell_{new} = \ell \cdot e^{m – m_{new}} + \sum_t e^{P_{ij,t} – m_{new}}
Onew=O⋅ℓ⋅em−mnewℓnew+∑tePij,t−mnewVjℓnewO_{new} = O \cdot \frac{\ell \cdot e^{m – m_{new}}}{\ell_{new}} + \frac{\sum_t e^{P_{ij,t} – m_{new}} V_j}{\ell_{new}}

3.2.3 前向流程

Card 0 (S_0)     Card 1 (S_1)     ...     Card N-1 (S_{N-1})
    |                 |                     |
    v                 v                     v
  Q0,K0,V0          Q1,K1,V1          Q_{N-1},K_{N-1},V_{N-1}

Step 0: 每卡用本地 K,V 计算本地 Attention 部分
         |--- P2P Send K_i,V_i -> Card (i+1)%N --->

Step 1: 每卡收到 Card (i-1) 的 K,V, 计算与本地 Q 的 Attention
        同时 P2P Send 当前 K,V 到下一卡

Step 2 ~ N-1: 重复上述过程, 直到每卡都见过所有 K,V

最终: 每卡的 O_i 包含全序列信息

通信细节:

  • 通信组:CP 组内 NcpN_{cp} 张卡
  • 通信方式:P2P Send-Recv(环形),共 Ncp−1N_{cp}-1 轮
  • 每轮通信数据量:发送 KjK_j 和 VjV_j,各 B×Slocal×H×2B \times S_{local} \times H \times 2 bytes
  • 单轮通信量:2×B×Slocal×H×2=2×2×4096×8192×2=268,435,4562 \times B \times S_{local} \times H \times 2 = 2 \times 2 \times 4096 \times 8192 \times 2 = 268{,}435{,}456 bytes = 256 MiB
  • 总通信量(前向):(Ncp−1)×256(N_{cp}-1) \times 256 MiB
  • Ncp=2N_{cp}=2:256 MiB
  • Ncp=4N_{cp}=4:768 MiB

3.2.4 反向流程

反向传播需计算 Q,K,VQ, K, V 的梯度。由于前向 K,VK, V 在卡间循环,反向需相应的梯度回传。

反向传播(Ring-Attention 反向):

Card i 持有 dO_i (来自上一层的梯度) [B, S_local, H]

Step 0: 用本地 Q_i, K_i, V_i 计算本地部分的 dQ_i, dK_i, dV_i
        |--- P2P Send dK_j, dV_j -> 前一卡(逆向环) --->

Step 1 ~ N-1:
  每卡收到来自 "下一卡" 的 dK_j, dV_j
  用本地 Q_i 和缓存的 K_j, V_j 计算:
    dQ_i  += dP_{ij}^T @ K_j  (累加到本地 dQ)
    dK_j  += dP_{ij}^T @ Q_i  (需传回)
    dV_j  += softmax(P_{ij})^T @ dO_i  (需传回)

  P2P Send dK_j, dV_j 回前一卡

反向通信量:与前向相同,(Ncp−1)×256(N_{cp}-1) \times 256 MiB

总通信量(前向 + 反向):2×(Ncp−1)×2562 \times (N_{cp}-1) \times 256 MiB

  • Ncp=2N_{cp}=2:512 MiB

3.2.5 优化手段

  1. 计算 - 通信重叠:Ring-Attention 的核心优势——计算当前块时预取下一块
  2. Flash-Attention 集成:将分块计算融合为 Flash-Attention kernel,减少 HBM 读写
  3. KV 缓存复用:前向时缓存收到的 K,VK, V,避免反向时重新通信(代价是显存增加)

3.3 变体二:AllGather 方案

3.3.1 原理

AllGather 方案更直接:每卡持有 SlocalS_{local} 的 Q,K,VQ, K, V,通过 All-Gather 收集全序列的 K,VK, V,然后各卡独立计算完整的 Attention。

3.3.2 数学推导

每卡 i 持有 Qi,Ki,Vi∈[B,Slocal,H]Q_i, K_i, V_i \in [B, S_{local}, H]。

All-Gather K, V:
Kfull=AllGather(K0,K1,…,KNcp−1)∈[B,S,H]K_{full} = \text{AllGather}(K_0, K_1, \ldots, K_{N_{cp}-1}) \in [B, S, H]
Vfull=AllGather(V0,V1,…,VNcp−1)∈[B,S,H]V_{full} = \text{AllGather}(V_0, V_1, \ldots, V_{N_{cp}-1}) \in [B, S, H]

计算完整 Attention:
Oi=softmax(QiKfullTdh)Vfull∈[B,Slocal,H]O_i = \text{softmax}\left(\frac{Q_i K_{full}^T}{\sqrt{d_h}}\right) V_{full} \in [B, S_{local}, H]

每卡只需计算自己 SlocalS_{local} 行的 Attention,但需全序列的 K,VK, V。

3.3.3 前向流程

Card 0 (S_0)       Card 1 (S_1)      ...    Card N-1 (S_{N-1})
    |                   |                      |
    v                   v                      v
  Q0,K0,V0            Q1,K1,V1             Q_{N-1},K_{N-1},V_{N-1}

Step 1: AllGather(K) -> 每卡获得 K_full [B, S, H]
        AllGather(V) -> 每卡获得 V_full [B, S, H]

        通信量: 每卡发送 B*S_local*H*2 bytes (K), 接收 B*S*H*2 bytes (K)
        同理 V -> 总通信量 = 2 * B * S * H * 2 = 512 MiB

Step 2: 每卡独立计算 O_i = Attn(Q_i, K_full, V_full)
        Attention 矩阵: [B, N_h, S_local, S]

Step 3: 输出 O_i [B, S_local, H], 过 W_o 投影(各卡独立)

通信细节:

  • 通信组:CP 组内 NcpN_{cp} 张卡
  • 通信方式:All-Gather(收集 K 和 V)
  • 总通信量(前向 K+V):2×B×S×H×2=2×2×8192×8192×2=536,870,9122 \times B \times S \times H \times 2 = 2 \times 2 \times 8192 \times 8192 \times 2 = 536{,}870{,}912 bytes = 512 MiB

3.3.4 反向流程

反向传播:

Card i 持有 dO_i [B, S_local, H]

Step 1: 用前向缓存的 K_full, V_full 计算:
  dQ_i     = dP_i @ V_full         [B, S_local, H]
  dK_full  = dP_i^T @ Q_i          [B, S, H] (每卡只算了 S_local 行的梯度)
  dV_full  = softmax(P_i)^T @ dO_i [B, S, H]

Step 2: dK_full 和 dV_full 需在 CP 组内 Reduce-Scatter:
  - Reduce-Scatter(dK_full) -> 每卡得到 dK_i [B, S_local, H]
  - Reduce-Scatter(dV_full) -> 每卡得到 dV_i [B, S_local, H]

  通信量: 2 * B * S * H * 2 = 512 MiB

总通信量(前向 + 反向):2×512=10242 \times 512 = 1024 MiB

比 Ring-Attention(512 MiB)多一倍,但实现更简单,计算效率更高(不需要在线 softmax)。

3.3.5 优化手段

  1. 只 AllGather K,V,不 AllGather Q:减少通信量
  2. Flash-Attention:用 Flash-Attention kernel 计算 [Slocal×S][S_{local} \times S] 的 Attention
  3. QKV 融合 AllGather:一次 All-Gather 同时收集 K 和 V(拼接为 [K;V][K;V])
  4. 滑动窗口 Attention:局部注意力只需 AllGather 窗口内的 K,V

3.4 变体三:Ulysses All-to-All

3.4.1 原理

Ulysses 采用与上述两种方案完全不同的策略:通过 All-to-All 通信在「序列切分」和「注意力头切分」之间切换。

  • 前向前:每卡持有 SlocalS_{local} 的所有头的 Q,K,VQ, K, V
  • All-to-All 后:每卡持有 SS(全序列)的 Nh/NcpN_h / N_{cp} 个头的 Q,K,VQ, K, V
  • 每卡独立计算自己负责的那些头的完整 Attention(序列维度完整,头维度切分)
  • 计算完成后:再次 All-to-All 切回序列切分模式

3.4.2 数学推导

初始状态(序列切分):
Qi∈[B,Slocal,Nh,dh](卡  i  持有序列段  i  的所有头)Q_i \in [B, S_{local}, N_h, d_h] \quad \text{(卡} i \text{持有序列段} i \text{的所有头)}

All-to-All 重排:
Q′i∈[B,S,Nhlocal,dh](卡  i  持有全序列的  Nhlocal=Nh/Ncp  个头)Q’i \in [B, S, N_h^{local}, d_h] \quad \text{(卡} i \text{持有全序列的} N_h^{local} = N_h / N{cp} \text{个头)}

这是在序列维度和头维度之间的数据重排。

计算 Attention:
Oi′=softmax(Qi′Ki′Tdh)Vi′∈[B,S,Nhlocal,dh]O’_i = \text{softmax}\left(\frac{Q’_i K’^{T}_i}{\sqrt{d_h}}\right) V’_i \in [B, S, N_h^{local}, d_h]

第二次 All-to-All(切回):
Oi∈[B,Slocal,Nh,dh](恢复为序列切分)O_i \in [B, S_{local}, N_h, d_h] \quad \text{(恢复为序列切分)}

3.4.3 前向流程

Card 0               Card 1              ...    Card N-1
  |                    |                          |
  v                    v                          v
Q0,K0,V0             Q1,K1,V1                Q_{N-1},K_{N-1},V_{N-1}
[B, S_local, N_h, d_h]                       [B, S_local, N_h, d_h]

Step 1: All-to-All (Q)  -> 重排: 序列切分 -> 头切分
  Card 0 获得 Q'_0 = [B, S, N_h/N_cp, d_h]  (全序列, 头 0~31)
  Card 1 获得 Q'_1 = [B, S, N_h/N_cp, d_h]  (全序列, 头 32~63)

  通信量: 每卡发送 (N_cp-1)/N_cp * B * S_local * N_h * d_h * 2
        = (1/2) * 2 * 4096 * 64 * 128 * 2 = 33,554,432 bytes = 32 MiB

  同理对 K, V 各做一次 All-to-All
  总 All-to-All 通信量(Q+K+V) = 3 * 32 = 96 MiB

Step 2: 每卡独立计算完整 Attention (自己负责的头)
  O'_i = Attn(Q'_i, K'_i, V'_i)  [B, S, N_h_local, d_h]
  Attention 矩阵: [B, N_h_local, S, S]  (完全本地计算, 无需跨卡)

Step 3: All-to-All (O) -> 重排: 头切分 -> 序列切分
  Card 0 获得 O_0 = [B, S_local, N_h, d_h]  (序列段 0, 所有头)
  通信量: 同 Step 1, 32 MiB

Step 4: 过 W_o 投影 (每卡独立, 无通信)

总通信量(前向):4×32=1284 \times 32 = 128 MiB(3 次 All-to-All for Q,K,V + 1 次 for O)

3.4.4 反向流程

反向传播:

Card i 持有 dO_i [B, S_local, N_h, d_h]

Step 1: All-to-All(dO) -> 头切分模式
  dO'_i = [B, S, N_h_local, d_h]
  通信量: 32 MiB

Step 2: 计算本地 Attention 的梯度
  dQ'_i, dK'_i, dV'_i = Attn_backward(dO'_i, Q'_i, K'_i, V'_i)
  完全本地计算, 无通信

Step 3: All-to-All(dQ), All-to-All(dK), All-to-All(dV) -> 序列切分模式
  通信量: 3 * 32 = 96 MiB

Step 4: W_o 的反向 (每卡独立)

总通信量(反向):4×32=1284 \times 32 = 128 MiB

总通信量(前向 + 反向):128+128=256128 + 128 = 256 MiB

3.4.5 优化手段

  1. QKV 融合 All-to-All:将 Q,K,VQ, K, V 拼接后一次 All-to-All,减少通信次数从 3 到 1
  2. 标准 Flash-Attention:每卡持有全序列,可直接用标准 Flash-Attention kernel

3.4.6 三种 CP 变体对比

维度 Ring-Attention AllGather Ulysses All-to-All
通信方式 P2P 环形 Send-Recv All-Gather + Reduce-Scatter All-to-All
前向通信量 (Ncp=2N_{cp}=2) 256 MiB 512 MiB 128 MiB
反向通信量 (Ncp=2N_{cp}=2) 256 MiB 512 MiB 128 MiB
总通信量 512 MiB 1024 MiB 256 MiB
计算效率 中(在线 softmax 开销) 高(标准 Flash-Attn) 高(标准 Flash-Attn)
通信 - 计算重叠 强(核心设计) 弱 弱
实现复杂度 高 低 中
扩展性 好(P2P 带宽充分利用) 中(All-Gather 瓶颈) 好(All-to-All 对大组友好)
额外显存 低 高(全序列 K,V) 中(全序列但部分头)

4. Expert Parallelism (EP)

4.1 原理

Expert Parallelism(专家并行)是 MoE(Mixture of Experts)模型专用的并行策略。将多个专家(Expert)分布到不同卡上,每张卡只持有部分专家的参数。路由器(Gate/Router)根据输入 token 决定将 token 发送到哪些卡上的专家进行计算。

MoE 层结构:

  • Router/Gate:Wgate∈[H,E]W_{gate} \in [H, E],计算每个 token 对每个专家的分数
  • Top-k 选择:每 token 选 k 个专家
  • Expert FFN:E 个独立的 SwiGLU MLP,每个 Wgatee∈[H,4H],Wupe∈[H,4H],Wdowne∈[4H,H]W_{gate}^e \in [H, 4H], W_{up}^e \in [H, 4H], W_{down}^e \in [4H, H]

4.2 数学推导

Gate 计算:
G(x)=TopK(softmax(x⋅Wgate))G(x) = \text{TopK}(\text{softmax}(x \cdot W_{gate}))

其中 x∈[B⋅S,H]x \in [B \cdot S, H],G(x)∈[B⋅S,k]G(x) \in [B \cdot S, k](每 token 选 k 个专家)。

Expert 计算与加权:
y=∑e∈TopK(G(x))G(x)e⋅FFNe(x)y = \sum_{e \in \text{TopK}(G(x))} G(x)_e \cdot \text{FFN}_e(x)

EP 切分: E 个专家分布在 NepN_{ep} 张卡上,每卡持有 E/NepE / N_{ep} 个专家。

4.3 切分内容精确描述

  • 参数切分:
    • Router WgateW_{gate}:不切分,每卡完整持有(所有卡都需要为本地 token 计算路由)
    • Expert Wgatee,Wupe,WdowneW_{gate}^e, W_{up}^e, W_{down}^e:按专家切分。卡 i 持有专家 e∈[i⋅E/Nep,(i+1)⋅E/Nep)e \in [i \cdot E/N_{ep}, (i+1) \cdot E/N_{ep}) 的 Wgatee,Wupe,WdowneW_{gate}^e, W_{up}^e, W_{down}^e
    • 非 MoE 层(Attention、RMSNorm):不受 EP 影响
  • 梯度切分:与参数对应,每个专家的梯度只在持有该专家的卡上计算
  • 优化器状态:与参数对应
  • 激活值切分(精确到层内位置):
    • Router 输入 X∈[B,S,H]X \in [B, S, H]:每卡持有完整序列(由上游 DP/TP/CP 决定,EP 不额外切分序列维度)
    • Router 输出(dispatch 前):每卡为所有 token 计算路由分数
    • Expert 输入激活值 :被路由到卡 i 的 token 子集 X~i\tilde{X}_i,这是一个 不规则的、动态的 激活值分块(不是沿固定维度切分,而是 token 的动态子集)
    • Expert 输出激活值:Y~i=FFNe(X~i)\tilde{Y}_i = \text{FFN}_e(\tilde{X}_i),需合并回原位置

4.4 通信组构成

  • EP 通信组:NepN_{ep} 张卡,持有 E 个专家
  • 跨 EP 组的通信通过 All-to-All 实现 token dispatch 和 combine

4.5 通信方式

前向(2 次 All-to-All):

  1. Dispatch All-to-All:每卡将本地 token 按路由结果发送到目标专家所在卡
  2. Combine All-to-All:专家计算完成后,将结果按逆路由发回原卡

反向(2 次 All-to-All,与前向对称但方向互换):

  1. Dispatch 反向 = Combine:将 loss 对 expert 输出的梯度 dYedY_e 从原卡发送到对应专家所在卡(与前向 Combine 方向相同,传的是梯度而非激活值)
  2. Combine 反向 = Dispatch:expert 计算完输入梯度 dXdX 后,将其按逆路由发回原卡(与前向 Dispatch 方向相同,传的是梯度而非激活值)

4.6 通信数据量计算

Dispatch 阶段(均匀路由假设):

  • 每 token 发送 k×H×2k \times H \times 2 bytes 到目标卡
  • 总 token 数 = B×S=2×8192=16384B \times S = 2 \times 8192 = 16384
  • 均匀路由下,每卡发送给 (Nep−1)/Nep(N_{ep}-1)/N_{ep} 的 dispatch
每卡发送量=Nep−1Nep×B×S×k×H×2\text{每卡发送量} = \frac{N_{ep}-1}{N_{ep}} \times B \times S \times k \times H \times 2

数值计算(B=2,S=8192,H=8192,k=2,Nep=8B=2, S=8192, H=8192, k=2, N_{ep}=8):

=78×2×8192×2×8192×2=78×536,870,912=469,762,048 bytes≈448 MiB= \frac{7}{8} \times 2 \times 8192 \times 2 \times 8192 \times 2 = \frac{7}{8} \times 536{,}870{,}912 = 469{,}762{,}048 \text{bytes} \approx 448 \text{MiB}

Combine 阶段: 与 Dispatch 对称,每卡发送 ≈448\approx 448 MiB

总通信量(前向):2×448=8962 \times 448 = 896 MiB

反向通信量:与前向对称(Dispatch 和 Combine 各一次)

总通信量(前向 + 反向):2×896=1792MiB≈1.752 \times 896 = 1792 MiB \approx 1.75 GiB

4.7 前向与反向流程

前向流程

所有 EP 卡持有: X [B, S, H] (完整序列)
                W_gate [H, E] (完整, 所有卡都有)

Step 1: Router 计算 (每卡独立)
  gate_scores = softmax(X @ W_gate)  [B*S, E]
  top_k_indices, top_k_weights = TopK(gate_scores, k=2)  [B*S, 2]

Step 2: Dispatch All-to-All
  每卡根据 top_k_indices 将 token 发送到对应专家所在卡
  通信: All-to-All, 发送 ~448 MiB, 接收 ~448 MiB

  通信后: 卡 i 持有被路由到本地专家的 token 子集 X_i'
  (均匀假设下约 B*S*k/N_ep = 16384*2/8 = 4096 个 token)

Step 3: Expert 计算 (每卡独立)
  卡 i 对本地每个专家 e 计算: Y_e = W_{down}^e @ (SiLU(W_{gate}^e @ X_i'_e) * (W_{up}^e @ X_i'_e))

Step 4: Combine All-to-All (逆路由)
  将专家输出发回原始卡
  通信: All-to-All, 发送 ~448 MiB, 接收 ~448 MiB

  通信后: 每卡持有完整 Y [B, S, H], 按 top_k_weights 加权合并

Step 5: 加权合并 (每卡独立)
  Y_final = sum_e(top_k_weights[e] * Y_e)  逐 token 加权

反向流程

反向传播:

每卡持有 dY [B, S, H]

Step 1: 加权合并的反向
  d(top_k_weights[e] * Y_e) = top_k_weights[e] * dY  (逐 token)
  每卡需对每个 token 的每个被选专家产生 dY_e

Step 2: Combine All-to-All 的反向 = Dispatch All-to-All
  将 dY_e 发送到对应专家所在卡
  通信: All-to-All, ~448 MiB

Step 3: Expert 反向 (每卡独立)
  卡 i 对本地专家计算:
    dW_{gate}^e, dW_{up}^e, dW_{down}^e (参数梯度)
    dX_i'_e (输入梯度)

Step 4: Dispatch All-to-All 的反向 = Combine All-to-All
  将 dX_i'_e 发回原始卡
  通信: All-to-All, ~448 MiB

Step 5: Router 反向
  每卡计算 dW_gate, 累加到梯度

Step 6: EP 组内梯度同步
  dW_gate 在所有 EP 卡上 All-Reduce (因为 W_gate 是复制的)
  专家参数梯度不需 All-Reduce (各卡只持有自己的专家)

4.8 优化手段

  1. Capacity Factor:限制每个专家处理的 token 数上限 C=cf×B⋅S/NepC = \text{cf} \times B \cdot S / N_{ep},防止负载不均
  2. Token Dropping:超出容量的 token 被丢弃(或通过残差连接直接传递)
  3. Expert Balancing Loss:添加辅助损失鼓励均匀路由
  4. Grouped GEMM:将同一卡多个专家的计算合并为一次 Grouped GEMM
  5. 通信 - 计算重叠:Dispatch 通信可与 Router 计算重叠
  6. DeepEP/Expert-EP 优化:使用专用 All-to-All kernel 减少小消息延迟

4.9 专家负载均衡

MoE 的核心痛点:Router 可能将大量 token 路由到少数专家(hot expert),导致部分 EP 卡过载而其他卡空闲。这不只是性能问题——训练不稳定、梯度爆炸、专家退化(某些专家永远不被训练)都可能由此引发。

4.9.1 问题根源

Router 计算 G(x)=TopK(softmax(x⋅Wgate))G(x) = \text{TopK}(\text{softmax}(x \cdot W_{gate})),每个 token 独立选择 top-k 专家。在训练初期 WgateW_{gate} 随机初始化,路由分布天然不均匀;训练过程中还可能出现 马太效应——被选中的专家获得更多梯度更新,变得更强,吸引更多 token,形成正反馈。

负载不均的三个层面:

层面 表现 影响
专家级 某些专家处理 5x token,某些几乎为 0 EP 卡计算不均
卡级 持有 hot expert 的卡成为瓶颈 All-to-All 延迟
训练级 Cold expert 不被训练,参数停滞 模型质量下降

4.9.2 负载均衡 Loss(Auxiliary Loss)

最核心的方法。 添加辅助损失项,鼓励 token 在专家间均匀分布。

Switch Transformer 的负载均衡 Loss:

设 N=B⋅SN = B \cdot S 为总 token 数,EE 为专家数,fef_e 为实际路由到专家 ee 的 token 比例,PeP_e 为 Router 对专家 ee 的平均概率:

fe=1N∑i=1N𝟙[e∈TopK(G(xi))]f_e = \frac{1}{N} \sum_{i=1}^{N} \mathbb{1}[e \in \text{TopK}(G(x_i))]
Pe=1N∑i=1Nsoftmax(xi⋅Wgate)eP_e = \frac{1}{N} \sum_{i=1}^{N} \text{softmax}(x_i \cdot W_{gate})_e
ℒbalance=α⋅E∑e=1Efe⋅Pe\mathcal{L}{balance} = \alpha \cdot E \sum{e=1}^{E} f_e \cdot P_e
  • 当路由完全均匀时:fe=k/E,Pe=1/E,ℒbalance=α⋅kf_e = k/E, P_e = 1/E,\mathcal{L}_{balance} = \alpha \cdot k(最小值)
  • 当路由完全不均匀时:ℒbalance\mathcal{L}_{balance} 增大
  • α\alpha 是超参(通常 0.01),过大会影响主任务

直觉: fe⋅Pef_e \cdot P_e 在均匀分布时最小(柯西 - 施瓦茨不等式),这个 loss 惩罚的是“高概率 + 高选中率”的专家。

GShard 的 z-loss:

额外添加 z-loss 鼓励 logits 的绝对值不要过大,防止 softmax 饱和:

ℒz=αz⋅1N∑i=1N(log⁡∑e=1Eezi,e)2\mathcal{L}z = \alpha_z \cdot \frac{1}{N} \sum{i=1}^{N} \left(\log \sum_{e=1}^{E} e^{z_{i,e}} \right)^2

其中 zi,e=(xi⋅Wgate)ez_{i,e} = (x_i \cdot W_{gate})_e 是未归一化的 logits。

4.9.3 容量因子(Capacity Factor)

硬性限制 每个专家处理的 token 数上限,防止极端不均。

Ce=cf×N⋅kEC_e = \text{cf} \times \frac{N \cdot k}{E}

其中 cf\text{cf} 是容量因子(通常 1.0-1.5),N⋅k/EN \cdot k / E 是均匀路由下每专家的期望 token 数。

  • 超出 CeC_e 的 token 被 丢弃(通过残差连接直接传给下一层,不经过 expert)
  • cf=1.0:严格限制,但丢弃率高,信息损失大
  • cf=1.5:允许 50% 冗余,丢弃少,但显存增加
  • 训练时常从 cf=1.25 退火到 1.0

显存影响: 容量因子直接决定 expert 输入 buffer 大小:

buffer size=Ce×H=cf×N⋅kE×H\text{buffer size} = C_e \times H = \text{cf} \times \frac{N \cdot k}{E} \times H

4.9.4 Token Dropping 策略

超出容量的 token 处理方式:

策略 做法 优缺点
直接丢弃 超出的 token 不计算 FFN 最简单,但信息损失
残差直通 y=xy = x(跳过该层的 MoE FFN) 保持信号传播,主流做法
随机丢弃 从超出的 token 中随机选 CeC_e 个 减少系统性偏倚
优先级丢弃 按 Router 分数排序,保留高分 token 低分 token 被丢弃,可能偏倚

Megatron-LM 默认使用 残差直通——被丢弃的 token 直接通过 y=x+0y = x + 0(不经过任何 expert)。

4.9.5 Expert Bias / Router Noise

Router Noise(训练初期): 在 Router logits 上加高斯噪声,鼓励探索:

G(x)=TopK(softmax(x⋅Wgate+ϵ)),ϵ∼𝒩(0,σ2)G(x) = \text{TopK}(\text{softmax}(x \cdot W_{gate} + \epsilon)), \quad \epsilon \sim \mathcal{N}(0, \sigma^2)

σ\sigma 通常设为 1.0,随训练退火到 0。这防止 Router 早期固化到某些专家。

Expert Bias: 给每个专家添加可学习偏置 beb_e,用于补偿路由不均:

G(x)=TopK(softmax(x⋅Wgate+b))G(x) = \text{TopK}(\text{softmax}(x \cdot W_{gate} + b))

冷专家的 beb_e 会被梯度推向更大值,自动吸引更多 token。类似于 MaxViT 的“专家拉力”机制。

4.9.6 路由方案对比

路由方案 均衡机制 通信量 模型质量
Top-k 路由 + 负载均衡 Loss 软约束(Loss 惩罚) 标准 好
Top-k 路由 + 容量因子 硬约束(丢弃超限 token) 降低(有上限) 中(信息丢失)
Expert Choice 路由 专家选 token(反向) 降低(自动均衡) 中高
Sparse MLP 路由 Top-1 + 量化 最低 中
Hash 路由 固定哈希函数 无(完全确定) 差

Expert Choice 路由(Gale et al. 2023)反转了路由方向:每个专家从所有 token 中选 top-CeC_e 个,天然保证负载均衡(每个专家恰好处理 CeC_e 个 token)。代价是 token 可能被多个专家处理或被零个处理。

4.9.7 Megatron-LM 的实际做法

Megatron-LM 默认组合使用:

  1. 负载均衡 Loss(α=0.01\alpha=0.01)——软约束,鼓励均匀
  2. 容量因子 cf=1.0\text{cf}=1.0 ——硬约束,防止极端不均
  3. 残差直通——超出容量的 token 跳过 MoE 层
  4. Router Noise 退火——训练初期探索
  5. z-loss——防止 softmax 饱和

5. Data Parallelism with ZeRO (DP)

5.1 原理

标准 Data Parallelism(DP)中,每张卡持有完整的模型参数、梯度和优化器状态,各自处理不同数据 batch,前向后通过 All-Reduce 同步梯度。

ZeRO(Zero Redundancy Optimizer)逐步切分这些冗余副本:

阶段 切分内容 冗余消除
ZeRO-1 优化器状态(Adam 的 m, v, master weights) 优化器状态不再每卡复制
ZeRO-2 优化器状态 + 梯度 梯度也按卡切分
ZeRO-3 优化器状态 + 梯度 + 参数 参数也按卡切分(完全无冗余)

5.2 显存分析

设模型参数量为 Ψ\Psi,使用 Adam 优化器,bf16 训练 + fp32 优化器状态。

标准 DP 每卡显存:

  • 参数(bf16):2Ψ2\Psi bytes
  • 梯度(bf16):2Ψ2\Psi bytes
  • 优化器状态(fp32 m, v + fp32 master weights):12Ψ12\Psi bytes
  • 总计:16Ψ16\Psi bytes

ZeRO-1(切分优化器状态,$N_{dp}$ 卡):

  • 参数:2Ψ2\Psi(不切分)
  • 梯度:2Ψ2\Psi(不切分)
  • 优化器状态:12Ψ/Ndp12\Psi / N_{dp}
  • 总计:4Ψ+12Ψ/Ndp4\Psi + 12\Psi/N_{dp} bytes

ZeRO-2(切分优化器状态 + 梯度):

  • 参数:2Ψ2\Psi(不切分)
  • 梯度:2Ψ/Ndp2\Psi / N_{dp}
  • 优化器状态:12Ψ/Ndp12\Psi / N_{dp}
  • 总计:2Ψ+14Ψ/Ndp2\Psi + 14\Psi/N_{dp} bytes

ZeRO-3(切分优化器状态 + 梯度 + 参数):

  • 参数:2Ψ/Ndp2\Psi / N_{dp}
  • 梯度:2Ψ/Ndp2\Psi / N_{dp}
  • 优化器状态:12Ψ/Ndp12\Psi / N_{dp}
  • 总计:16Ψ/Ndp16\Psi / N_{dp} bytes

5.3 切分内容精确描述

ZeRO-1

  • 切分对象:优化器状态(Adam 的 m,vm, v, master weights)
  • 切分维度:按参数维度均匀切分到 NdpN_{dp} 张卡
  • 不切分:参数(每卡完整)、梯度(每卡完整)
  • 激活值:不涉及
  • 前向:标准前向传播,无额外通信
  • 反向:标准反向计算梯度后,All-Reduce 同步梯度(每卡获得完整梯度),然后各卡用本地那段优化器状态更新本地那段参数副本(注意:参数仍需 All-Reduce 同步更新后的参数)

ZeRO-2

  • 切分对象:优化器状态 + 梯度
  • 切分维度:按参数维度均匀切分
  • 不切分:参数(每卡完整)
  • 激活值:不涉及
  • 前向:标准前向传播,无额外通信
  • 反向:计算梯度后,用 Reduce-Scatter 将梯度按卡切分(每卡只保留自己那段的梯度),然后各卡用本地梯度更新本地优化器状态和参数。更新后,参数需通过 All-Gather 广播给其他卡

ZeRO-3

  • 切分对象:优化器状态 + 梯度 + 参数
  • 切分维度:按参数维度均匀切分
  • 激活值:不直接切分激活值,但参数不完整需按需收集
  • 前向:每层计算前,All-Gather 收集当前层完整参数;计算后丢弃(释放显存)
  • 反向:每层反向前,All-Gather 收集当前层完整参数;计算梯度后,Reduce-Scatter 将梯度按卡切分;更新本地参数段

5.4 通信量分析

阶段 前向通信 反向通信 总通信量
标准 DP 0 All-Reduce(梯度) = 2Ψ×22\Psi \times 2 4Ψ4\Psi
ZeRO-1 0 All-Reduce(梯度) = 4Ψ4\Psi 4Ψ4\Psi
ZeRO-2 0 Reduce-Scatter(梯度) + All-Gather(参数) = 2Ψ+2Ψ2\Psi + 2\Psi 4Ψ4\Psi
ZeRO-3 All-Gather(参数) 每层 = 2Ψ2\Psi Reduce-Scatter(梯度) + All-Gather(参数) = 2Ψ+2Ψ2\Psi + 2\Psi 6Ψ6\Psi

注意:ZeRO-2 和标准 DP 通信量相同(All-Reduce = All-Gather + Reduce-Scatter),只是拆分了时机。ZeRO-3 比标准 DP 多 2Ψ2\Psi 的前向 All-Gather 通信量。

数值示例(假设模型参数 Ψ=175B\Psi = 175B 参数,bf16):

  • ZeRO-3 每卡显存:16×175×109/Ndp16 \times 175 \times 10^9 / N_{dp} bytes
  • Ndp=4N_{dp}=4:700 GB / 4 = 175 GB/ 卡(仍然很大,需配合其他并行)
  • Ndp=32N_{dp}=32:700 GB / 32 = 21.9 GB/ 卡
  • ZeRO-3 前向 All-Gather 通信量:2×175×109×2=7002 \times 175 \times 10^9 \times 2 = 700 GB(全部参数每层收集一次,总计)
  • ZeRO-3 反向通信量:2×175×109×2=7002 \times 175 \times 10^9 \times 2 = 700 GB
  • ZeRO-3 总通信量:700+700=1400700 + 700 = 1400 GB

5.5 前向与反向流程

ZeRO-1 前向与反向

前向: 标准前向, 每卡用完整参数计算, 无额外通信

反向:
  Step 1: 每卡计算本地梯度 dW (完整)
  Step 2: All-Reduce(dW) -> 每卡获得完整梯度
          通信量: 2 * Psi * 2 = 4*Psi bytes
  Step 3: 每卡用本地优化器状态段更新对应参数段
  Step 4: All-Reduce(更新后的参数) -> 同步参数
          通信量: 2 * Psi * 2 = 4*Psi bytes

ZeRO-2 前向与反向

前向: 标准前向, 每卡用完整参数计算, 无额外通信

反向:
  Step 1: 每卡计算本地梯度 dW (完整)
  Step 2: Reduce-Scatter(dW) -> 每卡只持有自己那段的梯度
          通信量: 2 * Psi * 2 = 4*Psi bytes (等价于 All-Reduce)
  Step 3: 每卡用本地梯度段 + 本地优化器状态段 -> 更新本地参数段
  Step 4: All-Gather(更新后的参数) -> 广播给所有卡
          通信量: 2 * Psi * 2 = 4*Psi bytes

ZeRO-3 前向与反向

前向:
  对每一层 l (从 0 到 L-1):
    Step 1: All-Gather(W_l) -> 收集当前层完整参数
            通信量: 2 * Psi_per_layer * 2 (bf16)
            总计(全部层): 2 * Psi * 2 = 4*Psi bytes
    Step 2: 用完整参数计算前向 (标准计算)
    Step 3: 丢弃完整参数 (释放显存, 只保留本地段)

反向:
  对每一层 l (从 L-1 到 0):
    Step 1: All-Gather(W_l) -> 收集当前层完整参数
            通信量: 2 * Psi_per_layer * 2
    Step 2: 用完整参数计算本地梯度 dW_l
    Step 3: Reduce-Scatter(dW_l) -> 每卡只持有自己那段的梯度
            通信量: 2 * Psi_per_layer * 2
    Step 4: 丢弃完整参数
    Step 5: 每卡用本地梯度段 + 本地优化器状态段 -> 更新本地参数段

5.6 优化手段

  1. ZeRO-Infinity:将优化器状态 offload 到 CPU/NVMe,进一步降低显存
  2. ZeRO-R:切分激活值(与 ZeRO-3 的参数切分配合),减少激活值显存
  3. 预取重叠:在计算当前层时,预取下一层的参数(ZeRO-3)
  4. 梯度累积:减少 DP 通信频率(每 N 个 micro-batch 通信一次)
  5. 通信 - 计算重叠:All-Gather 和 Reduce-Scatter 可与计算重叠
  6. 分层通信:利用 NVLink(节点内)和 InfiniBand(节点间)的层次结构优化

5.7 梯度累积(Gradient Accumulation)

5.7.1 原理

梯度累积是一种通过 多次小 batch 前向 + 反向、累积梯度后统一更新 来模拟大 batch 训练的技术。它解决了显存不足以容纳大 batch 的问题,同时减少 DP 通信频率。

核心流程:

目标全局 batch size = B_global
每步 micro-batch size = B_micro
累积步数 N_acc = B_global / B_micro

for step = 0, 1, ..., N_acc-1:
    前向: loss_i = forward(micro_batch_i)
    反向: grad_i = backward(loss_i)  # 梯度累加到 .grad
    # 不做 All-Reduce, 不更新参数

# 累积 N_acc 步后:
grad_total = sum(grad_0, ..., grad_{N_acc-1})
All-Reduce(grad_total)  # DP 组内同步梯度
grad_total /= N_acc     # 除以累积步数 (平均梯度)
optimizer.step()        # 统一更新参数
optimizer.zero_grad()   # 清零梯度

5.7.2 数学等价性

设全局 batch 为 BglobalB_{global},拆为 NaccN_{acc} 个 micro-batch 每个 Bmicro=Bglobal/NaccB_{micro} = B_{global}/N_{acc}:

标准大  batch  梯度=1Bglobal∑i=1Bglobal∇ℒ(xi)\text{标准大 batch 梯度} = \frac{1}{B_{global}} \sum_{i=1}^{B_{global}} \nabla \mathcal{L}(x_i)
梯度累积=1Nacc∑s=1Nacc(1Bmicro∑j=1Bmicro∇ℒ(xj))\text{梯度累积} = \frac{1}{N_{acc}} \sum_{s=1}^{N_{acc}} \left(\frac{1}{B_{micro}} \sum_{j=1}^{B_{micro}} \nabla \mathcal{L}(x_j) \right)
=1Nacc⋅Bmicro∑i=1Bglobal∇ℒ(xi)=1Bglobal∑i=1Bglobal∇ℒ(xi)= \frac{1}{N_{acc} \cdot B_{micro}} \sum_{i=1}^{B_{global}} \nabla \mathcal{L}(x_i) = \frac{1}{B_{global}} \sum_{i=1}^{B_{global}} \nabla \mathcal{L}(x_i)

两者数学等价(前提:BatchNorm 等跨样本统计的层需特殊处理,但 LLM 用 RMSNorm/LayerNorm 不受影响)。

5.7.3 通信优化

无梯度累积时(每步 All-Reduce):

  • 每 step 一次 All-Reduce,通信量 =2×Ψ×2= 2 \times \Psi \times 2 bytes(bf16)

有梯度累积时(每 NaccN_{acc} 步 All-Reduce):

  • 每 NaccN_{acc} step 一次 All-Reduce,通信量相同但频率降低 NaccN_{acc} 倍
  • 等效效果:通信占比从 TcommTcomm+Tcomp降 低到Tcomm/NaccTcomm/Nacc+Tcomp\frac{T_{comm}}{T_{comm}+T_{comp}} \text 降低到 \frac{T_{comm}/N_{acc}}{T_{comm}/N_{acc}+T_{comp}}

5.7.4 与 PP 的协同

梯度累积与 Pipeline Parallelism 天然协同——PP 本身就使用 micro-batch:

  • PP 的 M 个 micro-batch 就是 M 步梯度累积
  • PP 中每个 micro-batch 做完反向后,梯度累积在本地
  • 所有 micro-batch 完成后,DP 组做一次 All-Reduce

Megatron-LM 的默认做法:

全局 batch = B_global
PP micro-batches = M (由 PP 调度决定)
DP 梯度累积 = N_acc = B_global / (M * B_micro * N_dp)

每个 PP stage 完成所有 M 个 micro-batch 的反向后:
  → DP 组 All-Reduce 梯度
  → optimizer.step()
  → zero_grad()

5.7.5 显存影响

梯度累积 不增加 峰值显存——每次只处理 1 个 micro-batch 的激活值,反向完成后梯度累加到 .grad(已有显存),激活值可释放。

但需注意:

  • PP 的 1F1B 调度中,stage 0 最多缓存 NppN_{pp} 份 micro-batch 的激活值
  • 梯度累积不减少这个缓存——它影响的是 optimizer 更新频率,不是 PP 激活值缓存

5.7.6 太大的问题

1. 收敛质量下降

梯度累积的数学等价性成立的前提是:loss 在 micro-batch 间是独立同分布的。但 NaccN_{acc}​​ 过大时:

  • 全局 batch size Bglobal=Nacc×Bmicro×NdpB_{global} = N_{acc} \times B_{micro} \times N_{dp}​​ 膨胀过大
  • 大 batch 训练的泛化 gap 增大——模型容易收敛到 sharp minima
  • 学习率需要相应放大(线性 / 平方根 scaling rule),但不总能补偿

2. 训练有效速度变慢

  • NaccN_{acc}越大 → optimizer.step() 频率越低 → 参数更新次数减少
  • 同样的 wall-clock 时间,梯度更新步数少了 NaccN_{acc}​​ 倍
  • 虽然每步的 batch 更大、梯度更稳,但 收敛所需的总 token 数可能增加

3. Pipeline Bubble 比例变化

PP 中 micro-batch 数 MMM 与梯度累积直接相关。MM越大 bubble 越小:

bubble ratio=Npp−1M+Npp−1\text{bubble ratio} = \frac{N_{pp} – 1}{M + N_{pp} – 1}

但 MM 也不能无限大——MM 受限于全局 batch size 和显存(stage 0 缓存 MM 份激活值)。

5.7.6 太小的问题

1. 通信占比高——每步都要 All-Reduce,通信没被计算掩盖

2. 显存不够 ——NaccN_{acc} 太小意味着 BmicroB_{micro}​​ 要大才能达到目标 batch size,激活值显存撑不住

3. PP bubble 大——MM 小时 pipeline 利用率低

5.7.7 实际选择

因素 倾向大的 NaccN_{acc}N​acc​​ 倾向小的 NaccN_{acc}N​acc​​
显存不够 ✅ 用小 micro-batch + 大累积
通信瓶颈 ✅ 降低 All-Reduce 频率
PP bubble ✅ 增加 micro-batch 数
泛化 gap ✅ 别让全局 batch 太大
收敛速度 ✅ 更频繁更新参数
BatchNorm 统计 ✅ 小 batch 统计更准(LLM 用 RMSNorm 不受影响)

经验值: 大模型训练通常全局 batch 在 4M-8M tokens(如 Llama 2 用 4M tokens/batch),通过 Bmicro×Ndp×Nacc×SB_{micro} \times N_{dp} \times N_{acc} \times S 凑出这个数。NaccN_{acc}​​ 一般在 4-32 之间,不是越大越好,而是 刚好够用 ——在显存允许的 BmicroB_{micro} 下,NaccN_{acc} 取能达到目标全局 batch 的最小值。


6. 综合实例:TP+SP+PP+CP+EP+DP/ZeRO 全并行配置

6.1 模型配置

参数 值 说明
LL 80 Transformer 层数
HH 8192 hidden size
NhN_h 64 注意力头数
dhd_h 128 每头维度
EE 64 专家数
top-k\text {top-k} 2 每 token 选 2 个专家
dtype bf16 2 bytes
架构 MoE Transformer 每 4 层有 1 层 MoE,共 20 个 MoE 层,60 个 dense Attention 层

6.2 并行配置

并行维度 度数 说明
TP 4 Tensor Parallel(含 SP=4)
PP 4 Pipeline Parallel(4 个 stage,每 stage 20 层)
CP 2 Context Parallel
EP 8 Expert Parallel(64 专家 / 8 = 8 专家 / 卡)
DP 4 Data Parallel(ZeRO-3)
总 GPU 1024 4×4×2×8×4=10244 \times 4 \times 2 \times 8 \times 4 = 1024

6.3 GPU 拓扑与通信组

全局 GPU 编号: rank = f(dp, ep, cp, pp, tp)

通信组:
├── TP 组: 4 张卡 (同一节点内, NVLink 互联)
│   例: rank 0,1,2,3 (dp=0, ep=0, cp=0, pp=0, tp=0~3)
│
├── DP 组: 4 张卡 (跨节点, InfiniBand)
│   例: rank 0, 256, 512, 768 (dp=0~3, ep=0, cp=0, pp=0, tp=0)
│   ZeRO-3 在此组内切分参数 / 梯度 / 优化器状态
│
├── EP 组: 8 张卡 (跨节点)
│   例: rank 0, 32, 64, 96, 128, 160, 192, 224
│   (dp=0, ep=0~7, cp=0, pp=0, tp=0)
│
├── CP 组: 2 张卡
│   例: rank 0, 16 (dp=0, ep=0, cp=0~1, pp=0, tp=0)
│
└── PP 组: 4 张卡 (跨节点)
    例: rank 0, 64, 128, 192 (dp=0, ep=0, cp=0, pp=0~3, tp=0)

6.4 数据准备

每个 DP 卡处理不同的 micro-batch。

  • 全局 batch = Bglobal=B×Ndp×Ncp=2×4×2=16B_{global} = B \times N_{dp} \times N_{cp} = 2 \times 4 \times 2 = 16
  • 序列长度 S=8192S = 8192
  • 每 GPU 实际处理:B=2,Scp_local=S/Ncp=4096B=2, S_{cp\_local} = S/N_{cp} = 4096

经过 TP/SP 切分后,每 GPU 持有的激活值(进入 CP 前):

  • SP 区域激活值:[B,Slocal_sp,H]=[2,2048,8192][B, S_{local\_sp}, H] = [2, 2048, 8192](RMSNorm 区域,SP 按 4 切,S/4=2048S/4=2048)
  • 但 CP 进一步切分序列:[B,Slocal_sp/Ncp,H][B, S_{local\_sp}/N_{cp}, H]…

重要:SP 和 CP 都切分序列维度。SP 在 TP 组内切分(Nsp=4N_{sp}=4),CP 在 CP 组内进一步切分。实际每卡序列长度:

Spergpu=SNsp×Ncp=81924×2=1024S_{per_gpu} = \frac{S}{N_{sp} \times N_{cp}} = \frac{8192}{4 \times 2} = 1024

但实际实现中 SP 和 CP 的切分方式不同:

  • SP:在 RMSNorm 区域切分,在 Attention/MLP 的核心计算区域通过 All-Gather 恢复完整序列
  • CP:在整个 Attention 计算期间保持序列切分

这里我们采用 SP 切分后 Ssp_local=S/Nsp=2048S_{sp\_local} = S/N_{sp} = 2048,CP 在此基础上进一步切分 Scp_local=Ssp_local/Ncp=1024S_{cp\_local} = S_{sp\_local}/N_{cp} = 1024。

6.5 前向传播全过程

以下描述一个 micro-batch 从输入到输出的完整前向流程,聚焦于 GPU rank 0(dp=0, ep=0, cp=0, pp=0, tp=0)的视角。

========== 前向传播开始 ==========

--- Stage 0 (Layer 1~20), PP Stage 0 ---

输入: X [B=2, S_local=2048 (SP 后), H=8192]
  实际上 CP 进一步切分: X_cp [B=2, S_cp=1024, H=8192]

对每一层 l (1~20):

  === [1] RMSNorm 前置 ===
  - SP 区域: 每卡持有 [B=2, S_sp=2048, H=8192]
    但 CP 再切: 实际每卡 [B=2, S=1024, H=8192]
  - RMSNorm 在序列维度独立, 各卡独立计算
  - All-Gather(SP 组): 恢复完整序列 [B=2, 2048, 8192]
    通信量: 3 * B * S_sp/N_sp * H * 2 = 3 * 2 * 512 * 8192 * 2 = 50,331,648 bytes = 48 MiB
    (SP All-Gather, 4 张 TP 卡参与)

  === [2] QKV 投影 (TP) ===
  - TP 组内: W_q, W_k, W_v 按列切分, 每卡持有 1/4 的头 (16 个头)
  - 每卡计算: Q_i, K_i, V_i = X @ W_qkv_i  [B=2, 2048, 2048]
    (2048 = 16 头 * 128 维)
  - 无通信 (TP 区域内各卡独立计算自己那份)

  === [3] CP: 序列切分 Attention 计算 ===
  - CP 将 S=2048 切为 2 段, 每卡 S_cp=1024
  - 以 Ulysses 方案为例:

    Step 3a: All-to-All (Q, K, V) [CP 组, 2 张卡]
      Q: [B=2, S_cp=1024, N_h=64, d_h=128] -> [B=2, S=2048, N_h=32, d_h=128]
      通信量: 3 * (1/2) * B * S_cp * N_h * d_h * 2 = 3 * 0.5 * 2 * 1024 * 64 * 128 * 2
             = 3 * 16,777,216 = 50,331,648 bytes = 48 MiB

    Step 3b: Attention 计算 (完全本地)
      O'= softmax(Q' K'^T / sqrt(d_h)) V'  [B=2, S=2048, 32, 128]
      Attention 矩阵: [B=2, 32, 2048, 2048]

    Step 3c: All-to-All (O') [CP 组]
      O': [B=2, S=2048, 32, 128] -> [B=2, S_cp=1024, 64, 128]
      通信量: 16,777,216 bytes = 16 MiB

    CP 总通信量: 48 + 16 = 64 MiB

  === [4] W_o 投影 (TP) ===
  - W_o 按行切分, 每卡计算部分输出
  - Reduce-Scatter(TP 组): 合并 + 切分(SP 模式)
    通信量: B * S_sp * H * 2 = 2 * 2048 * 8192 * 2 = 67,108,864 bytes = 64 MiB
    输出: [B=2, S_sp/N_sp=512, H=8192] (SP 区域的序列切分)

  === [5] 残差连接 ===
  - Y = X + SP 的输出 (需要 All-Gather 恢复完整序列做残差)
  - 或者在 SP 模式下直接在切分状态做残差

  --- 如果是 MoE 层 (每 4 层 1 次) ---

  === [6a] Router 计算 (EP) ===
  - 每卡持有完整 W_gate
  - gate_scores = softmax(X @ W_gate) [B*S_cp_local, E=64]
  - top_k = TopK(gate_scores, k=2) [B*S_cp_local, 2]

  === [6b] Dispatch All-to-All (EP 组, 8 张卡) ===
  - 将 token 按 top_k 发送到对应专家卡
  - 通信量: (7/8) * B * S_cp_local * k * H * 2
          = (7/8) * 2 * 1024 * 2 * 8192 * 2 = 58,720,256 bytes = 56 MiB
  - 接收: 被路由到本地 8 个专家的 token

  === [6c] Expert FFN 计算 ===
  - 卡 i 持有 8 个专家 (64/8)
  - 对每个接收到的 token, 用对应专家计算:
    Y = W_{down}^e @ (SiLU(W_{gate}^e @ X) * (W_{up}^e @ X))
    W_{gate}^e: [H=8192, 4H=32768], W_{up}^e: [H=8192, 4H=32768], W_{down}^e: [4H=32768, H=8192]

  === [6d] Combine All-to-All (EP 组) ===
  - 将结果发回原卡
  - 通信量: 56 MiB
  - 加权合并: Y_final = sum(gate_weight * Y_e)

  --- 如果是 Dense 层 ---

  === [6'] MLP 计算 (TP) ===
  - W_gate 按列切分: 每卡 [H, 4H/4] = [8192, 8192]
  - W_up 按列切分: 每卡 [H, 4H/4] = [8192, 8192]
  - W_down 按行切分: 每卡 [4H/4, H] = [8192, 8192]
  - 前向: Z = (SiLU(X @ W_gate_i) * (X @ W_up_i)) @ W_down_i
  - Reduce-Scatter(TP 组): 合并 + SP 切分
    通信量: B * S_sp * H * 2 = 64 MiB

  === [7] 残差连接 + RMSNorm 后置 ===
  - 类似前述

--- Stage 间通信 (PP) ---

Layer 20 输出 -> P2P Send 到 Stage 1 (Layer 21~40)
  通信量: B * S_sp/N_sp * H * 2 = 2 * 512 * 8192 * 2 = 16,777,216 bytes = 16 MiB
  (SP 切分后的激活值大小)

--- Stage 1~3 (Layer 21~80) ---
  重复上述流程, 每个 stage 处理 20 层

--- 最终输出 ---
  Stage 3 输出 -> Loss 计算 (每卡独立计算)
  Cross-Entropy Loss: 需要 logits [B, S, vocab_size]
  vocab_size 可能很大, 通常也做 TP 切分

========== 前向传播结束 ==========

6.6 反向传播全过程

========== 反向传播开始 ==========

--- Stage 3 (Layer 80~61), 从最后 stage 开始 ---

Loss -> dLoss/dLogits
  每卡计算本地 logits 的梯度

对每一层 l (从 80 到 61):

  === [1] 反向 RMSNorm 后置 ===
  - 各卡独立计算 (序列维度独立)

  === [2] MLP/MoE 反向 ===

  Dense 层:
  - Reduce-Scatter 的反向 = All-Gather(TP 组)
    通信量: 64 MiB
  - 计算 dW_gate, dW_up, dW_down (每卡对自己那份参数的梯度)
  - dX = dZ @ W_down^T (传播给前一层)

  MoE 层:
  - Combine All-to-All 的反向 = Dispatch All-to-All
    通信量: 56 MiB
  - Expert 反向: dW_{gate}^e, dW_{up}^e, dW_{down}^e, dX_i'
  - Dispatch All-to-All 的反向 = Combine All-to-All
    通信量: 56 MiB
  - Router 反向: dW_gate

  === [3] Attention 反向 ===
  - W_o 反向: All-Gather(TP 组) 的反向
    通信量: 64 MiB

  - CP 反向 (Ulysses 方案):
    All-to-All(dO) -> 头切分
      通信量: 16 MiB
    Attention 反向计算 (本地)
    All-to-All(dQ, dK, dV) -> 序列切分
      通信量: 48 MiB

  - QKV 反向: 各卡独立 (TP 区域内)

  === [4] RMSNorm 前置反向 ===
  - All-Gather 的反向 = Reduce-Scatter(TP 组)
    通信量: 48 MiB

--- PP Stage 间反向通信 ---

  P2P Send: dX (梯度) -> Stage 2 (Layer 60~41)
  通信量: 16 MiB

--- Stage 2~0 (Layer 60~1) ---
  重复上述反向流程

--- DP/ZeRO-3 梯度同步 ---

  在所有 PP stage 反向完成后:

  Step 1: 每层的梯度需要 Reduce-Scatter(DP 组)
    通信量(每层): 2 * Psi_per_layer * 2 (bf16)
    总计: 2 * (全部参数) * 2 = 4 * Psi bytes

  Step 2: All-Gather(DP 组) 更新后的参数
    通信量: 2 * Psi * 2 = 4 * Psi bytes

  ZeRO-3 总通信量: 8 * Psi bytes
  (参数量 ~175B params for this model size:
   约 80 * 4 * 8192^2 + MoE params ~ 几百 B 参数)

========== 反向传播结束 ==========

6.7 单层前向 + 反向通信量汇总

以一个 Dense Attention + MLP 层 为例(不含 MoE):

通信 通信组 方式 数据量 前向 / 反向
SP All-Gather (RMSNorm 前) TP=4 All-Gather 48 MiB 前向
CP All-to-All (QKV) CP=2 All-to-All 48 MiB 前向
CP All-to-All (O) CP=2 All-to-All 16 MiB 前向
TP Reduce-Scatter (W_o) TP=4 Reduce-Scatter 64 MiB 前向
TP Reduce-Scatter (MLP) TP=4 Reduce-Scatter 64 MiB 前向
TP All-Gather (MLP 反向) TP=4 All-Gather 64 MiB 反向
TP All-Gather (W_o 反向) TP=4 All-Gather 64 MiB 反向
CP All-to-All (dO) CP=2 All-to-All 16 MiB 反向
CP All-to-All (dQ,dK,dV) CP=2 All-to-All 48 MiB 反向
SP Reduce-Scatter (RMSNorm 反向) TP=4 Reduce-Scatter 48 MiB 反向
单层总计 480 MiB

以一个 MoE 层 为例(替换 MLP 为 MoE):

通信 通信组 方式 数据量 前向 / 反向
SP All-Gather (RMSNorm 前) TP=4 All-Gather 48 MiB 前向
CP All-to-All (QKV) CP=2 All-to-All 48 MiB 前向
CP All-to-All (O) CP=2 All-to-All 16 MiB 前向
TP Reduce-Scatter (W_o) TP=4 Reduce-Scatter 64 MiB 前向
EP Dispatch All-to-All EP=8 All-to-All 56 MiB 前向
EP Combine All-to-All EP=8 All-to-All 56 MiB 前向
EP Dispatch (反向) EP=8 All-to-All 56 MiB 反向
EP Combine (反向) EP=8 All-to-All 56 MiB 反向
CP All-to-All (dO) CP=2 All-to-All 16 MiB 反向
CP All-to-All (dQ,dK,dV) CP=2 All-to-All 48 MiB 反向
TP All-Gather (W_o 反向) TP=4 All-Gather 64 MiB 反向
SP Reduce-Scatter (RMSNorm 反向) TP=4 Reduce-Scatter 48 MiB 反向
MoE 层总计 640 MiB

6.8 全局通信量估算

前向(全部 80 层):

  • 60 个 Dense 层: 60×240 MiB=14,040 MiB≈13.760 \times 240 \text{MiB} = 14{,}040 \text{MiB} \approx 13.7 GiB
  • 20 个 MoE 层: 20×352 MiB=7,040 MiB≈6.920 \times 352 \text{MiB} = 7{,}040 \text{MiB} \approx 6.9 GiB
  • PP P2P (3 次 stage 间): 3×16 MiB=48 MiB3 \times 16 \text{MiB} = 48 \text{MiB}
  • 前向总通信量: ≈21\approx 21 GiB

反向(全部 80 层):

  • 60 个 Dense 层: 60×240 MiB=14,040 MiB≈13.760 \times 240 \text{MiB} = 14{,}040 \text{MiB} \approx 13.7 GiB
  • 20 个 MoE 层: 20×288 MiB=5,760 MiB≈5.620 \times 288 \text{MiB} = 5{,}760 \text{MiB} \approx 5.6 GiB
  • PP P2P (3 次): 3×16 MiB=48 MiB3 \times 16 \text{MiB} = 48 \text{MiB}
  • 反向总通信量: ≈19.4\approx 19.4 GiB

ZeRO-3 DP 通信(每次 step):

  • All-Gather(参数, 前向): 2Ψ×22\Psi \times 2 bytes
  • Reduce-Scatter(梯度, 反向): 2Ψ×22\Psi \times 2 bytes
  • All-Gather(参数, 反向): 2Ψ×22\Psi \times 2 bytes
  • 总计: 8Ψ8\Psi bytes

6.9 显存估算

假设模型总参数 Ψ≈350B\Psi \approx 350B(含 MoE 专家参数):

每卡显存(ZeRO-3, Ndp=4N_{dp}=4):

  • 参数: 2×350B/4=1752 \times 350B / 4 = 175 GB … 需要更多 DP 度数或 TP 分担

实际中 TP=4 也会切分 Attention 和 MLP 参数:

  • TP 切分参数量: 约占总参数的 1/2(Attention + Dense MLP)
  • TP 后每卡参数: Ψtp≈Ψ/2+Ψ/2/4=5Ψ/8\Psi_{tp} \approx \Psi/2 + \Psi/2/4 = 5\Psi/8
  • ZeRO-3 进一步切分: 5Ψ/8/4=5Ψ/325\Psi/8 / 4 = 5\Psi/32
  • 参数: 2×350B×5/32≈1092 \times 350B \times 5/32 \approx 109 GB(仍较大)
  • 梯度: 109109 GB
  • 优化器: 6×350B×5/32≈3286 \times 350B \times 5/32 \approx 328 GB … 需要 offload 或更大 DP

实际中还需要 EP 分担 MoE 参数:

  • EP=8 切分了 64 个专家, 每卡 8 个专家
  • 专家参数约占总参数的 70%(典型 MoE 模型)
  • Dense 参数(30%): 350B×0.3=105B350B \times 0.3 = 105B
  • 每卡专家参数: 350B×0.7/8=30.6B350B \times 0.7 / 8 = 30.6B

最终每卡参数(TP+EP 切分后):

  • Dense 部分(TP=4 切): 105B×3/4≈79B105B \times 3/4 \approx 79B (TP 切走 1/4)
  • Expert 部分(EP=8 切): 30.6B30.6B
  • 合计每卡: ≈110B\approx 110B

ZeRO-3 (Ndp=4N_{dp}=4) 切分后:

  • 参数: 110B/4=27.5B→2×27.5=55110B / 4 = 27.5B → 2 \times 27.5 = 55 GB
  • 梯度: 55 GB
  • 优化器: 6×27.5=1656 \times 27.5 = 165 GB
  • 激活值(含重计算): ≈20\approx 20 GB
  • 总计: ≈295\approx 295 GB → 需要多节点或进一步优化

6.10 全流程时序图

时间轴 (从左到右) →

Stage 0:  [F1][F2]...[F_m] [B_m]...[B2][B1]
Stage 1:         [F1][F2]...[F_m] [B_m]...[B2][B1]
Stage 2:              [F1]...[F_m] [B_m]...[B1]
Stage 3:                  [F1]...[F_m][B_m]...[B1]
                                     ^
                                     |
                            反向从此开始

其中:
- F_i = 第 i 个 micro-batch 的前向
- B_i = 第 i 个 micro-batch 的反向
- m = micro-batch 数量 (1F1B 调度)

每个 F_i 内部:
  |RN|AG_sp|QKV|A2A_cp|Attn|A2A_cp|Wo|RS_tp|Res|MLP|RS_tp|Res|

每个 B_i 内部:
  |RN_b|AG_tp|MLP_b|RS_tp|Res_b|Wo_b|AG_tp|A2A_cp|Attn_b|A2A_cp|QKV_b|RS_sp|RN_b|

通信标记:
  AG_sp = All-Gather (SP, TP 组)
  RS_tp = Reduce-Scatter (TP 组)
  A2A_cp = All-to-All (CP 组)
  A2A_ep = All-to-All (EP 组, MoE 层)
  P2P_pp = P2P Send-Recv (PP 组)

7. 总结

本文档全面覆盖了 Megatron 框架下五种核心并行策略:

并行策略 切分对象 核心通信 前向通信量 反向通信量
TP+SP 参数(按头 / 列 / 行)、激活值(RN 区域按序列) All-Gather + Reduce-Scatter (TP 组) 112 MiB/ 层 112 MiB/ 层
PP 层(按层切分) P2P Send-Recv (PP 组) 16 MiB/stage 16 MiB/stage
CP (Ring) 激活值(Attention 的 Q,K,V 按序列) P2P 环形 (CP 组) 256 MiB 256 MiB
CP (AllGather) 同上 All-Gather + RS (CP 组) 512 MiB 512 MiB
CP (Ulysses) 同上 All-to-All (CP 组) 128 MiB 128 MiB
EP 专家参数(按专家) All-to-All (EP 组) 896 MiB/MoE 层 896 MiB/MoE 层
DP/ZeRO-1 优化器状态 All-Reduce (DP 组) 0 4Ψ4\Psi
DP/ZeRO-2 优化器状态 + 梯度 RS + AG (DP 组) 0 4Ψ4\Psi
DP/ZeRO-3 优化器状态 + 梯度 + 参数 AG + RS (DP 组) $4\Psi$ 4Ψ4\Psi

这些并行策略可以正交组合,通过合理的通信组划分和调度,在 1024+ 卡的规模上实现接近线性的扩展效率。

 0