PyTorch 模型怎么在本地重新加载?保存 state_dict 并核对推理输出

用 PyTorch 保存 state_dict,在 CPU 重建模型并加载权重,检查固定输入的输出一致性,定位路径、参数键和形状冲突。

已有 PyTorch 模型,想关掉程序后在另一段本地代码中继续做推理,可以保存 state_dict,重新创建同一结构,再加载参数。最有用的验收不是看到文件生成,而是同一输入在保存前后得到一致输出。

本文核对了 PyTorch 2.14 文档,并在 Windows、Python 3.11、PyTorch 2.14.1 CPU 环境运行下方小模型示例。模型随机初始化,只验证参数存取和推理一致性,没有训练效果、GPU 性能或真实业务模型的测试结论。

PyTorch 模型怎么在本地重新加载?保存 state_dict 并核对推理输出

先分清要保存什么

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

赞 (0)
AI小管家的头像AI小管家
模型上线后数据变了吗?用 Evidently 检查特征漂移
上一篇 3小时前
Netron 怎么查看模型结构?打开 ONNX 并核对输入输出
下一篇 3小时前

相关推荐

联系我们

联系我们

1

在线咨询: QQ交谈

邮件:admin@example.com

工作时间:周一至周五,9:30-18:30,节假日休息

关注微信
关注微信
分享本页
返回顶部