已有 PyTorch 模型,想关掉程序后在另一段本地代码中继续做推理,可以保存 state_dict,重新创建同一结构,再加载参数。最有用的验收不是看到文件生成,而是同一输入在保存前后得到一致输出。
本文核对了 PyTorch 2.14 文档,并在 Windows、Python 3.11、PyTorch 2.14.1 CPU 环境运行下方小模型示例。模型随机初始化,只验证参数存取和推理一致性,没有训练效果、GPU 性能或真实业务模型的测试结论。

先分清要保存什么
PyTorch 的 模型保存与加载教程推荐为推理保存 model.state_dict()。这个字典包含参数和已注册的缓冲区,例如 BatchNorm 的运行统计;它不替你保存模型类的 Python 定义。
- 换一个程序做推理:保留权重、模型定义、输入预处理和类别对应表。
- 继续训练:还需要优化器及其他训练状态。本文只处理推理加载,不能据此声称训练可无缝恢复。
- 向别人提供 HTTP 服务:还要做接口、访问控制和部署验证。加载权重只是这条流程中的一步。
.pth 是常见文件后缀,不能仅凭后缀判断文件内部是参数字典、完整模型还是训练检查点。请先确认保存方的格式。
运行最小保存、重建、加载示例
在独立 Python 环境安装 CPU 依赖,保存下面代码为 reload_demo.py,再运行 python reload_demo.py。它只在当前目录写入专用的 demo_model/weights.pth,再次运行会覆盖这个演示文件。
python -m pip install torch==2.14.1 numpy --index-url https://download.pytorch.org/whl/cpu
from pathlib import Path
import torch
from torch import nn
class TinyModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(2, 1)
def forward(self, x):
return self.fc(x)
torch.manual_seed(7)
model = TinyModel().cpu().eval()
x = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
with torch.no_grad():
expected = model(x)
folder = Path('demo_model')
folder.mkdir(exist_ok=True)
weight_path = folder / 'weights.pth'
torch.save(model.state_dict(), weight_path)
restored = TinyModel().cpu()
state = torch.load(weight_path, map_location='cpu', weights_only=True)
restored.load_state_dict(state, strict=True)
restored.eval()
with torch.no_grad():
actual = restored(x)
assert expected.shape == actual.shape == (2, 1)
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
print('keys:', sorted(state))
print('output_equal:', torch.equal(actual, expected))
print('max_abs_diff:', (actual - expected).abs().max().item())
TinyModel() 重建结构,torch.load() 读文件得到字典,load_state_dict() 才把字典内容放到模型中。不要把文件路径直接传给 load_state_dict()。
map_location='cpu' 指定加载位置。torch.load 官方接口说明还说明了 weights_only 的受限加载机制;示例明确设为 True。该选项不代表任何来源的文件都安全,只加载自己生成或可信来源的权重,不要为绕过报错随意改成 False。
看输出,确认真的复现
上述 CPU 示例实际得到以下输出:
keys: ['fc.bias', 'fc.weight']
output_equal: True
max_abs_diff: 0.0
两个参数键与模型定义对应;两个输出形状都是 (2, 1),逐元素相等,最大绝对差为零。这里采用同一 CPU 环境和同一小模型,因而用零容差检查。跨设备、不同精度或不同算子实现时应根据任务设定容差,并用代表性输入核对,不能把这次结果当成所有模型均逐位一致的保证。
换成自己的模型时,先保存一小批固定输入和保存前的输出,再在新的进程中重建模型、加载权重并逐项对比。输入的列顺序、归一化规则、图像尺寸、分词器或类别映射也要保持一致,文件成功加载不能发现所有预处理错误。
加载失败时按这张表排查
| 现象 | 先核对 | 处理后如何确认 |
|---|---|---|
| FileNotFoundError | 脚本运行目录与权重路径是否一致 | 用 Path.resolve() 查看实际路径,再读取确切文件 |
| Missing key / Unexpected key | 模块名称、是否保存了 DataParallel 外层,文件是否其实是训练检查点 | 打印 state.keys() 与 model.state_dict().keys(),恢复保存时的结构或取正确字段 |
| size mismatch | 输入维数、隐藏层大小、输出类别数 | 逐项比对参数形状,使用与保存时一致的模型定义 |
| 加载成功但输出变了 | 输入预处理、模型模式、精度和设备 | 用同一固定输入比对输出形状与误差,再扩展样本 |
Module 官方说明规定,strict=True 要求参数键匹配。不要把 strict=False 当通用修复:缺失参数可能仍是新实例的初始化值,同名参数的形状冲突也不是简单忽略键名就能解决。
推理前调用 eval() 控制 Dropout、BatchNorm 等具有训练与评估差异的层;用 no_grad() 关闭此次前向的反向图记录。这两个动作各自承担不同职责。确认固定样本复现后,再把加载逻辑接回自己的推理入口。
Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/30457.html