TensorFlow Lite安卓图像分类报错Label number 3 mismatch the shape on axis 1
问题修复方案
你遇到的java.lang.IllegalArgumentException: Label number 3 mismatch the shape on axis 1错误,核心是模型输出张量形状、声明的输出缓冲形状、标签列表长度三者不匹配,具体修复步骤如下:
1. 修正输出概率缓冲的形状配置
你代码中错误将输入图片的张量形状配置给了输出缓冲,分类模型的输出是对应类别的概率数组,和输入图片尺寸无关:
- 你当前共有3个分类类别,输出张量的形状应为
[1, 3](1代表单批次推理,3对应类别总数) - 找到代码中
probabilityBuffer初始化的位置,替换为以下代码:
// 替换原错误的输出缓冲声明 TensorBuffer probabilityBuffer = TensorBuffer.createFixedSize(new int[]{1, 3}, DataType.UINT8);
如果修改后仍有类型报错,可以打印模型输出的实际数据类型做对应调整:在Interpreter初始化完成后添加日志Log.d("OutputType", tflite.getOutputTensor(0).dataType().toString());,替换DataType.UINT8为打印出来的实际类型即可。
2. 清理标签文件冗余内容
你的labels.txt末尾存在多余空行,FileUtil.loadLabels会将空行识别为一个独立标签,导致标签列表长度超过3,需要将标签文件末尾的空行完全删除,确保文件仅保留3行有效标签内容。
可选校验步骤
如果修改后仍有形状不匹配报错,可以添加日志打印模型实际的输出形状,和你配置的缓冲形状做对齐:
// 在Interpreter初始化完成后添加以下代码 int[] outputShape = tflite.getOutputTensor(0).shape(); Log.d("ModelOutputShape", Arrays.toString(outputShape));
打印出来的数组就是你需要配置给probabilityBuffer的实际形状。
内容的提问来源于stack exchange,提问作者Michael Paulinus
相关产品推荐
相关产品推荐

