使用torch.onnx.export遇NotImplementedError及权重加载不匹配问题求助
PyTorch转ONNX:两个错误的解决方法
1. NotImplementedError 错误处理
原因:自定义的Model类继承自nn.Module,但未实现必须的forward方法。PyTorch要求所有自定义模块必须重写该方法定义前向传播逻辑,否则会触发此错误。
解决步骤:在Model类中添加forward方法,直接调用内部封装的self.model完成前向计算:
class Model(nn.Module): def __init__(self, n_h_layers, n_h_neurons, dim_in, dim_out, in_bound, out_bound): # 保留原__init__中的所有代码 super(Model,self).__init__() self.n_h_layers=n_h_layers self.n_h_neurons=n_h_neurons self.dim_in=dim_in self.dim_out=dim_out self.in_bound=in_bound self.out_bound=out_bound layer_input = [nn.Linear(dim_in, n_h_neurons, bias=True)] layer_output = [nn.ReLU(), nn.Linear(n_h_neurons, dim_out, bias=True), nn.Hardtanh(in_bound, out_bound)] # hidden layer module_hidden = [[nn.ReLU(), nn.Linear(n_h_neurons, n_h_neurons, bias=True)] for _ in range(n_h_layers - 1)] layer_hidden = list(np.array(module_hidden).flatten()) # nn model layers = layer_input + layer_hidden + layer_output self.model = nn.Sequential(*layers) print(self.model) # 添加以下forward方法 def forward(self, x): return self.model(x)
2. 权重加载后state_dict不匹配的问题
原因:加载的预训练权重字典trained_nn的键名(如'0.weight'),与当前trained_model.state_dict()的键名(如'model.0.weight')结构不一致。因为训练时保存的是直接的Sequential模型权重,而现在你将Sequential封装到了Model类的self.model属性下,导致键名多了model.前缀。加上strict=False后,PyTorch会跳过不匹配的键,权重仍为初始化的随机值,并未成功加载。
解决方法(二选一):
方法一:修改权重字典的键名
给预训练权重的每个键添加model.前缀,使其与当前模型的键名匹配:
from collections import OrderedDict trained_nn = torch.load('path') # 重新构建权重字典,修改键名 new_state_dict = OrderedDict() for k, v in trained_nn.items(): new_key = f'model.{k}' new_state_dict[new_key] = v trained_model = Model(1,5,2,1,-1,1) trained_model.load_state_dict(new_state_dict, strict=True) # 此时可以用strict=True确保所有权重加载成功
方法二:直接使用Sequential模型(更简洁)
如果训练时的模型就是直接的Sequential结构,无需封装到自定义Model类中,直接构建相同结构并加载权重:
import torch.nn as nn import torch # 构建与训练时一致的Sequential模型 model = nn.Sequential( nn.Linear(2, 5, bias=True), nn.ReLU(), nn.Linear(5, 1, bias=True), nn.Hardtanh(-1, 1) ) trained_nn = torch.load('path') model.load_state_dict(trained_nn) # PyTorch新版本无需使用Variable,直接用张量作为dummy输入 dummy_input = torch.randn(1, 2) torch.onnx.export(model, dummy_input, 'file.onnx', verbose=True)
内容的提问来源于stack exchange,提问作者Acad
相关产品推荐
相关产品推荐

