AI 工程基础体系 · 第 64/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。

分布式训练:数据并行、模型并行、梯度同步与故障恢复

分布式训练不是简单地“把一个训练脚本启动多次”。它改变了训练过程中的数据分片、模型状态存放位置、梯度聚合方式、通信时序和故障语义。一个可恢复的生产训练系统至少需要回答四个问题:

  1. 每个进程处理哪些数据?
  2. 每个进程持有哪些模型参数和中间激活?
  3. 多个进程如何得到与单机训练一致或可解释的梯度?
  4. 某个进程、节点或存储系统失败后,从哪个一致状态继续训练?

本文以 PyTorch 分布式训练为主要背景,覆盖机器学习、深度学习和生成式 AI 中常见的数据并行、模型并行、梯度同步以及故障恢复机制。


一、先建立分布式训练的基本模型

设训练集为

D={(xi,yi)}i=1ND=\{(x_i,y_i)\}_{i=1}^{N}

模型参数为 θ\theta,单个样本的损失为

(fθ(xi),yi)\ell(f_\theta(x_i),y_i)

一个大小为 BB 的 mini-batch 上,通常使用平均损失:

L(θ)=1Bi=1Bi(θ)L(\theta)=\frac{1}{B}\sum_{i=1}^{B}\ell_i(\theta)

对应梯度为:

g=θL(θ)=1Bi=1Bθi(θ)g=\nabla_\theta L(\theta) =\frac{1}{B}\sum_{i=1}^{B}\nabla_\theta \ell_i(\theta)

单机训练的一步通常是:

  1. 从数据集中取一个 batch;
  2. 前向计算损失;
  3. 反向计算梯度;
  4. 优化器根据梯度更新参数;
  5. 清空梯度,进入下一步。

分布式训练引入了 WW 个训练进程,通常每个 GPU 对应一个进程。这里的 world size 是进程总数,rank 是进程编号,范围为 00W1W-1。如果一个节点有多个 GPU,那么每个节点通常运行多个进程;某个进程绑定的 GPU 编号称为 local rank

在最常见的数据并行场景中,第 rr 个进程拥有本地 batch:

BrB_r

全局 batch 是所有本地 batch 的并集:

Bglobal=r=0W1BrB_{\text{global}}=\bigcup_{r=0}^{W-1}B_r

如果每个本地 batch 大小为 bb,且没有梯度累积,则:

Bglobal=WbB_{\text{global}}=Wb

但这只是“一个优化步骤处理多少样本”的定义。数据并行还必须解决参数副本一致、梯度聚合、数据不重复和故障恢复问题。


二、数据并行:每个进程处理不同数据,但维护完整模型副本

2.1 数据并行的定义

**数据并行(Data Parallelism)**是指:

  • 每个训练进程保存一份完整模型参数;
  • 不同进程处理同一全局 batch 的不同数据切片;
  • 每个进程独立完成前向和反向;
  • 反向结束后,将各进程梯度聚合;
  • 聚合后,每个进程使用相同梯度更新自己的参数副本。

因此,数据并行复制的是模型,切分的是数据。

基本执行时序如下:

sequenceDiagram
    participant D0 as Rank 0
    participant D1 as Rank 1
    participant C as 通信组
    participant S as 数据集

    S->>D0: 本地 batch B0
    S->>D1: 本地 batch B1
    D0->>D0: 前向与反向,得到 g0
    D1->>D1: 前向与反向,得到 g1
    D0->>C: AllReduce(g0)
    D1->>C: AllReduce(g1)
    C-->>D0: 平均梯度 (g0 + g1) / 2
    C-->>D1: 平均梯度 (g0 + g1) / 2
    D0->>D0: optimizer.step()
    D1->>D1: optimizer.step()

通信操作通常通过 torch.distributed 提供的进程组完成。NCCL 常用于 NVIDIA GPU 之间的通信,Gloo 可用于 CPU,具体后端可用性取决于安装方式、设备和 PyTorch 版本。

2.2 为什么梯度平均后参数仍然一致

假设有两个进程,每个进程本地 batch 大小为 bb,本地平均梯度分别为:

g0=1biB0ig_0=\frac{1}{b}\sum_{i\in B_0}\nabla\ell_i

g1=1biB1ig_1=\frac{1}{b}\sum_{i\in B_1}\nabla\ell_i

如果 B0B1 不重叠,且全局 batch 大小为 2b2b,则全局平均梯度为:

gglobal=12b(iB0i+iB1i)g_{\text{global}} =\frac{1}{2b} \left( \sum_{i\in B_0}\nabla\ell_i+ \sum_{i\in B_1}\nabla\ell_i \right)

代入两个本地平均梯度:

gglobal=g0+g12g_{\text{global}}=\frac{g_0+g_1}{2}

因此,执行 all_reduce 的求和结果再除以进程数,就得到全局平均梯度。

一个具体例子如下。设某个参数只有一维梯度:

  • Rank 0 的本地 batch 梯度为 g0=2g_0=2
  • Rank 1 的本地 batch 梯度为 g1=6g_1=6

平均后:

gglobal=2+62=4g_{\text{global}}=\frac{2+6}{2}=4

如果当前参数为 θ=10\theta=10,学习率为 η=0.1\eta=0.1,使用最简单的 SGD:

