自定义目标检测模型输出顺序异常致Android应用报错
问题
使用TF 2.10按官方教程训练自定义目标检测模型,转换为TFLite模型后在Android Java应用部署时出现错误:
EXCEPTION: Failed on interpreter inference -> Cannot copy from a TensorFlowLite tensor (StatefulPartionedCall:1) with shape [1,10] to a Java object with shape [1,10,4].
TF 2.6之前模型输出元数据顺序为boxes、classes、scores、检测数量,TF 2.6及之后变为scores、boxes、检测数量、classes。已尝试两种方案:
- 降级至TF 2.5可解决,但会引发其他库兼容问题,不优先考虑;
- 使用元数据写入器显式声明输出序列,仍出现相同异常。处理后模型输出详情如下:
[{'name': 'StatefulPartitionedCall:1', 'index': 249, 'shape': array([ 1, 10]), 'shape_signature': array([ 1, 10]), 'dtype': <class 'numpy.float32'>, 'quantization': (0.0, 0), 'quantization_parameters': {'scales': array([], dtype=float32), 'zero_points': array([], dtype=int32), 'quantized_dimension': 0}, 'sparsity_parameters': {}}, {'name': 'StatefulPartitionedCall:3', 'index': 247, 'shape': array([ 1, 10, 4]), 'shape_signature': array([ 1, 10, 4]), 'dtype': <class 'numpy.float32'>, 'quantization': (0.0, 0), 'quantization_parameters': {'scales': array([], dtype=float32), 'zero_points': array([], dtype=int32), 'quantized_dimension': 0}, 'sparsity_parameters': {}}, {'name': 'StatefulPartitionedCall:0', 'index': 250, 'shape': array([1]), 'shape_signature': array([1]), 'dtype': <class 'numpy.float32'>, 'quantization': (0.0, 0), 'quantization_parameters': {'scales': array([], dtype=float32), 'zero_points': array([], dtype=int32), 'quantized_dimension': 0}, 'sparsity_parameters': {}}, {'name': 'StatefulPartitionedCall:2', 'index': 248, 'shape': array([ 1, 10]), 'shape_signature': array([ 1, 10]), 'dtype': <class 'numpy.float32'>, 'quantization': (0.0, 0), 'quantization_parameters': {'scales': array([], dtype=float32), 'zero_points': array([], dtype=int32), 'quantized_dimension': 0}, 'sparsity_parameters': {}}]
输出形状顺序仍与旧版不匹配,现寻求无需修改Android应用代码的前提下,在TFLite转换阶段解决输出形状错位的方法。
附原转换脚本:
import tensorflow as tf import argparse parser = argparse.ArgumentParser( description="tfLite Converter") parser.add_argument("--saved_model_path", help="", type=str) parser.add_argument("--tflite_model_path", help="", type=str) args = parser.parse_args() converter = tf.lite.TFLiteConverter.from_saved_model(args.saved_model_path) tflite_model = converter.convert() with open(args.tflite_model_path, 'wb') as f: f.write(tflite_model)
解决方案
方法1:用包装模型重排输出顺序
创建一个包装函数,将原模型的输出按旧版顺序(boxes、classes、scores、检测数量)重新排列后再转换为TFLite。修改后的转换脚本如下:
import tensorflow as tf import argparse parser = argparse.ArgumentParser(description="tfLite Converter with output reordering") parser.add_argument("--saved_model_path", help="Path to saved model", type=str) parser.add_argument("--tflite_model_path", help="Path to output TFLite model", type=str) args = parser.parse_args() # 加载原SavedModel loaded_model = tf.saved_model.load(args.saved_model_path) infer_func = loaded_model.signatures["serving_default"] # 定义包装函数,按旧顺序重排输出 @tf.function(input_signature=infer_func.input_signature) def wrapped_inference(input_tensor): outputs = infer_func(input_tensor) # 对应关系:原输出张量 → 旧版语义 return { "detection_boxes": outputs["StatefulPartitionedCall:3"], "detection_classes": outputs["StatefulPartitionedCall:2"], "detection_scores": outputs["StatefulPartitionedCall:1"], "num_detections": outputs["StatefulPartitionedCall:0"] } # 保存包装后的模型 tf.saved_model.save(loaded_model, args.saved_model_path + "_wrapped", signatures={"serving_default": wrapped_inference}) # 转换为TFLite converter = tf.lite.TFLiteConverter.from_saved_model(args.saved_model_path + "_wrapped") tflite_model = converter.convert() with open(args.tflite_model_path, 'wb') as f: f.write(tflite_model)
方法2:转换时显式指定输出张量顺序
利用TFLiteConverter的output_arrays参数,直接指定输出张量的顺序为旧版要求的序列,修改转换脚本如下:
import tensorflow as tf import argparse parser = argparse.ArgumentParser(description="tfLite Converter") parser.add_argument("--saved_model_path", help="", type=str) parser.add_argument("--tflite_model_path", help="", type=str) args = parser.parse_args() converter = tf.lite.TFLiteConverter.from_saved_model(args.saved_model_path) # 按旧版顺序指定输出张量名称:boxes、classes、scores、检测数量 converter.output_arrays = ["StatefulPartitionedCall:3", "StatefulPartitionedCall:2", "StatefulPartitionedCall:1", "StatefulPartitionedCall:0"] tflite_model = converter.convert() with open(args.tflite_model_path, 'wb') as f: f.write(tflite_model)
方法3:修正元数据写入时的输出映射
之前使用元数据写入器未生效,是因为未正确将输出张量映射到旧版语义标签。重新编写元数据脚本,明确指定输出顺序和对应角色:
from tflite_support import metadata_writers from tflite_support.metadata_writers import object_detector from tflite_support.metadata_writers import writer_utils _MODEL_PATH = "your_tflite_model.tflite" # 替换为你的模型路径 _LABEL_FILE = "path/to/labelmap.txt" # 替换为你的标签文件路径 _SAVE_TO_PATH = "model_with_correct_metadata.tflite" writer = object_detector.MetadataWriter.create_for_inference( writer_utils.load_file(_MODEL_PATH), input_norm_mean=[127.5], input_norm_std=[127.5], label_file_paths=[_LABEL_FILE], # 按旧版顺序指定输出张量名称 output_names=["StatefulPartitionedCall:3", "StatefulPartitionedCall:2", "StatefulPartitionedCall:1", "StatefulPartitionedCall:0"] ) writer_utils.save_file(writer.populate(), _SAVE_TO_PATH)
内容的提问来源于stack exchange,提问作者peru_45
相关产品推荐
相关产品推荐

