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,467 @@
|
||||
package batchstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
type LogFunc func(string)
|
||||
|
||||
type Store struct {
|
||||
dbPath string
|
||||
stateFile string
|
||||
cameraName string
|
||||
objectLabel string
|
||||
cutoffTime string
|
||||
|
||||
batchTimeout time.Duration
|
||||
ignoreBatchLabelTimeout time.Duration
|
||||
minObjectPerBatch int
|
||||
minDurationPerBatch int
|
||||
carryIDs int
|
||||
|
||||
log LogFunc
|
||||
|
||||
db *sql.DB
|
||||
mu sync.Mutex
|
||||
currentState *batchState
|
||||
previousState *batchState
|
||||
|
||||
batchTimer *time.Timer
|
||||
ignoreBatchLabel bool
|
||||
ignoreBatchLabelTimer *time.Timer
|
||||
|
||||
shutdownCh chan struct{}
|
||||
stopped bool
|
||||
}
|
||||
|
||||
type batchState struct {
|
||||
CountingDate string `json:"counting_date"`
|
||||
BatchNumber int `json:"batch_number"`
|
||||
Count int `json:"count"`
|
||||
StartTime string `json:"start_time"`
|
||||
LastDetection string `json:"last_detection_time"`
|
||||
CountedEventIDs []string `json:"counted_event_ids"`
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
DBPath string
|
||||
StateFile string
|
||||
CameraName string
|
||||
ObjectLabel string
|
||||
CutoffTime string
|
||||
BatchTimeoutSec float64
|
||||
IgnoreBatchLabelTimeout float64
|
||||
MinObjectPerBatch int
|
||||
MinDurationPerBatch int
|
||||
}
|
||||
|
||||
func New(cfg Config, logFn LogFunc) (*Store, error) {
|
||||
if logFn == nil {
|
||||
logFn = func(s string) { fmt.Println(s) }
|
||||
}
|
||||
|
||||
s := &Store{
|
||||
dbPath: cfg.DBPath,
|
||||
stateFile: cfg.StateFile,
|
||||
cameraName: cfg.CameraName,
|
||||
objectLabel: cfg.ObjectLabel,
|
||||
cutoffTime: cfg.CutoffTime,
|
||||
batchTimeout: time.Duration(cfg.BatchTimeoutSec * float64(time.Second)),
|
||||
ignoreBatchLabelTimeout: time.Duration(cfg.IgnoreBatchLabelTimeout * float64(time.Second)),
|
||||
minObjectPerBatch: cfg.MinObjectPerBatch,
|
||||
minDurationPerBatch: cfg.MinDurationPerBatch,
|
||||
carryIDs: 50,
|
||||
log: logFn,
|
||||
shutdownCh: make(chan struct{}),
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(cfg.DBPath), 0755); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(cfg.StateFile), 0755); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var err error
|
||||
s.db, err = sql.Open("sqlite", cfg.DBPath+"?_journal_mode=WAL&_synchronous=NORMAL")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open db: %w", err)
|
||||
}
|
||||
|
||||
if err := s.initDB(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
s.currentState = s.loadState()
|
||||
s.previousState = s.currentState
|
||||
|
||||
return s, nil
|
||||
}
|
||||
|
||||
func (s *Store) initDB() error {
|
||||
_, err := s.db.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS batches (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
counting_date TEXT NOT NULL,
|
||||
batch_number INTEGER NOT NULL,
|
||||
camera_name TEXT NOT NULL,
|
||||
object_label TEXT NOT NULL,
|
||||
count INTEGER NOT NULL,
|
||||
start_time TEXT NOT NULL,
|
||||
end_time TEXT NOT NULL,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
UNIQUE(counting_date, batch_number, camera_name, object_label)
|
||||
)
|
||||
`)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = s.db.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS daily_summaries (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
counting_date TEXT NOT NULL,
|
||||
camera_name TEXT NOT NULL,
|
||||
object_label TEXT NOT NULL,
|
||||
total_count INTEGER NOT NULL DEFAULT 0,
|
||||
total_batches INTEGER NOT NULL DEFAULT 0,
|
||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
UNIQUE(counting_date, camera_name, object_label)
|
||||
)
|
||||
`)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *Store) countingDate(t time.Time) string {
|
||||
cutoff, _ := time.Parse("15:04", s.cutoffTime)
|
||||
ct := time.Date(t.Year(), t.Month(), t.Day(), cutoff.Hour(), cutoff.Minute(), 0, 0, t.Location())
|
||||
if t.Before(ct) {
|
||||
return t.Format("2006-01-02")
|
||||
}
|
||||
return t.Add(24 * time.Hour).Format("2006-01-02")
|
||||
}
|
||||
|
||||
func (s *Store) loadState() *batchState {
|
||||
data, err := os.ReadFile(s.stateFile)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
var st batchState
|
||||
if err := json.Unmarshal(data, &st); err != nil {
|
||||
s.log(fmt.Sprintf("Failed to load state file: %v", err))
|
||||
return nil
|
||||
}
|
||||
|
||||
currentDate := s.countingDate(time.Now())
|
||||
if st.CountingDate != currentDate {
|
||||
s.log(fmt.Sprintf("State file belongs to previous counting day (%s). Finalizing before fresh start.", st.CountingDate))
|
||||
s.insertBatch(st.CountingDate, st.BatchNumber, st.Count, st.StartTime, time.Now().Format(time.RFC3339))
|
||||
os.Remove(s.stateFile)
|
||||
return nil
|
||||
}
|
||||
|
||||
s.log(fmt.Sprintf("Resumed batch #%d from %s with count=%d", st.BatchNumber, st.StartTime, st.Count))
|
||||
s.resetBatchTimer()
|
||||
return &st
|
||||
}
|
||||
|
||||
func (s *Store) saveState() {
|
||||
if s.currentState == nil {
|
||||
os.Remove(s.stateFile)
|
||||
return
|
||||
}
|
||||
data, err := json.MarshalIndent(s.currentState, "", " ")
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if err := os.WriteFile(s.stateFile, data, 0644); err != nil {
|
||||
s.log(fmt.Sprintf("Failed to save state: %v", err))
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Store) nextBatchNumber(countingDate string) int {
|
||||
var maxNum sql.NullInt64
|
||||
err := s.db.QueryRow(
|
||||
`SELECT MAX(batch_number) FROM batches WHERE counting_date = ? AND camera_name = ? AND object_label = ?`,
|
||||
countingDate, s.cameraName, s.objectLabel,
|
||||
).Scan(&maxNum)
|
||||
if err != nil || !maxNum.Valid {
|
||||
return 1
|
||||
}
|
||||
return int(maxNum.Int64) + 1
|
||||
}
|
||||
|
||||
func (s *Store) startNewBatch(countingDate string) {
|
||||
batchNum := s.nextBatchNumber(countingDate)
|
||||
now := time.Now().Format(time.RFC3339)
|
||||
|
||||
var countedIDs []string
|
||||
if s.previousState != nil && len(s.previousState.CountedEventIDs) > 0 {
|
||||
ids := s.previousState.CountedEventIDs
|
||||
start := 0
|
||||
if len(ids) > s.carryIDs {
|
||||
start = len(ids) - s.carryIDs
|
||||
}
|
||||
countedIDs = ids[start:]
|
||||
}
|
||||
|
||||
s.currentState = &batchState{
|
||||
CountingDate: countingDate,
|
||||
BatchNumber: batchNum,
|
||||
Count: 0,
|
||||
StartTime: now,
|
||||
LastDetection: now,
|
||||
CountedEventIDs: countedIDs,
|
||||
}
|
||||
s.saveState()
|
||||
s.log(fmt.Sprintf("Started batch #%d for %s (%s)", batchNum, countingDate, s.objectLabel))
|
||||
}
|
||||
|
||||
func (s *Store) resetBatchTimer() {
|
||||
if s.batchTimer != nil {
|
||||
s.batchTimer.Stop()
|
||||
}
|
||||
s.batchTimer = time.AfterFunc(s.batchTimeout, func() {
|
||||
s.log(fmt.Sprintf("Batch inactivity timeout (%.0fs) reached", s.batchTimeout.Seconds()))
|
||||
s.endBatch("timeout")
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Store) startIgnoreBatchLabel() {
|
||||
if s.ignoreBatchLabelTimer != nil {
|
||||
return
|
||||
}
|
||||
s.ignoreBatchLabel = true
|
||||
s.ignoreBatchLabelTimer = time.AfterFunc(s.ignoreBatchLabelTimeout, func() {
|
||||
s.ignoreBatchLabel = false
|
||||
s.ignoreBatchLabelTimer = nil
|
||||
s.log("Ignore batch label cooldown finished")
|
||||
})
|
||||
s.log(fmt.Sprintf("Ignore batch label for %.0fs", s.ignoreBatchLabelTimeout.Seconds()))
|
||||
}
|
||||
|
||||
func (s *Store) RecordAyamCrossing(trackID int) (int, bool) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
countingDate := s.countingDate(time.Now())
|
||||
startedNew := false
|
||||
|
||||
if s.currentState == nil {
|
||||
s.startNewBatch(countingDate)
|
||||
startedNew = true
|
||||
} else if s.currentState.CountingDate != countingDate {
|
||||
s.endBatchLocked("cutoff")
|
||||
s.startNewBatch(countingDate)
|
||||
startedNew = true
|
||||
}
|
||||
|
||||
eventKey := fmt.Sprintf("%d", trackID)
|
||||
found := false
|
||||
for _, id := range s.currentState.CountedEventIDs {
|
||||
if id == eventKey {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
s.currentState.Count++
|
||||
s.currentState.CountedEventIDs = append(s.currentState.CountedEventIDs, eventKey)
|
||||
s.log(fmt.Sprintf("Counted ayam (track %d) | batch #%d total: %d", trackID, s.currentState.BatchNumber, s.currentState.Count))
|
||||
}
|
||||
|
||||
s.currentState.LastDetection = time.Now().Format(time.RFC3339)
|
||||
s.saveState()
|
||||
s.resetBatchTimer()
|
||||
return s.currentState.Count, startedNew
|
||||
}
|
||||
|
||||
func (s *Store) RecordTalenanCrossing(trackID int) bool {
|
||||
if s.ignoreBatchLabel {
|
||||
return false
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
s.startIgnoreBatchLabel()
|
||||
s.endBatchLocked("talenan")
|
||||
s.log(fmt.Sprintf("Batch closed by talenan (track %d)", trackID))
|
||||
|
||||
if s.batchTimer != nil {
|
||||
s.batchTimer.Stop()
|
||||
s.batchTimer = nil
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *Store) endBatch(closedBy string) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.endBatchLocked(closedBy)
|
||||
}
|
||||
|
||||
func (s *Store) endBatchLocked(closedBy string) bool {
|
||||
if s.currentState == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
s.previousState = s.currentState
|
||||
state := s.currentState
|
||||
|
||||
startTime, _ := time.Parse(time.RFC3339, state.StartTime)
|
||||
endTime := time.Now()
|
||||
durationSec := endTime.Sub(startTime).Seconds()
|
||||
|
||||
if state.Count < s.minObjectPerBatch || int(durationSec) < s.minDurationPerBatch {
|
||||
s.currentState = nil
|
||||
s.saveState()
|
||||
if s.batchTimer != nil {
|
||||
s.batchTimer.Stop()
|
||||
s.batchTimer = nil
|
||||
}
|
||||
s.log(fmt.Sprintf("Batch #%d discarded (count=%d, duration=%.0fs)", state.BatchNumber, state.Count, durationSec))
|
||||
return false
|
||||
}
|
||||
|
||||
endTimeStr := endTime.Format(time.RFC3339)
|
||||
if err := s.insertBatch(state.CountingDate, state.BatchNumber, state.Count, state.StartTime, endTimeStr); err != nil {
|
||||
s.log(fmt.Sprintf("Failed to persist batch: %v", err))
|
||||
return false
|
||||
}
|
||||
|
||||
cps := float64(state.Count) / durationSec
|
||||
if durationSec == 0 {
|
||||
cps = 0
|
||||
}
|
||||
s.log(fmt.Sprintf("Batch #%d ended | count=%d | duration=%.0fs | cps=%.3f | closed_by=%s",
|
||||
state.BatchNumber, state.Count, durationSec, cps, closedBy))
|
||||
|
||||
s.currentState = nil
|
||||
s.saveState()
|
||||
if s.batchTimer != nil {
|
||||
s.batchTimer.Stop()
|
||||
s.batchTimer = nil
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *Store) insertBatch(countingDate string, batchNumber, count int, startTime, endTime string) error {
|
||||
tx, err := s.db.Begin()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
|
||||
_, err = tx.Exec(
|
||||
`INSERT INTO batches (counting_date, batch_number, camera_name, object_label, count, start_time, end_time)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)`,
|
||||
countingDate, batchNumber, s.cameraName, s.objectLabel, count, startTime, endTime,
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, err = tx.Exec(
|
||||
`INSERT INTO daily_summaries (counting_date, camera_name, object_label, total_count, total_batches)
|
||||
VALUES (?, ?, ?, ?, 1)
|
||||
ON CONFLICT(counting_date, camera_name, object_label)
|
||||
DO UPDATE SET total_count = total_count + excluded.total_count,
|
||||
total_batches = total_batches + excluded.total_batches,
|
||||
updated_at = CURRENT_TIMESTAMP`,
|
||||
countingDate, s.cameraName, s.objectLabel, count,
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var totalCount, totalBatches int
|
||||
_ = s.db.QueryRow(
|
||||
`SELECT total_count, total_batches FROM daily_summaries
|
||||
WHERE counting_date = ? AND camera_name = ? AND object_label = ?`,
|
||||
countingDate, s.cameraName, s.objectLabel,
|
||||
).Scan(&totalCount, &totalBatches)
|
||||
|
||||
s.log(fmt.Sprintf("Daily totals for %s: %d objects across %d batch(es)", countingDate, totalCount, totalBatches))
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Store) cutoffWatcher() {
|
||||
ticker := time.NewTicker(60 * time.Second)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-s.shutdownCh:
|
||||
return
|
||||
case <-ticker.C:
|
||||
s.mu.Lock()
|
||||
if s.currentState != nil && s.currentState.CountingDate != s.countingDate(time.Now()) {
|
||||
s.log("Daily cutoff reached – finalizing batch")
|
||||
s.endBatchLocked("cutoff")
|
||||
}
|
||||
s.mu.Unlock()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Store) StartCutoffWatcher() {
|
||||
go s.cutoffWatcher()
|
||||
}
|
||||
|
||||
func (s *Store) CurrentBatchNumber() int {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.currentState == nil {
|
||||
return 0
|
||||
}
|
||||
return s.currentState.BatchNumber
|
||||
}
|
||||
|
||||
func (s *Store) CurrentBatchCount() int {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.currentState == nil {
|
||||
return 0
|
||||
}
|
||||
return s.currentState.Count
|
||||
}
|
||||
|
||||
func (s *Store) ClosedTotalForDay() int {
|
||||
countingDate := s.countingDate(time.Now())
|
||||
var total int
|
||||
_ = s.db.QueryRow(
|
||||
`SELECT COALESCE(total_count, 0) FROM daily_summaries
|
||||
WHERE counting_date = ? AND camera_name = ? AND object_label = ?`,
|
||||
countingDate, s.cameraName, s.objectLabel,
|
||||
).Scan(&total)
|
||||
return total
|
||||
}
|
||||
|
||||
func (s *Store) DisplayTotal() int {
|
||||
return s.ClosedTotalForDay() + s.CurrentBatchCount()
|
||||
}
|
||||
|
||||
func (s *Store) Shutdown() {
|
||||
if s.stopped {
|
||||
return
|
||||
}
|
||||
s.stopped = true
|
||||
close(s.shutdownCh)
|
||||
s.endBatch("shutdown")
|
||||
if s.batchTimer != nil {
|
||||
s.batchTimer.Stop()
|
||||
}
|
||||
if s.db != nil {
|
||||
s.db.Close()
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
OutputDir string
|
||||
DBPath string
|
||||
StateFile string
|
||||
Source string
|
||||
ModelPath string
|
||||
CameraName string
|
||||
ObjectLabel string
|
||||
ClassAyam string
|
||||
ClassTalenan string
|
||||
|
||||
LineX *int
|
||||
LineXFrac float64
|
||||
CrossDirection string
|
||||
|
||||
ImgSz int
|
||||
Half bool
|
||||
Conf float64
|
||||
CoreMask int
|
||||
NumClasses int
|
||||
ScoreSigmoid bool
|
||||
|
||||
TrackHighThresh float64
|
||||
TrackLowThresh float64
|
||||
TrackMatchThresh float64
|
||||
TrackBuffer int
|
||||
TrackMinHits int
|
||||
|
||||
DailyCutoffTime string
|
||||
BatchTimeoutSec float64
|
||||
IgnoreBatchLabelTimeoutSec float64
|
||||
MinObjectPerBatch int
|
||||
MinDurationPerBatch int
|
||||
|
||||
ExportCSV bool
|
||||
CrossCSV string
|
||||
RecordVideo bool
|
||||
|
||||
WarmupFrames int
|
||||
ReconnectDelaySec int
|
||||
MaxReconnectAttempts int
|
||||
FlushEveryNFrames int
|
||||
TrackedPruneSec float64
|
||||
VideoSegmentSec int
|
||||
OutputFPS int
|
||||
|
||||
LiveStreamEnabled bool
|
||||
LiveStreamFramePath string
|
||||
LiveStreamQuality int
|
||||
LiveStreamEveryN int
|
||||
RTSPFFmpegOptions string
|
||||
|
||||
IsLive bool
|
||||
}
|
||||
|
||||
func Load() *Config {
|
||||
cfg := &Config{
|
||||
OutputDir: getEnv("OUTPUT_DIR", "/opt/jetson-counter"),
|
||||
DBPath: getEnv("DB_PATH", ""),
|
||||
StateFile: getEnv("STATE_FILE", ""),
|
||||
Source: getEnv("SOURCE", "rtsp://user:pass@192.168.0.100:554/stream1"),
|
||||
ModelPath: getEnv("MODEL_PATH", "/opt/jetson-counter/yolo11n.rknn"),
|
||||
CameraName: getEnv("CAMERA_NAME", "CC1"),
|
||||
ObjectLabel: getEnv("OBJECT_LABEL", "ayam-potong"),
|
||||
ClassAyam: getEnv("CLASS_AYAM", "ayam"),
|
||||
ClassTalenan: getEnv("CLASS_TALENAN", "talenan"),
|
||||
|
||||
LineXFrac: getEnvFloat("LINE_X_FRAC", 0.5),
|
||||
CrossDirection: strings.ToLower(getEnv("CROSS_DIRECTION", "rtl")),
|
||||
|
||||
ImgSz: getEnvInt("IMGSZ", 320),
|
||||
Half: getEnvBool("HALF"),
|
||||
Conf: getEnvFloat("CONF", 0.3),
|
||||
CoreMask: getEnvInt("CORE_MASK", 1),
|
||||
NumClasses: getEnvInt("NUM_CLASSES", 2),
|
||||
ScoreSigmoid: getEnvBool("SCORE_SIGMOID"),
|
||||
|
||||
TrackHighThresh: getEnvFloat("TRACK_HIGH_THRESH", 0.5),
|
||||
TrackLowThresh: getEnvFloat("TRACK_LOW_THRESH", 0.1),
|
||||
TrackMatchThresh: getEnvFloat("TRACK_MATCH_THRESH", 0.8),
|
||||
TrackBuffer: getEnvInt("TRACK_BUFFER", 30),
|
||||
TrackMinHits: getEnvInt("TRACK_MIN_HITS", 3),
|
||||
|
||||
DailyCutoffTime: getEnv("DAILY_CUTOFF_TIME", "20:00"),
|
||||
BatchTimeoutSec: getEnvFloat("BATCH_TIMEOUT_SECONDS", 300),
|
||||
IgnoreBatchLabelTimeoutSec: getEnvFloat("IGNORE_BATCH_LABEL_TIMEOUT_SECONDS", 30),
|
||||
MinObjectPerBatch: getEnvInt("MIN_OBJECT_PER_BATCH", 60),
|
||||
MinDurationPerBatch: getEnvInt("MIN_DURATION_PER_BATCH", 60),
|
||||
|
||||
ExportCSV: getEnvBool("EXPORT_CSV"),
|
||||
CrossCSV: getEnv("CROSS_CSV", ""),
|
||||
RecordVideo: getEnvBool("RECORD_VIDEO"),
|
||||
|
||||
WarmupFrames: getEnvInt("WARMUP_FRAMES", 30),
|
||||
ReconnectDelaySec: getEnvInt("RECONNECT_DELAY_SEC", 3),
|
||||
MaxReconnectAttempts: getEnvInt("MAX_RECONNECT_ATTEMPTS", 0),
|
||||
FlushEveryNFrames: getEnvInt("FLUSH_EVERY_N_FRAMES", 100),
|
||||
TrackedPruneSec: getEnvFloat("TRACKED_PRUNE_SEC", 300),
|
||||
VideoSegmentSec: getEnvInt("VIDEO_SEGMENT_SEC", 3600),
|
||||
OutputFPS: getEnvInt("OUTPUT_FPS", 15),
|
||||
|
||||
LiveStreamEnabled: getEnvBool("LIVE_STREAM_ENABLED"),
|
||||
LiveStreamFramePath: getEnv("LIVE_STREAM_FRAME_PATH", "/dev/shm/jetson-counter/live_frame.jpg"),
|
||||
LiveStreamQuality: getEnvInt("LIVE_STREAM_QUALITY", 75),
|
||||
LiveStreamEveryN: getEnvInt("LIVE_STREAM_EVERY_N", 2),
|
||||
RTSPFFmpegOptions: getEnv("OPENCV_FFMPEG_CAPTURE_OPTIONS", "rtsp_transport;tcp|fflags;nobuffer|flags;low_delay"),
|
||||
}
|
||||
|
||||
if cfg.DBPath == "" {
|
||||
cfg.DBPath = cfg.OutputDir + "/jetson_counter.db"
|
||||
}
|
||||
if cfg.StateFile == "" {
|
||||
cfg.StateFile = cfg.OutputDir + "/current_batch.json"
|
||||
}
|
||||
if cfg.CrossCSV == "" {
|
||||
cfg.CrossCSV = cfg.OutputDir + "/batch_crossings.csv"
|
||||
}
|
||||
|
||||
if lx := os.Getenv("LINE_X"); lx != "" {
|
||||
v, err := strconv.Atoi(lx)
|
||||
if err == nil {
|
||||
cfg.LineX = &v
|
||||
}
|
||||
}
|
||||
|
||||
cfg.IsLive = strings.HasPrefix(strings.ToLower(cfg.Source), "rtsp://") ||
|
||||
strings.HasPrefix(strings.ToLower(cfg.Source), "http://")
|
||||
|
||||
return cfg
|
||||
}
|
||||
|
||||
func getEnv(key, def string) string {
|
||||
if v := os.Getenv(key); v != "" {
|
||||
return v
|
||||
}
|
||||
return def
|
||||
}
|
||||
|
||||
func getEnvInt(key string, def int) int {
|
||||
s := os.Getenv(key)
|
||||
if s == "" {
|
||||
return def
|
||||
}
|
||||
v, err := strconv.Atoi(s)
|
||||
if err != nil {
|
||||
return def
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func getEnvFloat(key string, def float64) float64 {
|
||||
s := os.Getenv(key)
|
||||
if s == "" {
|
||||
return def
|
||||
}
|
||||
v, err := strconv.ParseFloat(s, 64)
|
||||
if err != nil {
|
||||
return def
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func getEnvBool(key string) bool {
|
||||
return strings.ToLower(os.Getenv(key)) == "true"
|
||||
}
|
||||
@@ -0,0 +1,241 @@
|
||||
package drawing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"image"
|
||||
"image/color"
|
||||
"math"
|
||||
"time"
|
||||
|
||||
"gocv.io/x/gocv"
|
||||
)
|
||||
|
||||
const (
|
||||
CrossFlashFrames = 12
|
||||
PopupLifetime = 20
|
||||
LinePulseFrames = 12
|
||||
CountPulseFrames = 15
|
||||
BatchPulseFrames = 20
|
||||
)
|
||||
|
||||
var skeleton = [][2]int{{0, 1}, {4, 3}, {1, 2}, {3, 2}, {2, 6}, {2, 5}, {2, 7}, {7, 8}}
|
||||
var skColors = []color.RGBA{
|
||||
{0, 255, 255, 255},
|
||||
{0, 255, 255, 255},
|
||||
{255, 0, 255, 255},
|
||||
{255, 0, 255, 255},
|
||||
{0, 255, 0, 255},
|
||||
{255, 255, 0, 255},
|
||||
{0, 0, 255, 255},
|
||||
{200, 200, 0, 255},
|
||||
}
|
||||
|
||||
var cPanel = color.RGBA{28, 24, 18, 255}
|
||||
var cBorder = color.RGBA{90, 85, 75, 255}
|
||||
var cAccent = color.RGBA{255, 200, 60, 255}
|
||||
var cGreen = color.RGBA{80, 220, 100, 255}
|
||||
var cText = color.RGBA{220, 220, 220, 255}
|
||||
var cMuted = color.RGBA{150, 150, 150, 255}
|
||||
var cAyamBox = color.RGBA{0, 165, 255, 255}
|
||||
var cTalenanBox = color.RGBA{220, 120, 60, 255}
|
||||
var cLineCore = color.RGBA{180, 220, 255, 255}
|
||||
var cLineGlow = color.RGBA{100, 160, 220, 255}
|
||||
|
||||
type Popup struct {
|
||||
X, Y int
|
||||
Born int
|
||||
Text string
|
||||
}
|
||||
|
||||
func ResolveLineX(w int, lineX *int, lineXFrac float64) int {
|
||||
if lineX != nil {
|
||||
return *lineX
|
||||
}
|
||||
if lineXFrac != 0.5 {
|
||||
return int(float64(w) * lineXFrac)
|
||||
}
|
||||
return w / 2
|
||||
}
|
||||
|
||||
func CrossedLine(prevCX, cx float64, lineX int, direction string) bool {
|
||||
lx := float64(lineX)
|
||||
switch direction {
|
||||
case "ltr":
|
||||
return prevCX < lx && lx <= cx
|
||||
case "both":
|
||||
return (prevCX > lx && lx >= cx) || (prevCX < lx && lx <= cx)
|
||||
default:
|
||||
return prevCX > lx && lx >= cx
|
||||
}
|
||||
}
|
||||
|
||||
func OverlayRect(img *gocv.Mat, x1, y1, x2, y2 int, clr color.RGBA, alpha float64) {
|
||||
x1 = clampI(x1, 0, img.Cols())
|
||||
y1 = clampI(y1, 0, img.Rows())
|
||||
x2 = clampI(x2, 0, img.Cols())
|
||||
y2 = clampI(y2, 0, img.Rows())
|
||||
if x2 <= x1 || y2 <= y1 {
|
||||
return
|
||||
}
|
||||
roi := img.Region(image.Rect(x1, y1, x2, y2))
|
||||
defer roi.Close()
|
||||
overlay := gocv.NewMatWithSize(y2-y1, x2-x1, gocv.MatTypeCV8UC3)
|
||||
defer overlay.Close()
|
||||
overlay.SetTo(gocv.NewScalar(float64(clr.B), float64(clr.G), float64(clr.R), 0))
|
||||
gocv.AddWeighted(overlay, alpha, roi, 1-alpha, 0, &roi)
|
||||
}
|
||||
|
||||
func DrawPill(img *gocv.Mat, text string, x, y int, bg color.RGBA) {
|
||||
fontScale := 0.45
|
||||
size := gocv.GetTextSize(text, gocv.FontHersheySimplex, fontScale, 1)
|
||||
padX, padY := 6, 4
|
||||
|
||||
x1 := x
|
||||
y1 := y - size.Y - padY
|
||||
x2 := x + size.X + padX*2
|
||||
y2 := y + padY - 1
|
||||
|
||||
gocv.Rectangle(img, image.Rect(x1, y1, x2, y2), bg, -1)
|
||||
gocv.Rectangle(img, image.Rect(x1, y1, x2, y2), cBorder, 1)
|
||||
gocv.PutText(img, text, image.Point{X: x + padX, Y: y}, gocv.FontHersheySimplex, fontScale, cText, 1)
|
||||
}
|
||||
|
||||
func DrawCountingLine(img *gocv.Mat, lineX, h, pulse int) {
|
||||
strength := float64(pulse) / float64(max(LinePulseFrames, 1))
|
||||
glowAlpha := 0.12 + 0.18*strength
|
||||
for _, offset := range []int{14, 9, 5} {
|
||||
alpha := uint8(glowAlpha * 255)
|
||||
clr := color.RGBA{cLineGlow.R, cLineGlow.G, cLineGlow.B, alpha}
|
||||
gocv.Line(img, image.Point{X: lineX - offset, Y: 0}, image.Point{X: lineX - offset, Y: h}, clr, 1)
|
||||
gocv.Line(img, image.Point{X: lineX + offset, Y: 0}, image.Point{X: lineX + offset, Y: h}, clr, 1)
|
||||
}
|
||||
|
||||
dashLen, gap := 18, 12
|
||||
for y := 0; y < h; y += dashLen + gap {
|
||||
yEnd := y + dashLen
|
||||
if yEnd > h {
|
||||
yEnd = h
|
||||
}
|
||||
gocv.Line(img, image.Point{X: lineX, Y: y}, image.Point{X: lineX, Y: yEnd}, cLineCore, 2)
|
||||
}
|
||||
|
||||
gocv.PutText(img, "COUNT LINE", image.Point{X: lineX - 46, Y: 24}, gocv.FontHersheySimplex, 0.42, cLineCore, 1)
|
||||
}
|
||||
|
||||
func DrawHeroCount(img *gocv.Mat, lineX, h, count, pulse int) {
|
||||
text := fmt.Sprintf("%d", count)
|
||||
boost := 0.35 * float64(pulse) / float64(max(CountPulseFrames, 1))
|
||||
fontScale := 1.6 + boost
|
||||
thickness := 3
|
||||
|
||||
size := gocv.GetTextSize(text, gocv.FontHersheySimplex, fontScale, thickness)
|
||||
pad := 14
|
||||
tx := lineX - size.X/2
|
||||
ty := h/2 + size.Y/2
|
||||
|
||||
OverlayRect(img, tx-pad, ty-size.Y-pad, tx+size.X+pad, ty+pad/2, cPanel, 0.78)
|
||||
gocv.Rectangle(img, image.Rect(tx-pad, ty-size.Y-pad, tx+size.X+pad, ty+pad/2), cLineCore, 2)
|
||||
gocv.PutText(img, text, image.Point{X: tx, Y: ty}, gocv.FontHersheySimplex, fontScale, cGreen, thickness)
|
||||
}
|
||||
|
||||
func DrawHUD(img *gocv.Mat, w int, batchNum, batchCount, totalAyam int, elapsedSec float64, rate float64, cameraID string) {
|
||||
barH := 52
|
||||
OverlayRect(img, 0, 0, w, barH, cPanel, 0.72)
|
||||
gocv.Line(img, image.Point{X: 0, Y: barH}, image.Point{X: w, Y: barH}, cBorder, 1)
|
||||
|
||||
gocv.PutText(img, "BATCH", image.Point{X: 16, Y: 20}, gocv.FontHersheySimplex, 0.45, cMuted, 1)
|
||||
batchLabel := "--"
|
||||
if batchNum > 0 {
|
||||
batchLabel = fmt.Sprintf("%d", batchNum)
|
||||
}
|
||||
gocv.PutText(img, batchLabel, image.Point{X: 16, Y: 44}, gocv.FontHersheySimplex, 0.9, cAccent, 2)
|
||||
|
||||
gocv.PutText(img, "COUNT", image.Point{X: 100, Y: 20}, gocv.FontHersheySimplex, 0.45, cMuted, 1)
|
||||
gocv.PutText(img, fmt.Sprintf("%d", batchCount), image.Point{X: 100, Y: 44}, gocv.FontHersheySimplex, 0.9, cGreen, 2)
|
||||
|
||||
gocv.PutText(img, "TOTAL", image.Point{X: 190, Y: 20}, gocv.FontHersheySimplex, 0.45, cMuted, 1)
|
||||
gocv.PutText(img, fmt.Sprintf("%d", totalAyam), image.Point{X: 190, Y: 44}, gocv.FontHersheySimplex, 0.7, cText, 1)
|
||||
|
||||
gocv.PutText(img, "UPTIME", image.Point{X: 280, Y: 20}, gocv.FontHersheySimplex, 0.45, cMuted, 1)
|
||||
gocv.PutText(img, fmt.Sprintf("%.1fh", elapsedSec/3600), image.Point{X: 280, Y: 44}, gocv.FontHersheySimplex, 0.7, cText, 1)
|
||||
|
||||
gocv.PutText(img, "RATE", image.Point{X: 380, Y: 20}, gocv.FontHersheySimplex, 0.45, cMuted, 1)
|
||||
gocv.PutText(img, fmt.Sprintf("%.1f/min", rate), image.Point{X: 380, Y: 44}, gocv.FontHersheySimplex, 0.7, cAccent, 1)
|
||||
|
||||
gocv.PutText(img, fmt.Sprintf("CAM %s", cameraID), image.Point{X: w - 180, Y: 20}, gocv.FontHersheySimplex, 0.45, cMuted, 1)
|
||||
}
|
||||
|
||||
func DrawFooter(img *gocv.Mat, w, h, frameIdx int, liveTag string) {
|
||||
barH := 28
|
||||
OverlayRect(img, 0, h-barH, w, h, cPanel, 0.55)
|
||||
gocv.PutText(img, fmt.Sprintf("%s | Frame %d", liveTag, frameIdx), image.Point{X: 12, Y: h - 9}, gocv.FontHersheySimplex, 0.45, cMuted, 1)
|
||||
}
|
||||
|
||||
func DrawSkeleton(img *gocv.Mat, kpts [][2]float64) {
|
||||
for i, pair := range skeleton {
|
||||
a, b := pair[0], pair[1]
|
||||
if a >= len(kpts) || b >= len(kpts) {
|
||||
continue
|
||||
}
|
||||
xa, ya := int(kpts[a][0]), int(kpts[a][1])
|
||||
xb, yb := int(kpts[b][0]), int(kpts[b][1])
|
||||
if xa > 0 && ya > 0 && xb > 0 && yb > 0 {
|
||||
clr := cGreen
|
||||
if i < len(skColors) {
|
||||
clr = skColors[i]
|
||||
}
|
||||
gocv.Line(img, image.Point{X: xa, Y: ya}, image.Point{X: xb, Y: yb}, clr, 3)
|
||||
}
|
||||
}
|
||||
for _, kp := range kpts {
|
||||
x, y := int(kp[0]), int(kp[1])
|
||||
if x > 0 && y > 0 {
|
||||
gocv.Circle(img, image.Point{X: x, Y: y}, 6, color.RGBA{255, 255, 255, 255}, -1)
|
||||
gocv.Circle(img, image.Point{X: x, Y: y}, 6, color.RGBA{40, 40, 40, 255}, 2)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func DrawPopups(img *gocv.Mat, popups []Popup, frameIdx int) []Popup {
|
||||
var alive []Popup
|
||||
for _, pop := range popups {
|
||||
age := frameIdx - pop.Born
|
||||
if age > PopupLifetime {
|
||||
continue
|
||||
}
|
||||
alive = append(alive, pop)
|
||||
fade := 1.0 - float64(age)/PopupLifetime
|
||||
yy := pop.Y - int(float64(age)*1.8)
|
||||
clr := color.RGBA{
|
||||
uint8(float64(cGreen.R) * fade),
|
||||
uint8(float64(cGreen.G) * fade),
|
||||
uint8(float64(cGreen.B) * fade),
|
||||
255,
|
||||
}
|
||||
gocv.PutText(img, pop.Text, image.Point{X: pop.X, Y: yy}, gocv.FontHersheySimplex, 0.7, clr, 2)
|
||||
}
|
||||
return alive
|
||||
}
|
||||
|
||||
func DrawBatchBanner(img *gocv.Mat, w, batchNum, pulse int) {
|
||||
if pulse <= 0 {
|
||||
return
|
||||
}
|
||||
text := fmt.Sprintf("NEW BATCH %d", batchNum)
|
||||
size := gocv.GetTextSize(text, gocv.FontHersheySimplex, 0.8, 2)
|
||||
x1 := w/2 - size.X/2 - 16
|
||||
y1 := 62
|
||||
x2 := w/2 + size.X/2 + 16
|
||||
y2 := 62 + size.Y + 20
|
||||
OverlayRect(img, x1, y1, x2, y2, cPanel, 0.7)
|
||||
gocv.Rectangle(img, image.Rect(x1, y1, x2, y2), cAccent, 2)
|
||||
gocv.PutText(img, text, image.Point{X: w/2 - size.X/2, Y: 62 + size.Y + 4}, gocv.FontHersheySimplex, 0.8, cAccent, 2)
|
||||
}
|
||||
|
||||
func NowStr() string {
|
||||
return time.Now().Format("2006-01-02 15:04:05")
|
||||
}
|
||||
|
||||
func clampI(v, lo, hi int) int {
|
||||
return int(math.Max(float64(lo), math.Min(float64(hi), float64(v))))
|
||||
}
|
||||
@@ -0,0 +1,396 @@
|
||||
package kalman
|
||||
|
||||
import "math"
|
||||
|
||||
const (
|
||||
weightPosition = 1.0 / 20
|
||||
weightVelocity = 1.0 / 160
|
||||
)
|
||||
|
||||
type Filter struct {
|
||||
X [8]float64
|
||||
P [8][8]float64
|
||||
}
|
||||
|
||||
func NewFilter() *Filter {
|
||||
kf := &Filter{}
|
||||
for i := 0; i < 8; i++ {
|
||||
kf.P[i][i] = 10.0
|
||||
}
|
||||
return kf
|
||||
}
|
||||
|
||||
func (kf *Filter) Predict() {
|
||||
motionMat := [8][8]float64{
|
||||
{1, 0, 0, 0, 1, 0, 0, 0},
|
||||
{0, 1, 0, 0, 0, 1, 0, 0},
|
||||
{0, 0, 1, 0, 0, 0, 1, 0},
|
||||
{0, 0, 0, 1, 0, 0, 0, 1},
|
||||
{0, 0, 0, 0, 1, 0, 0, 0},
|
||||
{0, 0, 0, 0, 0, 1, 0, 0},
|
||||
{0, 0, 0, 0, 0, 0, 1, 0},
|
||||
{0, 0, 0, 0, 0, 0, 0, 1},
|
||||
}
|
||||
|
||||
stdPos := [4]float64{
|
||||
weightPosition * kf.X[2],
|
||||
weightPosition * kf.X[3],
|
||||
weightPosition * kf.X[2],
|
||||
weightPosition * kf.X[3],
|
||||
}
|
||||
stdVel := [4]float64{
|
||||
weightVelocity * kf.X[2],
|
||||
weightVelocity * kf.X[3],
|
||||
weightVelocity * kf.X[2],
|
||||
weightVelocity * kf.X[3],
|
||||
}
|
||||
|
||||
var Q [8][8]float64
|
||||
full := [8]float64{stdPos[0], stdPos[1], stdPos[2], stdPos[3], stdVel[0], stdVel[1], stdVel[2], stdVel[3]}
|
||||
for i := 0; i < 8; i++ {
|
||||
Q[i][i] = full[i] * full[i]
|
||||
}
|
||||
|
||||
kf.X = mul8x8_8x1(motionMat, kf.X)
|
||||
kf.P = add8x8(mul8x8_8x8(mul8x8_8x8(motionMat, kf.P), transpose8(motionMat)), Q)
|
||||
}
|
||||
|
||||
func (kf *Filter) Update(z [4]float64) {
|
||||
updateMat := [4][8]float64{
|
||||
{1, 0, 0, 0, 0, 0, 0, 0},
|
||||
{0, 1, 0, 0, 0, 0, 0, 0},
|
||||
{0, 0, 1, 0, 0, 0, 0, 0},
|
||||
{0, 0, 0, 1, 0, 0, 0, 0},
|
||||
}
|
||||
|
||||
Rdiag := [4]float64{
|
||||
weightPosition * z[2],
|
||||
weightPosition * z[3],
|
||||
weightPosition * z[2],
|
||||
weightPosition * z[3],
|
||||
}
|
||||
var R [4][4]float64
|
||||
for i := 0; i < 4; i++ {
|
||||
R[i][i] = Rdiag[i] * Rdiag[i]
|
||||
}
|
||||
|
||||
H := updateMat
|
||||
|
||||
HP := mul4x8_8x8(H, kf.P)
|
||||
|
||||
Ht := transpose4x8(H)
|
||||
HPHt := add4x4(mul4x8_8x4(HP, Ht), R)
|
||||
|
||||
sinv := inv4x4(HPHt)
|
||||
|
||||
PHt := mul8x8_8x4(kf.P, Ht)
|
||||
|
||||
K := mul8x4_4x4(PHt, sinv)
|
||||
|
||||
y := [4]float64{
|
||||
z[0] - dot8(H[0][:], kf.X[:]),
|
||||
z[1] - dot8(H[1][:], kf.X[:]),
|
||||
z[2] - dot8(H[2][:], kf.X[:]),
|
||||
z[3] - dot8(H[3][:], kf.X[:]),
|
||||
}
|
||||
|
||||
var Ky [8]float64
|
||||
for i := 0; i < 8; i++ {
|
||||
for j := 0; j < 4; j++ {
|
||||
Ky[i] += K[i][j] * y[j]
|
||||
}
|
||||
}
|
||||
for i := 0; i < 8; i++ {
|
||||
kf.X[i] += Ky[i]
|
||||
}
|
||||
|
||||
var KH [8][8]float64
|
||||
for i := 0; i < 8; i++ {
|
||||
for k := 0; k < 4; k++ {
|
||||
for j := 0; j < 8; j++ {
|
||||
KH[i][j] += K[i][k] * H[k][j]
|
||||
}
|
||||
}
|
||||
}
|
||||
var IKH [8][8]float64
|
||||
for i := 0; i < 8; i++ {
|
||||
IKH[i][i] = 1.0
|
||||
for j := 0; j < 8; j++ {
|
||||
IKH[i][j] -= KH[i][j]
|
||||
}
|
||||
}
|
||||
|
||||
IKHP := mul8x8_8x8(IKH, kf.P)
|
||||
IKHKt := mul8x8_8x8(IKHP, transpose8(IKH))
|
||||
|
||||
KR := mul8x4_4x4(K, R)
|
||||
KRKt := mul8x4_4x8(KR, transpose8x4(K))
|
||||
|
||||
kf.P = add8x8(IKHKt, KRKt)
|
||||
}
|
||||
|
||||
type Tracker struct {
|
||||
ID int
|
||||
Filter *Filter
|
||||
TimeSinceUpd int
|
||||
Hits int
|
||||
HitStreak int
|
||||
Age int
|
||||
}
|
||||
|
||||
var nextID int
|
||||
|
||||
func NewTracker(bbox [4]float64) *Tracker {
|
||||
nextID++
|
||||
x := (bbox[0] + bbox[2]) / 2
|
||||
y := (bbox[1] + bbox[3]) / 2
|
||||
w := bbox[2] - bbox[0]
|
||||
h := bbox[3] - bbox[1]
|
||||
|
||||
trk := &Tracker{
|
||||
ID: nextID,
|
||||
Filter: NewFilter(),
|
||||
Hits: 1,
|
||||
}
|
||||
trk.Filter.X = [8]float64{x, y, w, h, 0, 0, 0, 0}
|
||||
return trk
|
||||
}
|
||||
|
||||
func (t *Tracker) Predict() {
|
||||
if t.Filter.X[6]+t.Filter.X[2] <= 0 {
|
||||
t.Filter.X[6] = 0
|
||||
}
|
||||
t.Filter.Predict()
|
||||
t.Age++
|
||||
t.TimeSinceUpd++
|
||||
}
|
||||
|
||||
func (t *Tracker) Update(bbox [4]float64) {
|
||||
t.TimeSinceUpd = 0
|
||||
t.Hits++
|
||||
t.HitStreak++
|
||||
|
||||
x := (bbox[0] + bbox[2]) / 2
|
||||
y := (bbox[1] + bbox[3]) / 2
|
||||
w := bbox[2] - bbox[0]
|
||||
h := bbox[3] - bbox[1]
|
||||
|
||||
t.Filter.Update([4]float64{x, y, w, h})
|
||||
}
|
||||
|
||||
func (t *Tracker) GetState() [4]float64 {
|
||||
xx := t.Filter.X
|
||||
x, y, w, h := xx[0], xx[1], xx[2], xx[3]
|
||||
return [4]float64{
|
||||
x - w/2,
|
||||
y - h/2,
|
||||
x + w/2,
|
||||
y + h/2,
|
||||
}
|
||||
}
|
||||
|
||||
func (t *Tracker) GetCX() float64 {
|
||||
return t.Filter.X[0]
|
||||
}
|
||||
|
||||
func mul8x8_8x1(m [8][8]float64, v [8]float64) [8]float64 {
|
||||
var r [8]float64
|
||||
for i := 0; i < 8; i++ {
|
||||
for j := 0; j < 8; j++ {
|
||||
r[i] += m[i][j] * v[j]
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func mul8x8_8x8(a, b [8][8]float64) [8][8]float64 {
|
||||
var r [8][8]float64
|
||||
for i := 0; i < 8; i++ {
|
||||
for k := 0; k < 8; k++ {
|
||||
aik := a[i][k]
|
||||
for j := 0; j < 8; j++ {
|
||||
r[i][j] += aik * b[k][j]
|
||||
}
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func mul8x8_8x4(a [8][8]float64, b [8][4]float64) [8][4]float64 {
|
||||
var r [8][4]float64
|
||||
for i := 0; i < 8; i++ {
|
||||
for k := 0; k < 8; k++ {
|
||||
aik := a[i][k]
|
||||
for j := 0; j < 4; j++ {
|
||||
r[i][j] += aik * b[k][j]
|
||||
}
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func mul8x4_4x4(a [8][4]float64, b [4][4]float64) [8][4]float64 {
|
||||
var r [8][4]float64
|
||||
for i := 0; i < 8; i++ {
|
||||
for k := 0; k < 4; k++ {
|
||||
aik := a[i][k]
|
||||
for j := 0; j < 4; j++ {
|
||||
r[i][j] += aik * b[k][j]
|
||||
}
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func mul8x4_4x8(a [8][4]float64, b [4][8]float64) [8][8]float64 {
|
||||
var r [8][8]float64
|
||||
for i := 0; i < 8; i++ {
|
||||
for k := 0; k < 4; k++ {
|
||||
aik := a[i][k]
|
||||
for j := 0; j < 8; j++ {
|
||||
r[i][j] += aik * b[k][j]
|
||||
}
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func mul4x8_8x8(a [4][8]float64, b [8][8]float64) [4][8]float64 {
|
||||
var r [4][8]float64
|
||||
for i := 0; i < 4; i++ {
|
||||
for k := 0; k < 8; k++ {
|
||||
aik := a[i][k]
|
||||
for j := 0; j < 8; j++ {
|
||||
r[i][j] += aik * b[k][j]
|
||||
}
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func mul4x8_8x4(a [4][8]float64, b [8][4]float64) [4][4]float64 {
|
||||
var r [4][4]float64
|
||||
for i := 0; i < 4; i++ {
|
||||
for k := 0; k < 8; k++ {
|
||||
aik := a[i][k]
|
||||
for j := 0; j < 4; j++ {
|
||||
r[i][j] += aik * b[k][j]
|
||||
}
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func add8x8(a, b [8][8]float64) [8][8]float64 {
|
||||
var r [8][8]float64
|
||||
for i := 0; i < 8; i++ {
|
||||
for j := 0; j < 8; j++ {
|
||||
r[i][j] = a[i][j] + b[i][j]
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func add4x4(a, b [4][4]float64) [4][4]float64 {
|
||||
var r [4][4]float64
|
||||
for i := 0; i < 4; i++ {
|
||||
for j := 0; j < 4; j++ {
|
||||
r[i][j] = a[i][j] + b[i][j]
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func transpose8(a [8][8]float64) [8][8]float64 {
|
||||
var r [8][8]float64
|
||||
for i := 0; i < 8; i++ {
|
||||
for j := 0; j < 8; j++ {
|
||||
r[i][j] = a[j][i]
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func transpose4x8(a [4][8]float64) [8][4]float64 {
|
||||
var r [8][4]float64
|
||||
for i := 0; i < 4; i++ {
|
||||
for j := 0; j < 8; j++ {
|
||||
r[j][i] = a[i][j]
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func transpose8x4(a [8][4]float64) [4][8]float64 {
|
||||
var r [4][8]float64
|
||||
for i := 0; i < 8; i++ {
|
||||
for j := 0; j < 4; j++ {
|
||||
r[j][i] = a[i][j]
|
||||
}
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func inv4x4(a [4][4]float64) [4][4]float64 {
|
||||
var inv [4][4]float64
|
||||
|
||||
m := [4][4]float64{
|
||||
{a[0][0], a[0][1], a[0][2], a[0][3]},
|
||||
{a[1][0], a[1][1], a[1][2], a[1][3]},
|
||||
{a[2][0], a[2][1], a[2][2], a[2][3]},
|
||||
{a[3][0], a[3][1], a[3][2], a[3][3]},
|
||||
}
|
||||
|
||||
col := [4]int{0, 1, 2, 3}
|
||||
row := [4]int{0, 1, 2, 3}
|
||||
|
||||
for i := 0; i < 4; i++ {
|
||||
maxVal := math.Abs(m[row[i]][col[i]])
|
||||
pi, pj := i, i
|
||||
for r := i; r < 4; r++ {
|
||||
for c := i; c < 4; c++ {
|
||||
v := math.Abs(m[row[r]][col[c]])
|
||||
if v > maxVal {
|
||||
maxVal = v
|
||||
pi, pj = r, c
|
||||
}
|
||||
}
|
||||
}
|
||||
row[i], row[pi] = row[pi], row[i]
|
||||
col[i], col[pj] = col[pj], col[i]
|
||||
|
||||
pivot := m[row[i]][col[i]]
|
||||
if math.Abs(pivot) < 1e-12 {
|
||||
pivot = 1e-12
|
||||
}
|
||||
m[row[i]][col[i]] = 1.0
|
||||
for j := 0; j < 4; j++ {
|
||||
m[row[i]][j] /= pivot
|
||||
}
|
||||
|
||||
for r := 0; r < 4; r++ {
|
||||
if r != i {
|
||||
factor := m[row[r]][col[i]]
|
||||
m[row[r]][col[i]] = 0
|
||||
for j := 0; j < 4; j++ {
|
||||
m[row[r]][j] -= factor * m[row[i]][j]
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for i := 0; i < 4; i++ {
|
||||
for j := 0; j < 4; j++ {
|
||||
inv[col[i]][row[j]] = m[row[i]][j]
|
||||
}
|
||||
}
|
||||
return inv
|
||||
}
|
||||
|
||||
func dot8(a, b []float64) float64 {
|
||||
var s float64
|
||||
for i := 0; i < 8; i++ {
|
||||
s += a[i] * b[i]
|
||||
}
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
package rknn
|
||||
|
||||
/*
|
||||
#cgo linux LDFLAGS: -lrknnrt
|
||||
#cgo linux,arm64 LDFLAGS: -lrknnrt
|
||||
|
||||
#include <stdint.h>
|
||||
#include <stdlib.h>
|
||||
|
||||
typedef void* rknn_context;
|
||||
|
||||
typedef enum {
|
||||
RKNN_QUERY_IN_OUT_NUM = 0,
|
||||
RKNN_QUERY_INPUT_ATTR = 1,
|
||||
RKNN_QUERY_OUTPUT_ATTR = 2,
|
||||
RKNN_QUERY_SDK_VERSION = 5,
|
||||
} rknn_query_cmd;
|
||||
|
||||
typedef enum {
|
||||
RKNN_TENSOR_UINT8 = 0,
|
||||
RKNN_TENSOR_FLOAT32 = 2,
|
||||
} rknn_tensor_type;
|
||||
|
||||
typedef enum {
|
||||
RKNN_TENSOR_NCHW = 0,
|
||||
RKNN_TENSOR_NHWC = 1,
|
||||
} rknn_tensor_format;
|
||||
|
||||
typedef struct {
|
||||
uint32_t index;
|
||||
int32_t type;
|
||||
int32_t fmt;
|
||||
void* buf;
|
||||
uint32_t size;
|
||||
uint8_t want_float;
|
||||
uint8_t is_prealloc;
|
||||
} rknn_output;
|
||||
|
||||
typedef struct {
|
||||
uint32_t n_input;
|
||||
uint32_t n_output;
|
||||
} rknn_input_output_num;
|
||||
|
||||
typedef struct {
|
||||
void* buf;
|
||||
uint32_t size;
|
||||
uint8_t pass_through;
|
||||
uint32_t type;
|
||||
uint32_t fmt;
|
||||
} rknn_input;
|
||||
|
||||
extern int rknn_init(rknn_context* ctx, void* model, uint32_t size, uint32_t flag, rknn_input_output_num* io_num);
|
||||
extern int rknn_query(rknn_context ctx, rknn_query_cmd cmd, void* info, uint32_t size);
|
||||
extern int rknn_inputs_set(rknn_context ctx, uint32_t n_inputs, rknn_input inputs[]);
|
||||
extern int rknn_run(rknn_context ctx, void* ext);
|
||||
extern int rknn_outputs_get(rknn_context ctx, uint32_t n_outputs, rknn_output outputs[], void* ext);
|
||||
extern int rknn_outputs_release(rknn_context ctx, uint32_t n_outputs, rknn_output outputs[]);
|
||||
extern int rknn_destroy(rknn_context ctx);
|
||||
*/
|
||||
import "C"
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"unsafe"
|
||||
)
|
||||
|
||||
type Context struct {
|
||||
ctx C.rknn_context
|
||||
nInput uint32
|
||||
nOutput uint32
|
||||
}
|
||||
|
||||
func LoadModel(modelPath string) (*Context, error) {
|
||||
data, err := os.ReadFile(modelPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("rknn: read model file: %w", err)
|
||||
}
|
||||
|
||||
var ctx C.rknn_context
|
||||
var ioNum C.rknn_input_output_num
|
||||
|
||||
ret := C.rknn_init(
|
||||
&ctx,
|
||||
unsafe.Pointer(&data[0]),
|
||||
C.uint32_t(len(data)),
|
||||
0,
|
||||
&ioNum,
|
||||
)
|
||||
if ret != 0 {
|
||||
return nil, fmt.Errorf("rknn_init: error %d", int(ret))
|
||||
}
|
||||
|
||||
return &Context{
|
||||
ctx: ctx,
|
||||
nInput: uint32(ioNum.n_input),
|
||||
nOutput: uint32(ioNum.n_output),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *Context) InferenceRGB(inputData []uint8, height, width int) ([][]float32, error) {
|
||||
rknnIn := C.rknn_input{
|
||||
buf: unsafe.Pointer(&inputData[0]),
|
||||
size: C.uint32_t(len(inputData)),
|
||||
pass_through: 0,
|
||||
_type: C.RKNN_TENSOR_UINT8,
|
||||
fmt: C.RKNN_TENSOR_NHWC,
|
||||
}
|
||||
|
||||
ret := C.rknn_inputs_set(c.ctx, 1, &rknnIn)
|
||||
if ret < 0 {
|
||||
return nil, fmt.Errorf("rknn_inputs_set: error %d", int(ret))
|
||||
}
|
||||
|
||||
ret = C.rknn_run(c.ctx, nil)
|
||||
if ret < 0 {
|
||||
return nil, fmt.Errorf("rknn_run: error %d", int(ret))
|
||||
}
|
||||
|
||||
outputs := make([]C.rknn_output, c.nOutput)
|
||||
for i := range outputs {
|
||||
outputs[i].want_float = 1
|
||||
}
|
||||
|
||||
ret = C.rknn_outputs_get(c.ctx, C.uint32_t(c.nOutput), &outputs[0], nil)
|
||||
if ret < 0 {
|
||||
return nil, fmt.Errorf("rknn_outputs_get: error %d", int(ret))
|
||||
}
|
||||
|
||||
results := make([][]float32, c.nOutput)
|
||||
for i := range outputs {
|
||||
nVals := int(outputs[i].size) / 4
|
||||
vals := make([]float32, nVals)
|
||||
src := unsafe.Slice((*float32)(outputs[i].buf), nVals)
|
||||
copy(vals, src)
|
||||
results[i] = vals
|
||||
}
|
||||
|
||||
C.rknn_outputs_release(c.ctx, C.uint32_t(c.nOutput), &outputs[0])
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (c *Context) Release() {
|
||||
C.rknn_destroy(c.ctx)
|
||||
}
|
||||
@@ -0,0 +1,236 @@
|
||||
package yolo
|
||||
|
||||
import (
|
||||
"image"
|
||||
"math"
|
||||
|
||||
"github.com/anomalyco/bytetrack-counter-go/pkg/rknn"
|
||||
"gocv.io/x/gocv"
|
||||
)
|
||||
|
||||
type Detection struct {
|
||||
BBox [4]float64
|
||||
Score float64
|
||||
Class int
|
||||
Keypoints [][2]float64
|
||||
}
|
||||
|
||||
type Detector struct {
|
||||
rknn *rknn.Context
|
||||
imgSz int
|
||||
conf float64
|
||||
iouThr float64
|
||||
numClasses int
|
||||
scoreSigmoid bool
|
||||
}
|
||||
|
||||
func NewDetector(modelPath string, imgSz int, conf float64, numClasses int, scoreSigmoid bool) (*Detector, error) {
|
||||
ctx, err := rknn.LoadModel(modelPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Detector{
|
||||
rknn: ctx,
|
||||
imgSz: imgSz,
|
||||
conf: conf,
|
||||
iouThr: 0.45,
|
||||
numClasses: numClasses,
|
||||
scoreSigmoid: scoreSigmoid,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (d *Detector) Detect(frame gocv.Mat) ([]Detection, error) {
|
||||
h0, w0 := frame.Rows(), frame.Cols()
|
||||
|
||||
scale := math.Min(float64(d.imgSz)/float64(h0), float64(d.imgSz)/float64(w0))
|
||||
nh, nw := int(float64(h0)*scale), int(float64(w0)*scale)
|
||||
|
||||
resized := gocv.NewMat()
|
||||
defer resized.Close()
|
||||
gocv.Resize(frame, &resized, image.Point{X: nw, Y: nh}, 0, 0, gocv.InterpolationLinear)
|
||||
|
||||
letterbox := gocv.NewMatWithSize(d.imgSz, d.imgSz, gocv.MatTypeCV8UC3)
|
||||
defer letterbox.Close()
|
||||
letterbox.SetTo(gocv.NewScalar(114, 114, 114, 0))
|
||||
|
||||
dy := (d.imgSz - nh) / 2
|
||||
dx := (d.imgSz - nw) / 2
|
||||
roi := letterbox.Region(image.Rect(dx, dy, dx+nw, dy+nh))
|
||||
resized.CopyTo(&roi)
|
||||
|
||||
rgb := gocv.NewMat()
|
||||
defer rgb.Close()
|
||||
gocv.CvtColor(letterbox, &rgb, gocv.ColorBGRToRGB)
|
||||
|
||||
data := rgb.ToBytes()
|
||||
|
||||
outputs, err := d.rknn.InferenceRGB(data, d.imgSz, d.imgSz)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(outputs) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
out := outputs[0]
|
||||
return d.decodeOutput(out, h0, w0, scale, float64(dy), float64(dx))
|
||||
}
|
||||
|
||||
func (d *Detector) decodeOutput(out []float32, h0, w0 int, scale, padY, padX float64) ([]Detection, error) {
|
||||
if len(out) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
stride := d.numClasses + 4
|
||||
if len(out)%stride != 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
numDet := len(out) / stride
|
||||
if numDet == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
var detections []Detection
|
||||
|
||||
for i := 0; i < numDet; i++ {
|
||||
off := i * stride
|
||||
|
||||
cx := float64(out[off+0])
|
||||
cy := float64(out[off+1])
|
||||
w := float64(out[off+2])
|
||||
h := float64(out[off+3])
|
||||
|
||||
x1 := cx - w/2
|
||||
y1 := cy - h/2
|
||||
x2 := cx + w/2
|
||||
y2 := cy + h/2
|
||||
|
||||
var maxScore float64
|
||||
var bestClass int
|
||||
for c := 0; c < d.numClasses; c++ {
|
||||
score := float64(out[off+4+c])
|
||||
if d.scoreSigmoid {
|
||||
score = sigmoid(clamp(score, -10, 10))
|
||||
}
|
||||
if score > maxScore {
|
||||
maxScore = score
|
||||
bestClass = c
|
||||
}
|
||||
}
|
||||
|
||||
if maxScore <= d.conf {
|
||||
continue
|
||||
}
|
||||
|
||||
x1 = (x1 - padX) / scale
|
||||
y1 = (y1 - padY) / scale
|
||||
x2 = (x2 - padX) / scale
|
||||
y2 = (y2 - padY) / scale
|
||||
|
||||
x1 = clamp(x1, 0, float64(w0))
|
||||
y1 = clamp(y1, 0, float64(h0))
|
||||
x2 = clamp(x2, 0, float64(w0))
|
||||
y2 = clamp(y2, 0, float64(h0))
|
||||
|
||||
detections = append(detections, Detection{
|
||||
BBox: [4]float64{x1, y1, x2, y2},
|
||||
Score: maxScore,
|
||||
Class: bestClass,
|
||||
})
|
||||
}
|
||||
|
||||
return nmsByClass(detections, d.iouThr, d.numClasses), nil
|
||||
}
|
||||
|
||||
func nmsByClass(dets []Detection, iouThr float64, numClasses int) []Detection {
|
||||
if len(dets) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
byClass := make([][]int, numClasses)
|
||||
for i, d := range dets {
|
||||
byClass[d.Class] = append(byClass[d.Class], i)
|
||||
}
|
||||
|
||||
var result []Detection
|
||||
for cls := 0; cls < numClasses; cls++ {
|
||||
idxs := byClass[cls]
|
||||
if len(idxs) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
scores := make([]float64, len(idxs))
|
||||
boxes := make([][4]float64, len(idxs))
|
||||
for i, idx := range idxs {
|
||||
scores[i] = dets[idx].Score
|
||||
boxes[i] = dets[idx].BBox
|
||||
}
|
||||
|
||||
keep := nmsIndices(boxes, scores, iouThr)
|
||||
for _, k := range keep {
|
||||
result = append(result, dets[idxs[k]])
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func nmsIndices(boxes [][4]float64, scores []float64, iouThr float64) []int {
|
||||
if len(scores) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
order := make([]int, len(scores))
|
||||
for i := range order {
|
||||
order[i] = i
|
||||
}
|
||||
|
||||
for i := 0; i < len(scores); i++ {
|
||||
for j := i + 1; j < len(scores); j++ {
|
||||
if scores[order[i]] < scores[order[j]] {
|
||||
order[i], order[j] = order[j], order[i]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var keep []int
|
||||
for len(order) > 0 {
|
||||
idx := order[0]
|
||||
keep = append(keep, idx)
|
||||
if len(order) == 1 {
|
||||
break
|
||||
}
|
||||
|
||||
rest := order[1:]
|
||||
var newOrder []int
|
||||
for _, r := range rest {
|
||||
xx1 := math.Max(boxes[idx][0], boxes[r][0])
|
||||
yy1 := math.Max(boxes[idx][1], boxes[r][1])
|
||||
xx2 := math.Min(boxes[idx][2], boxes[r][2])
|
||||
yy2 := math.Min(boxes[idx][3], boxes[r][3])
|
||||
w := math.Max(0, xx2-xx1)
|
||||
h := math.Max(0, yy2-yy1)
|
||||
inter := w * h
|
||||
areaI := (boxes[idx][2] - boxes[idx][0]) * (boxes[idx][3] - boxes[idx][1])
|
||||
areaR := (boxes[r][2] - boxes[r][0]) * (boxes[r][3] - boxes[r][1])
|
||||
iou := inter / (areaI + areaR - inter + 1e-16)
|
||||
if iou < iouThr {
|
||||
newOrder = append(newOrder, r)
|
||||
}
|
||||
}
|
||||
order = newOrder
|
||||
}
|
||||
return keep
|
||||
}
|
||||
|
||||
func sigmoid(x float64) float64 {
|
||||
return 1.0 / (1.0 + math.Exp(-x))
|
||||
}
|
||||
|
||||
func clamp(x, lo, hi float64) float64 {
|
||||
return math.Max(lo, math.Min(hi, x))
|
||||
}
|
||||
|
||||
func (d *Detector) Release() {
|
||||
d.rknn.Release()
|
||||
}
|
||||
Reference in new issue
Block a user