AI 工程基础体系 · 第 67/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。
模型检查点与训练恢复:状态、随机数、分片、容错和一致性
模型检查点(checkpoint)不是“把模型参数保存到文件”。它是训练过程在某个时间点的可恢复状态快照。若只保存 model.state_dict(),通常只能完成“加载参数后继续训练”,不能保证训练轨迹连续,更不能保证分布式训练、混合精度、数据顺序和随机增强在故障后正确恢复。
训练恢复(resume)也有不同强度:
- 参数恢复:恢复模型权重,用于推理或从已有权重开始新的训练。
- 优化器恢复:同时恢复优化器状态,使动量、二阶矩等历史信息连续。
- 训练进度恢复:恢复 epoch、step、学习率调度器、梯度累积位置等控制状态。
- 随机性恢复:恢复 Python、NumPy、PyTorch CPU/CUDA 等随机数状态。
- 执行轨迹恢复:在相同环境、相同数据顺序和相同确定性条件下,尽量从故障点继续产生相同结果。
最后一种要求最强。即使保存了所有显式状态,也不一定能跨硬件、CUDA 内核、通信拓扑或数据加载实现获得逐位相同的结果。
一、先定义“训练状态”而不是只定义“模型”
设第 个训练更新前的完整状态为:
其中:
- :模型参数和必要的缓冲区,例如 BatchNorm 的运行均值与方差;
- :优化器状态,例如 Adam 的一阶矩、二阶矩和步计数;
- :学习率调度器状态;
- :自动混合精度的梯度缩放器状态;
- :随机数生成器状态;
- :数据迭代状态,包括采样顺序、已消费位置和分布式 sampler 状态;
- :训练控制状态,例如 epoch、global step、梯度累积位置;
- :环境与实验元数据,例如代码版本、配置、词表版本和数据版本。
一次训练更新可以抽象为:
这里 是本次更新取到的训练输入,可能还包括随机数据增强、dropout 掩码和负样本采样结果。要让恢复后的训练得到同一个 ,至少需要满足:
- 恢复的状态等于保存时的状态;
- 恢复后取到同一个 ;
- 随机数生成器在需要随机数时处于相同位置;
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_mean、running_var 就是典型缓冲区。只保存参数而不保存缓冲区,模型在评估模式下可能表现不同。
模型结构本身通常不在 state_dict() 中。加载时必须先构造兼容的模型对象:
model = MyModel(config)
model.load_state_dict(checkpoint["model"])
因此,模型配置、词表、分词器、类别映射和输出头定义也必须被版本化。一个权重文件无法独立证明“它应该加载到哪种结构上”。
2.2 优化器状态决定更新规则的历史
以 Adam 为例,对参数 和梯度 ,其状态包括:
参数更新使用 、 以及步数 。如果只恢复 ,却重新创建一个没有历史的 Adam,那么下一步更新会使用新的 ,这不是原训练轨迹的延续。
这会产生一个常见误解:
“模型参数一样,所以恢复训练应该一样。”
对于 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
假设每 个 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 单文件写入和原子替换
单进程中常见的安全写法是:
- 写入临时文件;
flush;fsync;- 使用
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,尤其是按参数分片的状态,必须使用支持重分片的加载逻辑,不能手工按文件名拼接。
八、故障路径与容错策略
训练作业可能在以下时刻失败:
- forward 期间进程被终止;
- backward 已完成但 optimizer step 尚未完成;
- optimizer step 已完成但 checkpoint 尚未提交;
- 部分 rank 已写完,另一些 rank 仍在写;
- 对象存储上传完成了一部分;
- 主进程成功退出,但 checkpoint 元数据未更新。
因此恢复系统需要定义“最后一个可提交点”,而不是简单地每隔几分钟调用一次保存。
8.1 保存频率的实际含义
如果每 个 optimizer step 保存一次,故障后最多重算约 个 step,但保存成本会减少。保存成本包括:
- 序列化;
- GPU 到 CPU 或主机内存拷贝;
- 多 rank 通信;
- 文件系统写入;
- 对象存储上传;
- 校验和计算。
对于大模型,保存期间阻塞训练会降低吞吐。异步保存可以减少主训练线程等待,但不能直接把正在变化的张量对象交给后台线程并假设它们稳定。必须在快照时形成稳定副本,例如:
- 在一致性边界复制到 CPU;
- 使用框架提供的安全 snapshot 机制;
- 采用双缓冲或写时复制;
- 明确后台保存期间不得修改快照缓冲区。
否则后台线程可能读到“模型参数已更新、优化器状态尚未更新”的混合状态。
8.2 preemption 与信号处理
云实例抢占、集群时间限制或作业调度器终止进程时,可以捕获终止信号并触发保存,但不能假设一定有足够时间完成大文件写入。信号处理程序应尽量:
- 设置“请求保存”标志;
- 让训练主循环在安全边界处理;
- 由主循环执行正常的 barrier 和快照;
- 设置超时,避免无限等待;
- 保存失败时保留旧的有效 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 key 或 Unexpected 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 成本取舍
检查点成本可以粗略写成:
只保存模型权重最便宜,但恢复能力最弱;保存完整 optimizer、RNG 和 sampler 状态成本更高,却能减少重算和故障后的不确定性。
实际系统常采用分层策略:
- 高频保存轻量的恢复点;
- 低频保存完整训练点;
- 更低频保存用于部署的模型权重;
- 对 optimizer state 使用压缩、分片或独立存储;
- 通过保留策略删除过期版本,但保留回滚窗口。
不能只看文件大小。频繁保存导致的训练停顿、对象存储请求数量和恢复读取时间,也属于总成本。
十二、参数检查点、训练检查点和发布制品应分开
三类产物的目标不同:
| 类型 | 主要用途 | 典型内容 |
|---|---|---|
| 参数检查点 | 推理、微调起点 | 模型权重、配置、词表 |
| 训练检查点 | 故障恢复 | 参数、优化器、调度器、scaler、RNG、数据进度 |
| 发布制品 | 部署和审计 | 权重、推理配置、模型签名、评测报告、依赖锁定 |
把三者混为一个文件,会导致权限过大、下载过慢或恢复信息不足。生成式 AI 模型还需要特别记录 tokenizer、special tokens、上下文长度、量化配置和 generation 配置;仅有 Transformer 参数无法保证推理行为一致。
十三、一个可操作的恢复契约
生产系统应明确写出检查点的恢复契约,而不是只在代码中隐含约定。契约至少回答:
- 检查点保存在哪个训练边界?
global_step统计 micro-batch 还是 optimizer step?- 恢复后下一个数据样本如何确定?
- 是否恢复 optimizer、scheduler 和 scaler?
- 是否保证 rank-local RNG?
- 是否支持改变 world size?
- 是否支持改变模型配置、词表或数据版本?
- 哪些版本不兼容时必须失败?
- 如何判断文件完整?
- 最新文件损坏时回滚到哪个版本?
- 恢复成功后如何验证?
- 检查点由谁读取、谁可以删除?
一个可靠的检查点系统最终应满足这样的因果链:
明确的更新边界
→ 一致地冻结相关状态
→ 完整写入并校验
→ 原子提交可见版本
→ 按兼容拓扑加载
→ 恢复数据进度和随机状态
→ 通过小规模验证
→ 再继续训练
检查点的核心不是序列化 API,而是对训练状态、数据流、随机性、分布式布局和故障提交点建立统一定义。只有这些状态在同一个逻辑时刻被保存,恢复才真正具有可解释性;只有把恢复保证的边界写清楚,工程团队才能区分“从权重重新开始”“从训练进度继续”和“尽量复现原执行轨迹”。
系列导航与关联阅读
- 系列入口:AI 工程完整学习路线:从机器学习与 Transformer 到 RAG、Agent 和生产治理
- 上一篇:混合精度训练:FP16、BF16、Loss Scaling、溢出与精度验证
- 下一篇:TensorFlow 工程基础:Tensor、GradientTape、tf.data、训练与导出
官方资料
本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。

评论
0 条讨论