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

预训练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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 18:32:28