提取ONNX模型中间层输出遇RuntimeError,求修复方案
解决ONNX Runtime提取中间层输出时的输入格式错误
问题场景
需要提取ONNX模型(如squeezenet.onnx)不同层的输出,参考相关代码实现后,尽管输入尺寸符合要求,但执行推理时出现以下RuntimeError:
---> 40 ort_outs = ort_session.run(outputs, {'data': img} ) 41 ort_outs = OrderedDict(zip(outputs, ort_outs)) /usr/local/lib/python3.7/dist-packages/onnxruntime/capi/onnxruntime_inference_collection.py in run(self, output_names, input_feed, run_options) 198 output_names = [output.name for output in self._outputs_meta] 199 try: ---> 200 return self._sess.run(output_names, input_feed, run_options) 201 except C.EPFail as err: 202 if self._enable_fallback: RuntimeError: Input must be a list of dictionaries or a single numpy array for input 'data'.
问题原因
ONNX Runtime仅支持接收numpy数组作为输入,而代码中传入的img是PyTorch张量,不符合输入格式要求。
修复方案
将PyTorch张量转换为numpy数组,同时确保输入键名与模型实际输入节点名一致(可通过ort_session.get_inputs()[0].name确认)。
修复后的完整代码
# add all intermediate outputs to onnx net import onnx import onnxruntime as ort from PIL import Image from torchvision import transforms from collections import OrderedDict ort_session = ort.InferenceSession('<your path>/model.onnx') org_outputs = [x.name for x in ort_session.get_outputs()] model = onnx.load('<your path>/model.onnx') for node in model.graph.node: for output in node.output: if output not in org_outputs: model.graph.output.extend([onnx.ValueInfoProto(name=output)]) # execute onnx ort_session = ort.InferenceSession(model.SerializeToString()) outputs = [x.name for x in ort_session.get_outputs()] # 确认模型输入节点名,避免键名错误 input_name = ort_session.get_inputs()[0].name img_path = '<your path>/input_img.raw' img = Image.open(img_path).convert('RGB') # 替换为标准图片加载逻辑 transform_fn = transforms.Compose([ transforms.Resize(224), transforms.ToTensor(), ]) img = transform_fn(img) img = img.expand_dims(axis=0) # 关键修复:将PyTorch张量转为numpy数组 img = img.numpy() ort_outs = ort_session.run(outputs, {input_name: img} ) ort_outs = OrderedDict(zip(outputs, ort_outs))
额外注意事项
- 部分ONNX模型要求输入为BGR通道顺序,而PyTorch的
ToTensor()会将RGB转为CHW格式,若模型需要BGR,可添加通道反转:img = img[:, [2,1,0], :, :] - 确保numpy数组的数据类型与模型输入要求一致(通常为float32,
ToTensor()输出的张量默认是float32,转numpy后无需额外转换)
内容的提问来源于stack exchange,提问作者andreadaou
相关产品推荐
相关产品推荐

