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

如何在Kotlin中用单张图片运行训练好的TFLite关键点检测模型

Android TFLite 11关键点检测实现代码

1. 依赖配置(build.gradle)

确保添加TFLite核心依赖:

dependencies {
    implementation 'org.tensorflow:tensorflow-lite:2.15.0'
    implementation 'org.tensorflow:tensorflow-lite-support:0.4.4'
}

2. 核心实现代码

import android.graphics.Bitmap
import android.graphics.Canvas
import android.graphics.Color
import android.graphics.Paint
import org.tensorflow.lite.Interpreter
import java.io.FileInputStream
import java.nio.ByteBuffer
import java.nio.ByteOrder
import java.nio.MappedByteBuffer
import java.nio.channels.FileChannel

class KeyPointDetector(private val modelPath: String) {
    private lateinit var interpreter: Interpreter
    private val INPUT_SIZE = 640
    private val NUM_KEYPOINTS = 11
    private val OUTPUT_COLUMNS = 38
    private val OUTPUT_ROWS = 8400
    private val CONFIDENCE_THRESHOLD = 0.5f

    init {
        loadModel()
    }

    private fun loadModel() {
        val modelBuffer: MappedByteBuffer = FileInputStream(modelPath).channel.map(
            FileChannel.MapMode.READ_ONLY, 0, FileInputStream(modelPath).channel.size()
        )
        val options = Interpreter.Options()
        interpreter = Interpreter(modelBuffer, options)
    }

    // 预处理:将Bitmap转为模型要求的输入格式
    private fun preprocessBitmap(bitmap: Bitmap): ByteBuffer {
        val inputBuffer = ByteBuffer.allocateDirect(1 * INPUT_SIZE * INPUT_SIZE * 3 * 4)
        inputBuffer.order(ByteOrder.nativeOrder())
        val scaledBitmap = Bitmap.createScaledBitmap(bitmap, INPUT_SIZE, INPUT_SIZE, true)

        val intValues = IntArray(INPUT_SIZE * INPUT_SIZE)
        scaledBitmap.getPixels(intValues, 0, scaledBitmap.width, 0, 0, scaledBitmap.width, scaledBitmap.height)

        var pixelIndex = 0
        for (y in 0 until INPUT_SIZE) {
            for (x in 0 until INPUT_SIZE) {
                val pixel = intValues[pixelIndex++]
                // 归一化到0-1(若模型训练时用0-255输入,删除除以255.0f的操作)
                inputBuffer.putFloat((Color.red(pixel) / 255.0f))
                inputBuffer.putFloat((Color.green(pixel) / 255.0f))
                inputBuffer.putFloat((Color.blue(pixel) / 255.0f))
            }
        }
        return inputBuffer
    }

    // 推理并解析关键点结果
    fun detectKeyPoints(originalBitmap: Bitmap): List<Pair<Float, Float>> {
        val inputBuffer = preprocessBitmap(originalBitmap)
        val outputArray = Array(1) { Array(OUTPUT_COLUMNS) { FloatArray(OUTPUT_ROWS) } }
        interpreter.run(inputBuffer, outputArray)

        val keyPoints = mutableListOf<Pair<Float, Float>>()
        // 遍历所有检测结果,筛选高置信度目标
        for (i in 0 until OUTPUT_ROWS) {
            val confidence = outputArray[0][4][i] // 假设第5列是置信度(索引从0开始)
            if (confidence >= CONFIDENCE_THRESHOLD) {
                // 提取11个关键点坐标(假设从第5列后,每2列对应一个关键点的x/y)
                for (kpIndex in 0 until NUM_KEYPOINTS) {
                    val x = outputArray[0][5 + kpIndex * 2][i] * originalBitmap.width
                    val y = outputArray[0][6 + kpIndex * 2][i] * originalBitmap.height
                    keyPoints.add(Pair(x, y))
                }
                // 若支持多目标检测,可移除break遍历所有符合条件的结果
                break
            }
        }
        return keyPoints
    }

    // 在原图像上绘制关键点
    fun drawKeyPoints(originalBitmap: Bitmap, keyPoints: List<Pair<Float, Float>>): Bitmap {
        val resultBitmap = originalBitmap.copy(Bitmap.Config.ARGB_8888, true)
        val canvas = Canvas(resultBitmap)
        val paint = Paint().apply {
            color = Color.RED
            strokeWidth = 8f
            style = Paint.Style.FILL
        }

        keyPoints.forEach { (x, y) ->
            canvas.drawCircle(x, y, 10f, paint)
        }
        return resultBitmap
    }

    fun close() {
        interpreter.close()
    }
}

3. 使用示例

在Activity/Fragment中调用:

// 模型文件需先从assets复制到本地路径(示例路径)
val modelPath = filesDir.path + "/your_model.tflite"
val detector = KeyPointDetector(modelPath)

// 传入待检测的Bitmap(从相册/相机获取)
val originalBitmap = ... 
val keyPoints = detector.detectKeyPoints(originalBitmap)
val resultBitmap = detector.drawKeyPoints(originalBitmap, keyPoints)

// 显示结果到ImageView
imageView.setImageBitmap(resultBitmap)

// 使用完毕后释放资源
detector.close()

注意事项

  • 预处理匹配:若模型训练时用BGR通道或其他归一化规则,需调整preprocessBitmap中的通道顺序和系数。
  • 输出索引调整:根据Netron查看的输出结构,确认置信度、关键点的列索引,修改代码中对应的数组下标。
  • 多目标支持:若模型输出多个检测框,移除代码中的break语句即可遍历所有符合置信度要求的目标。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 03:24:52