126 lines
3.8 KiB
Go
126 lines
3.8 KiB
Go
package pose
|
|||
|
|
|
||
|
|
import (
|
||
|
|
"fmt"
|
||
|
|
"sort"
|
||
|
|
)
|
||
|
|
|
||
|
|
// ParseYOLOv8Pose turns the fixed [1,56,8400] YOLO Pose output into source
|
||
|
|
// pixel coordinates, applying confidence filtering and class-agnostic NMS.
|
||
|
|
func ParseYOLOv8Pose(output []float32, transform LetterboxTransform, confidenceThreshold, iouThreshold float32) ([]PersonPose, error) {
|
||
|
|
if len(output) != YOLOPoseOutputValues {
|
||
|
|
return nil, fmt.Errorf("YOLO Pose output length = %d, want %d", len(output), YOLOPoseOutputValues)
|
||
|
|
}
|
||
|
|
if transform.Scale <= 0 {
|
||
|
|
return nil, fmt.Errorf("letterbox scale must be positive")
|
||
|
|
}
|
||
|
|
if confidenceThreshold < 0 || confidenceThreshold > 1 || iouThreshold < 0 || iouThreshold > 1 {
|
||
|
|
return nil, fmt.Errorf("confidence and IoU thresholds must be between 0 and 1")
|
||
|
|
}
|
||
|
|
|
||
|
|
candidates := make([]PersonPose, 0)
|
||
|
|
for candidate := 0; candidate < YOLOPoseCandidateCount; candidate++ {
|
||
|
|
confidence := valueAt(output, 4, candidate)
|
||
|
|
if confidence < confidenceThreshold {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
centerX := valueAt(output, 0, candidate)
|
||
|
|
centerY := valueAt(output, 1, candidate)
|
||
|
|
width := valueAt(output, 2, candidate)
|
||
|
|
height := valueAt(output, 3, candidate)
|
||
|
|
person := PersonPose{
|
||
|
|
Box: Box{
|
||
|
|
Left: transform.restoreX(centerX - width/2),
|
||
|
|
Top: transform.restoreY(centerY - height/2),
|
||
|
|
Right: transform.restoreX(centerX + width/2),
|
||
|
|
Bottom: transform.restoreY(centerY + height/2),
|
||
|
|
},
|
||
|
|
Confidence: confidence,
|
||
|
|
}
|
||
|
|
for keypoint := 0; keypoint < YOLOPoseKeypointCount; keypoint++ {
|
||
|
|
channel := 5 + keypoint*3
|
||
|
|
person.Keypoints[keypoint] = Keypoint{
|
||
|
|
X: transform.restoreX(valueAt(output, channel, candidate)),
|
||
|
|
Y: transform.restoreY(valueAt(output, channel+1, candidate)),
|
||
|
|
Confidence: valueAt(output, channel+2, candidate),
|
||
|
|
}
|
||
|
|
}
|
||
|
|
candidates = append(candidates, person)
|
||
|
|
}
|
||
|
|
|
||
|
|
sort.SliceStable(candidates, func(left, right int) bool {
|
||
|
|
return candidates[left].Confidence > candidates[right].Confidence
|
||
|
|
})
|
||
|
|
selected := make([]PersonPose, 0, len(candidates))
|
||
|
|
for _, candidate := range candidates {
|
||
|
|
overlaps := false
|
||
|
|
for _, accepted := range selected {
|
||
|
|
if intersectionOverUnion(candidate.Box, accepted.Box) > iouThreshold {
|
||
|
|
overlaps = true
|
||
|
|
break
|
||
|
|
}
|
||
|
|
}
|
||
|
|
if !overlaps {
|
||
|
|
selected = append(selected, candidate)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return selected, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func valueAt(output []float32, channel, candidate int) float32 {
|
||
|
|
return output[channel*YOLOPoseCandidateCount+candidate]
|
||
|
|
}
|
||
|
|
|
||
|
|
func (transform LetterboxTransform) restoreX(value float32) float32 {
|
||
|
|
return clampCoordinate((value-float32(transform.PadLeft))/transform.Scale, transform.OriginalWidth)
|
||
|
|
}
|
||
|
|
|
||
|
|
func (transform LetterboxTransform) restoreY(value float32) float32 {
|
||
|
|
return clampCoordinate((value-float32(transform.PadTop))/transform.Scale, transform.OriginalHeight)
|
||
|
|
}
|
||
|
|
|
||
|
|
func clampCoordinate(value float32, size int) float32 {
|
||
|
|
if size <= 0 {
|
||
|
|
return value
|
||
|
|
}
|
||
|
|
if value < 0 {
|
||
|
|
return 0
|
||
|
|
}
|
||
|
|
maximum := float32(size)
|
||
|
|
if value > maximum {
|
||
|
|
return maximum
|
||
|
|
}
|
||
|
|
return value
|
||
|
|
}
|
||
|
|
|
||
|
|
func intersectionOverUnion(left, right Box) float32 {
|
||
|
|
intersectionLeft := maxFloat(left.Left, right.Left)
|
||
|
|
intersectionTop := maxFloat(left.Top, right.Top)
|
||
|
|
intersectionRight := minFloat(left.Right, right.Right)
|
||
|
|
intersectionBottom := minFloat(left.Bottom, right.Bottom)
|
||
|
|
intersectionWidth := maxFloat(0, intersectionRight-intersectionLeft)
|
||
|
|
intersectionHeight := maxFloat(0, intersectionBottom-intersectionTop)
|
||
|
|
intersection := intersectionWidth * intersectionHeight
|
||
|
|
leftArea := maxFloat(0, left.Right-left.Left) * maxFloat(0, left.Bottom-left.Top)
|
||
|
|
rightArea := maxFloat(0, right.Right-right.Left) * maxFloat(0, right.Bottom-right.Top)
|
||
|
|
union := leftArea + rightArea - intersection
|
||
|
|
if union <= 0 {
|
||
|
|
return 0
|
||
|
|
}
|
||
|
|
return intersection / union
|
||
|
|
}
|
||
|
|
|
||
|
|
func minFloat(left, right float32) float32 {
|
||
|
|
if left < right {
|
||
|
|
return left
|
||
|
|
}
|
||
|
|
return right
|
||
|
|
}
|
||
|
|
|
||
|
|
func maxFloat(left, right float32) float32 {
|
||
|
|
if left > right {
|
||
|
|
return left
|
||
|
|
}
|
||
|
|
return right
|
||
|
|
}
|