预训练ONNX图像识别模型运行时的输入适配问题
ONNX模型输入适配问题解决
问题背景
我尝试运行第三方标注工具训练的ONNX图像识别预训练模型(基于工具预定义标签训练),目标是在工具外部调用该模型。用样本图片测试时,期望输出识别标签,但遇到输入适配障碍。模型所需输入包含4个参数(分别为image、scale_factor、image_shape、ratio,对应形状为[1,3,640,640]、[1,2]、[1,2]、[1,2])。
当前代码运行时抛出错误:
ValueError: Model requires 4 inputs. Input Feed contains 1
代码调整方案
修改代码中的模型输入部分,补充模型要求的4个输入参数,调整后的完整代码如下:
import cv2 import numpy as np import onnxruntime import pytesseract import PyPDF2 # Load the image image = cv2.imread("example.jpg") # Check if the image has been loaded successfully if image is None: raise ValueError("Failed to load the image") # Get the original image shape orig_height, orig_width = image.shape[:2] # Make sure the height and width are positive if orig_height <= 0 or orig_width <= 0: raise ValueError("Invalid image size") # Set the desired size of the resized image target_size = (640, 640) # Resize the image using cv2.resize resized_image = cv2.resize(image, target_size) # Convert BGR (cv2 default) to RGB, then transpose to [C, H, W] format # ONNX model usually expects CHW format instead of HWC input_image = cv2.cvtColor(resized_image, cv2.COLOR_BGR2RGB) input_image = input_image.transpose(2, 0, 1) # Add batch dimension to get [1, 3, 640, 640] input_image = np.expand_dims(input_image, axis=0).astype(np.float32) # Prepare other required inputs # scale_factor: [1,2] - ratio of resized size to original size scale_factor = np.array([[target_size[0]/orig_width, target_size[1]/orig_height]], dtype=np.float32) # image_shape: [1,2] - original image shape (height, width) image_shape = np.array([[orig_height, orig_width]], dtype=np.float32) # ratio: [1,2] - usually same as scale_factor for resize cases, adjust if needed ratio = np.array([[target_size[0]/orig_width, target_size[1]/orig_height]], dtype=np.float32) # Load the ONNX model session = onnxruntime.InferenceSession("ic/model.onnx") # Check if the model has been loaded successfully if session is None: raise ValueError("Failed to load the model") # Get the input names and shapes of the model (verify with your model's actual input names) inputs = session.get_inputs() for i, input_info in enumerate(inputs): print(f"Input {i}: name = {input_info.name}, shape = {input_info.shape}") # Run the ONNX model - match input names with your model's actual input names input_dict = {} for input_info in inputs: if input_info.name == "image": input_dict[input_info.name] = input_image elif input_info.name == "scale_factor": input_dict[input_info.name] = scale_factor elif input_info.name == "image_shape": input_dict[input_info.name] = image_shape elif input_info.name == "ratio": input_dict[input_info.name] = ratio output_name = session.get_outputs()[0].name prediction = session.run([output_name], input_dict)[0] # Postprocess the prediction to obtain the labels (ensure postprocess function is defined) # labels = postprocess(prediction) # Use PyTesseract to extract the text from the image text = pytesseract.image_to_string(image) # Print the labels and the text # print("Labels:", labels) print("Text:", text)
关键调整说明
- 图像格式转换:OpenCV默认读取为BGR格式,需转换为RGB;同时将HWC格式转为ONNX常用的CHW格式,并添加batch维度,得到符合要求的
[1,3,640,640]输入张量。 - 补充其他输入参数:根据模型要求,生成
scale_factor(缩放比例)、image_shape(原图尺寸)、ratio(通常与缩放比例一致,若模型有特殊要求可调整)三个张量,形状均为[1,2]。 - 输入字典匹配:通过模型的输入名称匹配对应的张量,确保所有4个输入都传入模型。
注意:如果模型的输入名称与上述示例不同,请根据
session.get_inputs()打印的实际名称调整input_dict中的键名。
内容的提问来源于stack exchange,提问作者learningPython
相关产品推荐
相关产品推荐

