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

模型检查点与训练恢复:状态、随机数、分片、容错和一致性

模型检查点(checkpoint)不是“把模型参数保存到文件”。它是训练过程在某个时间点的可恢复状态快照。若只保存 model.state_dict(),通常只能完成“加载参数后继续训练”,不能保证训练轨迹连续,更不能保证分布式训练、混合精度、数据顺序和随机增强在故障后正确恢复。

训练恢复(resume)也有不同强度:

  1. 参数恢复:恢复模型权重,用于推理或从已有权重开始新的训练。
  2. 优化器恢复:同时恢复优化器状态,使动量、二阶矩等历史信息连续。
  3. 训练进度恢复:恢复 epoch、step、学习率调度器、梯度累积位置等控制状态。
  4. 随机性恢复:恢复 Python、NumPy、PyTorch CPU/CUDA 等随机数状态。
  5. 执行轨迹恢复:在相同环境、相同数据顺序和相同确定性条件下,尽量从故障点继续产生相同结果。

最后一种要求最强。即使保存了所有显式状态,也不一定能跨硬件、CUDA 内核、通信拓扑或数据加载实现获得逐位相同的结果。


一、先定义“训练状态”而不是只定义“模型”

设第 tt 个训练更新前的完整状态为:

St=(Wt,Ot,Lt,At,Rt,Dt,Ct,Et)S_t = (W_t, O_t, L_t, A_t, R_t, D_t, C_t, E_t)

其中:

  • WtW_t:模型参数和必要的缓冲区,例如 BatchNorm 的运行均值与方差;
  • OtO_t:优化器状态,例如 Adam 的一阶矩、二阶矩和步计数;
  • LtL_t:学习率调度器状态;
  • AtA_t:自动混合精度的梯度缩放器状态;
  • RtR_t:随机数生成器状态;
  • DtD_t:数据迭代状态,包括采样顺序、已消费位置和分布式 sampler 状态;
  • CtC_t:训练控制状态,例如 epoch、global step、梯度累积位置;
  • EtE_t:环境与实验元数据,例如代码版本、配置、词表版本和数据版本。

一次训练更新可以抽象为:

St+1=F(St,xt)S_{t+1} = F(S_t, x_t)

这里 xtx_t 是本次更新取到的训练输入,可能还包括随机数据增强、dropout 掩码和负样本采样结果。要让恢复后的训练得到同一个 St+1S_{t+1},至少需要满足:

  1. 恢复的状态等于保存时的状态;
  2. 恢复后取到同一个 xtx_t
  3. 随机数生成器在需要随机数时处于相同位置;
  4. F 的实现、数值精度和并行归约顺序等价。

因此,下面两种文件的含义不同:

torch.save(model.state_dict(), "weights.pt")

它表达的是:

这组模型参数可以被加载。

而下面这种检查点表达的是:

训练在某个明确的更新边界上暂停,之后可以尝试从该边界继续。

torch.save({
    "model": model.state_dict(),
    "optimizer": optimizer.state_dict(),
    "scheduler": scheduler.state_dict(),
    "scaler": scaler.state_dict(),
    "step": step,
}, "checkpoint.pt")

二、参数、缓冲区和优化器状态分别决定什么

2.1 模型状态不只包含可训练参数

PyTorch 的 state_dict() 通常包含两类内容:

  • named_parameters() 返回的参数;
  • register_buffer() 注册的缓冲区。

BatchNorm 的 running_meanrunning_var 就是典型缓冲区。只保存参数而不保存缓冲区,模型在评估模式下可能表现不同。

模型结构本身通常不在 state_dict() 中。加载时必须先构造兼容的模型对象:

model = MyModel(config)
model.load_state_dict(checkpoint["model"])

因此,模型配置、词表、分词器、类别映射和输出头定义也必须被版本化。一个权重文件无法独立证明“它应该加载到哪种结构上”。

2.2 优化器状态决定更新规则的历史

以 Adam 为例,对参数 ww 和梯度 gtg_t,其状态包括:

mt=β1mt1+(1β1)gtm_t = \beta_1 m_{t-1} + (1-\beta_1)g_t

vt=β2vt1+(1β2)gt2v_t = \beta_2 v_{t-1} + (1-\beta_2)g_t^2

