分布式训练:并行策略¶
选择并行策略的第一问是:模型、激活与优化器状态能否在目标 batch/sequence 下放进单卡? 第二问才是如何扩吞吐。
并行方式对比¶
| 方式 | 核心做法 | 主要通信 | 适用条件 | 主要代价 |
|---|---|---|---|---|
| DDP / 数据并行 | 每卡完整模型,不同数据 | 梯度 all-reduce | 模型单卡可放下 | 每卡复制模型与优化器状态 |
| FSDP / ZeRO | 跨 rank 分片参数、梯度、优化器 | all-gather + reduce-scatter | 状态单卡放不下 | 通信更多、状态管理更复杂 |
| Tensor Parallel | 单层矩阵沿维度切分 | 高频 all-reduce/all-gather | 单层或模型单卡放不下 | 强依赖低时延高带宽互连 |
| Pipeline Parallel | 不同层放不同 stage | stage 间激活/梯度 | 模型纵向可切分 | pipeline bubble、调度复杂 |
| Expert Parallel | MoE experts 分布到不同 rank | all-to-all token dispatch | 稀疏 MoE | 负载不均与网络压力 |
PyTorch 官方建议的基本路径¶
PyTorch Distributed Overview 给出的起点很清楚:
- 模型能放进单 GPU、想扩吞吐:先用 DDP。
- 模型放不进单 GPU:考虑 FSDP。
- FSDP 仍触及扩展边界:组合 Tensor Parallel / Pipeline Parallel,形成多维并行。
这是起点而不是自动答案;真实选择还要考虑拓扑、序列长度、容错、checkpoint 和团队运维复杂度。
DDP 的关键路径¶
sequenceDiagram
participant R0 as Rank 0
participant R1 as Rank 1
participant R2 as Rank 2
R0->>R0: forward + backward
R1->>R1: forward + backward
R2->>R2: forward + backward
R0<<->>R1: gradient buckets all-reduce
R1<<->>R2: gradient buckets all-reduce
Note over R0,R2: 通信可与后续梯度计算重叠
R0->>R0: optimizer step
R1->>R1: optimizer step
R2->>R2: optimizer step DDP 在各进程保留模型副本;backward 过程中 reducer 把梯度组织成 bucket 并进行同步。调 bucket 大小的本质,是在“更早发起通信”和“更大消息效率”之间权衡。
FSDP 的显存换通信¶
FSDP 的 FULL_SHARD 会分片参数、梯度和优化器状态;计算某层前通过 all-gather 临时恢复所需参数,反向后用 reduce-scatter 同步并重新分片。
收益是显著减少单 rank 常驻状态,代价包括:
- forward/backward 引入更多 collective;
- wrap 粒度影响峰值显存、通信频率和 overlap;
- full state dict 的收集可能让 rank 0 CPU 内存成为瓶颈;
- checkpoint 格式、world size 变化与恢复流程更复杂。
多维并行的拓扑映射¶
优先把通信最频繁、最敏感的维度放在最快链路:
- Tensor Parallel 常放在单节点 NVLink/NVSwitch 域内。
- Data Parallel 可跨节点,但梯度通信仍需足够网络带宽。
- Pipeline Parallel 的 stage 边界需平衡计算与激活传输。
- MoE Expert Parallel 对 all-to-all 非常敏感,要关注负载与 fabric。
flowchart TB
subgraph Node0[Node 0 · 快互连]
A0[TP 0] --- A1[TP 1]
A1 --- A2[TP 2]
A2 --- A3[TP 3]
end
subgraph Node1[Node 1 · 快互连]
B0[TP 0] --- B1[TP 1]
B1 --- B2[TP 2]
B2 --- B3[TP 3]
end
A0 <--> B0
A1 <--> B1
A2 <--> B2
A3 <--> B3 图中节点内是 TP group,跨节点同位置 rank 组成 DP group。这只是常见映射,最终应依据拓扑和 profiler 验证。
Global batch 与收敛语义¶
常见近似:
\[ B_{global}=B_{micro}\times N_{data\ parallel}\times N_{accumulation} \]
改变 GPU 数、micro batch 或梯度累积会改变 global batch;学习率、warmup、数据 sampler 与随机种子策略也可能需要调整。扩容后吞吐上升但收敛变差,不能只当 Infra 性能问题。
启动作业前的检查¶
- world size、rank 与 local rank 映射正确。
- 每个进程绑定唯一设备,CPU/NUMA 亲和合理。
- 所有 rank 的模型初始化与数据 sampler 语义一致。
- rendezvous 地址可达,超时与失败策略明确。
- checkpoint 同时包含模型、优化器、scheduler、scaler、随机状态和训练进度。
- 使用小规模 smoke test 验证 loss、梯度与恢复,再扩大集群。