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模型对象使用。
步骤:
将模型转为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")加载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模型。
步骤:
保存为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'])加载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
相关产品推荐
相关产品推荐

