196 lines
4.6 KiB
Go
196 lines
4.6 KiB
Go
package bytetrack
|
|
|
|
import (
|
|
"git.proit.id/dsutanto/bytetrack-counter-go/internal/iou"
|
|
"git.proit.id/dsutanto/bytetrack-counter-go/pkg/kalman"
|
|
)
|
|
|
|
type ByteTracker struct {
|
|
HighThresh float64
|
|
LowThresh float64
|
|
MatchThresh float64
|
|
Buffer int
|
|
MinHits int
|
|
|
|
tracked []*kalman.Tracker
|
|
lost []*kalman.Tracker
|
|
removed []*kalman.Tracker
|
|
frameID int
|
|
}
|
|
|
|
func New(highThresh, lowThresh, matchThresh float64, buffer, minHits int) *ByteTracker {
|
|
return &ByteTracker{
|
|
HighThresh: highThresh,
|
|
LowThresh: lowThresh,
|
|
MatchThresh: matchThresh,
|
|
Buffer: buffer,
|
|
MinHits: minHits,
|
|
}
|
|
}
|
|
|
|
type UpdateResult struct {
|
|
Tracked map[int]float64
|
|
DetToTrack map[int]int
|
|
Lost map[int]float64
|
|
}
|
|
|
|
func (bt *ByteTracker) Update(boxes [][4]float64, scores []float64) UpdateResult {
|
|
bt.frameID++
|
|
|
|
result := UpdateResult{
|
|
Tracked: make(map[int]float64),
|
|
DetToTrack: make(map[int]int),
|
|
Lost: make(map[int]float64),
|
|
}
|
|
|
|
n := len(boxes)
|
|
|
|
var dets [][4]float64
|
|
var filtToGlobal []int
|
|
var highFiltIdx []int
|
|
var lowFiltIdx []int
|
|
|
|
for i := 0; i < n; i++ {
|
|
if scores[i] > bt.LowThresh {
|
|
filtToGlobal = append(filtToGlobal, i)
|
|
dets = append(dets, boxes[i])
|
|
if scores[i] > bt.HighThresh {
|
|
highFiltIdx = append(highFiltIdx, len(dets)-1)
|
|
} else {
|
|
lowFiltIdx = append(lowFiltIdx, len(dets)-1)
|
|
}
|
|
}
|
|
}
|
|
|
|
pool := make([]*kalman.Tracker, 0, len(bt.tracked)+len(bt.lost))
|
|
pool = append(pool, bt.tracked...)
|
|
pool = append(pool, bt.lost...)
|
|
|
|
numPool := len(pool)
|
|
var poolBoxes [][4]float64
|
|
if numPool > 0 {
|
|
poolBoxes = make([][4]float64, numPool)
|
|
for i, trk := range pool {
|
|
trk.Predict()
|
|
poolBoxes[i] = trk.GetState()
|
|
}
|
|
}
|
|
|
|
matchedPool := make(map[int]bool)
|
|
|
|
if numPool > 0 && len(highFiltIdx) > 0 {
|
|
highDets := make([][4]float64, len(highFiltIdx))
|
|
for i, hi := range highFiltIdx {
|
|
highDets[i] = dets[hi]
|
|
}
|
|
iouMat := iou.PairwiseIoU(highDets, poolBoxes)
|
|
cost := make([][]float64, len(iouMat))
|
|
for i := range cost {
|
|
cost[i] = make([]float64, len(iouMat[i]))
|
|
for j := range cost[i] {
|
|
cost[i][j] = 1.0 - iouMat[i][j]
|
|
}
|
|
}
|
|
matches := iou.GreedyMatchRowCol(cost, 1.0-bt.MatchThresh)
|
|
for _, m := range matches {
|
|
filtIdx := highFiltIdx[m.Row]
|
|
globalIdx := filtToGlobal[filtIdx]
|
|
ti := m.Col
|
|
pool[ti].Update(dets[filtIdx])
|
|
if pool[ti].HitStreak < 1 {
|
|
pool[ti].HitStreak = 1
|
|
}
|
|
matchedPool[ti] = true
|
|
result.DetToTrack[globalIdx] = pool[ti].ID
|
|
result.Tracked[pool[ti].ID] = pool[ti].GetCX()
|
|
}
|
|
}
|
|
|
|
var unmatchedPool []int
|
|
for i := 0; i < numPool; i++ {
|
|
if !matchedPool[i] {
|
|
unmatchedPool = append(unmatchedPool, i)
|
|
}
|
|
}
|
|
|
|
if len(lowFiltIdx) > 0 && len(unmatchedPool) > 0 {
|
|
lowDets := make([][4]float64, len(lowFiltIdx))
|
|
for i, li := range lowFiltIdx {
|
|
lowDets[i] = dets[li]
|
|
}
|
|
unmatchedBoxes := make([][4]float64, len(unmatchedPool))
|
|
for i, pi := range unmatchedPool {
|
|
unmatchedBoxes[i] = poolBoxes[pi]
|
|
}
|
|
iouMat2 := iou.PairwiseIoU(lowDets, unmatchedBoxes)
|
|
cost2 := make([][]float64, len(iouMat2))
|
|
for i := range cost2 {
|
|
cost2[i] = make([]float64, len(iouMat2[i]))
|
|
for j := range cost2[i] {
|
|
cost2[i][j] = 1.0 - iouMat2[i][j]
|
|
}
|
|
}
|
|
matches2 := iou.GreedyMatchRowCol(cost2, 0.5)
|
|
for _, m := range matches2 {
|
|
filtIdx := lowFiltIdx[m.Row]
|
|
globalIdx := filtToGlobal[filtIdx]
|
|
ti := unmatchedPool[m.Col]
|
|
pool[ti].Update(dets[filtIdx])
|
|
if pool[ti].HitStreak < 1 {
|
|
pool[ti].HitStreak = 1
|
|
}
|
|
matchedPool[ti] = true
|
|
result.DetToTrack[globalIdx] = pool[ti].ID
|
|
result.Tracked[pool[ti].ID] = pool[ti].GetCX()
|
|
}
|
|
}
|
|
|
|
for i, trk := range pool {
|
|
if !matchedPool[i] {
|
|
trk.HitStreak = 0
|
|
}
|
|
}
|
|
|
|
newTracked := make([]*kalman.Tracker, 0)
|
|
newLost := make([]*kalman.Tracker, 0)
|
|
for _, trk := range pool {
|
|
if trk.TimeSinceUpd > bt.Buffer {
|
|
bt.removed = append(bt.removed, trk)
|
|
} else if trk.TimeSinceUpd > 0 {
|
|
newLost = append(newLost, trk)
|
|
} else {
|
|
newTracked = append(newTracked, trk)
|
|
}
|
|
}
|
|
bt.tracked = newTracked
|
|
bt.lost = newLost
|
|
|
|
for _, trk := range bt.tracked {
|
|
if trk.HitStreak >= bt.MinHits || trk.Hits >= bt.MinHits {
|
|
result.Tracked[trk.ID] = trk.GetCX()
|
|
}
|
|
}
|
|
for _, trk := range bt.lost {
|
|
if trk.HitStreak >= bt.MinHits || trk.Hits >= bt.MinHits {
|
|
result.Tracked[trk.ID] = trk.GetCX()
|
|
result.Lost[trk.ID] = trk.GetCX()
|
|
}
|
|
}
|
|
|
|
matchedDets := make(map[int]bool)
|
|
for idx := range result.DetToTrack {
|
|
matchedDets[idx] = true
|
|
}
|
|
for _, fi := range highFiltIdx {
|
|
globalIdx := filtToGlobal[fi]
|
|
if !matchedDets[globalIdx] {
|
|
trk := kalman.NewTracker(dets[fi])
|
|
bt.tracked = append(bt.tracked, trk)
|
|
result.DetToTrack[globalIdx] = trk.ID
|
|
result.Tracked[trk.ID] = trk.GetCX()
|
|
}
|
|
}
|
|
|
|
return result
|
|
}
|