PyTorch 训练中断后,想接着原来的训练状态继续,通常需要恢复模型参数、优化器状态和已经完成的位置。只加载模型权重后新建优化器,可能让 Adam 的动量统计重新开始;程序能够继续运行,不代表训练轨迹接上了。PyTorch 保存与加载教程给出了模型与优化器一起保存、重新加载并切回训练模式的做法。
先确定检查点里保存什么
下面让一个线性层学习人为设定的 y=2x+0.5。总共更新 10 次,在第 4 次更新完成后保存检查点。一条路径继续训练 6 次;另一条路径新建模型与 Adam、加载检查点,再训练 6 次。两条路径使用相同输入和更新次数,才能核对恢复是否正确。

| 字段 | 用途 | 遗漏后会怎样 |
|---|---|---|
| model | 网络参数及模块缓冲 | 重新从随机权重开始 |
| optimizer | 优化器参数组与历史统计 | Adam 等优化器状态重置 |
| completed_steps | 已完成的更新次数 | 重复或跳过训练更新 |
这三个字段只够本文的固定输入演示。实际任务使用学习率调度器、AMP、乱序数据或 Dropout 时,还应保存与恢复调度器、GradScaler、随机数生成器和数据采样进度;应根据训练循环实际用到的状态补齐。
运行一个能对照的恢复示例
以下完整脚本在 Windows、Python 3.11.15、PyTorch 2.8.0+cpu、NumPy 2.3.3 环境实际运行通过,只验证这个 CPU 玩具任务。它不下载模型权重,也不代表真实模型的精度、速度或显存表现。在已启用的独立 Python 环境中可安装相同依赖:
python -m pip install torch==2.8.0 --index-url https://download.pytorch.org/whl/cpu
python -m pip install numpy==2.3.3
将下面代码保存为 resume_demo.py,在一个可写目录运行 python resume_demo.py。脚本会在当前目录创建或覆盖自己的 resume-demo.pt 演示文件;正式检查点请使用独立路径并保留备份。
import torch
from pathlib import Path
torch.set_num_threads(1)
torch.manual_seed(7)
x = torch.linspace(-1, 1, 32).reshape(-1, 1)
y = 2 * x + 0.5
model = torch.nn.Linear(1, 1)
optimizer = torch.optim.Adam(model.parameters(), lr=0.03)
loss_fn = torch.nn.MSELoss()
def train_steps(net, opt, count):
net.train()
for _ in range(count):
opt.zero_grad(set_to_none=True)
loss = loss_fn(net(x), y)
if not torch.isfinite(loss):
raise RuntimeError("non-finite loss")
loss.backward()
opt.step()
train_steps(model, optimizer, 4)
checkpoint_path = Path("resume-demo.pt")
torch.save({
"completed_steps": 4,
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
}, checkpoint_path)
# Keep the original run going as a reference.
train_steps(model, optimizer, 6)
# New objects simulate starting a new training process.
resumed = torch.nn.Linear(1, 1)
resumed_optimizer = torch.optim.Adam(resumed.parameters(), lr=0.03)
saved = torch.load(checkpoint_path, map_location="cpu", weights_only=True)
resumed.load_state_dict(saved["model"])
resumed_optimizer.load_state_dict(saved["optimizer"])
train_steps(resumed, resumed_optimizer, 10 - saved["completed_steps"])
max_diff = max((a - b).abs().max().item()
for a, b in zip(model.parameters(), resumed.parameters()))
print("completed_steps:", saved["completed_steps"])
print("max_parameter_diff:", max_diff)
assert max_diff == 0.0
# A negative control: loading weights alone loses Adam's running state.
weights_only_model = torch.nn.Linear(1, 1)
weights_only_model.load_state_dict(saved["model"])
new_optimizer = torch.optim.Adam(weights_only_model.parameters(), lr=0.03)
train_steps(weights_only_model, new_optimizer, 6)
weights_only_diff = max((a - b).abs().max().item()
for a, b in zip(model.parameters(), weights_only_model.parameters()))
print("weights_without_optimizer_diff_positive:", weights_only_diff > 0)
assert weights_only_diff > 0
核对输出,而不只看文件是否存在
本次运行结果为:
completed_steps: 4
max_parameter_diff: 0.0
weights_without_optimizer_diff_positive: True
恢复后的模型与不中断模型,最大参数差为 0.0;只加载权重却不恢复 Adam 的对照则产生了非零差值。读者应同时看到 completed_steps 为 4、参数差通过断言以及反例检查为 True。只有文件能读取,不能证明恢复过程正确。
这里的 weights_only=True 是 torch.load 的反序列化限制选项,和“只保存模型权重”不是同一个概念:演示文件还包含优化器状态和整数。PyTorch 2.8 序列化说明解释了这个选项的允许对象范围;不要为了读取来历不明的文件,直接关闭限制。
接入自己的训练任务时,按这个顺序检查
- 在完整 optimizer.step() 之后保存本轮状态,避免保存到一次更新中间。
- 用原模型结构重新创建模型,并以相同参数组结构创建优化器,再分别加载状态。
- 按保存的位置决定下一次从哪一轮、哪一步开始;completed_steps 表示已经做完的次数,不是接下来要重做的编号。
- 先在同一设备和数据顺序下做“连续跑”和“中断恢复”对照,再扩大训练。
如果 load_state_dict 报 missing keys、unexpected keys 或张量尺寸不匹配,先核对模型定义与检查点版本;如果 optimizer 状态加载失败,先检查参数组是否改变。不要用 strict=False 静默绕过结构错误并宣称完全恢复。
常见问题:为什么参数恢复了,后续损失仍不同?
先确认是否恢复了优化器,再检查样本乱序、随机数状态、Dropout、学习率调度器和 AMP 状态。本例只测试 CPU、固定整批输入、没有随机训练算子的一维回归;没有验证跨设备、分布式训练或精确恢复任意数据加载器到中途位置。
Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/30306.html