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

Android Kotlin中Yolo11n-Pose TFLite模型输出处理及关键点提取

YOLO11n-Pose TFLite输出解析为Person对象(Kotlin)

先明确输出格式

YOLO11 Pose的TFLite输出形状为[1, 56, 8400],各维度含义:

  • 1:固定批量大小(单张图像输入)
  • 56:每个候选框的参数集:
    • 前4个值:检测框的中心点x,y、宽高w,h(均为相对于输入图像的比例值)
    • 第5个值:该框为人物的置信度得分
    • 后续51个值:17个关键点的x,y,score(每个关键点占3个值,共17×3=51)
  • 8400:模型生成的候选锚框总数

完整解析实现

以下是填充后的processOutput函数,包含低置信度过滤、关键点解析、非极大值抑制(NMS,去除重复检测)逻辑:

// 可根据实际检测效果调整阈值
private const val PERSON_CONFIDENCE_THRESHOLD = 0.5f
private const val KEYPOINT_SCORE_THRESHOLD = 0.3f
private const val IOU_THRESHOLD = 0.5f

data class Person(
    var keyPoints: MutableList<KeyPoint>,
    val score: Float,
    // 可选:存储检测框参数,用于NMS计算
    val boxX: Float,
    val boxY: Float,
    val boxW: Float,
    val boxH: Float
)

enum class BodyPart {
    NOSE, LEFT_EYE, RIGHT_EYE, LEFT_EAR, RIGHT_EAR,
    LEFT_SHOULDER, RIGHT_SHOULDER, LEFT_ELBOW, RIGHT_ELBOW,
    LEFT_WRIST, RIGHT_WRIST, LEFT_HIP, RIGHT_HIP,
    LEFT_KNEE, RIGHT_KNEE, LEFT_ANKLE, RIGHT_ANKLE
}

data class KeyPoint(val bodyPart: BodyPart, var coordinate: PointF, val score: Float)

private fun processOutput(output: FloatArray, inputWidth: Int, inputHeight: Int): List<Person> {
    val persons = mutableListOf<Person>()
    val numAnchors = 8400
    val paramsPerAnchor = 56

    // 遍历所有候选锚框
    for (anchorIndex in 0 until numAnchors) {
        val baseIdx = anchorIndex * paramsPerAnchor

        // 1. 过滤低置信度的人物框
        val personScore = output[baseIdx + 4]
        if (personScore < PERSON_CONFIDENCE_THRESHOLD) continue

        // 2. 解析检测框参数(用于NMS)
        val boxX = output[baseIdx] * inputWidth
        val boxY = output[baseIdx + 1] * inputHeight
        val boxW = output[baseIdx + 2] * inputWidth
        val boxH = output[baseIdx + 3] * inputHeight

        // 3. 解析17个关键点
        val keyPoints = mutableListOf<KeyPoint>()
        for (kpIdx in 0 until BodyPart.values().size) {
            val kpBaseIdx = baseIdx + 5 + kpIdx * 3
            // 比例转像素坐标
            val kpX = output[kpBaseIdx] * inputWidth
            val kpY = output[kpBaseIdx + 1] * inputHeight
            val kpScore = output[kpBaseIdx + 2]

            if (kpScore >= KEYPOINT_SCORE_THRESHOLD) {
                val bodyPart = BodyPart.values()[kpIdx]
                keyPoints.add(KeyPoint(bodyPart, PointF(kpX, kpY), kpScore))
            }
        }

        // 仅保留有效关键点的人物
        if (keyPoints.isNotEmpty()) {
            persons.add(Person(keyPoints, personScore, boxX, boxY, boxW, boxH))
        }
    }

    // 4. 执行非极大值抑制,去除重复检测
    return applyNonMaxSuppression(persons)
}

// 非极大值抑制:保留得分最高、重叠度低的检测结果
private fun applyNonMaxSuppression(persons: List<Person>): List<Person> {
    // 按人物得分降序排序
    val sortedPersons = persons.sortedByDescending { it.score }
    val keptPersons = mutableListOf<Person>()

    for (current in sortedPersons) {
        var keep = true
        for (kept in keptPersons) {
            if (calculateIOU(current, kept) > IOU_THRESHOLD) {
                keep = false
                break
            }
        }
        if (keep) keptPersons.add(current)
    }
    return keptPersons
}

// 计算两个检测框的交并比(IOU)
private fun calculateIOU(p1: Person, p2: Person): Float {
    // 转换为框的左上角、右下角坐标
    val p1Left = p1.boxX - p1.boxW / 2
    val p1Top = p1.boxY - p1.boxH / 2
    val p1Right = p1.boxX + p1.boxW / 2
    val p1Bottom = p1.boxY + p1.boxH / 2

    val p2Left = p2.boxX - p2.boxW / 2
    val p2Top = p2.boxY - p2.boxH / 2
    val p2Right = p2.boxX + p2.boxW / 2
    val p2Bottom = p2.boxY + p2.boxH / 2

    // 计算交集区域
    val intersectLeft = max(p1Left, p2Left)
    val intersectTop = max(p1Top, p2Top)
    val intersectRight = min(p1Right, p2Right)
    val intersectBottom = min(p1Bottom, p2Bottom)

    val intersectWidth = max(0f, intersectRight - intersectLeft)
    val intersectHeight = max(0f, intersectBottom - intersectTop)
    val intersectArea = intersectWidth * intersectHeight

    // 计算两个框的总面积
    val p1Area = p1.boxW * p1.boxH
    val p2Area = p2.boxW * p2.boxH

    // 计算IOU
    return intersectArea / (p1Area + p2Area - intersectArea)
}

关键注意事项

  1. 阈值调整:PERSON_CONFIDENCE_THRESHOLD(人物置信度)和KEYPOINT_SCORE_THRESHOLD(关键点得分)可根据检测效果灵活调整,阈值越低检测数量越多,但误检概率也会上升。
  2. 坐标转换:模型输出的坐标是输入图像的比例值(0~1),必须乘以模型接收的输入图像宽高,才能得到实际像素坐标。
  3. 关键点顺序:BodyPart枚举的顺序必须与YOLO11 Pose的关键点定义完全一致,当前枚举顺序是正确的。
  4. NMS必要性:YOLO会生成大量重叠候选框,NMS是去除重复检测、保证结果准确性的必要步骤。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 12:34:55