参数更新使用 mtm_tvtv_t 以及步数 tt。如果只恢复 wtw_t,却重新创建一个没有历史的 Adam,那么下一步更新会使用新的 m0,v0m_0,v_0,这不是原训练轨迹的延续。

这会产生一个常见误解:

“模型参数一样,所以恢复训练应该一样。”

对于 SGD 且没有 momentum 时,这个近似更接近事实;对于 Adam、AdamW、带 momentum 的 SGD、LAMB 等优化器,通常不成立。

优化器状态还可能比模型参数占用更多空间。例如混合精度训练中,优化器常保存 FP32 master weights、momentum 和 variance。大模型的检查点成本往往主要来自优化器,而不是模型权重本身。

2.3 学习率调度器必须与更新次数对齐

调度器的关键不是当前学习率这一项,而是它如何计算下一次学习率。它通常保存:

  • last_epoch 或类似计数;
  • warmup 状态;
  • 周期位置;
  • 最小学习率、基准学习率等配置。

若训练在第 10,000 个 optimizer step 保存,恢复后又错误地调用一次 scheduler.step(),学习率就会提前前进一个位置。反过来,如果保存发生在 scheduler.step() 前后不同的边界,也会出现一个 step 的偏移。

应先明确更新顺序。例如:

读取 batch
→ forward
→ backward
→ optimizer.step()
→ scheduler.step()
→ global_step += 1
→ 保存检查点

那么检查点中的 global_step 应表示“已经完成的 optimizer 更新数”,而不是“刚要执行的更新数”。恢复后应从下一个 batch 开始,并避免再次执行已完成的 scheduler.step()


三、梯度累积和混合精度引入了额外边界

3.1 梯度累积不能只保存 global step

假设每 kk 个 micro-batch 执行一次参数更新:

micro-batch 0: backward,保留梯度
micro-batch 1: backward,保留梯度
...
micro-batch k-1: backward,optimizer.step()

global_step 通常统计 optimizer 更新次数,而不是 micro-batch 数。如果在一个累积周期的中间保存:

  • 模型参数还没有更新;
  • 优化器状态还没有更新;
  • 参数的 .grad 中已经存在部分累积梯度;
  • 数据游标已经消费了若干 micro-batch。

若不保存梯度和 micro_step,恢复时必须丢弃当前累积周期,回退数据游标,或者接受一次不同的更新。最稳妥的工程选择通常是只在累积边界保存:

if (micro_step + 1) % grad_accum_steps == 0:
    optimizer.step()
    optimizer.zero_grad(set_to_none=True)
    scheduler.step()
    global_step += 1
    save_checkpoint()

如果业务要求任意时刻容错,就必须明确保存:

  • 每个参数当前的 .grad
  • 当前 micro-batch 位置;
  • 梯度缩放器状态;
  • 是否已经执行过 unscale、梯度裁剪或 optimizer step。

3.2 AMP 的梯度缩放器也是训练状态

自动混合精度通常通过梯度缩放避免 FP16 梯度下溢。GradScaler 会根据溢出情况调整 scale:

  • 梯度溢出时跳过本次 optimizer 更新并降低 scale;
  • 连续稳定时提高 scale。

因此需要保存:

scaler.state_dict()

恢复时:

scaler.load_state_dict(checkpoint["scaler"])

如果漏掉它,恢复后的训练可能在相同梯度下采取不同的“跳过更新”决策。使用 BF16 时通常不需要同样的梯度缩放,但具体训练代码仍可能使用 scaler,不能仅根据模型数据类型推断状态是否存在。


四、随机数状态与数据顺序

4.1 随机数不是一个全局变量

一个进程中可能同时存在多个随机数来源:

  • Python 的 random
  • NumPy 的随机数生成器;
  • PyTorch CPU generator;
  • 每个 CUDA 设备上的 PyTorch generator;
  • 数据增强库或第三方库自己的 generator;
  • DataLoader worker 中独立的随机状态。

PyTorch 提供了常用的获取和恢复接口:

import random
import numpy as np
import torch

def capture_rng_state():
    state = {
        "python": random.getstate(),
        "numpy": np.random.get_state(),
        "torch_cpu": torch.get_rng_state(),
    }
    if torch.cuda.is_available():
        state["torch_cuda_all"] = torch.cuda.get_rng_state_all()
    return state

