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

如何修改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中注释的绘制逻辑无法正常工作,应用崩溃。

解决方案

核心修改点

  1. 正确解析Yolov5 TFLite输出:Yolov5的TFLite输出每个检测框包含9个元素,依次为x_center, y_center, width, height, confidence, 类别1得分, ..., 类别8得分,需要遍历所有25200个候选框并过滤低置信度结果。
  2. 修复坐标转换逻辑:处理相机预览的镜像问题,确保检测框位置与实际物体对齐。
  3. 优化性能避免崩溃:将模型推理和绘制逻辑移到后台线程,避免主线程高频阻塞。

修改后的完整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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 07:43:11