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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 10:10:29