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)
- 前4个值:检测框的中心点
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) }
关键注意事项
- 阈值调整:
PERSON_CONFIDENCE_THRESHOLD(人物置信度)和KEYPOINT_SCORE_THRESHOLD(关键点得分)可根据检测效果灵活调整,阈值越低检测数量越多,但误检概率也会上升。 - 坐标转换:模型输出的坐标是输入图像的比例值(0~1),必须乘以模型接收的输入图像宽高,才能得到实际像素坐标。
- 关键点顺序:
BodyPart枚举的顺序必须与YOLO11 Pose的关键点定义完全一致,当前枚举顺序是正确的。 - NMS必要性:YOLO会生成大量重叠候选框,NMS是去除重复检测、保证结果准确性的必要步骤。
内容的提问来源于stack exchange,提问作者samedhrmn
相关产品推荐
相关产品推荐

