一层 MLP 太宽而放不下时怎样切?列并行、行并行与块尾归约
从 FSDP2 仍需临时聚齐单层参数出发,手算 Transformer MLP 的列并行与行并行,解释局部张量形状、通信位置、反向传播与 PyTorch DTensor 落地。
上一篇用 FSDP2 把参数、梯度和 optimizer state 分片常驻,但某个模块计算前仍要 all-gather(全聚合)该模块的完整参数。如果一层超宽 MLP 本身就放不进单卡,或者反复聚齐它已经成为瓶颈,切数据轴不够,必须切开层内矩阵。
本篇只讲透一个配对:先用 Column-wise Parallel(列并行)切上投影的输出维,再用 Row-wise Parallel(行并行)切下投影的输入维。中间激活一直分片,直到下投影的局部部分和需要合并。
01 为什么两层必须一起设计?#
忽略 bias,Transformer MLP 可写成
其中 ,,,。 常为 的数倍,因此参数和中间激活都很大。
若只切 ,算完就 all-gather 完整 ,下一层又要重新切开,通信抵消了分片收益。正确配对让 的列 shard 恰好成为 的行 shard 输入。
02 两卡的数据流与张量形状#
设 Tensor Parallel(张量并行,TP)大小 :
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列并行写作
行并行把 沿第一维对应切开:
关键不是“每层各切一半”,而是两个切分轴能首尾相接。激活函数 是逐元素运算,各 rank 可直接在自己的 个通道上计算,不需要通信。
03 用四维隐藏层手算一次#
令一个 token 的输入 ,上投影为
两卡各取两列:
rank 0: x @ [[1,0],[0,1]] = [1,2]
rank 1: x @ [[1,2],[1,-1]] = [3,0]text用 ReLU 后仍为 、。令
则行分片得到
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直接算完整 也是 [4,5]。这个极小例子最适合作为多卡实现的 oracle(正确性基准)。
04 为什么只需在块尾合并?#
上投影输出是拼接关系,若必须恢复完整 ,需要 all-gather;下投影输出则是求和关系,需要 all-reduce。把两层配对后,前者被省掉,只保留块尾一次逻辑归约。
实际实现还能把 all-reduce 拆为 reduce-scatter,使残差流也保持分片,再在后续合适位置 all-gather。这会改变激活 layout(布局)契约,却不改变数学式:某处最终必须把各 rank 的部分和组合起来。
通信量不能只按“collective 次数”比较。对 个 token,块尾传输的逻辑张量为 ;当 很小或 太窄时,通信延迟与小矩阵低利用率可能让 TP 比单卡更慢。
05 Bias 为什么容易被加两次?#
若每个 rank 在部分和 上都加完整下投影 bias ,all-reduce 后会得到
正确做法是先归约部分和再加一次 bias,或让每卡只贡献 。框架的 RowwiseParallel 会管理兼容层的参数与通信;手写实现时必须把 bias 的所有权写进测试。上投影 bias 按输出列自然分片,不存在重复相加。
06 反向传播怎样沿相反方向流动?#
前向输出 在每卡复制,因此每卡拿到相同的 。行并行反向可在本地得到 与 ;后者继续穿过本地激活和上投影。
上投影对复制输入 的梯度是各列 shard 的贡献之和:
因此通信不会消失,只是前向与反向分别落在能保持中间分片的位置。分析性能时要同时查看两个方向,不能只数 forward collective。
07 当前 PyTorch 的最小落地#
PyTorch 2.14 的 Tensor Parallel API 构建在 DTensor(分布式张量)之上,parallelize_module 只接收一维 DeviceMesh。多维 mesh 必须先取出 tp 子 mesh。
import torch
from torch import nn
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor.parallel import (
ColwiseParallel,
RowwiseParallel,
parallelize_module,
)
class MLP(nn.Module):
def __init__(self, d_model: int, d_ff: int):
super().__init__()
self.up_proj = nn.Linear(d_model, d_ff)
self.act = nn.GELU()
self.down_proj = nn.Linear(d_ff, d_model)
def forward(self, x): # x: [B,L,D],各 TP rank 复制
return self.down_proj(self.act(self.up_proj(x)))
torch.distributed.init_process_group("nccl")
tp_mesh = init_device_mesh("cuda", (torch.distributed.get_world_size(),))
model = MLP(d_model=4096, d_ff=11008).cuda()
model = parallelize_module(
model,
tp_mesh,
{
# 默认输出沿最后一维分片;逐元素 GELU 可本地执行
"up_proj": ColwiseParallel(),
# 默认假设输入最后一维已分片,并返回复制输出
"down_proj": RowwiseParallel(),
},
)
x = torch.randn(2, 128, 4096, device="cuda")
y = model(x) # [2,128,4096],复制布局
loss = y.square().mean()
loss.backward()python当前 API 仍标记 experimental(实验性),应固定 PyTorch 版本。若列并行和行并行之间有依赖全局维度的 view、split 或自定义 kernel,普通本地 Tensor 可能把 F/2 错当成 ;此时使用 use_local_output=False 保留 DTensor 的全局 shape/placement 信息,并逐个审计中间算子。
08 Gated MLP 怎样扩展同一配对?#
SwiGLU 常写作
与 必须使用相同的列切分,确保每卡拿到相同通道范围,才能本地逐元素相乘; 再按对应输入行切分。若两个上投影的 shard 顺序不同,shape 完全正确,通道语义却已经错位。
tp_plan = {
"gate_proj": ColwiseParallel(),
"up_proj": ColwiseParallel(),
"down_proj": RowwiseParallel(),
}python09 初始化、保存与 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 如何证明实现没有静默算错?#
- 用本篇 的整数矩阵关闭 bias,比较每卡输出与完整矩阵答案。
- 再加入不同的 bias,专门检查下投影 bias 没被乘以 。
- 保存单卡 FP32 的输出、输入梯度和完整参数梯度;TP 后 gather shards,用
torch.testing.assert_close比较。 - 给 gate/up 投影设置可辨认的通道编号,检查两者 shard 对齐。
- profiler 中核对 collective 的张量元素数、process group 与调用顺序,再测 tokens/s。
11 常见错误与最短调试路径#
| 症状 | 常见原因 | 最短检查 |
|---|---|---|
| 中间 shape 变成预期一半后报错 | 本地 shard 被当作全局 shape | 打印 DTensor placements 与 local/global shape |
| 输出整体多一个常数 | 下投影 bias 在每卡先加后归约 | bias 置为已知非零数做手算 |
| 两卡正常、四卡结果改变 | 分母、bias 或 process group 依赖 | 对 跑等价性测试 |
| 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 今天真正需要记住什么?#
- 上投影列并行产生 通道,下投影按对应输入行并行;二者配对让中间激活无需 all-gather。
- 每卡下投影先得到完整输出形状的部分和,必须在正确位置归约;bias 只能逻辑上加一次。
- local shape 与 global shape 不同,逐元素算子通常安全,
view、分头与自定义 kernel 必须审计 DTensor layout。 - 多卡“能跑”不是正确性证据:输出、梯度、初始化、checkpoint 与 process group 都要对单卡 oracle。
14 思考题与小练习#
- 对 、、 做 8 路 TP,写出每卡两块权重与所有中间激活 shape。
- 若每个 rank 都给下投影部分和加 bias,证明 all-reduce 后 bias 被放大 倍,并给出两种修复。
- 把手算例改成带负数、使用 ReLU,分别算完整路径与两卡路径的输出和 。
相关工作#
- Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism ↗,系统化展示 Transformer MLP 与 attention 的层内张量并行。
- Shazeer et al., Mesh-TensorFlow: Deep Learning for Supercomputers ↗,用命名张量维度表达分布式切分。
- Xu et al., GSPMD: General and Scalable Parallelization for ML Computation Graphs ↗,讨论编译器传播 sharding annotation 的方法。
- Lian et al., Colossal-AI: A Unified Deep Learning System For Large-Scale Parallel Training ↗,总结多维并行的系统组合。
- Li et al., PyTorch Distributed: Experiences on Accelerating Data Parallel Training ↗,提供 collective 与 PyTorch 分布式运行时背景。
15 下一篇预告#
张量并行解决了单层太宽,却让每层都依赖高速 collective。下一篇转向按深度切分的 Pipeline Parallel:micro-batch 如何填满 stages,GPipe 与 1F1B 的气泡、激活驻留和梯度语义有何不同。