θ=θηgglobal=100.1×4=9.6\theta'=\theta-\eta g_{\text{global}} =10-0.1\times4=9.6

只要两个进程都收到梯度 44,并使用相同的优化器状态和学习率,它们都会得到 9.69.6

2.3 DDP 与单机多卡 DataParallel 的区别

PyTorch 中常见的两种数据并行方式是:

  • torch.nn.DataParallel
  • torch.nn.parallel.DistributedDataParallel,简称 DDP

DataParallel 通常在一个进程中控制多个 GPU,主设备需要聚合输入、收集输出和处理梯度。它使用方便,但 Python 线程调度、主设备聚合和扩展性会成为瓶颈。

DDP 通常采用“一进程一 GPU”:

  • 每个进程独立执行前向和反向;
  • 梯度同步通过进程间 collective 通信完成;
  • 通信可以与反向计算重叠;
  • 可以跨节点运行。

在多 GPU 和多节点训练中,通常应优先考虑 DDP,而不是把多个 GPU 都塞进一个 Python 进程。

2.4 DDP 的梯度同步不是在 optimizer.step() 时才发生

DDP 在构造模型时为参数注册梯度相关的钩子。反向传播过程中,当一组参数的梯度就绪后,DDP 可以把这些梯度放入通信 bucket,并发起 all-reduce

这意味着典型流程是:

  1. 前向计算;
  2. 调用 loss.backward()
  3. 反向图从后向前计算;
  4. 某些梯度先就绪;
  5. 对应 bucket 开始通信;
  6. 其余层继续反向;
  7. 通信与剩余反向尽可能重叠;
  8. backward() 返回时,必要的梯度同步已经完成;
  9. 调用 optimizer.step()

因此,默认情况下:

loss.backward()
optimizer.step()

中的 optimizer.step() 通常已经使用了全局同步后的梯度。

DDP 的 bucket 化是常见实现机制,不应理解为模型并行。模型参数仍然完整存在于每个进程中,只是梯度通信被分组和调度。


三、数据切分:DistributedSampler 解决什么问题,不能解决什么问题

如果每个进程都直接遍历同一个 DataLoader,那么所有进程可能读取相同样本。这样做不会自动产生有效的数据并行,全球 batch 只是重复计算同一批数据。

DistributedSampler 的作用是根据 num_replicasrank 为每个进程分配不同索引。典型用法如下:

sampler = torch.utils.data.distributed.DistributedSampler(
    dataset,
    num_replicas=dist.get_world_size(),
    rank=dist.get_rank(),
    shuffle=True,
    drop_last=True,
)

loader = torch.utils.data.DataLoader(
    dataset,
    batch_size=local_batch_size,
    sampler=sampler,
    num_workers=4,
    pin_memory=True,
)

每个 epoch 开始时还应调用:

sampler.set_epoch(epoch)

原因是 DistributedSampler 通常根据 epoch 和随机种子生成新的打乱顺序。如果不调用 set_epoch,不同 epoch 可能重复使用相同的顺序。

3.1 drop_last 与数据量不整除

假设数据集大小为 N=10N=10,进程数为 W=3W=3

如果需要每个进程拿到相同数量的样本,采样器可能:

  • 丢弃一部分样本;
  • 或补齐一些索引,使各进程长度相同。

这会影响每个 epoch 的实际样本数。drop_last=True 可以避免最后一个不完整 batch,但会丢弃数据;drop_last=False 可能保留数据,却需要注意不同进程的 batch 数量和最后一个 batch 的形状。

DDP 的默认同步假设各进程大致以相同顺序执行 collective。如果某个进程提前结束 DataLoader,而其他进程仍然进入梯度通信,就可能发生挂起或通信错误。因此,数据加载器长度、最后 batch 策略和输入数据过滤逻辑必须在各 rank 之间一致。

3.2 数据并行不等于样本严格无重复

分布式采样器通常保证一个 epoch 内索引分片的基本规则,但在以下情况下可能出现重复或有效样本数量变化:

  • 数据集长度不能被进程数整除;
  • 采样器为了对齐长度而补齐索引;
  • 使用加权随机采样;
  • 动态过滤样本;
  • 各 rank 的数据预处理结果不同;
  • 训练恢复时从错误的位置重新开始。

因此,需要明确“重复样本”是采样策略允许的行为,还是数据管线 bug。


四、全局 batch、学习率与梯度累积

4.1 全局 batch 的计算

设:

  • GPU 数量或训练进程数为 WW
  • 每个进程的本地 batch 为 bb
  • 梯度累积步数为 KK

若每个 micro-batch 都调用一次反向,并在累积 KK 次后更新参数,则一次优化更新处理的有效 batch 大小通常是:

Beffective=W×b×KB_{\text{effective}}=W\times b\times K

例如:

  • 8 个进程;
  • 每进程 batch 为 4;
  • 梯度累积 8 次;

则一次参数更新对应:

8×4×8=2568\times4\times8=256

个样本。

但这只在每次损失按本地 micro-batch 正确归一化,并且累积期间不提前更新参数时成立。

4.2 no_sync() 的作用

DDP 默认每次 backward() 都同步梯度。进行梯度累积时,如果每个 micro-batch 都同步,会产生不必要的通信。可以使用:

