导出自定义检测模型为TFLite时遇KeyError:1720求助
解决TFLite元数据写入时的KeyError:1720问题
错误原因
这个错误是因为你的自定义目标检测TFLite模型的输出张量索引,和ObjectDetectorWriter.create_for_inference()方法预设的标准SSD类模型输出索引不匹配。该快捷方法默认假设模型输出符合固定结构的张量索引,但你的模型导出后输出索引发生了变化,导致找不到索引为1720的张量。
解决步骤
1. 先确认模型实际的输出张量信息
运行以下代码,打印模型所有输出张量的索引、名称和形状,明确你的模型真实输出结构:
import tensorflow as tf interpreter = tf.lite.Interpreter(model_path="/content/detect.tflite") interpreter.allocate_tensors() # 遍历输出张量详情 output_details = interpreter.get_output_details() for detail in output_details: print(f"索引: {detail['index']}, 名称: {detail['name']}, 形状: {detail['shape']}")
执行后你会得到类似索引:0、索引:1这类真实的输出张量索引,替换掉预设的1720。
2. 手动指定输出张量索引构建元数据
放弃使用create_for_inference()快捷方法,改用手动指定参数的方式构建元数据writer,示例代码如下:
from tensorflow_lite_support.metadata.python.metadata_writers import object_detector from tensorflow_lite_support.metadata.python.metadata_writers import writer_utils from tensorflow_lite_support.metadata import schema_py_generated as _metadata_fb _MODEL_PATH = "/content/detect.tflite" _LABEL_FILE = "/content/labelmap.txt" _SAVE_TO_PATH = "/content/tflite_with_metadata/detect.tflite" # 加载模型和标签文件 model_buffer = writer_utils.load_file(_MODEL_PATH) label_buffer = writer_utils.load_file(_LABEL_FILE) # 替换为第一步查到的真实输出张量索引(示例为0、1、2,根据你的实际情况修改) output_tensor_indices = [0, 1, 2] # 手动构建元数据writer writer = object_detector.MetadataWriter( model_buffer=model_buffer, input_metadata=object_detector.InputMetadata( name="image", content=_metadata_fb.ContentT(contentProperties=_metadata_fb.ImagePropertiesT(colorSpace=_metadata_fb.ColorSpaceType.RGB)), normalization=_metadata_fb.ProcessUnitT(options=_metadata_fb.NormalizationOptionsT(mean=[127.5], std=[127.5])) ), output_metadata=[ # 检测框输出元数据(对应你的第一个输出张量) object_detector.OutputMetadata( name="bounding_boxes", content=_metadata_fb.ContentT(contentProperties=_metadata_fb.BoundingBoxPropertiesT(type=_metadata_fb.BoundingBoxType.BOUNDING_BOX, coordinateType=_metadata_fb.CoordinateType.RATIO)), associated_files=[object_detector.AssociatedFile(label_buffer, "labels.txt")] ), # 类别输出元数据(对应你的第二个输出张量) object_detector.OutputMetadata( name="classes", content=_metadata_fb.ContentT(contentProperties=_metadata_fb.ClassificationPropertiesT()), associated_files=[object_detector.AssociatedFile(label_buffer, "labels.txt")] ), # 置信度分数输出元数据(对应你的第三个输出张量) object_detector.OutputMetadata( name="scores", content=_metadata_fb.ContentT(contentProperties=_metadata_fb.ClassificationPropertiesT()), associated_files=[object_detector.AssociatedFile(label_buffer, "labels.txt")] ) ], output_tensor_indices=output_tensor_indices ) # 保存带元数据的模型 writer_utils.save_file(writer.populate(), _SAVE_TO_PATH) # 验证元数据 displayer = metadata.MetadataDisplayer.with_model_file(_SAVE_TO_PATH) print("Metadata populated:") print(displayer.get_metadata_json()) print("Associated file(s) populated:") print(displayer.get_packed_associated_file_list())
注意:需要根据第一步查到的输出张量的实际含义(哪个是检测框、哪个是类别、哪个是分数),对应调整output_metadata的顺序和内容。
3. 检查模型导出流程
确认你导出TFLite模型时,没有修改默认输出节点,或者使用TensorFlow Object Detection API导出时,参数配置正确,确保所有必要的输出张量都被包含在导出的模型中。
内容的提问来源于stack exchange,提问作者Grandpa
相关产品推荐
相关产品推荐

