PyTorch转ONNX再转TensorFlow导出图时遇KeyError: 'input.1'求助
PyTorch转ONNX再转TFLite时KeyError: 'input.1'的解决方法
我将PyTorch实现的LeNet5模型转成ONNX格式,再构建TensorFlow表示,最后尝试导出TFLite时,在导出图阶段碰到了KeyError: 'input.1'错误,相关代码如下:
torch.save(model.state_dict(), '/Plastic_dataset/MyDrive/model.pth') trained_model = LeNet5(activation='relu', conv_size=3, pooling='avg', use_batch_norm=True) trained_model.load_state_dict(torch.load('/Plastic_dataset/MyDrive/model.pth'), strict = False) dummy_input = torch.randn(1, 1, 64, 64) torch.onnx.export(trained_model, dummy_input, '/Plastic_dataset/MyDrive/model.onnx') onnx_model = onnx.load('/Plastic_dataset/MyDrive/model.onnx') tf_rep = prepare(onnx_model) tf_rep.export_graph("/Plastic_dataset/MyDrive/model.pb") converter = tf.lite.TFLiteConverter.from_frozen_graph( "/Plastic_dataset/MyDrive/model.pb", tf_rep.inputs, tf_rep.outputs) tflite_model = converter.convert() open("/Plastic_dataset/MyDrive/model.tflite", "wb").write(tflite_model)
错误原因
这个错误是因为TFLite转换器找不到你指定的输入节点名称。通常是ONNX转TensorFlow过程中,自动生成的输入节点名和你传入转换器的名称不匹配,或者导出的pb图中根本不存在input.1这个节点。
解决步骤
1. 导出ONNX时显式指定输入输出名称
避免PyTorch自动生成混乱的节点名称,在torch.onnx.export时通过参数固定输入输出名:
dummy_input = torch.randn(1, 1, 64, 64) # 显式指定输入输出节点名,比如'input'和'output' torch.onnx.export(trained_model, dummy_input, '/Plastic_dataset/MyDrive/model.onnx', input_names=['input'], output_names=['output'])
2. 确认TensorFlow图的实际输入输出节点名
转换为TensorFlow表示后,不要直接使用tf_rep.inputs,先打印确认实际的节点名称:
tf_rep = prepare(onnx_model) # 打印输入输出节点名,确认是否和预期一致 print("输入节点名:", tf_rep.inputs) print("输出节点名:", tf_rep.outputs) # 导出pb图 tf_rep.export_graph("/Plastic_dataset/MyDrive/model.pb")
3. 修正TFLite转换器的输入参数
用打印确认后的节点名替换原代码中的tf_rep.inputs和tf_rep.outputs:
# 假设打印出的输入是['input'],输出是['output'] converter = tf.lite.TFLiteConverter.from_frozen_graph( "/Plastic_dataset/MyDrive/model.pb", ['input'], ['output']) tflite_model = converter.convert() open("/Plastic_dataset/MyDrive/model.tflite", "wb").write(tflite_model)
额外排查:查看pb图所有节点
如果还是无法确定正确的节点名,可以用以下代码列出pb图中所有节点的名称:
import tensorflow as tf def print_graph_nodes(pb_path): with tf.io.gfile.GFile(pb_path, 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) for node in graph_def.node: print(node.name) print_graph_nodes("/Plastic_dataset/MyDrive/model.pb")
找到对应的输入节点名后,替换到转换器参数中即可。
内容的提问来源于stack exchange,提问作者Yaroslav
相关产品推荐
相关产品推荐

