观文听傑

返回

上一篇用 FSDP2 把参数、梯度和 optimizer state 分片常驻,但某个模块计算前仍要 all-gather(全聚合)该模块的完整参数。如果一层超宽 MLP 本身就放不进单卡,或者反复聚齐它已经成为瓶颈,切数据轴不够,必须切开层内矩阵。

本篇只讲透一个配对:先用 Column-wise Parallel(列并行)切上投影的输出维,再用 Row-wise Parallel(行并行)切下投影的输入维。中间激活一直分片,直到下投影的局部部分和需要合并。

01 为什么两层必须一起设计?#

忽略 bias,Transformer MLP 可写成

H=ϕ(XW1),Y=HW2H=\phi(XW_1),\qquad Y=HW_2

其中 XRB×L×DX\in\mathbb R^{B\times L\times D}W1RD×FW_1\in\mathbb R^{D\times F}HRB×L×FH\in\mathbb R^{B\times L\times F}W2RF×DW_2\in\mathbb R^{F\times D}FF 常为 DD 的数倍,因此参数和中间激活都很大。

若只切 W1W_1,算完就 all-gather 完整 HH,下一层又要重新切开,通信抵消了分片收益。正确配对让 W1W_1 的列 shard 恰好成为 W2W_2 的行 shard 输入。

02 两卡的数据流与张量形状#

设 Tensor Parallel(张量并行,TP)大小 T=2T=2

flowchart LR
  X[每卡复制 X<br/>B×L×D] --> W10[rank 0: W1⁰<br/>D×F/2]
  X --> W11[rank 1: W1¹<br/>D×F/2]
  W10 --> H0[H⁰<br/>B×L×F/2]
  W11 --> H1[H¹<br/>B×L×F/2]
  H0 --> W20[rank 0: W2⁰<br/>F/2×D]
  H1 --> W21[rank 1: W2¹<br/>F/2×D]
  W20 --> P0[部分和 P⁰<br/>B×L×D]
  W21 --> P1[部分和 P¹<br/>B×L×D]
  P0 --> R[all-reduce SUM]
  P1 --> R
  R --> Y[每卡复制 Y<br/>B×L×D]
mermaid

列并行写作

W1=[W1(0)  W1(1)],H(r)=ϕ(XW1(r))RB×L×F/2.W_1=[W_1^{(0)}\;W_1^{(1)}],\quad H^{(r)}=\phi(XW_1^{(r)})\in\mathbb R^{B\times L\times F/2}.

行并行把 W2W_2 沿第一维对应切开:

P(r)=H(r)W2(r)RB×L×D,Y=P(0)+P(1).P^{(r)}=H^{(r)}W_2^{(r)}\in\mathbb R^{B\times L\times D},\qquad Y=P^{(0)}+P^{(1)}.

关键不是“每层各切一半”,而是两个切分轴能首尾相接。激活函数 phiphi 是逐元素运算,各 rank 可直接在自己的 F/2F/2 个通道上计算,不需要通信。

03 用四维隐藏层手算一次#

令一个 token 的输入 x=[1,2]x=[1,2],上投影为

W1=[10120111].W_1=\begin{bmatrix}1&0&1&2\\0&1&1&-1\end{bmatrix}.

两卡各取两列:

rank 0: x @ [[1,0],[0,1]] = [1,2]
rank 1: x @ [[1,2],[1,-1]] = [3,0]
text

用 ReLU 后仍为 h(0)=[1,2]h^{(0)}=[1,2]h(1)=[3,0]h^{(1)}=[3,0]。令

W2=[10011121],W_2=\begin{bmatrix}1&0\\0&1\\1&1\\2&-1\end{bmatrix},

则行分片得到

rank 0 部分和: [1,2] @ [[1,0],[0,1]] = [1,2]
rank 1 部分和: [3,0] @ [[1,1],[2,-1]] = [3,3]
all-reduce SUM: [1,2] + [3,3] = [4,5]
text

直接算完整 ReLU(xW1)W2\operatorname{ReLU}(xW_1)W_2 也是 [4,5]。这个极小例子最适合作为多卡实现的 oracle(正确性基准)。

04 为什么只需在块尾合并?#