def restore_rng_state(state):
    random.setstate(state["python"])
    np.random.set_state(state["numpy"])
    torch.set_rng_state(state["torch_cpu"])
    if torch.cuda.is_available() and "torch_cuda_all" in state:
        torch.cuda.set_rng_state_all(state["torch_cuda_all"])

这些接口保存的是生成器的内部状态,而不是一个简单整数。恢复时必须在随机操作发生前调用,否则已经消耗的随机数无法追回。

4.2 保存 RNG 仍不等于保存数据迭代器

下面两种实现都可能使用随机数,但恢复难度不同:

for batch in DataLoader(dataset, shuffle=True):
    ...
indices = torch.randperm(len(dataset), generator=generator)
for i in range(0, len(indices), batch_size):
    batch = dataset[indices[i:i+batch_size]]

第一种把顺序隐藏在 sampler、DataLoader 和 worker 中;第二种可以显式保存 indices 和当前位置。对于需要精确恢复的数据管线,显式状态通常更容易验证。

一个 epoch 级别的 sampler 至少需要表达:

{
    "epoch": 3,
    "permutation_seed": 12345,
    "position": 64000,
    "dataset_version": "sha256:..."
}

只保存 epoch=3 不够,因为无法判断本 epoch 已经消费到哪一个样本。

4.3 DataLoader worker 是精确恢复的边界

多进程 DataLoader 可能有:

  • worker 自己的随机种子;
  • 预取队列中已经生成但尚未交给训练循环的 batch;
  • persistent worker 的长期状态;
  • 自定义 collate、增强和缓存。

故障发生时,队列中哪些 batch 已经生成、哪些 batch 已经被模型消费,可能难以精确重建。因此:

  • 追求可复现时,减少隐式 worker 状态,使用显式 generator;
  • 追求吞吐时,接受恢复后少量数据顺序变化;
  • 无论选择哪种,都要把“恢复保证”写清楚,而不是笼统地声称可复现。

4.4 分布式训练中每个 rank 都有自己的随机状态

在 DDP 中,每个 rank 是独立进程,拥有独立的 CPU/CUDA RNG。只保存 rank 0 的 RNG 状态,不能恢复其他 rank 的随机序列。

常见做法是每个 rank 写入自己的文件:

checkpoint/
├── manifest.json
├── model.pt
├── optimizer.pt
├── rng-rank-0.pt
├── rng-rank-1.pt
└── ...

恢复时由对应 rank 读取自己的 RNG。若 world size 改变,旧 rank 与新 rank 不再一一对应,原来的随机序列和数据分片通常不能直接延续。


五、一个可运行的单进程恢复示例

下面示例使用显式 batch 索引,演示以下边界:

  • 模型、优化器、调度器、AMP scaler;
  • Python、NumPy、PyTorch RNG;
  • epoch、batch 游标和 global step;
  • 先写临时文件,再原子替换;
  • 在“完成一个 optimizer step 后”保存。

它没有使用多进程 DataLoader,因此数据顺序状态是显式且可检查的。

from pathlib import Path
import os
import random
import numpy as np
import torch
from torch import nn
from torch.utils.data import TensorDataset

DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
CKPT = Path("run/checkpoint.pt")
BATCH_SIZE = 32
EPOCHS = 3
SEED = 1234


def seed_everything(seed: int):
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    if torch.cuda.is_available():
        torch.cuda.manual_seed_all(seed)


def capture_rng():
    result = {
        "python": random.getstate(),
        "numpy": np.random.get_state(),
        "torch_cpu": torch.get_rng_state(),
    }
    if torch.cuda.is_available():
        result["torch_cuda_all"] = torch.cuda.get_rng_state_all()
    return result


def restore_rng(state):
    random.setstate(state["python"])
    np.random.set_state(state["numpy"])
    torch.set_rng_state(state["torch_cpu"])
    if torch.cuda.is_available() and "torch_cuda_all" in state:
        torch.cuda.set_rng_state_all(state["torch_cuda_all"])


def atomic_torch_save(obj, path: Path):
    path.parent.mkdir(parents=True, exist_ok=True)
    tmp = path.with_name(path.name + ".tmp")
    with open(tmp, "wb") as f:
        torch.save(obj, f)
        f.flush()
        os.fsync(f.fileno())
    os.replace(tmp, path)


