跳转至

分布式训练:并行策略

选择并行策略的第一问是:模型、激活与优化器状态能否在目标 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 给出的起点很清楚:

  1. 模型能放进单 GPU、想扩吞吐:先用 DDP。
  2. 模型放不进单 GPU:考虑 FSDP。
  3. 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、梯度与恢复,再扩大集群。

延伸阅读