如何让自定义PyTorch模型保存后无需预定义类即可加载?
我用PyTorch定义了如下SimpleNN模型,并用torch.save(model, "my_model.pth")保存:
class SimpleNN(nn.Module): def __init__(self): super(SimpleNN, self).__init__() self.flatten = nn.Flatten() self.fc = nn.Sequential( nn.Linear(28 * 28, 128), nn.ReLU(), nn.Linear(128, 64), nn.ReLU(), nn.Linear(64, 10) ) def forward(self, x): x = self.flatten(x) x = self.fc(x) return x torch.save(model, "my_model.pth")
但加载时必须在当前文件中存在SimpleNN类,要是在另一文件直接执行:
model = torch.load("my_model.pth")
会报错:
AttributeError: Can't get attribute 'SimpleNN' on <module '__main__' from ····
而用torchvision.models里的resnet18时,执行以下代码保存:
from torchvision.models import resnet18 model = resnet18(pretrained=True) torch.save(model, "resnet_model.pth")
在另一文件直接执行:
resnet = torch.load("resnet_model.pth") print(resnet)
不用提前声明模型类就能成功加载,且pth文件包含模型结构和权重信息。
我的问题:
- 这种效果是怎么实现的?
- 自定义模型能不能实现这种效果,也就是pth文件包含模型结构与权重,加载时无需提前声明类?
1. torchvision预训练模型直接加载的原理
PyTorch用torch.save()保存整个模型时,会序列化模型对象,其中包含模型类的完整模块路径信息。对于torchvision.models下的模型(比如resnet18),它们的类定义在PyTorch官方维护的公开模块中(具体是torchvision.models.resnet)。
当调用torch.load()时,PyTorch会读取序列化文件里的类路径,自动从对应的官方模块中导入模型类,不需要手动提前声明。而你的自定义SimpleNN类默认在__main__模块(即当前运行的脚本),其他文件加载时找不到这个模块下的类,所以触发报错。
2. 自定义模型实现无需声明类加载的方法
有两种可行方案:
方案一:将自定义模型放入独立可导入模块
把SimpleNN类单独写在一个Python文件中,比如命名为my_custom_models.py:
# my_custom_models.py import torch.nn as nn class SimpleNN(nn.Module): def __init__(self): super(SimpleNN, self).__init__() self.flatten = nn.Flatten() self.fc = nn.Sequential( nn.Linear(28 * 28, 128), nn.ReLU(), nn.Linear(128, 64), nn.ReLU(), nn.Linear(64, 10) ) def forward(self, x): x = self.flatten(x) x = self.fc(x) return x
保存模型时,从该模块导入类并实例化:
from my_custom_models import SimpleNN import torch model = SimpleNN() torch.save(model, "my_model.pth")
加载时,只要你的Python环境能找到my_custom_models.py(比如放在同一目录),直接用torch.load()即可,无需手动导入SimpleNN。
方案二:保存模型配置+状态字典,加载时重建模型
如果不想依赖外部模块,可以把模型的结构配置(比如各层的参数)和权重字典一起保存:
import torch import torch.nn as nn # 保存时同时存储配置和状态字典 model = SimpleNN() model_config = { "input_dim": 28*28, "hidden_dims": [128, 64], "output_dim": 10 } save_data = { "config": model_config, "state_dict": model.state_dict() } torch.save(save_data, "my_model_with_config.pth")
加载时,根据配置重建模型再加载权重:
import torch.nn as nn def rebuild_model(config): return nn.Sequential( nn.Flatten(), nn.Linear(config["input_dim"], config["hidden_dims"][0]), nn.ReLU(), nn.Linear(config["hidden_dims"][0], config["hidden_dims"][1]), nn.ReLU(), nn.Linear(config["hidden_dims"][1], config["output_dim"]) ) save_data = torch.load("my_model_with_config.pth") model = rebuild_model(save_data["config"]) model.load_state_dict(save_data["state_dict"])
这种方式不需要提前定义模型类,但需要自己实现模型的重建逻辑。
内容的提问来源于stack exchange,提问作者Ldemon