for micro_step, (x, y) in enumerate(loader):
    should_sync = (micro_step + 1) % accumulation_steps == 0

    context = model.no_sync() if not should_sync else nullcontext()
    with context:
        output = model(x)
        loss = criterion(output, y) / accumulation_steps
        loss.backward()

    if should_sync:
        optimizer.step()
        optimizer.zero_grad(set_to_none=True)

这里除以 accumulation_steps 是为了让累积后的梯度近似于这些 micro-batch 梯度的平均值,而不是累加值。

no_sync() 只会抑制 DDP 梯度同步,不会阻止本地反向计算。最后一次 backward() 必须不在 no_sync() 中,否则各 rank 可能保留不同的本地梯度。

4.3 学习率线性缩放不是数学定律

当全局 batch 从 BB 增大到 kBkB 时,工程中常见学习率按 kk 线性放大。但这依赖优化器、模型、数据分布、warmup 和训练目标,并不保证严格等价。

对于普通 SGD,小 batch 多次更新与大 batch 少次更新本来就不是完全相同的轨迹:

θt+1=θtηgt\theta_{t+1} =\theta_t-\eta g_t

如果把多个小 batch 合并成一个大 batch,更新次数和每一步的参数位置都发生了变化。Adam、AdamW 等优化器还会维护一阶和二阶动量,batch 改变会影响动量统计。因此,改变进程数或本地 batch 后,需要重新验证学习率、warmup、训练步数和评测结果。


五、梯度同步:AllReduce、Reduce、Broadcast 和同步边界

5.1 Collective 通信原语

常用 collective 操作包括:

  • AllReduce:所有进程输入一个张量,执行求和、平均或最大值等归约,并把结果返回给所有进程;
  • Reduce:把所有进程的值归约到指定 root,其他进程不一定得到结果;
  • Broadcast:root 把数据发送给所有进程;
  • AllGather:收集所有进程的张量;
  • Barrier:等待所有进程到达同一同步点。

DDP 梯度同步通常使用 AllReduce。模型初始化时,参数副本也需要一致;常见实现会从某个进程广播参数和缓冲区,具体时机和细节由 DDP 实现负责。

5.2 梯度平均的正确性条件

“分布式训练等价于单机大 batch 训练”需要满足一组条件:

  1. 各进程本地样本集合组成目标全局 batch;
  2. 各本地损失的归一化方式一致;
  3. 梯度同步使用正确的求和或平均系数;
  4. 每个进程在更新前拥有相同参数;
  5. 优化器类型、超参数和优化器状态一致;
  6. 随机性不会引入未解释的差异;
  7. 模型中没有破坏同步假设的本地状态更新;
  8. 数据预处理、标签和有效 token 数在各 rank 上定义一致。

任何条件不满足,都不能简单声称“只是把 batch 放大了”。

5.3 反例:token 级损失的错误平均

语言模型训练中,一个 batch 可能包含不同数量的有效 token。假设:

  • Rank 0 有 100 个有效 token,平均 token loss 为 1;
  • Rank 1 有 10 个有效 token,平均 token loss 为 3。

如果直接对两个 rank 的平均 loss 再平均:

1+32=2\frac{1+3}{2}=2

但按有效 token 数计算的全局平均 loss 应为:

100×1+10×3100+10=1301101.182\frac{100\times1+10\times3}{100+10} =\frac{130}{110}\approx1.182

因此,对于带 padding 的生成式 AI 训练,若每个 rank 先对自己的有效 token 求平均,再对 rank 平均,可能得到错误的全局梯度权重。

一种更准确的做法是:

  1. 本地计算有效 token 的损失总和;
  2. 本地统计有效 token 数;
  3. 对损失总和和 token 数分别 AllReduce;
  4. 用全局损失总和除以全局 token 数进行归一化。

概念代码如下:

local_loss_sum = token_loss.masked_select(valid_mask).sum()
local_token_count = valid_mask.sum()

global_loss_sum = local_loss_sum.detach().clone()
global_token_count = local_token_count.detach().clone()

dist.all_reduce(global_loss_sum, op=dist.ReduceOp.SUM)
dist.all_reduce(global_token_count, op=dist.ReduceOp.SUM)

loss = local_loss_sum / global_token_count
loss.backward()

这里的关键是:用于反向的 local_loss_sum 保留本地计算图,而归一化分母 global_token_count 通常作为不需要梯度的标量参与除法。实际实现还需要确认张量设备、数据类型和混合精度行为。

5.4 梯度同步不是所有状态同步

DDP 主要同步参数梯度,不会自动让所有训练状态都“神奇地一致”。需要区分:

  • 参数 parameter
  • 非参数缓冲区 buffer,例如 BatchNorm 的 running mean;
  • 优化器状态,例如 Adam 的一阶、二阶矩;
  • 学习率调度器状态;
  • 随机数生成器状态;
  • 数据采样器状态;
  • 自定义缓存和计数器。

如果每个 rank 都从相同初始状态开始,并且每次使用相同的同步梯度更新,参数和优化器状态通常会保持一致。但自定义逻辑可能破坏这一点。

5.5 BatchNorm 是典型边界

普通 BatchNorm 使用当前进程本地 batch 计算均值和方差。数据并行后,每个 GPU 看到的是更小的本地 batch,因此其统计量与单机全局 batch 的统计量不同。

