TensorFlow 2 SSD FPN系列模型Android端运行输出张量维度不匹配问题
报错根因
该错误是因为自行转换的TFLite模型输出缺失首维batch维度,安卓端TensorFlow Lite ObjectDetector任务库要求0号输出张量(检测框坐标)维度为[1, 检测框数量, 4],而你转换得到的模型输出维度为[检测框数量, 4],维度不匹配导致初始化失败。
这一问题是TensorFlow 2.6及以上版本Object Detection API的默认导出逻辑导致的,导出时会默认压缩掉batch维度,和移动端任务库的输入输出要求不兼容。
修复步骤
- 重新导出SavedModel时添加配置参数
执行官方的exporter_main_v2.py导出脚本时,额外添加--use_regular_nms=true参数,同时固定输入shape的首维为1,示例命令如下:
如果你用的是1024*1024输入的模型,对应修改python exporter_main_v2.py \ --input_type=float_image \ --pipeline_config_path=你的pipeline配置文件路径 \ --trained_checkpoint_dir=训练checkpoint存储路径 \ --output_directory=SavedModel导出路径 \ --use_regular_nms=true \ --input_shape=1,640,640,3input_shape里的分辨率参数即可。 - TFLite转换时添加强制维度配置
转换脚本中添加以下配置,禁止压缩batch维度:import tensorflow as tf saved_model_dir = "你的SavedModel路径" converter = tf.lite.TFLiteConverter.from_saved_model(saved_model_dir) # 开启默认优化 converter.optimizations = [tf.lite.Optimize.DEFAULT] # 指定支持的算子集 converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS] # 关闭tensor list压缩逻辑,保留batch维度 converter._experimental_lower_tensor_list_ops = False converter.experimental_new_converter = True tflite_model = converter.convert() with open("model.tflite", "wb") as f: f.write(tflite_model) - 转换后验证输出维度
可以执行以下代码确认输出维度符合要求:
正常情况下0号输出的维度为import tensorflow as tf interpreter = tf.lite.Interpreter(model_path="model.tflite") interpreter.allocate_tensors() for out_info in interpreter.get_output_details(): print(f"输出名:{out_info['name']},维度:{out_info['shape']}")[1, N, 4],其中N为你配置的最大检测框数量,此时模型即可在安卓端正常加载运行。
内容的提问来源于stack exchange,提问作者puelo
相关产品推荐
相关产品推荐