def make_batch_order(n, batch_size, epoch, seed):
    # 每个 epoch 的排列由固定公式生成,因此 epoch 和 seed 足以重建完整顺序。
    generator = torch.Generator(device="cpu")
    generator.manual_seed(seed + epoch)
    permutation = torch.randperm(n, generator=generator)
    return [
        permutation[i:i + batch_size]
        for i in range(0, n, batch_size)
    ]


def build_objects():
    model = nn.Sequential(
        nn.Linear(20, 64),
        nn.ReLU(),
        nn.Linear(64, 2),
    ).to(DEVICE)

    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
        optimizer, T_max=EPOCHS * 100
    )
    scaler = torch.amp.GradScaler(
        "cuda", enabled=(DEVICE == "cuda")
    )
    return model, optimizer, scheduler, scaler


def save_state(model, optimizer, scheduler, scaler,
               epoch, batch_pos, global_step, data_seed):
    state = {
        "format_version": 1,
        "model": model.state_dict(),
        "optimizer": optimizer.state_dict(),
        "scheduler": scheduler.state_dict(),
        "scaler": scaler.state_dict(),
        "epoch": epoch,
        # batch_pos 指向下一个尚未消费的 batch
        "batch_pos": batch_pos,
        "global_step": global_step,
        "data_seed": data_seed,
        "rng": capture_rng(),
        "torch_version": torch.__version__,
    }
    atomic_torch_save(state, CKPT)


def load_state(model, optimizer, scheduler, scaler):
    # 只加载可信来源的 checkpoint;torch.load 可能反序列化 Python 对象。
    state = torch.load(CKPT, map_location=DEVICE, weights_only=False)

    if state.get("format_version") != 1:
        raise RuntimeError("不支持的 checkpoint 格式")

    model.load_state_dict(state["model"])
    optimizer.load_state_dict(state["optimizer"])
    scheduler.load_state_dict(state["scheduler"])
    scaler.load_state_dict(state["scaler"])

    restore_rng(state["rng"])
    return (
        state["epoch"],
        state["batch_pos"],
        state["global_step"],
        state["data_seed"],
    )


def train(resume=False):
    seed_everything(SEED)

    x = torch.randn(1000, 20)
    y = (x[:, :2].sum(dim=1) > 0).long()
    dataset = TensorDataset(x, y)

    model, optimizer, scheduler, scaler = build_objects()

    start_epoch = 0
    start_batch_pos = 0
    global_step = 0
    data_seed = 9000

    if resume and CKPT.exists():
        start_epoch, start_batch_pos, global_step, data_seed = load_state(
            model, optimizer, scheduler, scaler
        )
        print(
            f"resumed: epoch={start_epoch}, "
            f"batch_pos={start_batch_pos}, step={global_step}"
        )

    model.train()

    for epoch in range(start_epoch, EPOCHS):
        batches = make_batch_order(
            len(dataset), BATCH_SIZE, epoch, data_seed
        )
        batch_pos = start_batch_pos if epoch == start_epoch else 0

        for pos in range(batch_pos, len(batches)):
            indices = batches[pos]
            xb = dataset.tensors[0][indices].to(DEVICE)
            yb = dataset.tensors[1][indices].to(DEVICE)

            optimizer.zero_grad(set_to_none=True)

            with torch.autocast(
                device_type=DEVICE,
                enabled=(DEVICE == "cuda")
            ):
                logits = model(xb)
                loss = nn.functional.cross_entropy(logits, yb)

            scaler.scale(loss).backward()
            scaler.step(optimizer)
            scaler.update()
            scheduler.step()
            global_step += 1

            # 保存完成本次更新后的边界;下次从 pos + 1 开始。
            save_state(
                model, optimizer, scheduler, scaler,
                epoch=epoch,
                batch_pos=pos + 1,
                global_step=global_step,
                data_seed=data_seed,
            )

        start_batch_pos = 0

    print("finished:", global_step)


if __name__ == "__main__":
    train(resume=False)
    # 删除或保留 checkpoint 后,执行 train(resume=True) 可测试恢复路径。

这个示例为什么能避免重复 batch

保存发生在:

batch pos
→ forward/backward
→ optimizer.step
→ scheduler.step
→ global_step += 1
→ 保存 batch_pos = pos + 1

所以 batch_pos 表示“下一个尚未消费的位置”。恢复后循环从该位置开始,已经完成的 batch 不会再次训练。

如果把保存放到 optimizer step 之前,却仍然记录 batch_pos = pos + 1,恢复后就会跳过一个尚未完成的更新;如果保存后才递增游标,恢复后则会重复更新。游标和更新状态必须属于同一个一致性边界。

