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
相关产品推荐
相关产品推荐