如果需要跨进程统计 BatchNorm,PyTorch 提供 SyncBatchNorm 相关机制,但它会增加通信,且通常需要在包装 DDP 前完成转换。对于 Transformer,LayerNorm 更常见,不依赖 batch 统计,因此不会遇到同样的跨 rank BatchNorm 问题。


六、混合精度与梯度缩放对同步的影响

生成式 AI 训练通常使用 FP16 或 BF16 混合精度,以降低显存和通信成本。

典型结构是:

with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
    output = model(x)
    loss = criterion(output, y)

loss.backward()
optimizer.step()

使用 FP16 时,常见做法是配合梯度缩放器,避免小梯度下溢。无论是否使用缩放,都必须保证:

  • 所有 rank 对溢出或非法梯度采取一致的更新决定;
  • 跳过某次 optimizer.step() 时,各 rank 不能一部分更新、一部分不更新;
  • 梯度裁剪应在正确的反缩放之后执行;
  • 检查点需要保存缩放器状态,若缩放器包含可恢复状态。

梯度裁剪应作用于同步后的全局梯度。例如:

scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()

如果在梯度同步前只对本地梯度裁剪,各 rank 的裁剪操作与“先全局聚合、再裁剪”通常不等价。


七、模型并行:切分模型本身,而不是切分样本

7.1 模型并行的定义

**模型并行(Model Parallelism)**是指将一个模型的参数、计算或激活分布到多个设备,使单个设备不必保存完整模型或完整计算图。

它与数据并行的核心区别是:

方式 参数存放 数据处理
数据并行 每个进程通常有完整模型副本 不同进程处理不同样本
模型并行 参数或计算图分布在多个设备 一个样本可能经过多个设备
混合并行 两者同时使用 每个并行组有不同职责

当模型参数、优化器状态或中间激活超过单 GPU 显存时,仅增加数据并行副本不能解决问题,因为每个副本本身仍然放不下模型。

7.2 张量并行

**张量并行(Tensor Parallelism)**把某个层的矩阵运算拆分到多个设备。

设全连接层为:

Y=XWY=XW

其中:

  • XRm×dX\in\mathbb{R}^{m\times d}
  • WRd×hW\in\mathbb{R}^{d\times h}
  • YRm×hY\in\mathbb{R}^{m\times h}

如果按输出维度切分:

W=[W0,W1,,Wp1]W=[W_0,W_1,\ldots,W_{p-1}]

则每个设备计算:

Yj=XWjY_j=XW_j

最后需要把各设备的 YjY_j 拼接起来。这是列并行。

如果按输入维度切分:

W=[W0W1Wp1],X=[X0,X1,,Xp1]W= \begin{bmatrix} W_0\\ W_1\\ \vdots\\ W_{p-1} \end{bmatrix}, \qquad X=[X_0,X_1,\ldots,X_{p-1}]

则:

Y=j=0p1XjWjY=\sum_{j=0}^{p-1}X_jW_j

每个设备产生部分结果,之后需要 AllReduce 求和。这是行并行。

矩阵切分减少了单个设备持有的权重,但增加了设备间通信。通信是否划算取决于:

  • 设备间带宽;
  • batch 和序列长度;
  • 张量切分维度;
  • 层的计算量;
  • 通信是否能与计算重叠。

7.3 Pipeline Parallelism

**流水线并行(Pipeline Parallelism)**按网络层或模块切分模型。例如:

  • Stage 0:Embedding 和前若干 Transformer 层;
  • Stage 1:中间 Transformer 层;
  • Stage 2:后若干层和输出头。

一个样本的激活在 stage 之间传递。若一次只处理一个 batch,后面的设备必须等待前面的设备,设备利用率很低。流水线并行通常把 batch 切成多个 micro-batch,让不同 stage 同时处理不同 micro-batch。

时序可简化为:

时间 1: Stage 0 处理 M0
时间 2: Stage 0 处理 M1,Stage 1 处理 M0
时间 3: Stage 0 处理 M2,Stage 1 处理 M1,Stage 2 处理 M0
...

流水线调度需要处理:

  • warmup:前端先产生激活;
  • steady state:各 stage 并行工作;
  • cooldown:后端完成剩余反向;
  • 激活保存:反向时需要中间结果;
  • micro-batch 数量与 pipeline stage 数的关系;
  • stage 间发送和接收的阻塞关系。

流水线并行的常见问题是 pipeline bubble,即某些时间段设备没有有效计算。增加 micro-batch 通常可以降低气泡比例,但会增加调度复杂度和激活管理成本。

7.4 参数分片不应混同于传统模型并行

**参数分片(Parameter Sharding)**把参数、梯度或优化器状态分布到不同进程,而不是让每个进程保存完整副本。它可以显著降低内存占用,但其通信语义与传统张量并行不同。

PyTorch 的 FSDP(Fully Sharded Data Parallel)属于数据并行体系下的参数、梯度和优化器状态分片方案。它常被称为“分片数据并行”,不等同于把 Transformer 层按张量或 stage 切开的模型并行。

区分这几种方式很重要:

  • DDP:完整参数副本 + 数据切分;
  • FSDP:参数、梯度和优化器状态分片 + 数据切分;
  • 张量并行:同一层内部张量切分;
  • 流水线并行:不同层或模块切分;
  • 混合并行:组合上述策略。

八、混合并行的进程组设计

