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

Android Kotlin无法加载训练后.tflite模型及计算边界框问题

解决YOLOv5转TFLite模型在Android Kotlin中加载失败的问题

报错Input tensor has type kTfLiteFloat32: it requires specifying NormalizationOptions metadata to preprocess input images的核心原因是TensorFlow Lite Task Vision库的ObjectDetector要求Float32输入的模型必须附带图像归一化元数据,而YOLOv5导出的.tflite模型默认未添加该元数据,导致初始化失败。以下是两种可行解决方案:


方案一:重新导出带元数据的TFLite模型

1. 用YOLOv5导出脚本自动添加元数据

执行导出命令时,添加--metadata参数(YOLOv5输入通常将像素值除以255归一化到[0,1]):

python export.py --weights model-fp16.pt --include tflite --metadata --img-size 640 640

导出后的模型会自动包含NormalizationOptions元数据,直接用ObjectDetector.createFromFileAndOptions加载即可。

2. 手动给现有模型添加元数据

若不想重新导出,可通过TensorFlow Lite Metadata Writer工具手动添加。创建Python脚本如下:

from tflite_support import flatbuffers
from tflite_support import metadata as _metadata
from tflite_support import metadata_schema_py_generated as _metadata_fb

# 加载现有模型
model_file = "model-fp16.tflite"
with open(model_file, "rb") as f:
    model_buffer = f.read()

# 配置输入元数据与归一化参数
input_metadata = _metadata_fb.TensorMetadataT()
input_metadata.name = "image"
input_metadata.description = "Input image to be detected. Expected shape: [1, 640, 640, 3], pixels normalized to [0,1]."

normalization = _metadata_fb.ProcessUnitT()
normalization.optionsType = _metadata_fb.ProcessUnitOptions.NormalizationOptions
normalization.options = _metadata_fb.NormalizationOptionsT()
normalization.options.mean = [0.0, 0.0, 0.0]
normalization.options.std = [255.0, 255.0, 255.0]
input_metadata.processUnits = [normalization]

# 配置输出元数据
output_metadata = _metadata_fb.TensorMetadataT()
output_metadata.name = "detection_results"
output_metadata.description = "Bounding boxes, classes, scores."

# 组装模型元数据并写入
model_metadata = _metadata_fb.ModelMetadataT()
model_metadata.name = "YOLOv5 Object Detector"
model_metadata.description = "Detects objects from input images."
model_metadata.inputTensorMetadata = [input_metadata]
model_metadata.outputTensorMetadata = [output_metadata]

b = flatbuffers.Builder(0)
b.Finish(model_metadata.Pack(b), _metadata.MetadataPopulator.METADATA_FILE_IDENTIFIER)
metadata_buffer = b.Output()

populator = _metadata.MetadataPopulator.with_model_buffer(model_buffer)
populator.load_metadata_buffer(metadata_buffer)
populator.populate()

# 保存带元数据的模型
with open("model-fp16-with-metadata.tflite", "wb") as f:
    f.write(populator.get_model_buffer())

运行脚本后得到带元数据的模型,替换原模型即可正常加载。


方案二:绕过Task Library自动预处理,手动加载模型

若不想修改模型,可直接使用TfliteInterpreter手动处理图像输入与输出解析:

1. 添加依赖

确保build.gradle(app)中包含基础依赖:

implementation 'org.tensorflow:tensorflow-lite:2.15.0'
implementation 'org.tensorflow:tensorflow-lite-support:0.4.4'

2. 编写检测类与调用代码

import org.tensorflow.lite.Interpreter
import org.tensorflow.lite.support.common.FileUtil
import android.graphics.Bitmap
import android.graphics.Matrix

class YoloV5Detector(context: Context) {
    private val interpreter: Interpreter
    private val inputSize = 640 // 对应模型输入尺寸

    init {
        val modelFile = FileUtil.loadMappedFile(context, "model-fp16.tflite")
        val options = Interpreter.Options().apply {
            setNumThreads(4)
            setUseNNAPI(true) // 启用NNAPI加速
        }
        interpreter = Interpreter(modelFile, options)
    }

    fun detect(bitmap: Bitmap): List<DetectionResult> {
        // 预处理:缩放图像、转Float32数组并归一化到[0,1]
        val resizedBitmap = resizeBitmap(bitmap, inputSize, inputSize)
        val inputArray = convertBitmapToFloatArray(resizedBitmap)

        // 准备输出数组(需匹配YOLOv5模型输出形状,示例为[1,25200,85])
        val outputShape = interpreter.getOutputTensor(0).shape()
        val outputArray = Array(outputShape[0]) { Array(outputShape[1]) { FloatArray(outputShape[2]) } }

        // 执行推理
        interpreter.run(inputArray, outputArray)

        // 解析输出并过滤结果
        return parseOutput(outputArray)
    }