这个示例仍有明确限制:

  • 每次保存一个完整文件,频繁保存会影响吞吐;
  • 没有 DataLoader worker;
  • 没有分布式通信;
  • 不保证跨设备或跨 PyTorch/CUDA 版本逐位一致;
  • torch.load 对不可信文件有反序列化风险,应只加载可信检查点。具体 weights_only 行为随 PyTorch 版本变化,不能把它当成完整的安全边界。

六、检查点的一致性:必须保存一个“状态切面”

6.1 什么是一致性检查点

一致性检查点要求文件中的各个状态属于同一个训练边界。例如:

模型参数:第 1000 步之后
优化器状态:第 999 步之后
学习率调度器:第 1000 步之后
数据游标:第 1001 个 batch

这组状态无法共同描述一个合法的训练时刻。恢复后可能表现为:

  • optimizer 使用错误的动量;
  • scheduler 多走或少走一步;
  • 数据重复或跳过;
  • AMP 在错误的 scale 上继续;
  • 多个 rank 看到不同的 global step。

一致性不是“文件能否成功反序列化”,而是“所有相互关联的状态是否来自同一个逻辑提交点”。

6.2 单文件写入和原子替换

单进程中常见的安全写法是:

  1. 写入临时文件;
  2. flush
  3. fsync
  4. 使用 os.replace 替换目标文件。

os.replace 通常能保证同一文件系统内的目录项替换具有原子性:读者要么看到旧文件,要么看到新文件,而不是半个文件。但这不等于所有存储系统都提供相同语义。网络文件系统、对象存储挂载层和云端上传接口可能不支持真正的原子 rename。

更稳妥的目录式提交方式是:

checkpoint-0001000/
├── manifest.json
├── model.pt
├── optimizer.pt
├── rng-rank-0.pt
└── COMPLETE

写入顺序为:

创建临时目录
→ 写所有分片
→ 校验大小和哈希
→ 写 manifest.json
→ 写 COMPLETE
→ 将该目录登记为 latest

恢复程序只接受存在 COMPLETE 且 manifest 校验通过的目录。这样即使进程在写 optimizer 分片时被杀死,也不会把半成品当成可恢复版本。

6.3 manifest 应记录什么

一个实际的 manifest 可以包含:

{
  "format_version": 3,
  "global_step": 100000,
  "epoch": 4,
  "world_size": 8,
  "model_arch": "decoder-transformer-v2",
  "config_sha256": "...",
  "dataset_version": "...",
  "tokenizer_version": "...",
  "code_commit": "...",
  "torch_version": "...",
  "cuda_version": "...",
  "files": {
    "model-rank0.pt": {
      "bytes": 123456789,
      "sha256": "..."
    }
  }
}

这里的哈希用于检测传输损坏或截断,不用于证明文件内容“语义正确”。语义验证还需要实际加载、检查 key、形状、dtype、参数数量,并执行一次前向或小规模验证。


七、分布式训练中的复制、分片和保存布局

7.1 DDP 与 FSDP 的状态不同

在 DistributedDataParallel(DDP)中,每个 rank 通常持有一份完整模型副本。梯度通过通信同步,优化器一般也在各 rank 持有相应状态。因此保存时有两种思路:

  • 只让 rank 0 保存模型和优化器;
  • 每个 rank 保存自己的随机数、数据 sampler 和其他本地状态。

只保存 rank 0 的文件,通常不足以恢复 rank-local 的执行轨迹。

在 FullyShardedDataParallel(FSDP)中,参数、梯度或优化器状态可能被切分到不同 rank。此时不能把一个 rank 上的局部 state_dict() 当成完整模型检查点。需要明确选择:

  • full state dict:聚合成完整状态,便于单卡推理和跨规模加载,但可能产生显著 CPU/GPU 内存峰值;
  • sharded state dict:保持分片,保存和恢复更适合大模型,但必须使用兼容的分片布局和加载流程;
  • local state:只描述当前 rank 的局部状态,迁移和变更 world size 的能力最弱。

PyTorch 的 FSDP 和 torch.distributed.checkpoint 相关 API 在不同版本中持续演进,尤其是 state dict 配置、planner 和保存接口。生产代码应锁定 PyTorch 版本,直接按该版本官方文档验证 API,而不能复制旧版本示例后假定语义不变。

