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