    private fun resizeBitmap(bitmap: Bitmap, width: Int, height: Int): Bitmap {
        val matrix = Matrix()
        val scaleWidth = width.toFloat() / bitmap.width
        val scaleHeight = height.toFloat() / bitmap.height
        matrix.postScale(scaleWidth, scaleHeight)
        return Bitmap.createBitmap(bitmap, 0, 0, bitmap.width, bitmap.height, matrix, true)
    }

    private fun convertBitmapToFloatArray(bitmap: Bitmap): Array<Array<Array<FloatArray>>> {
        val intValues = IntArray(inputSize * inputSize)
        bitmap.getPixels(intValues, 0, bitmap.width, 0, 0, bitmap.width, bitmap.height)

        val inputArray = Array(1) { Array(inputSize) { Array(inputSize) { FloatArray(3) } } }
        for (y in 0 until inputSize) {
            for (x in 0 until inputSize) {
                val pixel = intValues[y * inputSize + x]
                inputArray[0][y][x][0] = ((pixel shr 16 and 0xFF) / 255.0f) // R通道
                inputArray[0][y][x][1] = ((pixel shr 8 and 0xFF) / 255.0f)  // G通道
                inputArray[0][y][x][2] = ((pixel and 0xFF) / 255.0f)        // B通道
            }
        }
        return inputArray
    }

    private fun parseOutput(outputArray: Array<Array<Array<FloatArray>>>): List<DetectionResult> {
        val results = mutableListOf<DetectionResult>()
        val detections = outputArray[0]
        val confidenceThreshold = 0.5f
        val iouThreshold = 0.5f

        for (detection in detections) {
            val confidence = detection[4]
            if (confidence < confidenceThreshold) continue

            // 获取最高置信度的类别
            var classIndex = 0
            var maxClassScore = 0.0f
            for (i in 5 until detection.size) {
                if (detection[i] > maxClassScore) {
                    maxClassScore = detection[i]
                    classIndex = i - 5
                }
            }
            if (maxClassScore < confidenceThreshold) continue

            // 将YOLO格式边界框转为左上右下坐标
            val xCenter = detection[0]
            val yCenter = detection[1]
            val width = detection[2]
            val height = detection[3]

            val left = (xCenter - width / 2) / inputSize
            val top = (yCenter - height / 2) / inputSize
            val right = (xCenter + width / 2) / inputSize
            val bottom = (yCenter + height / 2) / inputSize

            results.add(DetectionResult(left, top, right, bottom, confidence, classIndex))
        }

        // 执行非极大值抑制过滤重叠框
        return applyNMS(results, iouThreshold)
    }

    private fun applyNMS(results: List<DetectionResult>, iouThreshold: Float): List<DetectionResult> {
        val sortedResults = results.sortedByDescending { it.confidence }
        val keptResults = mutableListOf<DetectionResult>()

        for (result in sortedResults) {
            var overlap = false
            for (kept in keptResults) {
                val iou = calculateIOU(result, kept)
                if (iou > iouThreshold) {
                    overlap = true
                    break
                }
            }
            if (!overlap) keptResults.add(result)
        }
        return keptResults
    }

    private fun calculateIOU(a: DetectionResult, b: DetectionResult): Float {
        val intersectionLeft = maxOf(a.left, b.left)
        val intersectionTop = maxOf(a.top, b.top)
        val intersectionRight = minOf(a.right, b.right)
        val intersectionBottom = minOf(a.bottom, b.bottom)

        val intersectionArea = maxOf(0.0f, intersectionRight - intersectionLeft) * maxOf(0.0f, intersectionBottom - intersectionTop)
        val areaA = (a.right - a.left) * (a.bottom - a.top)
        val areaB = (b.right - b.left) * (b.bottom - b.top)

        return intersectionArea / (areaA + areaB - intersectionArea)
    }

    data class DetectionResult(
        val left: Float,
        val top: Float,
        val right: Float,
        val bottom: Float,
        val confidence: Float,
        val classIndex: Int
    )
}

调用示例:

val detector = YoloV5Detector(context)
val inputBitmap = ... // 你的输入图像Bitmap
val detectionResults = detector.detect(inputBitmap)
// 处理检测结果,如绘制边界框、显示类别等

内容的提问来源于stack exchange,提问作者Arslan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 09:53:14