TensorFlow Lite目标检测模型输出形状不匹配报错求助
嘿,我来帮你搞定这个报错!从你贴的Logcat日志来看,核心问题很明确:TFLite模型输出的detection_boxes张量形状是[1, 1917, 1, 4],但Demo代码里预期的是[1, 1917, 4],多了一个额外的维度导致不匹配。
下面是具体的解决步骤:
1. 理解问题根源
你基于ssd_mobilenet_v1_coco_2017_11_17训练的自定义模型,转换为TFLite后,输出的检测框张量保留了原始SSD模型的4维结构(批次数、检测框数量、单框维度、坐标值),但官方Demo的代码是按照3维结构来写的,所以才会抛出形状不匹配的异常。
2. 修改TFLiteObjectDetectionAPIModel.java的输出处理逻辑
找到recognizeImage方法中处理输出张量的部分,调整数组维度并提取有效数据:
// 原来的代码(预期3维数组) float[][][] outputLocations = new float[1][NUM_DETECTIONS][4]; tflite.getOutputTensor(0).copyTo(outputLocations); // 修改为适配4维输出的代码 float[][][][] outputLocations = new float[1][NUM_DETECTIONS][1][4]; tflite.getOutputTensor(0).copyTo(outputLocations); // 把4维数组转换成Demo需要的3维格式 float[][][] adjustedLocations = new float[1][NUM_DETECTIONS][4]; for (int i = 0; i < NUM_DETECTIONS; i++) { adjustedLocations[0][i] = outputLocations[0][i][0]; }
之后在后续的检测逻辑(比如解析检测框坐标)中,用adjustedLocations替代原来的outputLocations即可。
3. 同步检查其他输出张量的处理
除了detection_boxes,还要确认detection_scores、detection_classes这些输出的形状是否也有类似的维度差异,要是有的话,按照同样的方式调整数组结构。
4. 确认参数一致性
别忘了检查NUM_DETECTIONS这个常量的值,它需要和你的模型输出的检测框数量一致(SSD MobileNet V1默认是1917),如果训练时修改过这个参数,一定要同步更新Demo里的对应值。
5. 验证模型转换是否正确
如果上面的修改后还是有问题,可以先确认TFLite模型的输出形状是否正确。你可以用tflite_convert转换模型时,指定正确的输出节点,或者用TensorFlow的工具查看模型的输出结构:
saved_model_cli show --dir ./你的SavedModel路径 --tag_set serve --signature_def serving_default
这个命令会列出模型所有输出节点的形状,确保和你代码里处理的结构一致。
内容的提问来源于stack exchange,提问作者user 007