大型模型通常同时使用数据并行、张量并行和流水线并行。假设总进程数为:

W=D×T×PW=D\times T\times P

其中:

  • DD:数据并行组数量;
  • TT:张量并行组大小;
  • PP:流水线 stage 数量。

每个进程同时属于三个逻辑组:

  1. 张量并行组:在同一层的切分计算中通信;
  2. 流水线并行组:在相邻 stage 之间传递激活和梯度;
  3. 数据并行组:在相同模型分片之间同步梯度。

同一张卡上的计算可能需要不同通信组。如果组划分错误,可能出现:

  • 本应同步张量分片的进程没有加入同一组;
  • 数据并行 AllReduce 把不属于同一模型分片的进程混在一起;
  • rank 映射与物理拓扑不匹配,导致通信性能下降;
  • 某个组等待另一个组,形成死锁。

因此,混合并行的 rank 映射不是装饰性配置,而是训练正确性和性能的一部分。


九、一个可运行的 PyTorch DDP 示例

下面的示例展示一个最小但完整的数据并行训练生命周期。它使用:

  • torchrun 启动多个进程;
  • init_process_group 初始化通信;
  • DistributedSampler 切分数据;
  • DistributedDataParallel 同步梯度;
  • rank 0 保存检查点;
  • 所有 rank 从同一检查点恢复。

保存为 train_ddp.py

import os
import random
from pathlib import Path

import torch
import torch.distributed as dist
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, TensorDataset
from torch.utils.data.distributed import DistributedSampler


def seed_everything(seed: int) -> None:
    random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)


class MLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(32, 128),
            nn.ReLU(),
            nn.Linear(128, 2),
        )

    def forward(self, x):
        return self.net(x)


def atomic_save(obj, path: Path) -> None:
    tmp_path = path.with_suffix(path.suffix + ".tmp")
    torch.save(obj, tmp_path)
    os.replace(tmp_path, path)


def main():
    # torchrun 会注入这些环境变量
    rank = int(os.environ["RANK"])
    local_rank = int(os.environ["LOCAL_RANK"])
    world_size = int(os.environ["WORLD_SIZE"])

    torch.cuda.set_device(local_rank)
    device = torch.device("cuda", local_rank)

    # NCCL 用于 GPU 通信;初始化完成后,各 rank 才能进入 collective
    dist.init_process_group(backend="nccl")

    # 先固定初始随机性。真实项目还需根据数据管线保存和恢复 RNG 状态。
    seed_everything(1234)

    x = torch.randn(4096, 32)
    y = (x[:, :4].sum(dim=1) > 0).long()
    dataset = TensorDataset(x, y)

    sampler = DistributedSampler(
        dataset,
        num_replicas=world_size,
        rank=rank,
        shuffle=True,
        drop_last=True,
    )
    loader = DataLoader(
        dataset,
        batch_size=64,
        sampler=sampler,
        num_workers=2,
        pin_memory=True,
    )

    model = MLP().to(device)
    model = DDP(model, device_ids=[local_rank])
    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
    criterion = nn.CrossEntropyLoss()

    checkpoint_path = Path("checkpoint.pt")
    start_epoch = 0

    # 只由 rank 0 读取文件,避免多个进程同时争抢共享存储。
    if rank == 0 and checkpoint_path.exists():
        checkpoint = torch.load(checkpoint_path, map_location=device)
        model.module.load_state_dict(checkpoint["model"])
        optimizer.load_state_dict(checkpoint["optimizer"])
        start_epoch = checkpoint["epoch"] + 1

    # 将 rank 0 恢复的参数和优化器状态传给其他 rank。
    # 对大型优化器状态,直接 broadcast 可能成本较高;
    # 生产系统可让每个 rank 从共享存储读取同一检查点,或使用框架提供的分片状态恢复。
    dist.barrier()
    if world_size > 1:
        for tensor in model.module.state_dict().values():
            if torch.is_tensor(tensor):
                dist.broadcast(tensor, src=0)

    # 这里仅演示参数同步。实际 optimizer state 的跨 rank 同步应使用
    # 统一加载策略;不同版本/设备下直接手工广播 optimizer.state 可能不稳妥。
    # 最简单可靠的做法是让每个 rank 都从检查点加载 optimizer state。

    if rank != 0 and checkpoint_path.exists():
        checkpoint = torch.load(checkpoint_path, map_location=device)
        optimizer.load_state_dict(checkpoint["optimizer"])
        start_epoch = checkpoint["epoch"] + 1

    for epoch in range(start_epoch, 5):
        sampler.set_epoch(epoch)
        model.train()

        for step, (features, labels) in enumerate(loader):
            features = features.to(device, non_blocking=True)
            labels = labels.to(device, non_blocking=True)

            optimizer.zero_grad(set_to_none=True)
            logits = model(features)
            loss = criterion(logits, labels)
            loss.backward()       # DDP 在反向过程中同步梯度
            optimizer.step()

            if step % 20 == 0 and rank == 0:
                print(
                    f"epoch={epoch}, step={step}, "
                    f"loss={loss.item():.4f}"
                )

        # 保存前让各 rank 完成本 epoch,避免 rank 0 保存时其他 rank
        # 仍在使用可能被替换的外部状态。
        dist.barrier()

        if rank == 0:
            state = {
                "epoch": epoch,
                "model": model.module.state_dict(),
                "optimizer": optimizer.state_dict(),
            }
            atomic_save(state, checkpoint_path)
            print(f"saved {checkpoint_path}")

        dist.barrier()

    dist.destroy_process_group()


