使用ONNX Runtime时遇Invalid Argument Error,代码是否正确?
问题分析与解决
错误原因
报错提示输入float_input的第二个维度(索引1)实际值为1,但模型期望是4,核心问题是输入数据的特征维度与模型训练时的输入维度不匹配,同时你的代码在提取输入数据时存在逻辑错误。
代码问题定位
你想要提取数据集的首行作为输入,但X.iloc[:,[0]]的作用是提取所有行的第一列,得到的是形状为(183, 1)的数组,既不是首行数据,特征数也只有1,完全不符合模型的输入要求。
修正步骤
1. 正确提取并调整输入数据形状
将提取首行的代码替换为以下内容,确保取到首行的所有特征,并调整为模型期望的批量输入格式(通常ONNX模型要求输入为二维数组,形状为(批量大小, 特征数)):
# 提取首行的所有特征,转换为numpy数组 x = X.iloc[0, :].to_numpy() # 调整形状为(1, 特征数),适配模型的批量输入要求 x = x.reshape(1, -1) # 转换为模型要求的float32类型 x = x.astype(numpy.float32)
2. 确认模型输入要求
可以先打印模型期望的输入形状,明确特征维度是否正确:
input_shape = sess.get_inputs()[0].shape print("模型期望的输入形状:", input_shape)
如果输出类似[None, 4],说明模型要求输入的特征数为4,但你的当前X有20000列,这意味着你在数据转换环节可能遗漏了和训练时一致的特征选择/降维步骤,需要核对训练流程,确保当前输入的特征数量、顺序、预处理逻辑和训练阶段完全一致。
3. 完整修正后的代码片段
import pandas as pd import pyodbc import datetime import numpy import onnxruntime as rt data = pd.read_sql(data_sql,cnxn) # 确保此处的转换步骤和训练模型时完全一致 X = data.drop('Field_T',axis=1) sess = rt.InferenceSession("Importedrf_iris.onnx") input_name = sess.get_inputs()[0].name output_name = sess.get_outputs()[0].name # 修正输入数据提取与形状调整 x = X.iloc[0, :].to_numpy() x = x.reshape(1, -1) x = x.astype(numpy.float32) res = sess.run([output_name], {input_name: x})[0] print(res)
内容的提问来源于stack exchange,提问作者Bryan Socha
相关产品推荐
相关产品推荐