上投影输出是拼接关系,若必须恢复完整 HH,需要 all-gather;下投影输出则是求和关系,需要 all-reduce。把两层配对后,前者被省掉,只保留块尾一次逻辑归约。

实际实现还能把 all-reduce 拆为 reduce-scatter,使残差流也保持分片,再在后续合适位置 all-gather。这会改变激活 layout(布局)契约,却不改变数学式:某处最终必须把各 rank 的部分和组合起来。

通信量不能只按“collective 次数”比较。对 N=BLN=B L 个 token,块尾传输的逻辑张量为 N×DN\times D;当 NN 很小或 D/TD/T 太窄时,通信延迟与小矩阵低利用率可能让 TP 比单卡更慢。

05 Bias 为什么容易被加两次?#

若每个 rank 在部分和 P(r)P^{(r)} 上都加完整下投影 bias b2b_2,all-reduce 后会得到

r(P(r)+b2)=Y+Tb2.\sum_r(P^{(r)}+b_2)=Y+Tb_2.

正确做法是先归约部分和再加一次 bias,或让每卡只贡献 b2/Tb_2/T。框架的 RowwiseParallel 会管理兼容层的参数与通信;手写实现时必须把 bias 的所有权写进测试。上投影 bias 按输出列自然分片,不存在重复相加。

06 反向传播怎样沿相反方向流动?#

前向输出 YY 在每卡复制,因此每卡拿到相同的 L/Y\partial\mathcal L/\partial Y。行并行反向可在本地得到 L/W2(r)\partial\mathcal L/\partial W_2^{(r)}L/H(r)\partial\mathcal L/\partial H^{(r)};后者继续穿过本地激活和上投影。

上投影对复制输入 XX 的梯度是各列 shard 的贡献之和:

LX=rLZ(r)(W1(r)).\frac{\partial\mathcal L}{\partial X} =\sum_r \frac{\partial\mathcal L}{\partial Z^{(r)}}(W_1^{(r)})^\top.

因此通信不会消失,只是前向与反向分别落在能保持中间分片的位置。分析性能时要同时查看两个方向,不能只数 forward collective。

07 当前 PyTorch 的最小落地#

PyTorch 2.14 的 Tensor Parallel API 构建在 DTensor(分布式张量)之上,parallelize_module 只接收一维 DeviceMesh。多维 mesh 必须先取出 tp 子 mesh。

当前 API 仍标记 experimental(实验性),应固定 PyTorch 版本。若列并行和行并行之间有依赖全局维度的 view、split 或自定义 kernel,普通本地 Tensor 可能把 F/2 错当成 FF;此时使用 use_local_output=False 保留 DTensor 的全局 shape/placement 信息,并逐个审计中间算子。

08 Gated MLP 怎样扩展同一配对?#

SwiGLU 常写作

H=SiLU(XWg)(XWu),Y=HWd.H=\operatorname{SiLU}(XW_g)\odot(XW_u),\qquad Y=HW_d.

WgW_gWuW_u 必须使用相同的列切分,确保每卡拿到相同通道范围,才能本地逐元素相乘;WdW_d 再按对应输入行切分。若两个上投影的 shard 顺序不同,shape 完全正确,通道语义却已经错位。

tp_plan = {
    "gate_proj": ColwiseParallel(),
    "up_proj": ColwiseParallel(),
    "down_proj": RowwiseParallel(),
}
python

09 初始化、保存与 FSDP2 组合#

parallelize_module 会把参数转换为 DTensor。optimizer 应在并行化后创建,使 state 跟随本地 shards。保存 checkpoint 时使用理解 DTensor/sharded state 的分布式接口;不要每卡把本地 shard 当完整 state_dict

二维 [dp,tp] mesh 中,先在 mesh["tp"] 应用 TP,再在独立 mesh["dp"] 上应用 FSDP2。TP ranks 协作处理同一份样本,DP ranks 才处理不同 batch shards。若把两个 group 交换,程序可能在 collective 上挂起,或把 batch 错切两次。

初始化也要验证全局等价性:用固定完整权重切片,比“每卡同 seed 各初始化一个较小矩阵”更可靠,因为不同局部 shape 会改变随机数消费顺序。

