5.模型保存与恢复
第 5 章 — 模型保存与恢复(Checkpoint)
第 4 章结束时,我们学会了用 lr_scheduler 让学习率随训练动态变化,但训练产生的所有成果都还只存在于内存里:进程一退出,几小时(甚至几十小时)的进度就灰飞烟灭。本章引入 checkpoint(检查点/存档),把"训练现场"完整保存到磁盘——训练中断可以无缝续跑,训练结束后还能取出验证集最优的那个模型。
一、本章要解决的问题
- 第 4 章的代码一关终端就"失忆":模型、优化器、调度器的状态全部停在内存里,进程退出即清零 。训练到一半断电、报错、不小心关掉终端,全部进度归零,只能从头再来。
- 想选出"验证集最优的那个 epoch"的模型:第 4 章只是把 val_acc 打印出来,没有存档 ,跑完之后想拿最好的模型去做测试、去部署,手边什么凭据都没有。
- 想"接着上次继续"或"换个机器接着跑":没有持久化文件,这些需求一个都做不了。
二、核心概念速览
本章有 6 个新概念,先花 5 分钟读完,再看代码会轻松很多。
1. 模型权重的持久化(游戏存档类比)
训练是一段漫长的"游戏":参数在内存里不断更新,就像角色打怪练级。内存是易失的,进程退出、机器断电,一切归零——这相当于"没存档就关机"。把模型参数这一瞬间的"快照"写到磁盘上的文件里,就叫持久化(persistence)。torch.save 是"存档",torch.load 是"读档"。有了存档,中断了也不怕,读档接着玩。
2. state_dict(模型所有可学习参数的字典)
state_dict 是 PyTorch 内置的一张"参数清单":一个字典,键是层名(如 features.0.weight、classifier.0.weight),值是那个层的参数张量。注意只有"带可学习参数的层"才会出现——卷积、BatchNorm、全连接有,ReLU、池化没有(它们只是运算)。用 model.state_dict() 取出来,用 model.load_state_dict(...) 灌回去,模型结构不变,参数瞬间替换。它是存档里最核心的内容。
3. 为什么光存权重不够(存档要能"无缝续玩")
玩到一半的游戏,存档如果只记"等级"不记"道具、位置、剧情进度",读档后体验就变了。训练也一样:光存权重,续跑时会出问题——optimizer 里有 Adam 的一阶/二阶矩(相当于它的"惯性"),丢了它续跑几步就会蹦;scheduler 有当前衰减阶段,丢了它学习率曲线会断;epoch 决定从第几轮继续;best_acc 和 metrics 是历史成绩单。所以本章的 checkpoint 是一个大字典,把这些全部打包存进去,--resume 才能做到"无缝续玩"。
4. last.pth 与 best.pth 双文件策略(自动存档 vs 最优存档)
两个文件各司其职:last.pth 每轮结束都覆盖保存,是"自动存档"——崩溃后从最近的进度续跑,损失不超过一个 epoch;best.pth 只在验证集准确率刷新历史纪录时保存,是"最优存档"——训练结束时它就是"打得最好的那一次",用来做测试和部署。一个保证"不白跑",一个保证"拿最好",两者互不干扰。
5. torch.load 的 weights_only 参数与反序列化安全
torch.load 底层用的是 Python 的 pickle 反序列化,而 pickle 有一个著名特性:加载文件时会执行里面的任意代码。加载一个来历不明的 .pth 文件,等于让陌生代码在你的机器上裸奔。PyTorch 从 2.6 起默认 weights_only=True(只允许张量这类安全对象),而我们的 checkpoint 里含 optimizer 等非张量对象,必须显式传 weights_only=False。一句话规则:只对自己生成、可信的文件关闭这个保护。
6. map_location 的作用(跨设备加载)
保存时张量可能在 GPU 上,加载时可能换了机器、换了设备(比如从服务器拷到没 GPU 的笔记本)。map_location 负责"搬设备":map_location="cpu" 把 GPU 存档搬到 CPU(没 GPU 也能读);map_location=device 把存档搬到当前设备。没有它,GPU 上存的权重在 CPU 机器上直接报错"not on the same device"。
三、解决思路
- 统一存档结构:定义一个 dict,把"训练现场"完整打包——
model_state(权重)、optimizer_state(优化器状态)、scheduler_state(调度器状态)、epoch(进度)、best_acc(历史最优)、metrics(历史指标曲线)。缺哪一样,"续玩"都不完整。 - 双文件策略:每轮结束覆盖保存
last.pth(自动存档);仅当val_acc刷新历史纪录时额外保存best.pth(最优存档)。训练结束用best.pth做最终评估。 - 恢复路径:
--resume指定存档路径时,加载全部状态,epoch从ckpt["epoch"] + 1接着跑,优化器、调度器、历史指标全部无缝衔接——继续训练就像"中途没停过"。 - 安全与兼容:加载时
map_location=device对齐设备;weights_only=False配合中文注释讲清"这只用于可信文件"。
trade-off:每轮多一次写盘,有少量 IO 开销(对 CIFAR-10 这种小模型可以忽略);只存权重更省空间,但牺牲"无缝续玩"。本章选择存完整状态——工程上这是绝大多数训练框架(含 Lightning)的默认做法。
四、代码变更
相对第 4 章的改动:
+ import os # 新增:save_checkpoint 里创建目录要用
import argparse
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms
from collections import defaultdict
+ from tqdm import tqdm # tqdm 自第 3 章引入,本章完整代码中沿用(恢复 batch 级进度条)
...
parser.add_argument("--device", type=str, default="auto",
help="运行设备:auto/cuda/cpu")
+ parser.add_argument("--ckpt-dir", type=str, default="./checkpoints",
+ help="checkpoint 保存目录")
+ parser.add_argument("--resume", type=str, default=None,
+ help="从指定 checkpoint 恢复训练")
...
+ # 新增:把完整训练状态写入磁盘;目录不存在时自动创建
+ def save_checkpoint(state, path):
+ os.makedirs(os.path.dirname(path), exist_ok=True)
+ torch.save(state, path)
+ print(f"[Checkpoint] 已保存到 {path}")
metrics = defaultdict(list)
- for epoch in range(1, args.epochs + 1):
+ start_epoch, best_acc = 1, 0.0
+ if args.resume: # 新增:恢复逻辑,从存档接续训练
+ ckpt = torch.load(args.resume, map_location=device, weights_only=False)
+ model.load_state_dict(ckpt["model_state"])
+ optimizer.load_state_dict(ckpt["optimizer_state"])
+ if scheduler is not None and ckpt.get("scheduler_state") is not None:
+ scheduler.load_state_dict(ckpt["scheduler_state"])
+ start_epoch = ckpt["epoch"] + 1
+ best_acc = ckpt.get("best_acc", 0.0)
+ metrics = defaultdict(list, ckpt.get("metrics", {}))
+ print(f"[Resume] 从 epoch {ckpt['epoch']} 恢复,历史最优 val_acc={best_acc:.4f}")
+
+ for epoch in range(start_epoch, args.epochs + 1):
train_loss, train_acc = train_one_epoch(...)
val_loss, val_acc = validate(...)
...
print(f"Epoch {epoch:02d}/{args.epochs} | "
f"Train loss {train_loss:.4f} acc {train_acc:.4f} | "
f"Val loss {val_loss:.4f} acc {val_acc:.4f} | "
f"lr {current_lr:.2e}")
+ # 新增:每轮结束存 last.pth;刷新纪录时另存 best.pth
+ state = {
+ "model_state": model.state_dict(),
+ "optimizer_state": optimizer.state_dict(),
+ "scheduler_state": scheduler.state_dict() if scheduler else None,
+ "epoch": epoch,
+ "best_acc": best_acc,
+ "metrics": dict(metrics),
+ }
+ save_checkpoint(state, os.path.join(args.ckpt_dir, "last.pth"))
+ if val_acc > best_acc:
+ best_acc = val_acc
+ state["best_acc"] = best_acc
+ save_checkpoint(state, os.path.join(args.ckpt_dir, "best.pth"))另外,train_one_epoch / validate 的循环由 for images, labels in loader: 改为 for images, labels in tqdm(loader, desc="Train"/"Val", leave=False):——tqdm 已在第 3 章引入,这里只是沿用恢复 batch 级进度条,不是本章新知识。
五、完整代码
创建 train.py(完整版):
import os
import argparse
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms
from collections import defaultdict
from tqdm import tqdm
def parse_args():
parser = argparse.ArgumentParser(description="CIFAR-10 图像分类训练")
parser.add_argument("--epochs", type=int, default=10, help="训练轮数")
parser.add_argument("--batch-size", type=int, default=64, help="每个 batch 的样本数")
parser.add_argument("--lr", type=float, default=1e-3, help="初始学习率")
parser.add_argument("--lr-scheduler", type=str, default="step",
choices=["step", "cosine"], help="学习率调度策略")
parser.add_argument("--lr-step-size", type=int, default=15, help="StepLR 每多少轮衰减一次")
parser.add_argument("--lr-gamma", type=float, default=0.1, help="StepLR 衰减系数")
parser.add_argument("--data-dir", type=str, default="./data", help="数据集存放目录")
parser.add_argument("--num-workers", type=int, default=2, help="DataLoader 数据加载进程数")
parser.add_argument("--device", type=str, default="auto", help="运行设备:auto/cuda/cpu")
# 本章新增的两个参数:存档目录 + 是否从旧存档恢复
parser.add_argument("--ckpt-dir", type=str, default="./checkpoints", help="checkpoint 保存目录")
parser.add_argument("--resume", type=str, default=None, help="从指定 checkpoint 恢复训练")
return parser.parse_args()
args = parse_args()
# 设备选择,与第 4 章完全一致
if args.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(args.device)
print(f"Using device: {device}")
class SimpleCNN(nn.Module):
def __init__(self, num_classes=10):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, padding=1),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(64, 128, kernel_size=3, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
)
self.classifier = nn.Sequential(
nn.AdaptiveAvgPool2d(1),
nn.Flatten(),
nn.Linear(128, num_classes),
)
def forward(self, x):
return self.classifier(self.features(x))
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
])
train_dataset = datasets.CIFAR10(root=args.data_dir, train=True, download=True, transform=transform)
test_dataset = datasets.CIFAR10(root=args.data_dir, train=False, download=True, transform=transform)
train_dataset, val_dataset = random_split(
train_dataset, [45000, len(train_dataset) - 45000]
)
train_loader = DataLoader(train_dataset, batch_size=args.batch_size,
shuffle=True, num_workers=args.num_workers)
val_loader = DataLoader(val_dataset, batch_size=args.batch_size,
shuffle=False, num_workers=args.num_workers)
test_loader = DataLoader(test_dataset, batch_size=args.batch_size,
shuffle=False, num_workers=args.num_workers)
model = SimpleCNN().to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=args.lr)
if args.lr_scheduler == "step":
scheduler = optim.lr_scheduler.StepLR(
optimizer, step_size=args.lr_step_size, gamma=args.lr_gamma)
elif args.lr_scheduler == "cosine":
scheduler = optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=args.epochs)
else:
scheduler = None
def train_one_epoch(model, loader, criterion, optimizer, device):
model.train()
total_loss, correct, total = 0.0, 0, 0
for images, labels in tqdm(loader, desc="Train", leave=False):
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def validate(model, loader, criterion, device):
model.eval()
total_loss, correct, total = 0.0, 0, 0
with torch.no_grad():
for images, labels in tqdm(loader, desc="Val", leave=False):
images, labels = images.to(device), labels.to(device)
outputs = model(images)
loss = criterion(outputs, labels)
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def save_checkpoint(state, path):
"""把完整训练状态写入磁盘;目录不存在时自动创建。"""
os.makedirs(os.path.dirname(path), exist_ok=True)
torch.save(state, path)
print(f"[Checkpoint] 已保存到 {path}")
# 恢复逻辑:--resume 存在时,从 checkpoint 接续训练
metrics = defaultdict(list)
start_epoch, best_acc = 1, 0.0
if args.resume:
# weights_only=False:checkpoint 含 optimizer 等非 tensor 对象,
# 只应加载自己生成、可信的文件
ckpt = torch.load(args.resume, map_location=device, weights_only=False)
model.load_state_dict(ckpt["model_state"])
optimizer.load_state_dict(ckpt["optimizer_state"])
if scheduler is not None and ckpt.get("scheduler_state") is not None:
scheduler.load_state_dict(ckpt["scheduler_state"])
start_epoch = ckpt["epoch"] + 1
best_acc = ckpt.get("best_acc", 0.0)
metrics = defaultdict(list, ckpt.get("metrics", {}))
print(f"[Resume] 从 epoch {ckpt['epoch']} 恢复,历史最优 val_acc={best_acc:.4f}")
print("=" * 60)
print("训练配置:")
for k, v in vars(args).items():
print(f" {k:12s} = {v}")
print("=" * 60)
for epoch in range(start_epoch, args.epochs + 1):
train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device)
val_loss, val_acc = validate(model, val_loader, criterion, device)
if scheduler is not None:
scheduler.step()
current_lr = optimizer.param_groups[0]["lr"]
metrics["train_loss"].append(train_loss)
metrics["train_acc"].append(train_acc)
metrics["val_loss"].append(val_loss)
metrics["val_acc"].append(val_acc)
print(f"Epoch {epoch:02d}/{args.epochs} | "
f"Train loss {train_loss:.4f} acc {train_acc:.4f} | "
f"Val loss {val_loss:.4f} acc {val_acc:.4f} | "
f"lr {current_lr:.2e}")
# --- 存档:last.pth 每轮都存,best.pth 只在刷新验证集最优时存 ---
state = {
"model_state": model.state_dict(),
"optimizer_state": optimizer.state_dict(),
"scheduler_state": scheduler.state_dict() if scheduler else None,
"epoch": epoch,
"best_acc": best_acc,
"metrics": dict(metrics),
}
save_checkpoint(state, os.path.join(args.ckpt_dir, "last.pth"))
if val_acc > best_acc:
best_acc = val_acc
state["best_acc"] = best_acc
save_checkpoint(state, os.path.join(args.ckpt_dir, "best.pth"))
test_loss, test_acc = validate(model, test_loader, criterion, device)
print(f"\n[Test] loss {test_loss:.4f} acc {test_acc:.4f}")六、运行示例
# 1) 第一次训练:每轮结束自动生成 checkpoints/last.pth,刷新纪录时另存 best.pth
python train.py --epochs 10
# 2) 训练到一半中断(Ctrl+C / 断电)后,接着上次的存档继续跑
# 注意:--epochs 必须大于已完成的轮数,否则循环不执行
python train.py --epochs 20 --resume ./checkpoints/last.pth
# 3) 用训练期间保留下来的最优模型 best.pth 在测试集上评估
# (临时内联脚本;正式的独立推理/评估脚本见第 14 章)
python - <<'PY'
import torch
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
class SimpleCNN(torch.nn.Module): # 结构必须与保存时完全一致
def __init__(self, num_classes=10):
super().__init__()
self.features = torch.nn.Sequential(
torch.nn.Conv2d(3, 32, kernel_size=3, padding=1), torch.nn.BatchNorm2d(32),
torch.nn.ReLU(inplace=True), torch.nn.MaxPool2d(2),
torch.nn.Conv2d(32, 64, kernel_size=3, padding=1), torch.nn.BatchNorm2d(64),
torch.nn.ReLU(inplace=True), torch.nn.MaxPool2d(2),
torch.nn.Conv2d(64, 128, kernel_size=3, padding=1), torch.nn.BatchNorm2d(128),
torch.nn.ReLU(inplace=True), torch.nn.MaxPool2d(2))
self.classifier = torch.nn.Sequential(
torch.nn.AdaptiveAvgPool2d(1), torch.nn.Flatten(),
torch.nn.Linear(128, num_classes))
def forward(self, x):
return self.classifier(self.features(x))
model = SimpleCNN().to("cpu") # 目标是 CPU,map_location="cpu" 兜底
ckpt = torch.load("./checkpoints/best.pth", map_location="cpu", weights_only=False)
model.load_state_dict(ckpt["model_state"]) # 只取权重,评估不需要 optimizer/epoch
transform = transforms.Compose([transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616))])
test_set = datasets.CIFAR10(root="./data", train=False, download=True, transform=transform)
loader = DataLoader(test_set, batch_size=64, shuffle=False)
model.eval() # 评估前务必切到 eval 模式
correct = total = 0
with torch.no_grad():
for images, labels in loader:
preds = model(images).argmax(dim=1)
correct += (preds == labels).sum().item()
total += images.size(0)
print(f"[Test] acc {correct / total:.4f}")
PY七、本章小结
- 学到了什么
- 一个 checkpoint 不只是"权重",而是训练现场的完整快照:模型、优化器、调度器、当前轮数、历史最优、指标曲线,缺一样续跑都不"无缝"。
- 双文件策略:
last.pth管"崩溃后不白跑",best.pth管"事后取最优",一自动一精选,各司其职。 map_location解决跨设备加载;weights_only是加载文件的"安检开关",只对可信文件关闭。
- 常见坑
- 只 load 权重、不 load optimizer/scheduler:续跑时 Adam 的动量丢了、学习率曲线断了,看似"恢复了",其实状态错乱。
- 加载模型的结构必须与保存时完全一致:改了层名/通道数再
load_state_dict,会报 key 不匹配。 - 加载后忘了
model.eval():直接拿存档做验证/测试,BatchNorm 还在用训练行为,指标不可信。 --resume继续时--epochs设小了:比如已跑到 epoch 8 却--epochs 5,循环直接跳过,什么都不跑。- 乱加载不可信的
.pth文件:weights_only=False意味着对方代码可以在你机器上执行,只用于自己生成的文件。
- 下一章预告:存档解决了"想拿最好的模型",但训练本身还在"闷头跑满全部 epoch"——如果第 5 轮就过拟合、后面 25 轮全是浪费时间怎么办?第 6 章引入早停机制(Early Stopping),验证集连续 N 轮不提升就提前收工。
八、动手练习
- 双文件体检:训练 30 个 epoch 后,分别用
last.pth和best.pth在测试集上评估(复用运行示例第 3 条),对比两个文件对应的 test acc——思考:为什么best.pth一定不差于last.pth?在什么情况下两者完全相同? - 状态残缺实验:故意修改恢复逻辑,只
load_state_dict(model_state)、跳过 optimizer/scheduler 的恢复,然后--resume继续训练,观察日志里 val_acc / lr 曲线的变化,体会"只存权重"续跑的代价。 - 模拟中断续跑:训练 10 个 epoch 中途 Ctrl+C 打断,看
checkpoints/里生成了哪些文件;再用--resume ./checkpoints/last.pth --epochs 10继续,确认日志从Epoch 04/10(举例)无缝接上,历史指标没有被清空。 - 跨设备加载:把
best.pth复制到另一台机器(或本机强制--device cpu保存、map_location="cpu"加载),验证"GPU 上存的档、CPU 机器也能读"——然后试一下不加map_location加载,观察报错信息,加深理解。
