You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让自定义PyTorch模型保存后无需预定义类即可加载?

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文件包含模型结构和权重信息。

我的问题:

  1. 这种效果是怎么实现的?
  2. 自定义模型能不能实现这种效果,也就是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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.16 05:28:11