10 如何证明实现没有静默算错?#

  1. 用本篇 2422\to4\to2 的整数矩阵关闭 bias,比较每卡输出与完整矩阵答案。
  2. 再加入不同的 bias,专门检查下投影 bias 没被乘以 TT
  3. 保存单卡 FP32 的输出、输入梯度和完整参数梯度;TP 后 gather shards,用 torch.testing.assert_close 比较。
  4. 给 gate/up 投影设置可辨认的通道编号,检查两者 shard 对齐。
  5. profiler 中核对 collective 的张量元素数、process group 与调用顺序,再测 tokens/s。

11 常见错误与最短调试路径#

症状常见原因最短检查
中间 shape 变成预期一半后报错本地 shard 被当作全局 shape打印 DTensor placements 与 local/global shape
输出整体多一个常数下投影 bias 在每卡先加后归约bias 置为已知非零数做手算
两卡正常、四卡结果改变分母、bias 或 process group 依赖 TTT=1,2,4T=1,2,4 跑等价性测试
collective 永久等待ranks 走了不同控制流或 group 错给每次 collective 编号并对齐日志
显存下降但吞吐变差shard 矩阵太窄或通信未被隐藏同看 GEMM shape、带宽与端到端 tokens/s
checkpoint 恢复后 loss 跳变shard 元数据或 optimizer state 不完整做跨 world-size 恢复演练

12 它与相近并行方式有什么区别?#

  • FSDP2 在层外分片、层内临时使用完整参数;TP 让层内算子本身跨卡执行。
  • Sequence Parallel(序列并行)切激活的 token/sequence 维,常用于 LayerNorm、dropout 等算子,不能替代超宽权重的切分。
  • Expert Parallel(专家并行)把不同 MoE experts 放到不同 rank,通信核心是 token dispatch 的 all-to-all,而非稠密 MLP 的部分和。
  • Pipeline Parallel(流水线并行)按层深度切分,只在 stage 边界传激活;下一篇将讨论它的调度。

TP 的失败边界也很明确:高速互连不足、层太小、分片数不能整除 head/FFN 维、频繁动态控制流,都会让复杂度高于收益。先用单机高速互连建立基线,再扩到跨节点。

13 今天真正需要记住什么?#

  1. 上投影列并行产生 F/TF/T 通道,下投影按对应输入行并行;二者配对让中间激活无需 all-gather。
  2. 每卡下投影先得到完整输出形状的部分和,必须在正确位置归约;bias 只能逻辑上加一次。
  3. local shape 与 global shape 不同,逐元素算子通常安全,view、分头与自定义 kernel 必须审计 DTensor layout。
  4. 多卡“能跑”不是正确性证据:输出、梯度、初始化、checkpoint 与 process group 都要对单卡 oracle。

14 思考题与小练习#

  1. X[8,128,4096]X[8,128,4096]W1[4096,16384]W_1[4096,16384]W2[16384,4096]W_2[16384,4096] 做 8 路 TP,写出每卡两块权重与所有中间激活 shape。
  2. 若每个 rank 都给下投影部分和加 bias,证明 all-reduce 后 bias 被放大 TT 倍,并给出两种修复。
  3. 把手算例改成带负数、使用 ReLU,分别算完整路径与两卡路径的输出和 Y0/x\partial Y_0/\partial x

相关工作#

  1. Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism,系统化展示 Transformer MLP 与 attention 的层内张量并行。
  2. Shazeer et al., Mesh-TensorFlow: Deep Learning for Supercomputers,用命名张量维度表达分布式切分。
  3. Xu et al., GSPMD: General and Scalable Parallelization for ML Computation Graphs,讨论编译器传播 sharding annotation 的方法。
  4. Lian et al., Colossal-AI: A Unified Deep Learning System For Large-Scale Parallel Training,总结多维并行的系统组合。
  5. Li et al., PyTorch Distributed: Experiences on Accelerating Data Parallel Training,提供 collective 与 PyTorch 分布式运行时背景。

15 下一篇预告#

张量并行解决了单层太宽,却让每层都依赖高速 collective。下一篇转向按深度切分的 Pipeline Parallel:micro-batch 如何填满 stages,GPipe 与 1F1B 的气泡、激活驻留和梯度语义有何不同。

一层 MLP 太宽而放不下时怎样切?列并行、行并行与块尾归约
https://zwjcode.cn/blog/mlp-tensor-parallel-column-row-collective
作者
发布于 2026年9月16日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。