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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 00:38:33