First commit

This commit is contained in:
proitlab committed 2026-06-30 18:49:39 +07:00
1 parent 4f70c7f8d5
commit dc8dcca75e
17 files changed
+4382 -2

No files matched your search

+195
View File
@@ -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
}