模型保存与加载:state_dict 与检查点
本教程共 60 篇 · 第 25 篇 · 更新于 2026-08-17 · 约 4 分钟阅读
本节目标:搞懂 state_dict 是什么,学会保存和加载模型的正确姿势,并能写出可断点续训的检查点代码。
为什么要保存模型
训练一个模型少则几分钟,多则几天。练完直接关电脑,努力全白费。所以训练结束后的第一件事,就是把模型存到硬盘。
保存模型有这么几个实际用途:
- 训练中断恢复:断电、报错、机房维护,训练随时可能中断。有存档就能从断点接着练,不浪费已经烧掉的算力。
- 模型对比:训练过程中每过几轮存一份,事后把不同阶段的模型拿出来比一比,挑验证集表现最好的。
- 部署与共享:把训练好的模型拷到服务器、发给同事,别人不需要你的数据和训练代码,也能直接做推理。
PyTorch 里负责这件事的三个核心函数是:torch.save 负责存,torch.load 负责读,model.load_state_dict 负责把参数装回模型。
state_dict:模型的「参数户口本」
先认识一个关键概念:状态字典(state_dict)。
model.state_dict() 返回一个 Python 字典,键是每一层的名字,值是这一层的参数张量。看个最简单的例子:
import torch
import torch.nn as nn
model = nn.Sequential(
nn.Linear(4, 8),
nn.ReLU(),
nn.Linear(8, 2),
)
for k, v in model.state_dict().items():
print(k, v.shape)
输出:
0.weight torch.Size([8, 4])
0.bias torch.Size([8])
2.weight torch.Size([2, 8])
2.bias torch.Size([2])
注意两点:ReLU 没有可学习参数,不在字典里;权重张量的形状是「输出在前、输入在后」。正因为 state_dict 是普通字典,你可以随便改键名、删键、合并,灵活性很高。
Note优化器也有自己的 state_dict,存着学习率、动量等状态。断点续训时它和模型参数一样重要,马上会讲到。
推荐姿势:只存参数
最常用的做法是只保存 state_dict,文件小、最灵活:
torch.save(model.state_dict(), "model_weights.pth")
加载分三步:重建同样的模型结构,把参数装进去,最后切到评估模式:
model = nn.Sequential(
nn.Linear(4, 8),
nn.ReLU(),
nn.Linear(8, 2),
)
model.load_state_dict(torch.load("model_weights.pth"))
model.eval()
这里我踩过坑:load_state_dict 只认字典,不认文件路径。直接传 "model_weights.pth" 会报错,得先用 torch.load 读出来再传进去。
存完立刻验证一下,是个值得养成的好习惯:
test_x = torch.randn(16, 4)
with torch.no_grad():
before = model(test_x)
# 重新加载后输出应该一模一样
model.load_state_dict(torch.load("model_weights.pth"))
after = model(test_x)
print(torch.allclose(before, after)) # True
Tip扩展名
.pth和.pt都可以,社区里.pth更常见。加载完记得model.eval(),否则带 Dropout 和 BatchNorm 的模型,推理结果会不稳定。
另一条路:存整个模型
也可以把模型对象整个存进去:
torch.save(model, "whole_model.pth")
model = torch.load("whole_model.pth")
代码最短,但我不太推荐。它用 Python 的 pickle 序列化整个对象,加载时要求原始类的定义还在老地方。代码一重构、文件一挪窝,就可能加载失败。另外,整个模型的文件体积通常也比 state_dict 大一圈,里面塞了不少类结构等额外信息。反正参数才是核心,我建议你养成只存 state_dict 的习惯,长期保存和团队协作都更稳妥。
检查点:断点续训的秘密武器
只存参数只能做推理,没法继续训练。想中途断了接着练,要存检查点(checkpoint),把训练现场整个打包:
checkpoint = {
"epoch": epoch,
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"loss": loss.item(),
}
torch.save(checkpoint, "checkpoint.pth")
恢复训练时,把各部分原样装回去:
checkpoint = torch.load("checkpoint.pth")
model.load_state_dict(checkpoint["model_state_dict"])
optimizer.load_state_dict(checkpoint["optimizer_state_dict"])
start_epoch = checkpoint["epoch"]
model.train() # 继续训练切 train;只做推理切 eval
为什么优化器也要存?以 Adam 为例,它为每个参数维护了动量状态。丢了它,续训的前几步就像喝醉酒,模型表现会抖一下。
Note从 PyTorch 2.6 起,
torch.load默认weights_only=True,只反序列化张量、字典、列表等安全类型,防止恶意 pickle 文件偷偷执行代码。普通检查点不受影响,只有存了自定义对象的旧文件才会报错。
保存策略:像打游戏一样存进度
检查点怎么存,也是有讲究的,我的建议是:
- 定期存:每 N 个 epoch 存一份,文件名带上轮数,比如
resnet18_epoch50.pth。 - 只留最好的:训练时盯住验证集指标,刷新纪录就覆盖保存,避免磁盘被一堆中间文件塞满。
- 中断恢复靠检查点:训练循环开头先查有没有检查点文件,有就从那里接着跑。
- 顺手存一份最佳模型:验证集最好成绩对应的权重单独存一份,别等训练结束再后悔。
这套组合下来,训练就是可回放的:随时能看历史、随时能续跑、随时能回滚。
部分加载:热启动
有时候你只想加载一部分参数。比如迁移学习里,新模型的分类头和预训练模型对不上。这时加 strict=False,对不上的键会被跳过:
model.load_state_dict(torch.load("pretrained.pth"), strict=False)
它不会报错,但会打印缺了哪些键(missing keys)、多了哪些键(unexpected keys)。看一眼这个提示,确认忽略的正是你想忽略的。
跨设备加载
在 GPU 上训的模型,拿到没 GPU 的机器上直接读会报错。加个 map_location 就行:
state_dict = torch.load("model_weights.pth", map_location="cpu")
model.load_state_dict(state_dict)
反过来,CPU 上存的文件想上 GPU,先读进内存再搬:
model = model.to("cuda")
model.load_state_dict(torch.load("model_weights.pth", map_location="cuda"))
Warning显存紧张时加载大模型容易 OOM。保险做法是先
map_location="cpu"读进内存,再model.to("cuda")搬到显卡。
Note老项目里常看到
model.module.state_dict()的写法:那是nn.DataParallel多卡训练的遗产,参数被包了一层module。现在官方推荐分布式训练 DDP,DataParallel 已不推荐,新项目不必再和它纠缠。
常见报错速查
Missing key(s) in state_dict:模型结构和保存时不一致。检查网络定义,或加strict=FalseUnexpected key(s) in state_dict:字典里有模型不存在的键,常见于改过分类头,同样strict=False- 加载后效果差得离谱:八成是忘了
model.eval(),Dropout 和 BatchNorm 还停在训练模式 - 旧版本文件加载失败:PyTorch 升级后读老检查点偶尔报错。先确认版本兼容性,必要时用老环境重新保存一份
- 报 EOF 或 unpickling 错误:多半是文件没存完,或下载不完整。重新保存一次,并检查磁盘空间
进阶:大模型加载技巧
加载几十 GB 的大模型时,官方文档给了三个优化手段:torch.load(..., mmap=True) 把文件映射到内存、不整块读入;用 torch.device("meta") 上下文创建不占内存的「空壳模型」;再配 load_state_dict(assign=True) 直接把参数换过去。现在知道有这些招就行,真到大模型场景再回头查。
模型存取这门手艺你已经会了。下一章我们把它用在刀刃上:迁移学习。