First commit
This commit is contained in:
1 parent
4f70c7f8d5
commit
dc8dcca75e
17 files changed
+4382
-2
No files matched your search
@@ -0,0 +1,195 @@
|
||||
package bytetrack
|
||||
|
||||
import (
|
||||
"github.com/anomalyco/bytetrack-counter-go/internal/iou"
|
||||
"github.com/anomalyco/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
|
||||
}
|
||||
Reference in new issue
Block a user