I am trying to convert a custom network from PTH format to onnx format, but using the torch.onnx.export() method, TypeError occurs: Forward () missing 1 required positional argument: 'x_body' error.
import os.path as osp
import numpy as np
import onnx
import onnxruntime as ort
import torch
import torchvision
import torch.nn as nn
class Emotic(nn.Module):
def __init__(self, num_context_features, num_body_features):
super(Emotic, self).__init__()
self.num_context_features = num_context_features
self.num_body_features = num_body_features
self.fc1 = nn.Linear((self.num_context_features + num_body_features), 256)
self.bn1 = nn.BatchNorm1d(256)
self.d1 = nn.Dropout(p=0.5)
self.fc_cat = nn.Linear(256, 26)
self.fc_cont = nn.Linear(256, 3)
self.relu = nn.ReLU()
def forward(self, x_context, x_body): #定义前向传播
context_features = x_context.view(-1, self.num_context_features)
body_features = x_body.view(-1, self.num_body_features)
fuse_features = torch.cat((context_features, body_features),
1)
fuse_out = self.fc1(fuse_features)
fuse_out = self.bn1(fuse_out)
fuse_out = self.relu(fuse_out)
fuse_out = self.d1(fuse_out)
cat_out = self.fc_cat(fuse_out)
cont_out = self.fc_cont(fuse_out)
return cat_out, cont_out
test_arr=np.random.randn(10,3,224,224).astype(np.float32) #astype():数组的副本,转换为指定类型
dummy_input=torch.tensor(test_arr)
body_test_arr=np.random.randn(10,3,128,128).astype(np.float32)
body_dummy_input=torch.tensor(body_test_arr)
model_body=torch.load("./models/model_body1.pth")
model_context=torch.load("./models/model_context1.pth")
pred_context=model_context(torch.from_numpy(test_arr))
pred_body=model_body(torch.from_numpy(body_test_arr))
model_emotic=torch.load("./models/model_emotic1.pth")
model_emotic.eval()
torch_output=model_emotic(pred_context,pred_body) #torch.from_numpy():从numpy.ndarray创建Tensor
input_names=["input"]
output_names=["cat","cont"]
torch.onnx.export(model_emotic,
dummy_input,
"./models/model_emotic1.onnx",
verbose=False, #如果verbose指定了,将输出正在导出的跟踪的调试描述。默认:false
input_names=input_names, #按顺序分配给图的输入节点的名称
output_names=output_names) #按顺序分配给图的输入节点的名称
model=onnx.load("./models/model_emotic1.onnx")
ort_session=ort.InferenceSession("./models/model_emotic1.onnx")
onnx_outputs=ort_session.run(None,{'input':test_arr})
print('Export ONNX!')
Then the console reported an error: Traceback (most recent call last): File "E:/pythonProject/Pytorch2TFLite/PyTorch2ONNX.py", line 61, in torch.onnx.export(model_emotic, File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\onnx_init_.py", line 275, in export return utils.export(model, args, f, export_params, verbose, training, File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\onnx\utils.py", line 88, in export _export(model, args, f, export_params, verbose, training, input_names, output_names, File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\onnx\utils.py", line 689, in _export _model_to_graph(model, args, verbose, input_names, File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\onnx\utils.py", line 458, in _model_to_graph graph, params, torch_out, module = _create_jit_graph(model, args, File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\onnx\utils.py", line 422, in _create_jit_graph graph, torch_out = _trace_and_get_graph_from_model(model, args) File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\onnx\utils.py", line 373, in _trace_and_get_graph_from_model torch.jit._get_trace_graph(model, args, strict=False, _force_outplace=False, _return_inputs_states=True) File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\jit_trace.py", line 1160, in _get_trace_graph outs = ONNXTracedModule(f, strict, _force_outplace, return_inputs, _return_inputs_states)(*args, **kwargs) File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\nn\modules\module.py", line 1051, in _call_impl return forward_call(*input, **kwargs) File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\jit_trace.py", line 127, in forward graph, out = torch._C._create_graph_by_tracing( File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\jit_trace.py", line 118, in wrapper outs.append(self.inner(*trace_inputs)) File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\nn\modules\module.py", line 1051, in _call_impl return forward_call(*input, **kwargs) File "D:\Anaconda3\lib\site-packages\torch-1.9.0-py3.8-win-amd64.egg\torch\nn\modules\module.py", line 1039, in _slow_forward result = self.forward(*input, **kwargs) TypeError: forward() missing 1 required positional argument: 'x_body'