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

PyTorch模型(.nnet/.onnx格式)保存加载报错及解决方案求助

问题与解决方案

问题背景

本地训练了一个PyTorch模型,用torch.save(model,'theModel.nnet')保存后,尝试用torch.load('theModel.nnet')加载时出现AttributeError: Can't get attribute 'Net' on <module '__main__'>错误,希望加载时无需重复编写Net模型类代码就能直接使用PyTorch模型对象。

错误原因

torch.save(model)保存的是完整模型对象,它依赖模型类(即Net)的定义。加载时,Python解释器需要在当前运行环境中找到Net类的定义,否则会触发找不到类的报错。

解决方案

方案1:使用TorchScript(推荐,无需原模型类代码)

TorchScript可将PyTorch模型转为可序列化的脚本格式,加载时不需要原模型类的定义,直接就能作为PyTorch模型对象使用。

步骤:

  1. 将模型转为TorchScript并保存
    可通过**追踪(Trace)或脚本(Script)**两种方式转换,以下是追踪方式的示例:

    import torch
    import torch.nn as nn
    
    # 训练时的模型定义代码
    class Net(nn.Module):
        def __init__(self, input_size, hidden_size1, hidden_size2, output_size):
            super(Net, self).__init__()
            self.hidden1 = nn.Linear(input_size, hidden_size1)
            self.hidden2 = nn.Linear(hidden_size1, hidden_size2)
            self.output = nn.Linear(hidden_size2, output_size)
            self.relu = nn.ReLU()
    
        def forward(self, x):
            x = self.relu(self.hidden1(x))
            x = self.relu(self.hidden2(x))
            x = self.output(x)
            return x
    
    input_size =5
    hidden_size1 = 2
    hidden_size2 = 3
    output_size = 5
    model = Net(input_size, hidden_size1, hidden_size2, output_size)
    
    # 创建示例输入,用于追踪模型计算图
    example_input = torch.randn(1, input_size)
    # 转换为TorchScript模型
    traced_model = torch.jit.trace(model, example_input)
    # 保存模型(后缀可自定义,如.nnet/.pt/.pth)
    traced_model.save("theModel.nnet")
    
  2. 加载TorchScript模型
    加载时无需定义Net类,直接调用即可:

    import torch
    # 加载模型
    saved_model = torch.jit.load("theModel.nnet")
    # 验证模型可用性
    test_input = torch.randn(1, 5)
    output = saved_model(test_input)
    print(output)
    

方案2:保存为ONNX格式(跨框架兼容)

如果需要跨框架使用模型,可保存为ONNX格式,之后可用ONNX Runtime运行,或转换回PyTorch模型。

步骤:

  1. 保存为ONNX格式

    import torch
    # 假设model是训练完成的模型
    example_input = torch.randn(1, input_size)
    torch.onnx.export(model, example_input, "theModel.onnx", 
                      opset_version=11,
                      do_constant_folding=True,
                      input_names=['input'],
                      output_names=['output'])
    
  2. 加载ONNX模型并推理
    若仅需推理,可直接用ONNX Runtime:

    import onnxruntime as ort
    import torch
    
    sess = ort.InferenceSession("theModel.onnx")
    test_input = torch.randn(1, 5).numpy()
    output = sess.run(None, {'input': test_input})
    print(output)
    

方案3:保留原模型类定义(不推荐,不符合需求)

如果一定要用torch.save(model)的原生保存方式,加载时必须在当前环境中重新定义Net类,这需要重复编写模型代码,不符合需求,仅作参考:

import torch
import torch.nn as nn

# 必须重新定义Net类
class Net(nn.Module):
    def __init__(self, input_size, hidden_size1, hidden_size2, output_size):
        super(Net, self).__init__()
        self.hidden1 = nn.Linear(input_size, hidden_size1)
        self.hidden2 = nn.Linear(hidden_size1, hidden_size2)
        self.output = nn.Linear(hidden_size2, output_size)
        self.relu = nn.ReLU()

    def forward(self, x):
        x = self.relu(self.hidden1(x))
        x = self.relu(self.hidden2(x))
        x = self.output(x)
        return x

# 加载模型
saved_model = torch.load('theModel.nnet')

内容的提问来源于stack exchange,提问作者Coder Boy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 04:35:30