TFLite Model Maker训练的模型在Android物体检测App报张量维度错误
错误原因
- 核心诱因是Google Colab默认安装的
tflite-model-maker包版本发生了变更:新版本的Model Maker导出EfficientDet系列检测模型时,会默认压缩掉固定为1的batch维度,将原本3维的输出张量[1, 检测数量, 数据长度]压缩为2维[检测数量, 数据长度] - 你使用的TensorFlow官方Android物体检测示例代码是适配旧版3维输出格式开发的,解析输出张量时默认读取3维结构,就会抛出维度不匹配的错误
- 由于你训练时没有固定依赖包版本,每次Colab运行都会拉取最新版Model Maker,就出现了原有流程无修改但报错的情况
解决方案
你可以根据自身需求任选以下任意一种方案解决:
- 方案1(最快适配,无需修改原有逻辑):固定Model Maker版本为之前稳定运行的旧版
将训练代码中安装依赖的行替换为指定版本的安装命令即可,常用的兼容旧版为0.3.4:
重新训练导出的模型就和原有Android代码完全兼容。!pip install -q tflite-model-maker==0.3.4 - 方案2(无需降级依赖,保留新版功能):导出模型时关闭维度压缩逻辑
仅需修改model.export代码,添加disable_squeeze_output=True参数即可:
导出的模型会保留原有3维输出结构,直接兼容现有Android代码。model.export(export_dir='.', disable_squeeze_output=True) - 方案3(适配新版输出格式,无需修改训练逻辑):调整Android端张量解析代码
找到Android项目中解析检测输出张量的代码位置,将原本读取3维张量索引的逻辑调整为适配2维结构即可:
比如原逻辑读取边界框是outputTensors[0].getFloatArray()[0 * 4 * detCount + i *4],调整为outputTensors[0].getFloatArray()[i *4]即可,对应调整其余3个输出张量的读取索引即可。
内容的提问来源于stack exchange,提问作者danih1207
相关产品推荐
相关产品推荐

