ONNX Runtime中model.run()参数不兼容问题排查(MobileNet人脸检测)
问题:ONNX Runtime推理MobileNet 0.25人脸检测模型时参数不兼容错误
在人脸检测任务中,将MobileNet 0.25模型转为ONNX格式后,使用以下Python代码推理时出现参数错误:
import onnx import onnxruntime import cv2 import numpy as np import time import sys import onnx from onnxruntime import InferenceSession, RunOptions def input_output_layer(model_path): model = onnx.load(model_path) output =[node.name for node in model.graph.output] input_all = [node.name for node in model.graph.input] input_initializer = [node.name for node in model.graph.initializer] net_feed_input = list(set(input_all) - set(input_initializer)) print('Inputs: ', net_feed_input) print('Outputs: ', output) return net_feed_input, output model_path = "mnet.25.onnx" model = onnx.load(model_path) print(onnx.checker.check_model(model)) sess = InferenceSession(model_path) for t in sess.get_inputs(): print("input:", t.name, t.type, t.shape) for t in sess.get_outputs(): print("input:", t.name, t.type, t.shape) img_path = "Face.jpg" image = cv2.imread(img_path, cv2.IMREAD_COLOR) img_data = cv2.resize(image, (640, 640)).astype(np.float32) img_data = np.expand_dims(img_data, 0) print(f" onnx shapeee: {np.shape(img_data)}") img_data = np.transpose(img_data, [0, 3, 1, 2]) print(f" onnx shapeee: {np.shape(img_data)}, {type(img_data)}") session_option = onnxruntime.SessionOptions() session_option.log_severity_level = 4 model = onnxruntime.InferenceSession(model_path, sess_options=session_option, providers=['CPUExecutionProvider']) ort_inputs_name, ort_ouputs_names = input_output_layer(model_path) print(ort_inputs_name, ort_ouputs_names) start = time.time() ort_outs = model.run(ort_ouputs_names[0], {ort_inputs_name[0]: img_data.astype('float32')}) outputs = np.array(ort_outs[0]).astype("float32") print(outputs)
运行后抛出如下错误:
Traceback (most recent call last): File "onnx_test_1.py", line 56, in <module> ort_outs = model.run(output_x, {ort_inputs_name[0]: img_data.astype('float32')}, None) File "/home/mohammad/Documents/insightface/insightface.onnx.env/lib/python3.8/site-packages/onnxruntime/capi/onnxruntime_inference_collection.py", line 200, in run return self._sess.run(output_names, input_feed, run_options) TypeError: run(): incompatible function arguments. The following argument types are supported: 1. (self: onnxruntime.capi.onnxruntime_pybind11_state.InferenceSession, arg0: List[str], arg1: Dict[str, object], arg2: onnxruntime.capi.onnxruntime_pybind11_state.RunOptions) -> List[object] Invoked with: <onnxruntime.capi.onnxruntime_pybind11_state.InferenceSession object at 0x7f5e26b0d670>, 'output', {'data': array([[[[184., 184., 184., ..., 167., 167., 167.], [184., 184., 184., ..., 167., 167., 167.], [183., 183., 184., ..., 169., 169., 169.], ..., [243., 243., 243., ..., 254., 254., 254.], [243., 243., 243., ..., 255., 255., 255.], [243., 243., 243., ..., 255., 255., 255.]], [[196., 196., 196., ..., 178., 178., 178.], [196., 196., 196., ..., 178., 178., 178.], [196., 196., 196., ..., 179., 179., 179.], ..., [251., 251., 251., ..., 254., 255., 255.], [251., 251., 251., ..., 255., 255., 255.], [251., 251., 251., ..., 255., 255., 255.]], [[176., 176., 176., ..., 174., 175., 175.], [176., 176., 176., ..., 174., 175., 175.], [176., 176., 176., ..., 174., 175., 175.], ..., [249., 249., 249., ..., 254., 255., 255.], [250., 250., 250., ..., 255., 255., 255.], [250., 250., 250., ..., 255., 255., 255.]]]], dtype=float32)}, None
错误明确指出model.run()要求第一个参数为List[str]类型,但当前传入的是单个字符串。此前在图像质量评估模型上使用类似代码可正常运行,需要解决该参数问题。
解决方法
问题核心是ONNX Runtime的run()方法对第一个参数的要求:必须是包含输出节点名称的字符串列表,而非单个字符串。即便模型只有一个输出,也需要用列表包裹。
修改代码中调用model.run()的行:
# 原错误代码 ort_outs = model.run(ort_ouputs_names[0], {ort_inputs_name[0]: img_data.astype('float32')}) # 修改后代码 ort_outs = model.run([ort_ouputs_names[0]], {ort_inputs_name[0]: img_data.astype('float32')})
此外,后续处理输出时,因为run()返回的是列表,原代码中outputs = np.array(ort_outs[0]).astype("float32")无需修改,可正常提取第一个输出结果。
如果需要获取所有输出,直接传入ort_ouputs_names即可(本身就是列表):
ort_outs = model.run(ort_ouputs_names, {ort_inputs_name[0]: img_data.astype('float32')})
不同模型表现差异的原因是:部分ONNX Runtime版本可能对单个字符串参数做了兼容处理,但MobileNet 0.25对应的ONNX模型或当前使用的ONNX Runtime版本严格遵循API规范,要求必须传入列表类型。
内容的提问来源于stack exchange,提问作者BarzanHayati
相关产品推荐
相关产品推荐

