首页 / PyTorch 入门教程 / 模型保存与加载:state_dict 与检查点

PyTorch 入门教程

模型保存与加载:state_dict 与检查点

本教程共 60 篇 · 第 25 篇 · 更新于 2026-08-17 · 约 4 分钟阅读

PyTorch模型保存模型加载state_dict检查点深度学习入门

本节目标:搞懂 state_dict 是什么,学会保存和加载模型的正确姿势,并能写出可断点续训的检查点代码。

为什么要保存模型

训练一个模型少则几分钟,多则几天。练完直接关电脑,努力全白费。所以训练结束后的第一件事,就是把模型存到硬盘。

保存模型有这么几个实际用途:

  1. 训练中断恢复:断电、报错、机房维护,训练随时可能中断。有存档就能从断点接着练,不浪费已经烧掉的算力。
  2. 模型对比:训练过程中每过几轮存一份,事后把不同阶段的模型拿出来比一比,挑验证集表现最好的。
  3. 部署与共享:把训练好的模型拷到服务器、发给同事,别人不需要你的数据和训练代码,也能直接做推理。

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 文件偷偷执行代码。普通检查点不受影响,只有存了自定义对象的旧文件才会报错。

保存策略:像打游戏一样存进度

检查点怎么存,也是有讲究的,我的建议是:

  1. 定期存:每 N 个 epoch 存一份,文件名带上轮数,比如 resnet18_epoch50.pth
  2. 只留最好的:训练时盯住验证集指标,刷新纪录就覆盖保存,避免磁盘被一堆中间文件塞满。
  3. 中断恢复靠检查点:训练循环开头先查有没有检查点文件,有就从那里接着跑。
  4. 顺手存一份最佳模型:验证集最好成绩对应的权重单独存一份,别等训练结束再后悔。

这套组合下来,训练就是可回放的:随时能看历史、随时能续跑、随时能回滚。

部分加载:热启动

有时候你只想加载一部分参数。比如迁移学习里,新模型的分类头和预训练模型对不上。这时加 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 已不推荐,新项目不必再和它纠缠。

常见报错速查

  1. Missing key(s) in state_dict:模型结构和保存时不一致。检查网络定义,或加 strict=False
  2. Unexpected key(s) in state_dict:字典里有模型不存在的键,常见于改过分类头,同样 strict=False
  3. 加载后效果差得离谱:八成是忘了 model.eval(),Dropout 和 BatchNorm 还停在训练模式
  4. 旧版本文件加载失败:PyTorch 升级后读老检查点偶尔报错。先确认版本兼容性,必要时用老环境重新保存一份
  5. 报 EOF 或 unpickling 错误:多半是文件没存完,或下载不完整。重新保存一次,并检查磁盘空间

进阶:大模型加载技巧

加载几十 GB 的大模型时,官方文档给了三个优化手段:torch.load(..., mmap=True) 把文件映射到内存、不整块读入;用 torch.device("meta") 上下文创建不占内存的「空壳模型」;再配 load_state_dict(assign=True) 直接把参数换过去。现在知道有这些招就行,真到大模型场景再回头查。

模型存取这门手艺你已经会了。下一章我们把它用在刀刃上:迁移学习。