将ONNX模型转TFLite时执行tf_rep.export_graph报错KeyError: 'input.1'
ONNX转TensorFlow时export_graph报错KeyError: 'input.1'
问题概述
将PyTorch导出的ONNX模型转换为TensorFlow格式,执行tf_rep.export_graph(tf_model_path)时触发KeyError: 'input.1',此前未找到明确解决方法。
环境依赖版本
- tensorflow: 2.12.0
- onnx: 1.14.0
- onnx-tf: 1.10.0
- Python: 3.10.12
执行代码
import torch import onnx import tensorflow as tf import onnx_tf from torchvision.models import resnet50 # 加载PyTorch ResNet50模型 pytorch_model = resnet50(pretrained=True) pytorch_model.eval() # 将PyTorch模型导出为ONNX格式 input_shape = (1, 3, 224, 224) dummy_input = torch.randn(input_shape) onnx_model_path = 'resnet50.onnx' torch.onnx.export(pytorch_model, dummy_input, onnx_model_path, opset_version=12, verbose=False) # 加载ONNX模型 onnx_model = onnx.load(onnx_model_path) # 将ONNX模型转换为TensorFlow格式 tf_model_path = 'resnet50.pb' # 原代码此处缺失闭合引号,已修正 onnx_model = onnx.load(onnx_model_path) from onnx_tf.backend import prepare tf_rep = prepare(onnx_model) tf_rep.export_graph(tf_model_path) # 错误发生在此行
错误信息
WARNING:absl:`input.1` is not a valid tf.function parameter name. Sanitizing to `input_1`. --------------------------------------------------------------------------- KeyError Traceback (most recent call last) <ipython-input-4-f35b83c104b8> in <cell line: 8>() 6 tf_model_path = 'resnet50' 7 tf_rep = prepare(onnx_model) ----> 8 tf_rep.export_graph(tf_model_path) 35 frames /usr/local/lib/python3.10/dist-packages/onnx_tf/handlers/backend/conv_mixin.py in tf__conv(cls, node, input_dict, transpose) 17 do_return = False 18 retval_ = ag__.UndefinedReturnValue() ---> 19 x = ag__.ld(input_dict)[ag__.ld(node).inputs[0]] 20 x_rank = ag__.converted_call(ag__.ld(len), (ag__.converted_call(ag__.ld(x).get_shape, (), None, fscope),), None, fscope) 21 x_shape = ag__.converted_call(ag__.ld(tf_shape), (ag__.ld(x), ag__.ld(tf).int32), None, fscope) KeyError: in user code: File "/usr/local/lib/python3.10/dist-packages/onnx_tf/backend_tf_module.py", line 99, in __call__ * output_ops = self.backend._onnx_node_to_tensorflow_op(onnx_node, File "/usr/local/lib/python3.10/dist-packages/onnx_tf/backend.py", line 347, in _onnx_node_to_tensorflow_op * return handler.handle(node, tensor_dict=tensor_dict, strict=strict) File "/usr/local/lib/python3.10/dist-packages/onnx_tf/handlers/handler.py", line 59, in handle * return ver_handle(node, **kwargs) File "/usr/local/lib/python3.10/dist-packages/onnx_tf/handlers/backend/conv.py", line 15, in version_11 * return cls.conv(node, kwargs["tensor_dict"]) File "/usr/local/lib/python3.10/dist-packages/onnx_tf/handlers/backend/conv_mixin.py", line 29, in conv * x = input_dict[node.inputs[0]] KeyError: 'input.1'
解决方案
错误根源是PyTorch默认导出的ONNX输入节点名称为input.1,包含TensorFlow不允许的.符号,导致onnx-tf转换时出现名称映射不匹配。可通过以下方法解决:
1. 导出ONNX时指定合法输入名称
修改torch.onnx.export代码,显式指定输入节点名称为不含特殊字符的标识符,比如input:
torch.onnx.export(pytorch_model, dummy_input, onnx_model_path, opset_version=12, verbose=False, input_names=["input"]) # 自定义输入名称
2. 修复已存在的ONNX模型
如果已经有导出好的ONNX模型,可直接修改其输入节点名称:
onnx_model = onnx.load(onnx_model_path) # 修改第一个输入节点的名称 onnx_model.graph.input[0].name = "input" # 保存修复后的模型 onnx.save(onnx_model, "resnet50_fixed.onnx") # 后续转换使用修复后的模型 onnx_model = onnx.load("resnet50_fixed.onnx")
3. 修正代码中的语法错误
原代码中tf_model_path = 'resnet50.pb缺失闭合引号,需修正为tf_model_path = 'resnet50.pb',避免额外报错。
完成以上修改后,重新执行转换流程即可成功导出TensorFlow模型。
内容的提问来源于stack exchange,提问作者afsara_ben
相关产品推荐
相关产品推荐