7.2 分片检查点的读写过程

分片保存可以抽象为:

flowchart LR
    T[训练进程组] --> B[在一致性边界暂停更新]
    B --> S[收集模型/优化器/进度/RNG状态]
    S --> P[按参数或优化器状态切分]
    P --> W0[rank 0 分片]
    P --> W1[rank 1 分片]
    P --> WN[rank N 分片]
    W0 --> M[manifest 与校验信息]
    W1 --> M
    WN --> M
    M --> C[写 COMPLETE]
    C --> L[登记 latest]

恢复时则反向执行:

flowchart LR
    L[读取 latest] --> V[校验 manifest]
    V --> R[按当前或兼容拓扑规划读取]
    R --> G[重组模型/优化器状态]
    G --> X[恢复进度、数据游标和 rank RNG]
    X --> B[所有 rank barrier]
    B --> T[继续训练]

关键点不在于“每个 rank 写一个文件”,而在于分片文件、逻辑 key 和 manifest 共同描述一个完整状态。若只保留某些分片,模型可能仍能加载一部分 key,但这不是可恢复检查点。

7.3 world size 改变不是简单的参数加载

从 8 个 rank 保存,改为 4 个 rank 恢复,至少会影响:

  • 数据如何重新分片;
  • sampler 的 epoch 和位置如何解释;
  • 每个 rank 的 RNG;
  • FSDP 参数和 optimizer state 的重新分配;
  • 通信归约顺序;
  • 每个 rank 应读取哪些分片。

模型权重有时可以通过 full state dict 迁移到新规模,但这不代表训练轨迹可精确延续。对于 optimizer state,尤其是按参数分片的状态,必须使用支持重分片的加载逻辑,不能手工按文件名拼接。


八、故障路径与容错策略

训练作业可能在以下时刻失败:

  1. forward 期间进程被终止;
  2. backward 已完成但 optimizer step 尚未完成;
  3. optimizer step 已完成但 checkpoint 尚未提交;
  4. 部分 rank 已写完,另一些 rank 仍在写;
  5. 对象存储上传完成了一部分;
  6. 主进程成功退出,但 checkpoint 元数据未更新。

因此恢复系统需要定义“最后一个可提交点”,而不是简单地每隔几分钟调用一次保存。

8.1 保存频率的实际含义

如果每 KK 个 optimizer step 保存一次,故障后最多重算约 K1K-1 个 step,但保存成本会减少。保存成本包括:

  • 序列化;
  • GPU 到 CPU 或主机内存拷贝;
  • 多 rank 通信;
  • 文件系统写入;
  • 对象存储上传;
  • 校验和计算。

对于大模型,保存期间阻塞训练会降低吞吐。异步保存可以减少主训练线程等待,但不能直接把正在变化的张量对象交给后台线程并假设它们稳定。必须在快照时形成稳定副本,例如:

  • 在一致性边界复制到 CPU;
  • 使用框架提供的安全 snapshot 机制;
  • 采用双缓冲或写时复制;
  • 明确后台保存期间不得修改快照缓冲区。

否则后台线程可能读到“模型参数已更新、优化器状态尚未更新”的混合状态。

8.2 preemption 与信号处理

云实例抢占、集群时间限制或作业调度器终止进程时,可以捕获终止信号并触发保存,但不能假设一定有足够时间完成大文件写入。信号处理程序应尽量:

  1. 设置“请求保存”标志;
  2. 让训练主循环在安全边界处理;
  3. 由主循环执行正常的 barrier 和快照;
  4. 设置超时,避免无限等待;
  5. 保存失败时保留旧的有效 checkpoint。

不要在异步信号处理函数中直接调用复杂的序列化和分布式通信 API,这些操作通常不是异步信号安全的。

8.3 保留策略和回滚

生产训练一般同时保留:

  • 最新检查点;
  • 若干按 step 编号的历史检查点;
  • 验证指标最优的检查点;
  • 任务配置和数据版本对应的不可变归档。

“latest” 只是一个指针,不能覆盖历史文件。若最新版本损坏,应能回退到上一个 manifest 校验通过的版本。指标最优模型和最新训练状态也不是一回事:前者适合部署,后者适合继续训练。


九、恢复流程不是 load_state_dict 一行代码

完整恢复可以分为以下阶段:

