如何修改Yolov5s导出的TFLite模型输出适配Kotlin安卓应用
问题
将导出的Yolov5s TFLite模型加载到Kotlin开发的安卓目标检测应用时,模型输出为形状[1, 25200, 9]的数组,而官方目标检测示例预期输出是detection_boxes、detection_classes、detection_scores、num_detections四个独立数组。当前尝试的检测框绘制逻辑仅在屏幕左上角生成静态框,且应用几秒后崩溃。
模型示例代码:
val model = BestFp16.newInstance(context) // Creates inputs for reference. val inputFeature0 = TensorBuffer.createFixedSize(intArrayOf(1, 640, 640, 3), DataType.FLOAT32) inputFeature0.loadBuffer(byteBuffer) // Runs model inference and gets result. val outputs = model.process(inputFeature0) val outputFeature0 = outputs.outputFeature0AsTensorBuffer // Releases model resources if no longer used. model.close()
原MainActivity.kt中注释的绘制逻辑无法正常工作,应用崩溃。
解决方案
核心修改点
- 正确解析Yolov5 TFLite输出:Yolov5的TFLite输出每个检测框包含9个元素,依次为
x_center, y_center, width, height, confidence, 类别1得分, ..., 类别8得分,需要遍历所有25200个候选框并过滤低置信度结果。 - 修复坐标转换逻辑:处理相机预览的镜像问题,确保检测框位置与实际物体对齐。
- 优化性能避免崩溃:将模型推理和绘制逻辑移到后台线程,避免主线程高频阻塞。
修改后的完整MainActivity.kt代码
package com.example.sightfulkotlin import android.annotation.SuppressLint import android.content.Context import android.content.pm.PackageManager import android.graphics.* import android.hardware.camera2.CameraCaptureSession import android.hardware.camera2.CameraDevice import android.hardware.camera2.CameraManager import android.os.Bundle import android.os.Handler import android.os.HandlerThread import android.view.Surface import android.view.TextureView import android.widget.ImageView import androidx.appcompat.app.AppCompatActivity import androidx.core.content.ContextCompat import com.example.sightfulkotlin.ml.BestFp16 import org.tensorflow.lite.DataType import org.tensorflow.lite.support.common.FileUtil import org.tensorflow.lite.support.image.ImageProcessor import org.tensorflow.lite.support.image.TensorImage import org.tensorflow.lite.support.image.ops.ResizeOp import org.tensorflow.lite.support.tensorbuffer.TensorBuffer class MainActivity : AppCompatActivity() { var colors = listOf( Color.BLUE, Color.GREEN, Color.RED, Color.CYAN, Color.GRAY, Color.BLACK, Color.DKGRAY, Color.MAGENTA, Color.YELLOW, Color.LTGRAY, Color.WHITE ) val paint = Paint() private lateinit var labels: List<String> lateinit var imageView: ImageView lateinit var cameraDevice: CameraDevice lateinit var handler: Handler private lateinit var cameraManager: CameraManager lateinit var textureView: TextureView lateinit var model: BestFp16 private val inferenceHandler = Handler(HandlerThread("inferenceThread").apply { start() }.looper) override fun onCreate(savedInstanceState: Bundle?) { super.onCreate(savedInstanceState) setContentView(R.layout.activity_main) getPermission() labels = FileUtil.loadLabels(this, "labels.txt") model = BestFp16.newInstance(this) val imageProcessor = ImageProcessor.Builder().add(ResizeOp(640, 640, ResizeOp.ResizeMethod.BILINEAR)).build() val handlerThread = HandlerThread("videoThread") handlerThread.start() handler = Handler(handlerThread.looper) paint.style = Paint.Style.STROKE paint.strokeWidth = 5f paint.textSize = 40f paint.textAlign = Paint.Align.LEFT imageView = findViewById(R.id.imageView) textureView = findViewById(R.id.textureView) textureView.surfaceTextureListener = object : TextureView.SurfaceTextureListener { override fun onSurfaceTextureAvailable(p0: SurfaceTexture, p1: Int, p2: Int) { openCamera() } override fun onSurfaceTextureSizeChanged(p0: SurfaceTexture, p1: Int, p2: Int) {} override fun onSurfaceTextureDestroyed(p0: SurfaceTexture): Boolean { return false } override fun onSurfaceTextureUpdated(p0: SurfaceTexture) { val currentBitmap = textureView.bitmap ?: return // 后台线程处理推理与绘制 inferenceHandler.post { val tensorImage = TensorImage(DataType.FLOAT32).apply { load(currentBitmap) imageProcessor.process(this) } val inputFeature0 = TensorBuffer.createFixedSize(intArrayOf(1, 640, 640, 3), DataType.FLOAT32) inputFeature0.loadBuffer(tensorImage.buffer) val outputs = model.process(inputFeature0) val outputArray = outputs.outputFeature0AsTensorBuffer.floatArray val mutableBitmap = currentBitmap.copy(Bitmap.Config.ARGB_8888, true) val canvas = Canvas(mutableBitmap) val imgHeight = currentBitmap.height val imgWidth = currentBitmap.width val numDetections = 25200 val detectionSize = 9 val confidenceThreshold = 0.5f for (i in 0 until numDetections) { val offset = i * detectionSize val xCenter = outputArray[offset] val yCenter = outputArray[offset + 1] val width = outputArray[offset + 2] val height = outputArray[offset + 3] val confidence = outputArray[offset + 4] // 过滤低置信度结果 if (confidence < confidenceThreshold) continue // 识别最高得分的类别 var maxClassScore = 0f var classIndex = 0 for (j in 5 until detectionSize) { if (outputArray[offset + j] > maxClassScore) { maxClassScore = outputArray[offset + j] classIndex = j - 5 } } // 转换为屏幕坐标(处理相机镜像) val left = (1 - (xCenter + width / 2)) * imgWidth val right = (1 - (xCenter - width / 2)) * imgWidth val top = (yCenter - height / 2) * imgHeight val bottom = (yCenter + height / 2) * imgHeight // 设置类别对应颜色 paint.color = colors[classIndex % colors.size] // 绘制检测框 canvas.drawRect(left, top, right, bottom, paint) // 绘制标签与置信度 val label = "${labels[classIndex]}: %.2f".format(confidence) canvas.drawText(label, left, top - 10, paint) } // 主线程更新UI runOnUiThread { imageView.setImageBitmap(mutableBitmap) } } } } cameraManager = getSystemService(Context.CAMERA_SERVICE) as CameraManager } override fun onDestroy() { super.onDestroy() model.close() inferenceHandler.looper.quitSafely() } @SuppressLint("MissingPermission") fun openCamera() { cameraManager.openCamera(cameraManager.cameraIdList[0], object : CameraDevice.StateCallback() { @SuppressLint("MissingPermission") override fun onOpened(p0: CameraDevice) { cameraDevice = p0 val surfaceTexture = textureView.surfaceTexture val surface = Surface(surfaceTexture) val captureRequest = cameraDevice.createCaptureRequest(CameraDevice.TEMPLATE_PREVIEW) captureRequest.addTarget(surface) cameraDevice.createCaptureSession(listOf(surface), object : CameraCaptureSession.StateCallback() { override fun onConfigured(p0: CameraCaptureSession) { p0.setRepeatingRequest(captureRequest.build(), null, handler) } override fun onConfigureFailed(p0: CameraCaptureSession) {} }, handler) } override fun onDisconnected(p0: CameraDevice) {} @SuppressLint("MissingPermission") override fun onError(p0: CameraDevice, p1: Int) {} }, handler) } fun getPermission() { if (ContextCompat.checkSelfPermission(this, android.Manifest.permission.CAMERA) != PackageManager.PERMISSION_GRANTED) { requestPermissions(arrayOf(android.Manifest.permission.CAMERA), 101) } } override fun onRequestPermissionsResult( requestCode: Int, permissions: Array<out String>, grantResults: IntArray ) { super.onRequestPermissionsResult(requestCode, permissions, grantResults) if (grantResults[0] != PackageManager.PERMISSION_GRANTED) { getPermission() } } }
关键修改说明
- 新增
inferenceThread后台线程,避免主线程因高频推理阻塞崩溃 - 遍历所有25200个候选框,加入0.5的置信度过滤(可根据需求调整)
- 处理相机预览镜像问题,翻转x轴坐标使检测框位置准确
- 自动识别每个检测框的最高得分类别,绘制对应标签和置信度
- 优化Paint参数,确保检测框和文字清晰可见
内容的提问来源于stack exchange,提问作者Fatema Shawki
相关产品推荐
相关产品推荐

