PyTorch 训练中断后怎么继续?保存并恢复模型和优化器检查点

用一维回归演示训练检查点保存和恢复,核对模型、Adam 状态与已完成步数,并与连续训练及遗漏优化器状态的反例对照。

PyTorch 训练中断后,想接着原来的训练状态继续,通常需要恢复模型参数、优化器状态和已经完成的位置。只加载模型权重后新建优化器,可能让 Adam 的动量统计重新开始;程序能够继续运行,不代表训练轨迹接上了。PyTorch 保存与加载教程给出了模型与优化器一起保存、重新加载并切回训练模式的做法。

先确定检查点里保存什么

下面让一个线性层学习人为设定的 y=2x+0.5。总共更新 10 次,在第 4 次更新完成后保存检查点。一条路径继续训练 6 次;另一条路径新建模型与 Adam、加载检查点,再训练 6 次。两条路径使用相同输入和更新次数,才能核对恢复是否正确。

PyTorch 训练中断后怎么继续?保存并恢复模型和优化器检查点

字段 用途 遗漏后会怎样
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 序列化说明解释了这个选项的允许对象范围;不要为了读取来历不明的文件,直接关闭限制。

接入自己的训练任务时,按这个顺序检查

  1. 在完整 optimizer.step() 之后保存本轮状态,避免保存到一次更新中间。
  2. 用原模型结构重新创建模型,并以相同参数组结构创建优化器,再分别加载状态。
  3. 按保存的位置决定下一次从哪一轮、哪一步开始;completed_steps 表示已经做完的次数,不是接下来要重做的编号。
  4. 先在同一设备和数据顺序下做“连续跑”和“中断恢复”对照,再扩大训练。

如果 load_state_dict 报 missing keys、unexpected keys 或张量尺寸不匹配,先核对模型定义与检查点版本;如果 optimizer 状态加载失败,先检查参数组是否改变。不要用 strict=False 静默绕过结构错误并宣称完全恢复。

常见问题:为什么参数恢复了,后续损失仍不同?

先确认是否恢复了优化器,再检查样本乱序、随机数状态、Dropout、学习率调度器和 AMP 状态。本例只测试 CPU、固定整批输入、没有随机训练算子的一维回归;没有验证跨设备、分布式训练或精确恢复任意数据加载器到中途位置。

Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/30306.html

赞 (0)
AI小管家的头像AI小管家
决策树模型怎么生成?用 scikit-learn 训练并导出规则
上一篇 1天前
神经网络结构图怎么画?用 Graphviz 标出层与张量尺寸
下一篇 1天前

相关推荐

联系我们

联系我们

1

在线咨询: QQ交谈

邮件:admin@example.com

工作时间:周一至周五,9:30-18:30,节假日休息

关注微信
关注微信
分享本页
返回顶部