Android中TensorFlow Lite YOLOv8模型报错:Label轴1形状不匹配
YOLOv8转TensorFlow Lite模型在Android目标检测中的形状不匹配问题解决
问题场景
我正在开发一款使用TensorFlow Lite模型进行目标检测的Android应用,尝试处理图像时遇到错误:
Error classifying image: Label number 1 mismatch the shape on axis 1
相关代码片段
currentImageUri?.let { uri -> val imageStream: InputStream? = contentResolver.openInputStream(uri) val bitmap = BitmapFactory.decodeStream(imageStream) bitmap?.let { classifyImage(it) } } private fun classifyImage(bitmap: Bitmap) { try { val model = Model.newInstance(applicationContext) val resizedBitmap = Bitmap.createScaledBitmap(bitmap, 640, 640, true) val image = TensorImage.fromBitmap(resizedBitmap) val outputs = model.process(image) val output = outputs.outputAsCategoryList val resultLabels = output.filter { it.score > CONFIDENCE_THRESHOLD } .map { it.label } model.close() val intent = Intent(this, ListFoodActivity::class.java) if (resultLabels.isEmpty()) { intent.putExtra("noPrediction", true) } else { intent.putStringArrayListExtra("resultLabels", ArrayList(resultLabels)) } startActivity(intent) } catch (e: Exception) { e.printStackTrace() Log.e(TAG, "Error classifying image: ${e.message}") Toast.makeText(this, "Error classifying image", Toast.LENGTH_SHORT).show() } }
已尝试的操作
- 确保输入图像已调整为模型预期的640x640尺寸;
- 基于置信度阈值过滤输出类别。
问题解答
1. “Label number 1 mismatch the shape on axis 1”错误的含义
这个错误的核心是模型输出形状与代码解析逻辑不匹配:
- 你使用的是YOLOv8转TensorFlow Lite的目标检测模型,输出张量形状为
[1, 8400, 85](以COCO数据集为例:8400个候选检测框,每个框包含4个坐标值+1个框置信度+80个类别概率)。 - 代码中调用的
outputAsCategoryList方法是为图像分类模型设计的,它期望输出张量形状为[1, N_classes](单批次、N个类别的概率分布)。当该方法尝试解析YOLOv8的输出时,发现轴1(第二个维度)的长度是8400而非预期的类别数,因此抛出形状不匹配的错误。
2. 解决方法:正确解析YOLOv8的TFLite输出
需要放弃分类模型的封装方法,直接解析目标检测模型的原始输出张量,并添加非极大值抑制(NMS)过滤重复检测框。以下是修改后的代码示例:
// 定义检测结果数据类 data class Detection( val left: Float, val top: Float, val right: Float, val bottom: Float, val score: Float, val classIndex: Int ) private const val CONFIDENCE_THRESHOLD = 0.5f private const val IOU_THRESHOLD = 0.5f private fun classifyImage(bitmap: Bitmap) { try { val model = Model.newInstance(applicationContext) val resizedBitmap = Bitmap.createScaledBitmap(bitmap, 640, 640, true) val image = TensorImage.fromBitmap(resizedBitmap) // 获取模型原始输出张量(替换为你的模型实际输出特征名,可能是outputFeature0) val outputs = model.process(image) val outputTensor = outputs.outputFeature0AsTensor val outputBuffer = outputTensor.buffer outputBuffer.rewind() // YOLOv8 TFLite输出:[1, 8400, 85],固定处理8400个检测框 val numDetections = 8400 val detectionSize = 85 val rawDetections = ArrayList<Detection>() for (i in 0 until numDetections) { // 解析检测框坐标(中心x、中心y、宽度、高度) val cx = outputBuffer.float val cy = outputBuffer.float val width = outputBuffer.float val height = outputBuffer.float // 解析框置信度 val boxConfidence = outputBuffer.float // 找到当前框的最高类别概率及对应索引 var maxClassScore = 0f var classIndex = 0 for (j in 0 until 80) { // 80对应COCO数据集类别数,替换为你的模型类别数 val score = outputBuffer.float if (score > maxClassScore) { maxClassScore = score classIndex = j } } // 计算最终置信度(框置信度 × 类别概率),过滤低置信度结果 val finalScore = boxConfidence * maxClassScore if (finalScore > CONFIDENCE_THRESHOLD) { // 将YOLO格式的坐标转换为图像的绝对坐标(适配原始图像尺寸) val scaleX = bitmap.width / 640f val scaleY = bitmap.height / 640f val left = (cx - width / 2) * scaleX val top = (cy - height / 2) * scaleY val right = (cx + width / 2) * scaleX val bottom = (cy + height / 2) * scaleY rawDetections.add(Detection(left, top, right, bottom, finalScore, classIndex)) } } // 应用非极大值抑制,去除重复检测框 val filteredDetections = applyNMS(rawDetections, IOU_THRESHOLD) // 将类别索引映射为标签名称(替换为你的模型对应的标签列表) val resultLabels = filteredDetections.map { getLabelName(it.classIndex) } model.close() // 跳转结果页面 val intent = Intent(this, ListFoodActivity::class.java) if (resultLabels.isEmpty()) { intent.putExtra("noPrediction", true) } else { intent.putStringArrayListExtra("resultLabels", ArrayList(resultLabels)) } startActivity(intent) } catch (e: Exception) { e.printStackTrace() Log.e(TAG, "Error classifying image: ${e.message}") Toast.makeText(this, "Error classifying image", Toast.LENGTH_SHORT).show() } } // 非极大值抑制(NMS)实现:过滤重叠度过高的检测框 private fun applyNMS(detections: List<Detection>, iouThreshold: Float): List<Detection> { // 按置信度降序排序 val sortedDetections = detections.sortedByDescending { it.score } val keepDetections = mutableListOf<Detection>() for (current in sortedDetections) { var keep = true for (kept in keepDetections) { if (calculateIOU(current, kept) > iouThreshold) { keep = false break } } if (keep) { keepDetections.add(current) } } return keepDetections } // 计算两个检测框的交并比(IOU) private fun calculateIOU(a: Detection, b: Detection): Float { val intersectLeft = max(a.left, b.left) val intersectTop = max(a.top, b.top) val intersectRight = min(a.right, b.right) val intersectBottom = min(a.bottom, b.bottom) val intersectArea = max(0f, intersectRight - intersectLeft) * max(0f, intersectBottom - intersectTop) val areaA = (a.right - a.left) * (a.bottom - a.top) val areaB = (b.right - b.left) * (b.bottom - b.top) return intersectArea / (areaA + areaB - intersectArea) } // 替换为你的模型训练时使用的标签列表 private fun getLabelName(classIndex: Int): String { val labels = listOf( "person", "bicycle", "car", "motorcycle", "airplane", "bus", "train", "truck", "boat", "traffic light" // 补充完整标签 ) return labels.getOrElse(classIndex) { "Unknown" } }
关键注意事项
- 确认模型输出格式:部分YOLOv8转TFLite的模型可能将输出拆分为多个张量(如框坐标、置信度、类别概率),需根据模型实际输出调整解析逻辑。
- 标签列表一致性:
getLabelName中的标签顺序必须与训练模型时的标签顺序完全一致,否则会出现类别映射错误。 - 阈值调整:
CONFIDENCE_THRESHOLD和IOU_THRESHOLD可根据实际场景调整,平衡检测精度和召回率。
内容的提问来源于stack exchange,提问作者Alvin Fajar Permana
相关产品推荐
相关产品推荐