9.1 验证运行环境

先比较:

  • 模型配置;
  • 词表和 tokenizer;
  • 数据集版本;
  • 代码提交;
  • PyTorch、CUDA 和关键依赖版本;
  • world size、混合精度设置;
  • 参数 dtype 和设备布局。

版本不一致时有三种处理方式:

  • 严格拒绝;
  • 明确执行迁移;
  • 只加载模型权重,开始一个新的训练阶段。

最危险的是静默兼容:文件能加载,但调度器、数据处理或参数语义已经改变。

9.2 构造对象并加载状态

通常顺序是:

读取 manifest
→ 构造模型
→ 构造优化器并绑定同一组参数
→ 构造 scheduler/scaler
→ 加载 model
→ 加载 optimizer、scheduler、scaler
→ 恢复进度和 RNG
→ 恢复 sampler/data cursor
→ barrier
→ 继续训练

优化器必须在加载状态前以兼容的参数组构造。若参数组顺序、参数数量或 weight decay 分组发生变化,优化器状态可能无法正确对应。

9.3 恢复后立即做验证

恢复后不要直接跑几个小时再观察 loss。应执行:

  • 检查缺失和多余的 state dict key;
  • 检查参数形状、dtype、设备;
  • 检查 optimizer state 的数量和形状;
  • 检查 global_step 与 scheduler 计数;
  • 检查下一个 batch 的样本 ID 或哈希;
  • 执行一次 forward;
  • 必要时执行一次完整的 optimizer step;
  • 与保存前记录的 loss、梯度范数、学习率进行比较。

如果追求确定性,可以在测试中保存“恢复前下一步”的输入、loss 和参数摘要,再比较恢复后结果。不要只比较最终指标,因为最终指标无法定位是哪一个状态发生了偏差。


十、常见失败表现与诊断路径

10.1 恢复后 loss 突然跳变

可能原因包括:

  • 数据游标回退或跳过;
  • optimizer 状态未恢复;
  • scheduler 多走一步;
  • AMP scaler 重置;
  • dropout 或数据增强 RNG 不同;
  • 恢复在梯度累积中间,但没有恢复 .grad

诊断顺序应从最容易验证的状态开始:

下一个 batch 是否相同
→ 模型参数摘要是否相同
→ optimizer step 和学习率是否相同
→ scaler 状态是否相同
→ RNG 状态和随机算子是否相同
→ 分布式归约及硬件差异

10.2 只有 rank 0 能恢复,其他 rank 卡住

这通常不是模型文件问题,而是分布式控制流不一致。例如:

  • rank 0 已进入保存或加载,其他 rank 没有进入;
  • 某个 rank 在文件写入失败后退出;
  • barrier 两侧调用次数不一致;
  • 只有一个 rank 更新了恢复状态,其他 rank 仍使用旧值。

所有 rank 必须遵守相同的阶段顺序。文件 I/O 可以由部分 rank 执行,但状态广播、barrier 和错误传播必须有明确协议。

10.3 Missing keyUnexpected key

常见原因:

  • 保存的是 DDP 包装后的模型,key 带有 module. 前缀;
  • FSDP full/sharded/local state dict 类型不匹配;
  • 模型结构或配置变更;
  • 检查点被截断;
  • 误把 optimizer state 当成 model state 加载。

不应一律使用 strict=False 掩盖问题。它适合明确的迁移场景,例如新增一个任务头;对于意外缺失核心层,应该直接失败并输出具体 key 差异。

10.4 “随机种子一样”但结果仍不同

随机种子只是初始化随机数生成器的起点,不是完整执行历史。以下情况仍会造成差异:

  • checkpoint 保存后恢复时 RNG 已被额外消耗;
  • DataLoader worker 启动顺序不同;
  • CUDA 非确定性算子;
  • atomic 或通信归约顺序不同;
  • 不同 GPU 架构或 CUDA 内核;
  • 多线程 CPU 算法的执行顺序不同;
  • 只保存了 rank 0 的 RNG。

应区分:

  • 可复现:在约束环境中重复运行得到相同或近似结果;
  • 可恢复:故障后能继续训练;
  • 逐位一致:每个张量元素都相同。

这三者的要求逐级提高,不能混用。


十一、评测、权限和成本也属于检查点系统

11.1 评测状态不能污染训练状态

