模型放不下一张卡时究竟该切什么?数据并行、张量并行与流水线并行的边界
从显存瓶颈和通信位置出发,用四卡手算比较数据并行、张量并行与流水线并行,解释三种切分的数据流、张量形状、性能代价与组合原则。
上一篇解决了 Gradient Accumulation(梯度累积)如何在一张卡和 DistributedDataParallel(分布式数据并行,DDP)中保持大 batch 语义。但如果模型参数、激活或 optimizer state 本身已经放不进单卡,继续缩小 micro-batch 也无济于事。
此时不能只问“有几张 GPU”,而要问:复制什么、切开什么,以及被切开的张量何时必须重新通信? 本篇只建立 Data Parallel(数据并行,DP)、Tensor Parallel(张量并行,TP)与 Pipeline Parallel(流水线并行,PP)的选择框架;具体分片算法留到后文。
01 先把一轮训练拆成四类显存#
设模型有 个参数,每个参数的权重、梯度和优化器状态分别占 字节,激活峰值为 。单卡粗略峰值为
其中 是 micro-batch 大小, 是序列长度, 是 attention、通信和算子 workspace。以混合精度 AdamW 为例,若 FP16 权重 2 字节、FP16 梯度 2 字节、FP32 master weight、动量和二阶矩共 12 字节,则仅模型状态约为 字节;10 亿参数已经约 16 GB,还没有算激活和临时张量。
先测清谁占满显存,才知道该切哪个轴:
| 瓶颈 | 仅缩小 batch | 更直接的方向 |
|---|---|---|
| 激活随 增长 | 有效,但吞吐可能下降 | checkpoint、sequence/pipeline parallel |
| 参数与 optimizer state | 几乎无效 | FSDP/ZeRO、tensor parallel |
| 单层矩阵本身放不下 | 无效 | tensor parallel |
| 层数很多、跨机带宽较低 | 部分有效 | pipeline parallel |
02 三种并行究竟切哪一维?#
flowchart TB
M[完整训练计算图] --> DP[数据并行: 切 batch B]
M --> TP[张量并行: 切 hidden/head/FFN 维]
M --> PP[流水线并行: 切 layer 深度]
DP --> D1[每卡完整层<br/>不同样本]
TP --> T1[每卡一层的部分矩阵<br/>同一批样本]
PP --> P1[每卡连续若干层<br/>micro-batch 流过 stages]
D1 --> C1[反向同步梯度或分片状态]
T1 --> C2[层内 collective]
P1 --> C3[stage 边界点对点传激活/梯度]mermaid- DP 切 batch 轴:rank 看到 ,通常保存完整计算图。
- TP 切模型层内部:同一个 被多个 rank 协作处理,权重或激活沿 hidden/head 维分片。
- PP 切层轴:stage 0 保存前几层,stage 1 保存后几层,边界激活从前向后传,边界梯度反向传回。
它们不是三种互斥框架,而是三个近乎正交的 mesh 维度。大规模训练常把总卡数写成
但乘积相同不代表性能相同,因为三种通信发生的位置和频率不同。
03 用四张卡手算“复制与切分”#
假设模型状态 32 GB,激活峰值 8 GB,每卡只有 24 GB,暂忽略通信 buffer。
方案 A:4 路普通 DDP#
每卡都复制 32 GB 模型状态,再放自己的激活。峰值约 GB,仍然 OOM。DDP 增加了吞吐,却没有解决完整模型状态放不下的问题。
方案 B:4 路状态分片的数据并行#
若参数、梯度与 optimizer state 理想地均分,每卡模型状态约 GB;再加约 2 GB 局部激活,静态估算 10 GB。实际还会在计算某层前短暂 all-gather(全聚合)该层参数,因此峰值取决于分片边界。
方案 C:2 路 TP × 2 路 PP#
每个 pipeline stage 放一半层,每层又做 2 路 tensor shard。理想模型状态约 GB;但 TP 每层可能做 all-reduce 或 reduce-scatter,PP 还会有 pipeline bubble(流水线气泡)。它能处理“单层也放不下”的情况,代价是更密集的调度与通信。
这个例子说明:显存除法只是容量下界,collective 的时机才决定实际吞吐与峰值。
04 数据并行:算子完整,样本不同#
设线性层 ,,。两路 DP 得到:
rank 0: X0 [B/2,D] × W [D,H] -> Y0 [B/2,H]
rank 1: X1 [B/2,D] × W [D,H] -> Y1 [B/2,H]
backward: 对 dW0、dW1 做同步,保证副本下一步仍一致text优点是每张卡执行完整大矩阵,算子利用率通常好,模型代码改动少。限制是普通 DDP 复制完整权重、梯度和 optimizer state,单卡容量必须先容纳模型。上一篇的 no_sync() 只减少梯度累积期间的通信次数,不减少长期保存的模型状态。
数据采样也必须显式分片。DDP 不会替你切输入:训练集若在每个 rank 都按相同顺序遍历,等价于重复计算同一批样本。
05 张量并行:同一层由多卡合算#
仍看 。两路 Column-wise Parallel(列并行)把 的输出维切开:
若下一层可直接消费分片 ,无需马上拼回;若后续算子要求完整 ,就要 all-gather。Row-wise Parallel(行并行)则把输入维与权重行切开:
每卡先算部分和 ,再 all-reduce 求 。Transformer 常把相邻的列并行与行并行配对,使中间激活保持分片,只在块的合适边界通信。
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor.parallel import (
ColwiseParallel,
RowwiseParallel,
parallelize_module,
)
tp_mesh = init_device_mesh("cuda", (2,))
model = parallelize_module(
model,
tp_mesh,
{
"mlp.up_proj": ColwiseParallel(),
"mlp.down_proj": RowwiseParallel(),
},
)python这是 PyTorch 2.14 的核心接口形状,真实模型还要声明输入/输出 DTensor layout,并检查两层之间的 reshape、分头和残差是否理解“全局 shape”与“本地 shard”。
06 流水线并行:按层切成 stages#
若 12 层模型分到 3 个 stage,每个 stage 持有 4 层。一个完整 batch 若一次穿过所有 stage,后两张卡在开头空闲、前两张卡在结尾空闲。PP 因而把 batch 再切成 个 micro-batch 并交错执行。
时间 -> t0 t1 t2 t3 t4
stage 0: μ0 μ1 μ2 μ3 --
stage 1: -- μ0 μ1 μ2 μ3
stage 2: -- -- μ0 μ1 μ2 ...text对理想 GPipe 前向排程, 个 stage、 个 micro-batch 的气泡比例近似
例如 时约为 ;增至 后约为 。但更多 micro-batch 会增加调度开销,并影响激活驻留和梯度累积语义。
PP 的关键不是“平均分层数”,而是平衡每个 stage 的实际前后向时间与显存。Embedding、attention、MLP 和输出词表层成本不同,最慢 stage 决定流水线节拍。
07 通信原语比名称更重要#
| 方法 | 主要通信 | 典型发生时机 | 对网络的要求 |
|---|---|---|---|
| DDP | gradient all-reduce | backward bucket ready | 可与反向重叠 |
| FSDP 类 DP | parameter all-gather、gradient reduce-scatter | 每个分片单元前/后 | 依赖预取与分组 |
| TP | all-reduce、all-gather、reduce-scatter | 层内部,频率高 | 适合机内高速互连 |
| PP | send/recv activation 与 gradient | stage 边界 | 消息较大但频率较低 |
因此常见拓扑是:节点内用 TP,利用 NVLink 等高带宽互连;节点间用 FSDP/DP;层很多或跨机带宽有限时再加 PP。它是经验起点,不是无需测量的定律。
08 一个可执行的选择流程#
profile 单卡峰值
├─ 模型状态能放下?
│ ├─ 是:先 DP;若激活爆炸,考虑 checkpoint/sequence parallel
│ └─ 否:状态分片 DP
├─ 最大单层 all-gather 后仍放不下?
│ └─ 是:给该层加 TP
└─ 层可自然分段且网络跨节点较慢?
└─ 测 PP,并调 micro-batch 数与 stage balancetext选择时至少记录:每卡参数/梯度/optimizer/activation 峰值、每种 collective 的字节数与次数、通信计算重叠率、最慢 rank、tokens/s 和收敛等价性。仅看 GPU utilization 很容易把等待 collective 的时间误读为有效工作。
09 组合成二维 mesh 时 shape 如何理解?#
假设 8 张卡组成 [dp=4,tp=2]:同一个 TP 小组的两卡共同计算一份样本;四个 DP 小组处理四份不同样本。若全局 batch 为 32:
- 每个 DP replica 得到 8 个样本;
- replica 内两张 TP 卡都参与这 8 个样本,不是各拿 4 个;
- TP 切的是 hidden/head,不应再次除 batch;
- DP 同步发生在“持有相同参数 shard”的四个 rank 之间。
把两个 mesh 维的 process group 搞反,轻则 shape mismatch,重则 collective 顺序不一致而永久 hang。
10 最小调试方法:先验证语义,再测性能#
- 用两层小模型、关闭 dropout,保存单卡基准的 loss 与完整梯度。
- 每加一个并行维度,就 gather 出等价张量,用
torch.testing.assert_close比较输出和梯度。 - 打印每个 rank 的 mesh coordinate、输入样本 ID、局部参数 shape 和 collective 顺序。
- 在极小 batch 上跑 2–3 个 optimizer step,比较更新后的完整 state dict。
- 最后才扩大模型,使用 profiler 检查通信、气泡与峰值显存。
多卡程序“没有报错”不代表正确。样本重复、梯度多除一次 world size、某个 shard 没参与 loss,都可能稳定运行却优化错误目标。
11 常见错误与症状#
| 症状 | 常见原因 | 最短检查 |
|---|---|---|
| 加卡后仍 OOM | 使用复制模型的 DDP 解决参数容量 | 分开统计模型状态与激活 |
| TP 输出 shape 突然减半 | 本地 shard 当成全局 tensor 使用 | 打印 DTensor placement 与 local shape |
| PP 吞吐呈锯齿 | stage 不平衡或 micro-batch 太少 | 画每 stage 时间线 |
| 程序永久 hang | ranks 的 collective 顺序不同 | 为每个 collective 编号并比日志 |
| loss 与单卡差 world size 倍 | DP 与 TP group/归一化混淆 | 写出每个 reduction 的数学目标 |
| GPU 很忙但 tokens/s 低 | 频繁小 collective 或重算过多 | 同看 kernel、通信和端到端吞吐 |
12 三种方法各自会在哪里失败?#
DP 的扩展受全局 batch 和数据并行通信限制;模型若已完整放不下,普通 DDP 无法启动。TP 会把通信插进每层,跨低带宽节点常被延迟拖垮,小矩阵分片后也可能失去算子效率。PP 要求可切分且 shape 契约稳定的图,stage 不平衡、气泡和跨 stage 状态会增加复杂度。
混合并行也不会自动解决数据加载、checkpoint 保存、故障恢复或数值等价性。mesh 越多维,rank 映射与 state dict 越需要被当成正式接口测试。
13 今天真正需要记住什么?#
- DP 切 batch,TP 切层内张量维,PP 切网络深度;先定位显存成分再选轴。
- 容量估算只给下界,all-gather、all-reduce、reduce-scatter 和 pipeline bubble 决定真实性能。
- TP 适合高带宽域,PP 适合可平衡的层段,状态分片 DP 是模型状态放不下时的常见第一步。
- 每增加一个并行维,都要用单卡基准验证输出、梯度、样本覆盖与归一化。
14 思考题与小练习#
- 一个模型状态 48 GB、激活 12 GB,使用 4 张 24 GB 卡。分别估算 4 路 DDP 与理想 4 路状态分片的每卡静态占用,并指出估算遗漏了什么。
- 对 写出两路列并行与行并行的每卡权重、输入、局部输出 shape,以及需要的 collective。
- 画出 3 个 stage、6 个 micro-batch 的前向流水线,计算理想气泡比例,并解释为何实际比例可能更高。
相关工作#
- Dean et al., Large Scale Distributed Deep Networks ↗,系统讨论模型并行与数据并行的早期大规模实践。
- Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism ↗,展示 Transformer 层内张量并行。
- Huang et al., GPipe: Efficient Training of Giant Neural Networks using Pipeline Parallelism ↗,用 micro-batch pipeline 扩展深层模型。
- Narayanan et al., Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM ↗,研究数据、张量与流水线并行的组合。
- Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models ↗,通过分片数据并行状态降低副本内存。
15 下一篇预告#
三种切分轴已经定位清楚。下一篇将聚焦状态分片数据并行:FSDP2 为什么在算某一层前 all-gather 参数、反向后 reduce-scatter 梯度,以及分片边界怎样同时决定峰值显存和通信重叠。