5.5 模型保存与加载:断点续训 本节摘要:训练过夜、机器重启、指标突然变差需要回滚——不会保存加载,长训练就无从谈起。本节讲清 statedict 保存的是什么、为什么断点续训必须连优化器一起存、以及加载时类定义与参数名必须严格对上的机制。 存什么、怎么存 训练循环末尾悬着最后一块工程拼图。PyTorch 的答案分两层:statedict 是模型全部状态的字典(键是参数名,值是张量);保存加载就是把这个字典序列化与还原。推荐的保存姿势是"存字典而不是存整个模型"——只存 statedict 的文件不绑定代码路径,换机器换目录都能加载。 第二个常见误区是"只存模型"。断点续训要恢复的是完整训练现场:模型参数、优化器状态(动量、Adam 的两个矩都在里面)、epoch 计数、随机数状态。
本节摘要:训练过夜、机器重启、指标突然变差需要回滚——不会保存加载,长训练就无从谈起。本节讲清 state_dict 保存的是什么、为什么断点续训必须连优化器一起存、以及加载时类定义与参数名必须严格对上的机制。
训练循环末尾悬着最后一块工程拼图。PyTorch 的答案分两层:state_dict 是模型全部状态的字典(键是参数名,值是张量);保存加载就是把这个字典序列化与还原。推荐的保存姿势是"存字典而不是存整个模型"——只存 state_dict 的文件不绑定代码路径,换机器换目录都能加载。
第二个常见误区是"只存模型"。断点续训要恢复的是完整训练现场:模型参数、优化器状态(动量、Adam 的两个矩都在里面)、epoch 计数、随机数状态。少存优化器,恢复后 Adam 的矩从零热身,等于悄悄换了个优化器继续走,曲线会出现一次莫名的跳变。
import torch import torch.nn as nn torch.manual_seed(0) model = nn.Sequential(nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 10)) opt = torch.optim.Adam(model.parameters(), lr=0.001) # ---- 保存完整现场 ---- checkpoint = { "epoch": 12, "model_state": model.state_dict(), "opt_state": opt.state_dict(), "val_acc": 0.884, # 顺手记录成绩,便于挑检查点 } torch.save(checkpoint, "ckpt_step12.pt") # ---- 恢复完整现场 ---- model2 = nn.Sequential(nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 10)) opt2 = torch.optim.Adam(model2.parameters(), lr=0.001) ckpt = torch.load("ckpt_step12.pt", map_location="cpu") model2.load_state_dict(ckpt["model_state"]) opt2.load_state_dict(ckpt["opt_state"]) print("恢复的 epoch:", ckpt["epoch"], "| 恢复时验证成绩:", ckpt["val_acc"]) print("两份模型参数完全一致:", torch.equal(model[0].weight, model2[0].weight))
输出:
恢复的 epoch: 12 | 恢复时验证成绩: 0.884 两份模型参数完全一致: True
两个细节值得点破:map_location="cpu" 让在 GPU 上存的检查点能在无卡机器上加载——分发模型时的保命参数;load_state_dict 默认严格模式,键名对不上或多一个少一个都会当场报错,这看着烦,实际是防止"拿错模型装错参数"的守门员。
背景:30 epoch 的训练跑到第 12 轮机器断电。要求从第 12 轮无缝续跑,且训练曲线不能出现跳变。
操作:恢复后先验证现场完整(参数、优化器状态、计数),再继续训练,并与"不间断跑完"的对照组比较最终成绩。
import torch import torch.nn as nn from torch.utils.data import DataLoader, TensorDataset torch.manual_seed(42) templates = torch.randn(10, 64) X = torch.cat([templates[d].repeat(80, 1) + 0.3 * torch.randn(80, 64) for d in range(10)]) Y = torch.cat([torch.full((80,), d) for d in range(10)]) loader = DataLoader(TensorDataset(X[:640], Y[:640]), batch_size=64, shuffle=True) val_X, val_Y = X[640:], Y[640:] loss_fn = nn.CrossEntropyLoss() def val_acc(m): m.eval() with torch.no_grad(): return (m(val_X).argmax(1) == val_Y).float().mean().item() def run(epochs, start_epoch=0, model=None, opt=None): if model is None: torch.manual_seed(0) model = nn.Sequential(nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 10)) opt = torch.optim.SGD(model.parameters(), lr=0.5, momentum=0.9) model.train() for ep in range(start_epoch, epochs): for xb, yb in loader: opt.zero_grad(); loss_fn(model(xb), yb).backward(); opt.step() torch.save({"epoch": ep, "model_state": model.state_dict(), "opt_state": opt.state_dict()}, "latest.pt") # 每轮存现场 return val_acc(model) full = run(30) model, opt = nn.Sequential(nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 10)), None torch.manual_seed(0) model = nn.Sequential(nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 10)) opt = torch.optim.SGD(model.parameters(), lr=0.5, momentum=0.9) run(12, model=model, opt=opt) # 前段训练12轮,现场已存 model.load_state_dict(torch.load("latest.pt")["model_state"]) # 模拟断电恢复 opt.load_state_dict(torch.load("latest.pt")["opt_state"]) resumed = run(30, start_epoch=12, model=model, opt=opt) # 从13轮续跑 print(f"不间断跑完 30 轮: acc {full:.3f}") print(f"断电恢复续跑 : acc {resumed:.3f}")
输出:
不间断跑完 30 轮: acc 0.902 断电恢复续跑 : acc 0.898
结果:续跑与连续跑的最终成绩几乎一致(微小差异来自 DataLoader 打乱顺序的随机性)。
解读:现场完整的判据不只是"参数一样",更严格的是"续跑曲线与原曲线无跳变"。如果你只恢复了模型不恢复优化器,SGD 加动量的场景里动量从零开始,前几轮成绩会短暂回落再爬起——看到这种"恢复后先掉一段"的曲线,九成是优化器状态没跟上。
变式:把每轮全量保存改成"只保留验证成绩最好的检查点"(best-so-far 策略),对比训练结束时的 best 与 last 成绩差——早停(5.4 节)与 best 保存是一对搭档,部署时用 best、继续训练用 last。
部署场景(第 6 章的 ONNX 之前的第一站)只关心前向,加载后记得 model.eval() 并配合 no_grad。还有两个工程习惯值得养成:一是版本信息进检查点(存个字典多放一个键的事),三个月后你会感谢它;二是加载外部来源的模型时,严格模式的报错是你排查结构差异的第一线索——逐个对不上的键名往往直接指出"哪一层被你改过"。
import torch model.eval() # 部署前必须切换推理面孔(2.3 节) with torch.no_grad(): # 推理全程关账(4.1 节) sample = torch.randn(1, 64) logits = model(sample) print("单样本推理输出形状:", tuple(logits.shape), "预测类别:", logits.argmax(1).item())
输出:
单样本推理输出形状: (1, 10) 预测类别: 7
至此训练闭环连同工程外围全部走完。第 6 章扩编:换更强的装备走更远的路。