if __name__ == "__main__":
    main()

单机两张 GPU 上运行:

torchrun --standalone --nproc-per-node=2 train_ddp.py

预期现象是:

  • 启动两个训练进程;
  • 每个进程绑定一张 GPU;
  • rank 0 周期性打印 loss;
  • 每个 epoch 结束后生成 checkpoint.pt
  • 删除或保留检查点重新启动时,程序从检查点记录的下一个 epoch 继续。

9.1 这个示例中的恢复边界

上面的示例适合说明机制,但有几个生产边界:

  1. 检查点只在 epoch 结束保存,因此中途失败会丢失当前 epoch 的进度;
  2. 只保存了模型和优化器,没有保存学习率调度器、AMP scaler、随机状态和采样器游标;
  3. 如果 rank 0 在保存期间失败,临时文件和目标文件的处理需要检查;
  4. 恢复后数据顺序不一定与失败前完全相同;
  5. 如果 world size 变化,原有数据分片和优化器状态的语义可能变化;
  6. 示例中的状态广播仅用于解释参数同步,生产代码应采用明确、经过版本验证的统一状态加载策略。

十、故障恢复:从“重新启动进程”到“恢复一致训练状态”

10.1 故障的不同层次

分布式训练中的故障至少分为以下几类:

  • 进程故障:Python 进程崩溃、CUDA out of memory、未捕获异常;
  • 设备故障:GPU Xid、驱动异常、设备不可用;
  • 节点故障:机器断电、网络断开、容器被驱逐;
  • 通信故障:NCCL 超时、链路异常、某个 rank 没有进入 collective;
  • 存储故障:检查点写入中断、共享存储不可用、文件损坏;
  • 数据故障:某个样本损坏、数据读取阻塞、各 rank 数据长度不一致;
  • 逻辑故障:只有部分 rank 执行了保存、更新或跳过步骤。

这些故障的恢复方式不同。一个 rank 退出后,其他 rank 通常不能继续安全执行原进程组中的 collective;通信库可能报告错误,也可能在超时前表现为挂起。多数场景需要结束整个训练作业,再由作业管理器重新启动。

10.2 检查点必须保存哪些状态

一个训练检查点不是只有 model.state_dict()。至少应考虑:

checkpoint = {
    "model": model.state_dict(),
    "optimizer": optimizer.state_dict(),
    "scheduler": scheduler.state_dict(),
    "scaler": scaler.state_dict(),       # 使用 AMP 时
    "epoch": epoch,
    "global_step": global_step,
    "best_metric": best_metric,
}

为了更接近“从失败点继续”,还应考虑:

  • Python random 状态;
  • CPU RNG 状态;
  • CUDA RNG 状态;
  • 数据增强 RNG;
  • sampler 的 epoch 和位置;
  • 当前数据 shard 或数据流 offset;
  • 梯度累积中的 micro-step;
  • 混合并行的 stage 和并行组状态;
  • 训练配置、代码版本、数据版本;
  • 模型和优化器的 dtype、设备映射信息。

但“保存 RNG 状态”也不自动保证位级复现。DataLoader worker、异步预取、非确定性 CUDA 算子、通信顺序和数据读取顺序都可能引入差异。

10.3 一致检查点的写入流程

安全的检查点写入通常采用:

  1. 选择保存协调者,常见为 rank 0;
  2. 训练进程在逻辑安全点同步;
  3. 将状态写入临时文件;
  4. 刷新并关闭文件;
  5. 使用原子重命名替换正式文件;
  6. 写入校验信息或元数据;
  7. 通知其他进程检查点已完成;
  8. 保留最近若干个已验证版本。

os.replace(tmp, target) 在同一文件系统内通常具有原子替换语义,但不能把它当作跨对象、跨存储系统的完整事务。对象存储、网络文件系统和分布式文件系统的可见性及一致性行为可能不同,需要结合实际存储验证。

更稳妥的目录结构可以是:

checkpoints/
  step-000100/
    model.pt
    optimizer.pt
    metadata.json
    COMPLETE
  step-000200/
    model.pt
    optimizer.pt
    metadata.json
    COMPLETE
  latest -> step-000200

只有当所有必要文件写完并校验成功后,才创建 COMPLETE 标志。恢复时只选择存在完整标志且校验通过的目录,避免加载半写入的检查点。

10.4 中断恢复与精确恢复不是一回事

假设每 1000 个优化步骤保存一次检查点,训练在第 1750 步失败。重启后从第 1000 步恢复,会重做第 1001 到 1750 步。

这称为 回滚恢复(rollback recovery)。它通常可接受,但会产生:

  • 重复计算;
  • 训练数据可能重复读取;
  • 日志中的 global step 需要明确语义;
  • 外部副作用可能重复,例如写入评测结果或上传模型。

如果要求从第 1750 步精确继续,就必须保存更细粒度的训练游标和随机状态。保存频率越高,恢复点越近,但检查点的 I/O、存储成本和对训练的干扰也越大。

10.5 外部副作用必须幂等