生成式 AI 训练通常伴随:

  • 验证集指标;
  • 生成样本;
  • reward 或 preference 评测;
  • 安全评测;
  • 最优模型选择。

评测结果应记录对应的 checkpoint ID、代码版本、数据版本和评测配置。不能只保存一个“best score”,否则无法证明该分数对应哪组权重和哪次评测。

评测本身也可能消耗 RNG。如果训练和评测共享同一进程 RNG,评测调用会改变训练后续的随机序列。可选方案包括:

  • 使用独立的 torch.Generator
  • 在评测前后保存并恢复训练 RNG;
  • 将评测放到独立进程;
  • 接受评测会改变随机轨迹,并在可复现声明中说明。

11.2 权限和敏感数据

检查点可能包含:

  • 模型权重;
  • optimizer state 中的训练痕迹;
  • 训练配置;
  • 数据处理元数据;
  • 可能被意外写入的密钥或路径;
  • 生成模型中的敏感能力或内部信息。

对象存储权限应遵循最小权限原则:

  • 训练作业只获得目标前缀的读写权限;
  • 推理服务只读模型权重,不读取 optimizer 和训练数据元数据;
  • 归档与删除权限分离;
  • 对象存储启用服务端加密或应用层加密;
  • 不把 token、密码或临时凭据写入 checkpoint。

尤其要避免把不可信 checkpoint 直接交给 torch.load。PyTorch 检查点常依赖 Python 序列化机制,加载前应确认来源可信,并根据目标版本选择更受限的加载方式。

11.3 成本取舍

检查点成本可以粗略写成:

成本保存频率×(序列化时间+存储容量+传输容量)\text{成本} \approx \text{保存频率} \times (\text{序列化时间}+\text{存储容量}+\text{传输容量})

只保存模型权重最便宜,但恢复能力最弱;保存完整 optimizer、RNG 和 sampler 状态成本更高,却能减少重算和故障后的不确定性。

实际系统常采用分层策略:

  • 高频保存轻量的恢复点;
  • 低频保存完整训练点;
  • 更低频保存用于部署的模型权重;
  • 对 optimizer state 使用压缩、分片或独立存储;
  • 通过保留策略删除过期版本,但保留回滚窗口。

不能只看文件大小。频繁保存导致的训练停顿、对象存储请求数量和恢复读取时间,也属于总成本。


十二、参数检查点、训练检查点和发布制品应分开

三类产物的目标不同:

类型 主要用途 典型内容
参数检查点 推理、微调起点 模型权重、配置、词表
训练检查点 故障恢复 参数、优化器、调度器、scaler、RNG、数据进度
发布制品 部署和审计 权重、推理配置、模型签名、评测报告、依赖锁定

把三者混为一个文件,会导致权限过大、下载过慢或恢复信息不足。生成式 AI 模型还需要特别记录 tokenizer、special tokens、上下文长度、量化配置和 generation 配置;仅有 Transformer 参数无法保证推理行为一致。


十三、一个可操作的恢复契约

生产系统应明确写出检查点的恢复契约,而不是只在代码中隐含约定。契约至少回答:

  1. 检查点保存在哪个训练边界?
  2. global_step 统计 micro-batch 还是 optimizer step?
  3. 恢复后下一个数据样本如何确定?
  4. 是否恢复 optimizer、scheduler 和 scaler?
  5. 是否保证 rank-local RNG?
  6. 是否支持改变 world size?
  7. 是否支持改变模型配置、词表或数据版本?
  8. 哪些版本不兼容时必须失败?
  9. 如何判断文件完整?
  10. 最新文件损坏时回滚到哪个版本?
  11. 恢复成功后如何验证?
  12. 检查点由谁读取、谁可以删除?

一个可靠的检查点系统最终应满足这样的因果链:

明确的更新边界
→ 一致地冻结相关状态
→ 完整写入并校验
→ 原子提交可见版本
→ 按兼容拓扑加载
→ 恢复数据进度和随机状态
→ 通过小规模验证
→ 再继续训练

检查点的核心不是序列化 API,而是对训练状态、数据流、随机性、分布式布局和故障提交点建立统一定义。只有这些状态在同一个逻辑时刻被保存,恢复才真正具有可解释性;只有把恢复保证的边界写清楚,工程团队才能区分“从权重重新开始”“从训练进度继续”和“尽量复现原执行轨迹”。


系列导航与关联阅读

官方资料

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