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 }