训练不仅更新内存中的模型,也可能:

  • 写入指标数据库;
  • 上传模型;
  • 注册实验;
  • 更新模型仓库;
  • 发送通知;
  • 生成评测结果。

如果作业从旧检查点重放,外部操作可能重复执行。应使用由 run_id + global_step 或检查点版本组成的幂等键,并让下游系统能够识别重复提交。


十一、torchrun 与弹性恢复的边界

torchrun 负责启动和管理训练进程,提供诸如 RANKLOCAL_RANKWORLD_SIZE 等环境变量。使用固定 world size 时,它可以启动多进程训练;若训练进程退出,是否自动重启、重启多少次、如何分配节点,取决于启动参数和上层作业调度系统。

需要区分:

  • 进程重启:启动器重新拉起进程;
  • 训练状态恢复:新进程从检查点加载参数和优化器状态;
  • 弹性训练:允许成员加入或离开,并重新形成进程组;
  • 精确继续:数据游标、随机状态和优化器状态都按照指定位置继续。

进程被自动重启不代表训练自动恢复。如果代码每次启动都从随机初始化开始,启动器只能让它“重新跑”,而不能恢复之前的训练。

当 world size 变化时,训练语义也可能变化:

Bglobal=W×b×KB_{\text{global}}=W\times b\times K

如果 WW 改变而本地 batch 和累积步数不变,全局 batch 就改变。除此之外:

  • 数据分片方式改变;
  • 每个 epoch 的 step 数改变;
  • 学习率调度器的步数解释改变;
  • BatchNorm 本地统计改变;
  • 通信组拓扑改变;
  • FSDP 或混合并行的分片布局可能需要重新处理。

因此,弹性恢复必须明确采用哪种语义:

  1. 接受 world size 改变,并重新定义全局 batch 和调度;
  2. 固定 world size,只允许相同规模重启;
  3. 使用支持重分片的检查点格式;
  4. 把 global step 定义为优化更新次数,而不是进程本地循环次数。

十二、常见失败路径与诊断方法

12.1 程序卡在 loss.backward()

可能原因包括:

  • 某个 rank 的输入 batch 数量不同;
  • 某个 rank 在数据读取时异常或阻塞;
  • 条件分支导致不同 rank 使用了不同参数;
  • 某个 rank 跳过了反向;
  • DDP 的 unused parameter 配置不匹配;
  • NCCL 通信链路或设备异常。

诊断方法:

  1. 为每个 rank 打印进入和离开前向、反向、优化器更新的日志;
  2. 记录 epoch、step、batch shape 和样本数量;
  3. 使用超时初始化进程组,避免无限等待;
  4. 检查各 rank 是否都执行了同样的 collective;
  5. 在小数据集和单机多卡环境复现;
  6. 查看 NCCL 和 GPU 驱动日志。

分布式 collective 的一个重要规则是:参与同一进程组的进程必须以兼容顺序调用对应操作。Rank 0 进入了 all_reduce,而 rank 1 进入了 broadcast,即使两个操作分别都合法,也可能造成死锁。

12.2 loss 在不同 rank 上不一致

DDP 默认不要求每个 rank 的本地 loss 数值相同,因为每个 rank 处理不同数据。可以区分:

  • 本地 loss:当前 rank batch 的损失;
  • 全局 loss:跨 rank 汇总后的损失;
  • 日志 loss:用于展示的聚合结果。

如果要记录全局平均 loss,应按样本数或有效 token 数聚合,而不是简单平均不等长 batch 的本地平均值。

分类任务中,如果每个 rank 样本数相同,简单平均通常可以:

loss_sum = loss.detach().clone()
dist.all_reduce(loss_sum, op=dist.ReduceOp.SUM)
global_loss = loss_sum / dist.get_world_size()

如果样本数不同,则应同时 AllReduce 样本数,再计算加权平均。

12.3 参数在 rank 之间逐渐不一致

可能原因包括:

  • 某个 rank 跳过了 optimizer.step()
  • 优化器状态没有正确加载;
  • 使用了 rank-specific 的学习率或随机逻辑;
  • 自定义参数更新没有同步;
  • 混合精度溢出判断在 rank 之间不一致;
  • 某些参数没有参与反向;
  • 模型缓冲区被本地修改。

可以在调试阶段计算参数校验值,例如:

with torch.no_grad():
    checksum = torch.zeros(1, device=device)
    for p in model.parameters():
        checksum += p.float().sum()

    checksums = [torch.zeros_like(checksum)
                 for _ in range(dist.get_world_size())]
    dist.all_gather(checksums, checksum)

    if rank == 0:
        print([v.item() for v in checksums])

这不是严格的哈希校验,可能发生抵消,只适合快速诊断。更严格的检查需要对参数字节或分块摘要进行验证,但会增加开销。

12.4 显存不足不一定应该增加数据并行

如果单个模型副本已经无法放入一张 GPU,增加数据并行 GPU 数量仍然不能解决问题,因为每个 rank 依然需要完整模型。

应根据内存来源选择方案:

  • 参数太大:参数分片、张量并行或流水线并行;
  • 优化器状态太大:优化器状态分片、低精度状态或更节省内存的优化器;
  • 激活太大:activation checkpointing、减少序列长度、流水线并行;
  • 本地 batch 太大:减小 micro-batch,使用梯度累积;
  • 临时张量过多:检查算子、生命周期和显存碎片。

