You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 09:03:11