12.5 通信占用过高但 GPU 利用率不高

数据并行的通信量主要与梯度规模相关。模型参数越大,每一步同步的梯度越多。常见缓解方向包括:

  • 增大计算量与通信量的比例;
  • 使用通信与反向重叠;
  • 选择合适的 bucket 大小;
  • 使用梯度累积并减少同步频率;
  • 使用参数或梯度分片;
  • 优化节点拓扑和 GPU 亲和性;
  • 避免把跨节点通信放在高频张量并行路径上。

但减少同步频率会改变训练语义,不能只看吞吐量。例如梯度累积期间不同 rank 的参数仍保持不变,最后一次同步的梯度才用于更新;如果自定义代码在累积期间修改参数,就会破坏这一假设。


十三、评测、权限和成本也属于分布式训练系统

13.1 评测不能默认由所有 rank 重复执行

训练结束后,所有 rank 都调用完整评测,会重复读取数据和计算结果。常见做法是:

  • 只让 rank 0 执行评测;
  • 或让各 rank 分片评测,再聚合预测结果和指标;
  • 对生成式任务,按样本 ID 去重,避免采样器补齐造成重复样本影响指标。

指标聚合也要遵循与训练 loss 类似的加权原则。准确率应聚合正确样本数和总样本数,而不是简单平均每个 rank 的准确率。

13.2 检查点权限和敏感数据

模型检查点可能包含:

  • 训练数据记忆;
  • 优化器状态;
  • 访问令牌或实验配置;
  • 用户数据路径;
  • 生成模型的安全策略和词表。

因此,保存路径、对象存储权限、加密、生命周期策略和下载审计都需要纳入设计。共享文件系统上给所有训练进程写权限虽然方便,但会扩大误删和篡改风险。更合理的方式是限制写入角色,使用临时对象和经过校验的发布流程。

13.3 成本不只是 GPU 租用费

分布式训练成本包括:

总成本=计算成本+通信成本+存储成本+故障重算成本+评测与监控成本\text{总成本} = \text{计算成本} + \text{通信成本} + \text{存储成本} + \text{故障重算成本} + \text{评测与监控成本}

如果检查点间隔太长,单次故障可能导致数小时重算;如果检查点过于频繁,则保存本身会降低吞吐并增加对象存储费用。应通过实际故障恢复时间、检查点耗时和重算成本选择间隔,而不是固定套用某个数字。


十四、设计一个可解释的分布式训练方案

对于一个 Transformer 训练任务,可以按以下顺序决定并行策略:

第一步:确认单卡是否能放下模型

如果完整模型、优化器状态和必要激活都能放入单卡,优先使用 DDP 做数据并行。

第二步:识别内存瓶颈

如果模型参数能放下但优化器状态或激活放不下,可分别考虑:

  • 优化器状态分片;
  • FSDP;
  • activation checkpointing;
  • 减小 micro-batch;
  • 梯度累积。

如果模型参数本身就放不下,则需要模型切分或参数分片。

第三步:确认通信拓扑

张量并行需要高频、低延迟通信,通常更适合放在同一节点或高速互联设备之间。数据并行通信频率较低或数据量更容易聚合,但梯度规模很大时也会成为瓶颈。

第四步:确定训练等价关系

记录并验证:

Beffective=D×b×KB_{\text{effective}}=D\times b\times K

同时明确:

  • loss 按样本还是按 token 归一化;
  • 学习率是否随全局 batch 调整;
  • scheduler 按 epoch 还是按 optimizer step;
  • 最后一个 batch 如何处理;
  • world size 变化后是否允许继续训练。

第五步:定义恢复点

至少保存:

  • 模型参数;
  • 优化器状态;
  • scheduler;
  • AMP scaler;
  • global step;
  • 数据和代码版本;
  • 检查点完整性标记。

如果训练作业需要从节点故障中自动恢复,还要验证:

  1. 作业管理器能否重启所有进程;
  2. 新进程能否访问检查点;
  3. 检查点是否覆盖最近一次可接受的训练进度;
  4. 恢复后的并行组和 world size 是否兼容;
  5. 恢复后评测指标和日志是否不会重复污染。

十五、核心边界总结

数据并行的核心是“不同数据、相同模型副本”;模型并行的核心是“一个模型分布在多个设备”。两者可以组合,但它们的通信目标不同:

  • 数据并行同步梯度;
  • 张量并行同步层内部分结果;
  • 流水线并行传递激活和反向梯度;
  • 参数分片在需要时重新聚合或分发参数。

梯度同步默认只保证 DDP 管理范围内的梯度聚合。它不自动修复错误的数据切分、不自动统一自定义状态、不自动保存优化器,也不保证带不等长有效 token 的 loss 归一化正确。

故障恢复的关键不是“重新启动”,而是“从一个明确、一致、可验证的训练状态重新启动”。这个状态既包括模型,也包括优化器、调度器、随机性、数据位置、并行拓扑和外部副作用的幂等语义。

在 PyTorch 中,torch.distributed、DDP、DistributedSamplertorchrun 提供了分布式训练的基础构件;它们并不替代对 batch 语义、通信顺序、检查点一致性和生产故障路径的设计。只有把这些状态和因果关系明确下来,增加 GPU 数量才不会从“加速训练”变成“放大不确定性”。


系列导航与关联阅读

官方资料

本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。