From dc8dcca75ee22b824ddbc4f2ceea41fcb2ae7f71 Mon Sep 17 00:00:00 2001 From: dsutanto Date: Tue, 30 Jun 2026 18:49:39 +0700 Subject: [PATCH] First commit --- README.md | 2 - batch_store.py | 379 ++++++++++ bytetrack-counter-go.service | 32 + cmd/counter/main.go | 545 ++++++++++++++ config.env.example | 54 ++ counter_live_rknn_bytetrack.py | 1289 ++++++++++++++++++++++++++++++++ go.mod | 24 + go.sum | 51 ++ internal/iou/iou.go | 110 +++ internal/nms/nms.go | 45 ++ pkg/batchstore/store.go | 467 ++++++++++++ pkg/bytetrack/tracker.go | 195 +++++ pkg/config/config.go | 173 +++++ pkg/drawing/drawing.go | 241 ++++++ pkg/kalman/kalman.go | 396 ++++++++++ pkg/rknn/rknn.go | 145 ++++ pkg/yolo/yolo.go | 236 ++++++ 17 files changed, 4382 insertions(+), 2 deletions(-) delete mode 100644 README.md create mode 100644 batch_store.py create mode 100644 bytetrack-counter-go.service create mode 100644 cmd/counter/main.go create mode 100644 config.env.example create mode 100644 counter_live_rknn_bytetrack.py create mode 100644 go.mod create mode 100644 go.sum create mode 100644 internal/iou/iou.go create mode 100644 internal/nms/nms.go create mode 100644 pkg/batchstore/store.go create mode 100644 pkg/bytetrack/tracker.go create mode 100644 pkg/config/config.go create mode 100644 pkg/drawing/drawing.go create mode 100644 pkg/kalman/kalman.go create mode 100644 pkg/rknn/rknn.go create mode 100644 pkg/yolo/yolo.go diff --git a/README.md b/README.md deleted file mode 100644 index a7c519a..0000000 --- a/README.md +++ /dev/null @@ -1,2 +0,0 @@ -# bytetrack-counter-go - diff --git a/batch_store.py b/batch_store.py new file mode 100644 index 0000000..94ee985 --- /dev/null +++ b/batch_store.py @@ -0,0 +1,379 @@ +""" +Production batch persistence for edge Jetson counter. +Mirrors frigate-counter SQLite schema + current_batch.json contract. +""" +import json +import sqlite3 +import threading +import time +from datetime import datetime, timedelta +from pathlib import Path + + +class BatchStore: + def __init__( + self, + db_path, + state_file, + camera_name, + object_label='ayam-potong', + cutoff_time='20:00', + batch_timeout=300.0, + ignore_batch_label_timeout=30.0, + min_object_per_batch=60, + min_duration_per_batch=60, + carry_ids=50, + logger=print, + ): + self.db_path = db_path + self.state_file = Path(state_file) + self.camera_name = camera_name + self.object_label = object_label + self.cutoff_time_str = cutoff_time + datetime.strptime(cutoff_time, '%H:%M') + + self.batch_timeout = float(batch_timeout) + self.ignore_batch_label_timeout = float(ignore_batch_label_timeout) + self.min_object_per_batch = int(min_object_per_batch) + self.min_duration_per_batch = int(min_duration_per_batch) + self.carry_ids = int(carry_ids) + self.log = logger + + self.state_lock = threading.Lock() + self.batch_timer = None + self.ignore_batch_label = False + self.ignore_batch_label_timer = None + self.previous_state = None + self.shutdown_event = threading.Event() + + Path(db_path).parent.mkdir(parents=True, exist_ok=True) + self.state_file.parent.mkdir(parents=True, exist_ok=True) + + self.db = sqlite3.connect(db_path, check_same_thread=False) + self._init_db() + self.current_state = self._load_state() + self.previous_state = self.current_state + + def _init_db(self): + cur = self.db.cursor() + cur.execute( + """ + 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) + ) + """ + ) + cur.execute( + """ + 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) + ) + """ + ) + self.db.commit() + + def get_counting_date(self, dt=None): + if dt is None: + dt = datetime.now() + cutoff = datetime.strptime(self.cutoff_time_str, '%H:%M').time() + if dt.time() < cutoff: + return dt.date().isoformat() + return (dt.date() + timedelta(days=1)).isoformat() + + def _load_state(self): + if not self.state_file.exists(): + return None + try: + with open(self.state_file, 'r', encoding='utf-8') as f: + state = json.load(f) + current_date = self.get_counting_date() + if state.get('counting_date') != current_date: + self.log( + f"State file belongs to previous counting day ({state.get('counting_date')}). " + 'Finalizing before fresh start.' + ) + self._insert_batch( + state['counting_date'], + state['batch_number'], + state['count'], + state['start_time'], + datetime.now().isoformat(), + ) + self.state_file.unlink(missing_ok=True) + return None + self.log( + f"Resumed batch #{state['batch_number']} from {state['start_time']} " + f"with count={state['count']}" + ) + self._reset_batch_timer() + return state + except Exception as exc: + self.log(f'Failed to load state file: {exc}') + return None + + def save_state(self): + if self.current_state is None: + self.state_file.unlink(missing_ok=True) + return + with open(self.state_file, 'w', encoding='utf-8') as f: + json.dump(self.current_state, f, indent=2, ensure_ascii=False) + + def get_next_batch_number(self, counting_date): + cur = self.db.cursor() + cur.execute( + """ + SELECT COALESCE(MAX(batch_number), 0) + FROM batches + WHERE counting_date = ? AND camera_name = ? AND object_label = ? + """, + (counting_date, self.camera_name, self.object_label), + ) + return cur.fetchone()[0] + 1 + + def start_new_batch(self, counting_date): + batch_number = self.get_next_batch_number(counting_date) + now = datetime.now().isoformat() + counted_ids = [] + if self.previous_state is not None: + try: + counted_ids = self.previous_state['counted_event_ids'][-self.carry_ids:] + except (KeyError, TypeError): + counted_ids = [] + self.current_state = { + 'counting_date': counting_date, + 'batch_number': batch_number, + 'count': 0, + 'start_time': now, + 'last_detection_time': now, + 'counted_event_ids': counted_ids, + } + self.save_state() + self.log(f'Started batch #{batch_number} for {counting_date} ({self.object_label})') + + def _reset_batch_timer(self): + if self.batch_timer: + self.batch_timer.cancel() + self.batch_timer = threading.Timer(self.batch_timeout, self._on_batch_timeout) + self.batch_timer.daemon = True + self.batch_timer.start() + + def _on_batch_timeout(self): + self.log(f'Batch inactivity timeout ({self.batch_timeout}s) reached') + self.end_batch(closed_by='timeout') + + def _ignore_batch_label(self): + if not self.ignore_batch_label_timer: + self.ignore_batch_label = True + self.ignore_batch_label_timer = threading.Timer( + self.ignore_batch_label_timeout, self._on_ignore_batch_label_timeout + ) + self.ignore_batch_label_timer.daemon = True + self.ignore_batch_label_timer.start() + self.log( + f'Ignore batch label for {self.ignore_batch_label_timeout}s' + ) + + def _on_ignore_batch_label_timeout(self): + self.ignore_batch_label_timer = None + self.ignore_batch_label = False + self.log('Ignore batch label cooldown finished') + + def record_ayam_crossing(self, track_id): + """Line-cross equivalent of production ayam-potong MQTT event.""" + with self.state_lock: + counting_date = self.get_counting_date() + started_new = False + if self.current_state is None: + self.start_new_batch(counting_date) + started_new = True + elif self.current_state['counting_date'] != counting_date: + self._end_batch_locked(closed_by='cutoff') + self.start_new_batch(counting_date) + started_new = True + + event_key = str(track_id) + if event_key not in self.current_state['counted_event_ids']: + self.current_state['count'] += 1 + self.current_state['counted_event_ids'].append(event_key) + self.log( + f'Counted ayam (track {track_id}) | batch #{self.current_state["batch_number"]} ' + f'total: {self.current_state["count"]}' + ) + + self.current_state['last_detection_time'] = datetime.now().isoformat() + self.save_state() + self._reset_batch_timer() + return self.current_state['count'], started_new + + def record_talenan_crossing(self, track_id): + """Line-cross equivalent of production telenan MQTT batch close.""" + if self.ignore_batch_label: + return False + with self.state_lock: + self._ignore_batch_label() + self._end_batch_locked(closed_by='talenan') + self.log(f'Batch closed by talenan (track {track_id})') + if self.batch_timer: + self.batch_timer.cancel() + self.batch_timer = None + return True + + def end_batch(self, closed_by='manual'): + with self.state_lock: + self._end_batch_locked(closed_by=closed_by) + + def _end_batch_locked(self, closed_by='manual'): + if self.current_state is None: + return False + + self.previous_state = self.current_state + state = self.current_state + + start_time_obj = datetime.fromisoformat(state['start_time']) + end_time_obj = datetime.now() + duration_seconds = (end_time_obj - start_time_obj).total_seconds() + + if (state['count'] < self.min_object_per_batch + or duration_seconds < self.min_duration_per_batch): + self.current_state = None + self.save_state() + if self.batch_timer: + self.batch_timer.cancel() + self.batch_timer = None + self.log( + f'Batch #{state["batch_number"]} discarded ' + f'(count={state["count"]}, duration={duration_seconds:.0f}s)' + ) + return False + + end_time = end_time_obj.isoformat() + try: + self._insert_batch( + state['counting_date'], + state['batch_number'], + state['count'], + state['start_time'], + end_time, + ) + cps = state['count'] / duration_seconds if duration_seconds > 0 else 0 + self.log( + f'Batch #{state["batch_number"]} ended | count={state["count"]} | ' + f'duration={duration_seconds:.0f}s | cps={cps:.3f} | closed_by={closed_by}' + ) + except Exception as exc: + self.log(f'Failed to persist batch: {exc}') + return False + + self.current_state = None + self.save_state() + if self.batch_timer: + self.batch_timer.cancel() + self.batch_timer = None + return True + + def _insert_batch(self, counting_date, batch_number, count, start_time, end_time): + cur = self.db.cursor() + cur.execute( + """ + INSERT INTO batches + (counting_date, batch_number, camera_name, object_label, count, start_time, end_time) + VALUES (?, ?, ?, ?, ?, ?, ?) + """, + (counting_date, batch_number, self.camera_name, self.object_label, count, start_time, end_time), + ) + cur.execute( + """ + 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 + """, + (counting_date, self.camera_name, self.object_label, count), + ) + self.db.commit() + + cur.execute( + """ + SELECT total_count, total_batches + FROM daily_summaries + WHERE counting_date = ? AND camera_name = ? AND object_label = ? + """, + (counting_date, self.camera_name, self.object_label), + ) + row = cur.fetchone() + if row: + self.log( + f'Daily totals for {counting_date}: {row[0]} objects across {row[1]} batch(es)' + ) + + def cutoff_watcher_loop(self): + while not self.shutdown_event.is_set(): + time.sleep(60) + with self.state_lock: + if self.current_state is None: + continue + if self.current_state['counting_date'] != self.get_counting_date(): + self.log('Daily cutoff reached – finalizing batch') + self._end_batch_locked(closed_by='cutoff') + + def start_cutoff_watcher(self): + t = threading.Thread(target=self.cutoff_watcher_loop, daemon=True) + t.start() + return t + + @property + def current_batch_number(self): + if self.current_state is None: + return 0 + return self.current_state['batch_number'] + + @property + def current_batch_count(self): + if self.current_state is None: + return 0 + return self.current_state['count'] + + def get_closed_total_for_day(self, counting_date=None): + if counting_date is None: + counting_date = self.get_counting_date() + cur = self.db.cursor() + cur.execute( + """ + SELECT COALESCE(total_count, 0) + FROM daily_summaries + WHERE counting_date = ? AND camera_name = ? AND object_label = ? + """, + (counting_date, self.camera_name, self.object_label), + ) + row = cur.fetchone() + return row[0] if row else 0 + + def display_total(self): + return self.get_closed_total_for_day() + self.current_batch_count + + def shutdown(self): + self.shutdown_event.set() + self.end_batch(closed_by='shutdown') + if self.batch_timer: + self.batch_timer.cancel() + self.db.close() diff --git a/bytetrack-counter-go.service b/bytetrack-counter-go.service new file mode 100644 index 0000000..e04b0ff --- /dev/null +++ b/bytetrack-counter-go.service @@ -0,0 +1,32 @@ +[Unit] +Description=Edge ByteTrack Counter (RTSP + RKNN) — Go +Documentation=file:///opt/bytetrack-counter-go/DEPLOY.md +After=network-online.target +Wants=network-online.target + +[Service] +Type=simple +User=root +Group=root + +WorkingDirectory=/opt/bytetrack-counter-go + +EnvironmentFile=/opt/bytetrack-counter-go/.env +Environment=PATH=/usr/local/bin:/usr/bin:/bin + +ExecStart=/opt/bytetrack-counter-go/counter + +TimeoutStopSec=30 +KillSignal=SIGTERM + +Restart=on-failure +RestartSec=10 +StartLimitInterval=120s +StartLimitBurst=5 + +NoNewPrivileges=true +ProtectHome=true +PrivateTmp=false + +[Install] +WantedBy=multi-user.target diff --git a/cmd/counter/main.go b/cmd/counter/main.go new file mode 100644 index 0000000..f62a74e --- /dev/null +++ b/cmd/counter/main.go @@ -0,0 +1,545 @@ +package main + +import ( + "encoding/csv" + "fmt" + "image" + "image/color" + "os" + "os/signal" + "path/filepath" + "syscall" + "time" + + "github.com/anomalyco/bytetrack-counter-go/pkg/batchstore" + "github.com/anomalyco/bytetrack-counter-go/pkg/bytetrack" + "github.com/anomalyco/bytetrack-counter-go/pkg/config" + "github.com/anomalyco/bytetrack-counter-go/pkg/drawing" + "github.com/anomalyco/bytetrack-counter-go/pkg/yolo" + + "gocv.io/x/gocv" +) + +var cGreen = color.RGBA{80, 220, 100, 255} +var cAyamBox = color.RGBA{0, 165, 255, 255} +var cTalenanBox = color.RGBA{220, 120, 60, 255} + +type trackedEntry struct { + CX float64 + Mono float64 +} + +func main() { + cfg := config.Load() + + storeCfg := batchstore.Config{ + DBPath: cfg.DBPath, + StateFile: cfg.StateFile, + CameraName: cfg.CameraName, + ObjectLabel: cfg.ObjectLabel, + CutoffTime: cfg.DailyCutoffTime, + BatchTimeoutSec: cfg.BatchTimeoutSec, + IgnoreBatchLabelTimeout: cfg.IgnoreBatchLabelTimeoutSec, + MinObjectPerBatch: cfg.MinObjectPerBatch, + MinDurationPerBatch: cfg.MinDurationPerBatch, + } + + logFn := func(msg string) { + fmt.Printf("[%s] %s\n", drawing.NowStr(), msg) + } + + store, err := batchstore.New(storeCfg, logFn) + if err != nil { + fmt.Fprintf(os.Stderr, "Failed to init batch store: %v\n", err) + os.Exit(1) + } + defer store.Shutdown() + store.StartCutoffWatcher() + + classIDs := map[string]int{ + cfg.ClassAyam: 0, + cfg.ClassTalenan: 1, + } + talenanCls := classIDs[cfg.ClassTalenan] + ayamCls := classIDs[cfg.ClassAyam] + + detector, err := yolo.NewDetector(cfg.ModelPath, cfg.ImgSz, cfg.Conf, cfg.NumClasses, cfg.ScoreSigmoid) + if err != nil { + fmt.Fprintf(os.Stderr, "Failed to init YOLO detector: %v\n", err) + os.Exit(1) + } + defer detector.Release() + + ayamTracker := bytetrack.New(cfg.TrackHighThresh, cfg.TrackLowThresh, cfg.TrackMatchThresh, cfg.TrackBuffer, cfg.TrackMinHits) + talenanTracker := bytetrack.New(cfg.TrackHighThresh, cfg.TrackLowThresh, cfg.TrackMatchThresh, cfg.TrackBuffer, cfg.TrackMinHits) + + shutdownCh := make(chan os.Signal, 1) + signal.Notify(shutdownCh, syscall.SIGINT, syscall.SIGTERM) + + cap, w, h, fps := connectStream(cfg, cfg.WarmupFrames) + if cap == nil { + return + } + defer cap.Close() + + lineX := drawing.ResolveLineX(w, cfg.LineX, cfg.LineXFrac) + fmt.Printf("RKNN+ByteTrack counter | %dx%d @ %.1ffps | line x=%d | cross=%s\n", w, h, fps, lineX, cfg.CrossDirection) + fmt.Printf("Model: %s | imgsz=%d | core_mask=%d\n", cfg.ModelPath, cfg.ImgSz, cfg.CoreMask) + fmt.Printf("ByteTrack: high=%.2f low=%.2f match=%.2f buffer=%d\n", cfg.TrackHighThresh, cfg.TrackLowThresh, cfg.TrackMatchThresh, cfg.TrackBuffer) + fmt.Printf("DB: %s\nState: %s\n", cfg.DBPath, cfg.StateFile) + + var csvLogger *localCSVLogger + if cfg.ExportCSV { + csvLogger, err = newLocalCSVLogger(cfg.CrossCSV, []string{"batch", "frame", "timestamp", "chicken_id"}) + if err != nil { + fmt.Fprintf(os.Stderr, "Failed to init CSV logger: %v\n", err) + } + if csvLogger != nil { + defer csvLogger.close() + } + } + + var videoWriter *segmentWriter + if cfg.RecordVideo { + videoWriter = newSegmentWriter(cfg.OutputDir, w, h, fps, cfg.VideoSegmentSec) + defer videoWriter.release() + } + + ayamTracked := make(map[int]trackedEntry) + talenanTracked := make(map[int]trackedEntry) + ayamLineCrossed := make(map[int]bool) + talenanLineCrossed := make(map[int]bool) + ayamCrossFlash := make(map[int]int) + talenanCrossFlash := make(map[int]int) + + var linePulse, countPulse, batchPulse int + var popups []drawing.Popup + + sessionStart := time.Now() + frameIdx := 0 + reconnectCount := 0 + + frame := gocv.NewMat() + defer frame.Close() + + for { + select { + case <-shutdownCh: + fmt.Println("\nShutdown requested -- finishing current frame...") + goto shutdown + default: + } + + if ok := cap.Read(&frame); !ok { + if !cfg.IsLive { + break + } + reconnectCount++ + fmt.Printf("Stream dropped (attempt %d), reconnecting in %ds...\n", reconnectCount, cfg.ReconnectDelaySec) + cap.Close() + time.Sleep(time.Duration(cfg.ReconnectDelaySec) * time.Second) + newCap, nw, nh, nfps := connectStream(cfg, 0) + if newCap == nil { + break + } + cap = newCap + w, h, fps = nw, nh, nfps + lineX = drawing.ResolveLineX(w, cfg.LineX, cfg.LineXFrac) + continue + } + + now := time.Now() + elapsed := now.Sub(sessionStart).Seconds() + mono := float64(time.Now().UnixNano()) / 1e9 + + ayamCrossedFrame := false + batchClosedFrame := false + batchStartedFrame := false + + detections, err := detector.Detect(frame) + if err != nil { + fmt.Fprintf(os.Stderr, "Detection error: %v\n", err) + } + + if len(detections) > 0 { + ayamBoxes := make([][4]float64, 0) + ayamScores := make([]float64, 0) + ayamKptsList := make([][][2]float64, 0) + ayamCXList := make([]float64, 0) + + talenanBoxes := make([][4]float64, 0) + talenanScores := make([]float64, 0) + talenanKptsList := make([][][2]float64, 0) + talenanCXList := make([]float64, 0) + + for _, det := range detections { + cx := (det.BBox[0] + det.BBox[2]) / 2.0 + if det.Class == talenanCls { + talenanBoxes = append(talenanBoxes, det.BBox) + talenanScores = append(talenanScores, det.Score) + talenanKptsList = append(talenanKptsList, det.Keypoints) + talenanCXList = append(talenanCXList, cx) + } else if det.Class == ayamCls { + ayamBoxes = append(ayamBoxes, det.BBox) + ayamScores = append(ayamScores, det.Score) + ayamKptsList = append(ayamKptsList, det.Keypoints) + ayamCXList = append(ayamCXList, cx) + } + } + + ayamResult := ayamTracker.Update(ayamBoxes, ayamScores) + talenanResult := talenanTracker.Update(talenanBoxes, talenanScores) + + for di := 0; di < len(talenanBoxes); di++ { + tid, ok := talenanResult.DetToTrack[di] + if !ok { + continue + } + cx := talenanCXList[di] + bbox := talenanBoxes[di] + + if prev, ok := talenanTracked[tid]; ok { + if drawing.CrossedLine(prev.CX, cx, lineX, cfg.CrossDirection) && !talenanLineCrossed[tid] { + talenanLineCrossed[tid] = true + if store.RecordTalenanCrossing(tid) { + batchClosedFrame = true + talenanCrossFlash[tid] = drawing.CrossFlashFrames + popups = append(popups, drawing.Popup{ + X: int(cx) - 20, + Y: int((bbox[1] + bbox[3]) / 2), + Born: frameIdx, + Text: "BATCH CLOSED", + }) + } + } + } + talenanTracked[tid] = trackedEntry{CX: cx, Mono: mono} + } + + for di := 0; di < len(ayamBoxes); di++ { + tid, ok := ayamResult.DetToTrack[di] + if !ok { + continue + } + cx := ayamCXList[di] + + if prev, ok := ayamTracked[tid]; ok { + if drawing.CrossedLine(prev.CX, cx, lineX, cfg.CrossDirection) && !ayamLineCrossed[tid] { + ayamLineCrossed[tid] = true + _, startedNew := store.RecordAyamCrossing(tid) + if csvLogger != nil { + csvLogger.write([]string{ + fmt.Sprintf("%d", store.CurrentBatchNumber()), + fmt.Sprintf("%d", frameIdx), + time.Now().Format(time.RFC3339), + fmt.Sprintf("%d", tid), + }) + } + ayamCrossedFrame = true + if startedNew { + batchStartedFrame = true + } + ayamCrossFlash[tid] = drawing.CrossFlashFrames + popups = append(popups, drawing.Popup{ + X: int(cx) - 12, + Y: int((ayamBoxes[di][1] + ayamBoxes[di][3]) / 2), + Born: frameIdx, + Text: "+1", + }) + } + } + ayamTracked[tid] = trackedEntry{CX: cx, Mono: mono} + } + + for tid, cx := range ayamResult.Lost { + if _, ok := ayamTracked[tid]; !ok { + ayamTracked[tid] = trackedEntry{CX: cx, Mono: mono} + } + } + + for di := 0; di < len(talenanBoxes); di++ { + tid, ok := talenanResult.DetToTrack[di] + if !ok { + continue + } + bbox := talenanBoxes[di] + x1, y1, x2, y2 := int(bbox[0]), int(bbox[1]), int(bbox[2]), int(bbox[3]) + flash := talenanCrossFlash[tid] + clr := cTalenanBox + thick := 2 + if flash > 0 { + clr = cGreen + thick = 3 + } + gocv.Rectangle(&frame, image.Rect(x1, y1, x2, y2), clr, thick) + drawing.DrawPill(&frame, fmt.Sprintf("TALENAN %d", tid), x1, y1-4, clr) + } + + for di := 0; di < len(ayamBoxes); di++ { + tid, ok := ayamResult.DetToTrack[di] + if !ok { + continue + } + bbox := ayamBoxes[di] + x1, y1, x2, y2 := int(bbox[0]), int(bbox[1]), int(bbox[2]), int(bbox[3]) + flash := ayamCrossFlash[tid] + clr := cAyamBox + thick := 2 + if flash > 0 { + clr = cGreen + thick = 3 + } + gocv.Rectangle(&frame, image.Rect(x1, y1, x2, y2), clr, thick) + drawing.DrawPill(&frame, fmt.Sprintf("ID %d", tid), x1, y1-4, clr) + if di < len(ayamKptsList) && len(ayamKptsList[di]) > 0 { + drawing.DrawSkeleton(&frame, ayamKptsList[di]) + } + } + } + + if ayamCrossedFrame { + linePulse = drawing.LinePulseFrames + countPulse = drawing.CountPulseFrames + } + if batchClosedFrame { + linePulse = drawing.LinePulseFrames + } + if batchStartedFrame { + batchPulse = drawing.BatchPulseFrames + } + + batchNum := store.CurrentBatchNumber() + batchCount := store.CurrentBatchCount() + displayTotal := store.DisplayTotal() + rate := (float64(displayTotal) / elapsed) * 60 + if elapsed <= 0 { + rate = 0 + } + + drawing.DrawCountingLine(&frame, lineX, h, linePulse) + drawing.DrawHeroCount(&frame, lineX, h, batchCount, countPulse) + drawing.DrawHUD(&frame, w, batchNum, batchCount, displayTotal, elapsed, rate, cfg.CameraName) + drawing.DrawBatchBanner(&frame, w, batchNum, batchPulse) + + liveTag := "LIVE-RKNN-BT" + if !cfg.IsLive { + liveTag = "FILE-RKNN-BT" + } + drawing.DrawFooter(&frame, w, h, frameIdx, liveTag) + + popups = drawing.DrawPopups(&frame, popups, frameIdx) + + for tid := range ayamCrossFlash { + ayamCrossFlash[tid]-- + if ayamCrossFlash[tid] <= 0 { + delete(ayamCrossFlash, tid) + } + } + for tid := range talenanCrossFlash { + talenanCrossFlash[tid]-- + if talenanCrossFlash[tid] <= 0 { + delete(talenanCrossFlash, tid) + } + } + if linePulse > 0 { + linePulse-- + } + if countPulse > 0 { + countPulse-- + } + if batchPulse > 0 { + batchPulse-- + } + + if videoWriter != nil { + videoWriter.write(frame) + } + + if cfg.LiveStreamEnabled && frameIdx%cfg.LiveStreamEveryN == 0 { + writeLiveFrame(frame, cfg.LiveStreamFramePath) + } + + frameIdx++ + if frameIdx%cfg.FlushEveryNFrames == 0 { + fmt.Printf("[%s] Frame %d | Batch %d: %d | Total: %d | Uptime %.2fh\n", + drawing.NowStr(), frameIdx, batchNum, batchCount, displayTotal, elapsed/3600) + } + + pruneStaleTracks(ayamTracked, mono, cfg.TrackedPruneSec) + pruneStaleTracks(talenanTracked, mono, cfg.TrackedPruneSec) + } + +shutdown: + fmt.Println("\n=== Batch Summary (SQLite) ===") + fmt.Printf("Database: %s\n", cfg.DBPath) +} + +func connectStream(cfg *config.Config, warmup int) (*gocv.VideoCapture, int, int, float64) { + fps := float64(cfg.OutputFPS) + attempts := 0 + + for { + if isRTSP(cfg.Source) || isHTTP(cfg.Source) { + os.Setenv("OPENCV_FFMPEG_CAPTURE_OPTIONS", cfg.RTSPFFmpegOptions) + } + + cap, err := gocv.OpenVideoCapture(cfg.Source) + if err != nil || !cap.IsOpened() { + attempts++ + if cfg.MaxReconnectAttempts > 0 && attempts >= cfg.MaxReconnectAttempts { + fmt.Fprintf(os.Stderr, "Cannot open source after %d attempts: %s\n", attempts, cfg.Source) + return nil, 0, 0, fps + } + fmt.Printf("Cannot open source, retry in %ds...\n", cfg.ReconnectDelaySec) + time.Sleep(time.Duration(cfg.ReconnectDelaySec) * time.Second) + continue + } + + cap.Set(38, 1.0) + + if warmup > 0 && (isRTSP(cfg.Source) || isHTTP(cfg.Source)) { + fmt.Println("Warming up stream...") + mat := gocv.NewMat() + for i := 0; i < warmup; i++ { + cap.Read(&mat) + } + mat.Close() + fmt.Println("Stream ready!") + } + + w := int(cap.Get(gocv.VideoCaptureFrameWidth)) + h := int(cap.Get(gocv.VideoCaptureFrameHeight)) + gotFPS := cap.Get(gocv.VideoCaptureFPS) + if gotFPS > 1 { + fps = gotFPS + } + + return cap, w, h, fps + } +} + +func isRTSP(s string) bool { + return len(s) >= 7 && s[:7] == "rtsp://" +} + +func isHTTP(s string) bool { + return len(s) >= 7 && s[:7] == "http://" +} + +type localCSVLogger struct { + file *os.File + writer *csv.Writer +} + +func newLocalCSVLogger(path string, header []string) (*localCSVLogger, error) { + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { + return nil, err + } + + newFile := false + if _, err := os.Stat(path); os.IsNotExist(err) { + newFile = true + } else if fi, _ := os.Stat(path); fi.Size() == 0 { + newFile = true + } + + f, err := os.OpenFile(path, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644) + if err != nil { + return nil, err + } + + w := csv.NewWriter(f) + if newFile { + if err := w.Write(header); err != nil { + f.Close() + return nil, err + } + w.Flush() + } + + return &localCSVLogger{file: f, writer: w}, nil +} + +func (c *localCSVLogger) write(row []string) { + c.writer.Write(row) + c.writer.Flush() +} + +func (c *localCSVLogger) close() { + c.writer.Flush() + c.file.Close() +} + +type segmentWriter struct { + outputDir string + w, h, fps int + segmentSec int + segmentStart time.Time + writer *gocv.VideoWriter +} + +func newSegmentWriter(outputDir string, w, h int, fps float64, segmentSec int) *segmentWriter { + os.MkdirAll(outputDir, 0755) + sw := &segmentWriter{ + outputDir: outputDir, + w: w, + h: h, + fps: int(fps), + segmentSec: segmentSec, + } + sw.openNext() + return sw +} + +func (sw *segmentWriter) segmentPath() string { + ts := time.Now().Format("20060102_150405") + return filepath.Join(sw.outputDir, fmt.Sprintf("live_%s.mp4", ts)) +} + +func (sw *segmentWriter) openNext() { + if sw.writer != nil { + sw.writer.Close() + } + path := sw.segmentPath() + wr, err := gocv.VideoWriterFile(path, "avc1", float64(sw.fps), sw.w, sw.h, true) + if err != nil { + fmt.Fprintf(os.Stderr, "Failed to create video writer: %v\n", err) + return + } + sw.writer = wr + sw.segmentStart = time.Now() + fmt.Printf("Recording segment: %s\n", path) +} + +func (sw *segmentWriter) write(frame gocv.Mat) { + if time.Since(sw.segmentStart).Seconds() >= float64(sw.segmentSec) { + sw.openNext() + } + if sw.writer != nil { + sw.writer.Write(frame) + } +} + +func (sw *segmentWriter) release() { + if sw.writer != nil { + sw.writer.Close() + } +} + +func writeLiveFrame(frame gocv.Mat, path string) { + os.MkdirAll(filepath.Dir(path), 0755) + buf, err := gocv.IMEncode(".jpg", frame) + if err != nil { + return + } + defer buf.Close() + os.WriteFile(path, buf.GetBytes(), 0644) +} + +func pruneStaleTracks(tracked map[int]trackedEntry, nowMono, maxAgeSec float64) { + for tid, entry := range tracked { + if nowMono-entry.Mono > maxAgeSec { + delete(tracked, tid) + } + } +} diff --git a/config.env.example b/config.env.example new file mode 100644 index 0000000..ca6fcfd --- /dev/null +++ b/config.env.example @@ -0,0 +1,54 @@ +# Edge RK3588 production counter — copy to .env on device +# cp config.env.example .env && nano .env + +OUTPUT_DIR=/opt/jetson-counter +DB_PATH=/opt/jetson-counter/jetson_counter.db +STATE_FILE=/opt/jetson-counter/current_batch.json + +# Direct LAN camera RTSP (low latency) +SOURCE=rtsp://user:pass@192.168.0.100:554/stream1 +OPENCV_FFMPEG_CAPTURE_OPTIONS=rtsp_transport;tcp|fflags;nobuffer|flags;low_delay + +# RKNN model — export'd from YOLO9t with imgsz=320 +MODEL_PATH=/opt/jetson-counter/yolo9t.rknn +IMGSZ=320 +HALF=false +CONF=0.3 + +# RKNN NPU core mask: 1=core0, 2=core1, 3=core0+core1, 7=all three +CORE_MASK=1 + +# YOLO decoder params (must match model export) +NUM_CLASSES=2 +NUM_KEYPOINTS=9 +REG_MAX=16 +STRIDES=8,16,32 + +CAMERA_NAME=CC1 +OBJECT_LABEL=ayam-potong +CLASS_AYAM=ayam +CLASS_TALENAN=talenan + +# Line crossing: rtl (default) | ltr | both +CROSS_DIRECTION=rtl +LINE_X= +LINE_X_FRAC=0.5 + +DAILY_CUTOFF_TIME=20:00 +CUTOFF_TIME=20:00 +BATCH_TIMEOUT_SECONDS=300 +IGNORE_BATCH_LABEL_TIMEOUT_SECONDS=30 +MIN_OBJECT_PER_BATCH=60 +MIN_DURATION_PER_BATCH=60 + +DASHBOARD_HOST=0.0.0.0 +DASHBOARD_PORT=5000 +SECRET_KEY=change-me-in-production + +EXPORT_CSV=true +RECORD_VIDEO=false + +LIVE_STREAM_ENABLED=false +LIVE_STREAM_FRAME_PATH=/dev/shm/jetson-counter/live_frame.jpg +LIVE_STREAM_QUALITY=75 +LIVE_STREAM_EVERY_N=2 diff --git a/counter_live_rknn_bytetrack.py b/counter_live_rknn_bytetrack.py new file mode 100644 index 0000000..1ad5e62 --- /dev/null +++ b/counter_live_rknn_bytetrack.py @@ -0,0 +1,1289 @@ +""" +Edge production live counter — RTSP + YOLO RKNN + ByteTrack + line crossing. +Runs on RK3588 hardware with RKNN model (320×320 input). +Uses ByteTrack (Kalman filter + two-stage IoU association) for tracking. +""" + +import numpy as np +import cv2 +import csv +import os +import signal +import time +from datetime import datetime +from pathlib import Path + +from dotenv import load_dotenv + +load_dotenv() + +from rknnlite.api import RKNNLite +from batch_store import BatchStore + +# --- config (override via env / .env) --- +OUTPUT_DIR = os.getenv("OUTPUT_DIR", "/opt/jetson-counter") +DB_PATH = os.getenv("DB_PATH", f"{OUTPUT_DIR}/jetson_counter.db") +STATE_FILE = os.getenv("STATE_FILE", f"{OUTPUT_DIR}/current_batch.json") +SOURCE = os.getenv("SOURCE", "rtsp://user:pass@192.168.0.100:554/stream1") +MODEL_PATH = os.getenv("MODEL_PATH", "/opt/jetson-counter/yolo11n.rknn") +CAMERA_NAME = os.getenv("CAMERA_NAME", "CC1") +OBJECT_LABEL = os.getenv("OBJECT_LABEL", "ayam-potong") +CLASS_AYAM = os.getenv("CLASS_AYAM", "ayam") +CLASS_TALENAN = os.getenv("CLASS_TALENAN", "talenan") + +LINE_X = int(os.getenv("LINE_X")) if os.getenv("LINE_X") else None +LINE_X_FRAC = float(os.getenv("LINE_X_FRAC", "0.5")) +CROSS_DIRECTION = os.getenv("CROSS_DIRECTION", "rtl").lower() + +IMGSZ = int(os.getenv("IMGSZ", "320")) +HALF = os.getenv("HALF", "false").lower() == "true" +CONF = float(os.getenv("CONF", "0.3")) +DEVICE = int(os.getenv("DEVICE", "0")) + +# RKNN NPU core mask +CORE_MASK = int(os.getenv("CORE_MASK", "1")) + +# YOLO decoder config +NUM_CLASSES = int(os.getenv("NUM_CLASSES", "2")) +SCORE_SIGMOID = os.getenv("SCORE_SIGMOID", "false").lower() == "true" + +# ByteTrack settings +TRACK_HIGH_THRESH = float(os.getenv("TRACK_HIGH_THRESH", "0.5")) +TRACK_LOW_THRESH = float(os.getenv("TRACK_LOW_THRESH", "0.1")) +TRACK_MATCH_THRESH = float(os.getenv("TRACK_MATCH_THRESH", "0.8")) +TRACK_BUFFER = int(os.getenv("TRACK_BUFFER", "30")) +TRACK_MIN_HITS = int(os.getenv("TRACK_MIN_HITS", "3")) + +DAILY_CUTOFF_TIME = os.getenv("DAILY_CUTOFF_TIME", "20:00") +BATCH_TIMEOUT_SECONDS = float(os.getenv("BATCH_TIMEOUT_SECONDS", "300")) +IGNORE_BATCH_LABEL_TIMEOUT = float( + os.getenv("IGNORE_BATCH_LABEL_TIMEOUT_SECONDS", "30") +) +MIN_OBJECT_PER_BATCH = int(os.getenv("MIN_OBJECT_PER_BATCH", "60")) +MIN_DURATION_PER_BATCH = int(os.getenv("MIN_DURATION_PER_BATCH", "60")) + +EXPORT_CSV = os.getenv("EXPORT_CSV", "true").lower() == "true" +CROSS_CSV = os.getenv("CROSS_CSV", f"{OUTPUT_DIR}/batch_crossings.csv") + +WARMUP_FRAMES = int(os.getenv("WARMUP_FRAMES", "30")) +RECONNECT_DELAY_SEC = int(os.getenv("RECONNECT_DELAY_SEC", "3")) +MAX_RECONNECT_ATTEMPTS = int(os.getenv("MAX_RECONNECT_ATTEMPTS", "0")) +FLUSH_EVERY_N_FRAMES = int(os.getenv("FLUSH_EVERY_N_FRAMES", "100")) +TRACKED_PRUNE_SEC = int(os.getenv("TRACKED_PRUNE_SEC", "300")) +RECORD_VIDEO = os.getenv("RECORD_VIDEO", "false").lower() == "true" +VIDEO_SEGMENT_SEC = int(os.getenv("VIDEO_SEGMENT_SEC", "3600")) +OUTPUT_FPS = int(os.getenv("OUTPUT_FPS", "15")) + +LIVE_STREAM_ENABLED = os.getenv("LIVE_STREAM_ENABLED", "false").lower() == "true" +LIVE_STREAM_FRAME_PATH = os.getenv( + "LIVE_STREAM_FRAME_PATH", "/dev/shm/jetson-counter/live_frame.jpg" +) +LIVE_STREAM_QUALITY = int(os.getenv("LIVE_STREAM_QUALITY", "75")) +LIVE_STREAM_EVERY_N = int(os.getenv("LIVE_STREAM_EVERY_N", "2")) + +RTSP_FFMPEG_OPTIONS = os.getenv( + "OPENCV_FFMPEG_CAPTURE_OPTIONS", + "rtsp_transport;tcp|fflags;nobuffer|flags;low_delay", +) + +IS_LIVE = SOURCE.lower().startswith(("rtsp://", "http://")) + +CROSS_FLASH_FRAMES = 12 +POPUP_LIFETIME = 20 +LINE_PULSE_FRAMES = 12 +COUNT_PULSE_FRAMES = 15 +BATCH_PULSE_FRAMES = 20 + +SKELETON = [(0, 1), (4, 3), (1, 2), (3, 2), (2, 6), (2, 5), (2, 7), (7, 8)] +SK_COLORS = [ + (0, 255, 255), + (0, 255, 255), + (255, 0, 255), + (255, 0, 255), + (0, 255, 0), + (255, 255, 0), + (0, 0, 255), + (200, 200, 0), +] + +C_PANEL = (28, 24, 18) +C_BORDER = (90, 85, 75) +C_ACCENT = (255, 200, 60) +C_GREEN = (80, 220, 100) +C_TEXT = (235, 235, 235) +C_MUTED = (150, 150, 150) +C_AYAM_BOX = (0, 165, 255) +C_TALENAN_BOX = (220, 120, 60) +C_LINE_CORE = (180, 220, 255) +C_LINE_GLOW = (100, 160, 220) + +shutdown_requested = False + + +def request_shutdown(signum, frame): + global shutdown_requested + shutdown_requested = True + print("\nShutdown requested — finishing current frame...") + + +signal.signal(signal.SIGINT, request_shutdown) +signal.signal(signal.SIGTERM, request_shutdown) + + +# ============================================================================= +# YOLO output decoder (NMS only — boxes are pre-decoded by the model) +# ============================================================================= + + +def _nms(boxes, scores, iou_thr=0.45): + order = np.argsort(scores)[::-1] + keep = [] + while len(order) > 0: + idx = order[0] + keep.append(idx) + if len(order) == 1: + break + xx1 = np.maximum(boxes[idx, 0], boxes[order[1:], 0]) + yy1 = np.maximum(boxes[idx, 1], boxes[order[1:], 1]) + xx2 = np.minimum(boxes[idx, 2], boxes[order[1:], 2]) + yy2 = np.minimum(boxes[idx, 3], boxes[order[1:], 3]) + w = np.maximum(0.0, xx2 - xx1) + h = np.maximum(0.0, yy2 - yy1) + inter = w * h + area_i = (boxes[idx, 2] - boxes[idx, 0]) * (boxes[idx, 3] - boxes[idx, 1]) + area_o = (boxes[order[1:], 2] - boxes[order[1:], 0]) * ( + boxes[order[1:], 3] - boxes[order[1:], 1] + ) + iou = inter / (area_i + area_o - inter + 1e-16) + order = order[1:][iou < iou_thr] + return np.array(keep) + + +# ============================================================================= +# IoU helpers (xyxy format) +# ============================================================================= + + +def _ious_xyxy(boxes_a, boxes_b): + """Pairwise IoU: (N,4) vs (M,4) → (N,M) matrix.""" + n, m = len(boxes_a), len(boxes_b) + if n == 0 or m == 0: + return np.zeros((n, m), dtype=np.float32) + xx1 = np.maximum(boxes_a[:, None, 0], boxes_b[None, :, 0]) + yy1 = np.maximum(boxes_a[:, None, 1], boxes_b[None, :, 1]) + xx2 = np.minimum(boxes_a[:, None, 2], boxes_b[None, :, 2]) + yy2 = np.minimum(boxes_a[:, None, 3], boxes_b[None, :, 3]) + iw = np.maximum(0.0, xx2 - xx1) + ih = np.maximum(0.0, yy2 - yy1) + inter = iw * ih + area_a = (boxes_a[:, 2] - boxes_a[:, 0]) * (boxes_a[:, 3] - boxes_a[:, 1]) + area_b = (boxes_b[:, 2] - boxes_b[:, 0]) * (boxes_b[:, 3] - boxes_b[:, 1]) + return inter / (area_a[:, None] + area_b[None, :] - inter + 1e-16) + + +def _greedy_match(cost_matrix, threshold=0.3): + """Greedy linear assignment. Returns pairs (row_idx, col_idx).""" + if cost_matrix.size == 0: + return [] + n, m = cost_matrix.shape + flat = [(cost_matrix[i, j], i, j) for i in range(n) for j in range(m)] + flat.sort() + row_used = set() + col_used = set() + pairs = [] + for cost, i, j in flat: + if cost >= threshold: + break + if i in row_used or j in col_used: + continue + row_used.add(i) + col_used.add(j) + pairs.append((i, j)) + return pairs + + +# ============================================================================= +# Kalman filter box tracker (state: x, y, w, h, vx, vy, vw, vh) +# ============================================================================= + + +class KalmanBoxTracker: + count = 0 + + def __init__(self, bbox_xyxy): + KalmanBoxTracker.count += 1 + self.track_id = KalmanBoxTracker.count + + x1, y1, x2, y2 = bbox_xyxy + w, h = x2 - x1, y2 - y1 + x, y = x1 + w / 2, y1 + h / 2 + + self.kf = _KalmanFilter() + self.kf.x[:4, 0] = np.array([x, y, w, h], dtype=np.float32) + + self.time_since_update = 0 + self.hits = 1 + self.hit_streak = 1 + self.age = 1 + + def predict(self): + if self.kf.x[6] + self.kf.x[2] <= 0: + self.kf.x[6] *= 0.0 + self.kf.predict() + self.age += 1 + self.time_since_update += 1 + + def update(self, bbox_xyxy): + self.time_since_update = 0 + self.hits += 1 + self.hit_streak += 1 + + x1, y1, x2, y2 = bbox_xyxy + w, h = x2 - x1, y2 - y1 + x, y = x1 + w / 2, y1 + h / 2 + self.kf.update(np.array([x, y, w, h], dtype=np.float32)) + + def get_state(self): + """Returns xyxy bbox from Kalman state.""" + xx = self.kf.x[:4, 0] + x, y, w, h = xx[0], xx[1], xx[2], xx[3] + x1 = x - w / 2 + y1 = y - h / 2 + x2 = x + w / 2 + y2 = y + h / 2 + return np.array([x1, y1, x2, y2], dtype=np.float32) + + def get_cx(self): + return float(self.kf.x[0, 0]) + + +class _KalmanFilter: + """8-state constant-velocity Kalman filter for bounding box tracking.""" + + def __init__(self): + ndim, dt = 4, 1.0 + self.motion_mat = np.eye(2 * ndim, 2 * ndim, dtype=np.float32) + for i in range(ndim): + self.motion_mat[i, ndim + i] = dt + self.update_mat = np.eye(ndim, 2 * ndim, dtype=np.float32) + + self._std_weight_position = 1.0 / 20 + self._std_weight_velocity = 1.0 / 160 + + self.x = np.zeros((8, 1), dtype=np.float32) + self.P = np.eye(8, dtype=np.float32) * 10.0 + + def predict(self): + std_pos = [ + self._std_weight_position * self.x[2], + self._std_weight_position * self.x[3], + self._std_weight_position * self.x[2], + self._std_weight_position * self.x[3], + ] + std_vel = [ + self._std_weight_velocity * self.x[2], + self._std_weight_velocity * self.x[3], + self._std_weight_velocity * self.x[2], + self._std_weight_velocity * self.x[3], + ] + Q = np.diag(np.square(np.concatenate([std_pos, std_vel]))) + + self.x = self.motion_mat @ self.x + self.P = self.motion_mat @ self.P @ self.motion_mat.T + Q + + def update(self, z): + R = np.diag( + np.square( + [ + self._std_weight_position * z[2], + self._std_weight_position * z[3], + self._std_weight_position * z[2], + self._std_weight_position * z[3], + ] + ) + ) + + H = self.update_mat + S = H @ self.P @ H.T + R + K = self.P @ H.T @ np.linalg.inv(S) + + y = z.reshape(4, 1) - H @ self.x + self.x = self.x + K @ y + I_KH = np.eye(8) - K @ H + self.P = I_KH @ self.P @ I_KH.T + K @ R @ K.T + + +# ============================================================================= +# ByteTrack multi-object tracker +# ============================================================================= + + +class ByteTracker: + """ByteTrack: two-stage association with Kalman filter prediction.""" + + def __init__( + self, + track_high_thresh=0.5, + track_low_thresh=0.1, + match_thresh=0.8, + track_buffer=30, + min_hits=3, + ): + self.high_thresh = track_high_thresh + self.low_thresh = track_low_thresh + self.match_thresh = match_thresh + self.track_buffer = track_buffer + self.min_hits = min_hits + self.tracked_tracks = [] + self.lost_tracks = [] + self.removed_tracks = [] + self.frame_id = 0 + + def update(self, boxes_xyxy, scores): + self.frame_id += 1 + + # --- separate detections by score --- + if len(boxes_xyxy) > 0: + remain = scores > self.low_thresh + dets = boxes_xyxy[remain] + det_scores = scores[remain] + is_high = det_scores > self.high_thresh + is_low = ~is_high + else: + dets = np.zeros((0, 4), dtype=np.float32) + det_scores = np.zeros(0, dtype=np.float32) + is_high = np.zeros(0, dtype=bool) + is_low = np.zeros(0, dtype=bool) + + # --- Kalman predict all existing tracks --- + track_pool = self.tracked_tracks + self.lost_tracks + num_tracks = len(track_pool) + + # Per-frame tracking results + matched_track_idx = set() + det_to_track = {} + tracked_map = {} + lost_map = {} + + # Pre-allocate these for scoping + high_idx = np.array([], dtype=np.int64) + low_idx = np.array([], dtype=np.int64) + match_pairs_high = [] + + if num_tracks > 0: + track_boxes = np.zeros((num_tracks, 4), dtype=np.float32) + for ti, trk in enumerate(track_pool): + trk.predict() + track_boxes[ti] = trk.get_state() + + # --- first association: high-score ↔ all tracks --- + high_idx = np.where(is_high)[0] + high_dets = dets[is_high] + + unmatched_tracks = list(range(num_tracks)) + + if len(high_dets) > 0: + iou_mat = _ious_xyxy(high_dets, track_boxes) + cost_mat = 1.0 - iou_mat + matches = _greedy_match(cost_mat, threshold=1.0 - self.match_thresh) + + for di, ti in matches: + det_global = int(high_idx[di]) + track_pool[ti].update(dets[det_global]) + track_pool[ti].hit_streak = max(1, track_pool[ti].hit_streak) + matched_track_idx.add(ti) + det_to_track[det_global] = track_pool[ti].track_id + tracked_map[track_pool[ti].track_id] = track_pool[ti].get_cx() + match_pairs_high.append((det_global, ti)) + + unmatched_tracks = [ + t for t in range(num_tracks) if t not in matched_track_idx + ] + + # --- second association: low-score ↔ unmatched tracks --- + low_idx = np.where(is_low)[0] + low_dets = dets[is_low] + + if len(low_dets) > 0 and len(unmatched_tracks) > 0: + unmatched_boxes = track_boxes[unmatched_tracks] + iou_mat = _ious_xyxy(low_dets, unmatched_boxes) + cost_mat = 1.0 - iou_mat + matches2 = _greedy_match(cost_mat, threshold=0.5) + + for di, uti in matches2: + det_global = int(low_idx[di]) + pool_idx = unmatched_tracks[uti] + track_pool[pool_idx].update(dets[det_global]) + track_pool[pool_idx].hit_streak = max( + 1, track_pool[pool_idx].hit_streak + ) + matched_track_idx.add(pool_idx) + det_to_track[det_global] = track_pool[pool_idx].track_id + tracked_map[track_pool[pool_idx].track_id] = track_pool[ + pool_idx + ].get_cx() + + # --- reset hit_streak for unmatched tracks --- + for ti, trk in enumerate(track_pool): + if ti not in matched_track_idx: + trk.hit_streak = 0 + + # --- lifecycle management --- + new_tracked = [] + new_lost = [] + for trk in track_pool: + if trk.time_since_update > self.track_buffer: + self.removed_tracks.append(trk) + elif trk.time_since_update > 0: + new_lost.append(trk) + else: + new_tracked.append(trk) + self.tracked_tracks = new_tracked + self.lost_tracks = new_lost + + # --- confirmed tracks (both tracked and lost) --- + for trk in self.tracked_tracks + self.lost_tracks: + if trk.hit_streak >= self.min_hits or trk.hits >= self.min_hits: + tracked_map.setdefault(trk.track_id, trk.get_cx()) + + for trk in self.lost_tracks: + if trk.hit_streak >= self.min_hits or trk.hits >= self.min_hits: + lost_map[trk.track_id] = trk.get_cx() + + # --- new tracks from unmatched high-score dets --- + high_all = np.where(is_high)[0] + matched_det_ids = set(det_to_track.keys()) + for dg in high_all: + if int(dg) not in matched_det_ids: + trk = KalmanBoxTracker(dets[dg]) + self.tracked_tracks.append(trk) + det_to_track[int(dg)] = trk.track_id + tracked_map[trk.track_id] = trk.get_cx() + + return tracked_map, det_to_track, lost_map + + +# ============================================================================= +# RKNN YOLO wrapper (detect output format: (1, 4+num_classes, N)) +# ============================================================================= + + +class RKNNYOLO: + def __init__( + self, + model_path, + core_mask=1, + imgsz=320, + conf=0.3, + iou=0.45, + num_classes=2, + num_keypoints=0, + score_sigmoid=False, + ): + self.imgsz = imgsz + self.conf = conf + self.iou = iou + self.num_classes = num_classes + self.num_keypoints = num_keypoints + self.score_sigmoid = score_sigmoid + + self.rknn = RKNNLite(verbose=False) + ret = self.rknn.load_rknn(model_path) + if ret != 0: + raise RuntimeError(f"Failed to load RKNN model: {model_path}") + ret = self.rknn.init_runtime(core_mask=core_mask) + if ret != 0: + raise RuntimeError(f"Failed to init RKNN runtime (core_mask={core_mask})") + + try: + sdk_ver = self.rknn.get_sdk_version() + print(f"RKNN SDK version: {sdk_ver}") + except Exception: + pass + + print(f"RKNN model loaded: {model_path} imgsz={imgsz} core_mask={core_mask}") + + def _preprocess(self, frame): + h0, w0 = frame.shape[:2] + scale = min(self.imgsz / h0, self.imgsz / w0) + nh, nw = int(h0 * scale), int(w0 * scale) + resized = cv2.resize(frame, (nw, nh), interpolation=cv2.INTER_LINEAR) + + letterbox = np.full((self.imgsz, self.imgsz, 3), 114, dtype=np.uint8) + dy = (self.imgsz - nh) // 2 + dx = (self.imgsz - nw) // 2 + letterbox[dy : dy + nh, dx : dx + nw] = resized + + rgb = cv2.cvtColor(letterbox, cv2.COLOR_BGR2RGB) + gains = np.array([scale, scale, dy, dx], dtype=np.float32) + return rgb, gains + + def __call__(self, frame): + h0, w0 = frame.shape[:2] + rgb, gains = self._preprocess(frame) + scale, _, pad_y, pad_x = gains + + inp = np.expand_dims(rgb, axis=0) + inp = np.ascontiguousarray(inp.astype(np.uint8)) + + outputs = self.rknn.inference(inputs=[inp]) + + if len(outputs) == 0: + return [] + + out = outputs[0] + out = np.squeeze(out, axis=0) + + if out.shape[0] == self.num_classes + 4: + out = out.T + + boxes_cxcywh = out[:, :4].copy() + cls_raw = out[:, 4:].copy() + + if self.score_sigmoid: + cls_scores = 1.0 / (1.0 + np.exp(-np.clip(cls_raw, -10, 10))) + else: + cls_scores = cls_raw + + boxes_xyxy = np.stack( + [ + boxes_cxcywh[:, 0] - boxes_cxcywh[:, 2] / 2, + boxes_cxcywh[:, 1] - boxes_cxcywh[:, 3] / 2, + boxes_cxcywh[:, 0] + boxes_cxcywh[:, 2] / 2, + boxes_cxcywh[:, 1] + boxes_cxcywh[:, 3] / 2, + ], + axis=1, + ) + + max_scores = cls_scores.max(axis=1) + class_ids = cls_scores.argmax(axis=1) + mask = max_scores > self.conf + + if mask.sum() == 0: + return [] + + bboxes = boxes_xyxy[mask].astype(np.float32) + scores = max_scores[mask].astype(np.float32) + clses = class_ids[mask] + + bboxes[:, 0] = (bboxes[:, 0] - pad_x) / scale + bboxes[:, 1] = (bboxes[:, 1] - pad_y) / scale + bboxes[:, 2] = (bboxes[:, 2] - pad_x) / scale + bboxes[:, 3] = (bboxes[:, 3] - pad_y) / scale + + bboxes[:, 0] = np.clip(bboxes[:, 0], 0, w0) + bboxes[:, 1] = np.clip(bboxes[:, 1], 0, h0) + bboxes[:, 2] = np.clip(bboxes[:, 2], 0, w0) + bboxes[:, 3] = np.clip(bboxes[:, 3], 0, h0) + + detections = [] + for cls_id in range(self.num_classes): + idx = np.where(clses == cls_id)[0] + if len(idx) == 0: + continue + keep = _nms(bboxes[idx], scores[idx], iou_thr=self.iou) + for k in keep: + j = idx[k] + detections.append( + { + "bbox": bboxes[j].tolist(), + "score": float(scores[j]), + "cls": int(clses[j]), + "keypoints": None, + } + ) + return detections + + def release(self): + self.rknn.release() + + +# ============================================================================= +# Drawing helpers +# ============================================================================= + + +def resolve_line_x(frame_width): + if LINE_X is not None: + return LINE_X + if LINE_X_FRAC != 0.5: + return int(frame_width * LINE_X_FRAC) + return frame_width // 2 + + +def crossed_line(prev_cx, cx, line_x, direction=CROSS_DIRECTION): + if direction == "ltr": + return prev_cx < line_x <= cx + if direction == "both": + return (prev_cx > line_x >= cx) or (prev_cx < line_x <= cx) + return prev_cx > line_x >= cx + + +def now_str(): + return datetime.now().strftime("%Y-%m-%d %H:%M:%S") + + +def open_capture(source): + if source.lower().startswith(("rtsp://", "http://")): + os.environ["OPENCV_FFMPEG_CAPTURE_OPTIONS"] = RTSP_FFMPEG_OPTIONS + cap = cv2.VideoCapture(source, cv2.CAP_FFMPEG) + cap.set(cv2.CAP_PROP_BUFFERSIZE, 1) + return cap + + +def warmup_stream(cap, n=WARMUP_FRAMES): + print("Warming up stream...") + for _ in range(n): + cap.read() + print("Stream ready!") + + +def open_video_writer(path, w, h, fps): + return cv2.VideoWriter(path, cv2.VideoWriter_fourcc(*"avc1"), fps, (w, h)) + + +class CsvLogger: + def __init__(self, path, header): + Path(path).parent.mkdir(parents=True, exist_ok=True) + new_file = not Path(path).exists() or Path(path).stat().st_size == 0 + self.file = open(path, "a", newline="", buffering=1) + self.writer = csv.writer(self.file) + if new_file: + self.writer.writerow(header) + self.file.flush() + + def write_row(self, row): + self.writer.writerow(row) + self.file.flush() + + def close(self): + self.file.close() + + +class VideoSegmentWriter: + def __init__(self, output_dir, w, h, fps, segment_sec): + self.output_dir = Path(output_dir) + self.output_dir.mkdir(parents=True, exist_ok=True) + self.w, self.h, self.fps = w, h, fps + self.segment_sec = segment_sec + self.segment_start = time.monotonic() + self.writer = None + self._open_next() + + def _segment_path(self): + ts = datetime.now().strftime("%Y%m%d_%H%M%S") + return str(self.output_dir / f"live_{ts}.mp4") + + def _open_next(self): + if self.writer is not None: + self.writer.release() + path = self._segment_path() + self.writer = open_video_writer(path, self.w, self.h, self.fps) + self.segment_start = time.monotonic() + print(f"Recording segment: {path}") + + def write(self, frame): + if time.monotonic() - self.segment_start >= self.segment_sec: + self._open_next() + self.writer.write(frame) + + def release(self): + if self.writer is not None: + self.writer.release() + + +def prune_stale_tracks(tracked, now_mono): + stale = [ + tid for tid, (_, ts) in tracked.items() if now_mono - ts > TRACKED_PRUNE_SEC + ] + for tid in stale: + del tracked[tid] + + +def overlay_rect(img, x1, y1, x2, y2, color, alpha=0.65): + x1, y1 = max(0, x1), max(0, y1) + x2, y2 = min(img.shape[1], x2), min(img.shape[0], y2) + if x2 <= x1 or y2 <= y1: + return + roi = img[y1:y2, x1:x2] + patch = np.full_like(roi, color, dtype=np.uint8) + cv2.addWeighted(patch, alpha, roi, 1 - alpha, 0, roi) + + +def draw_pill(img, text, x, y, bg, fg=C_TEXT, font_scale=0.45, pad_x=6, pad_y=4): + font = cv2.FONT_HERSHEY_SIMPLEX + (tw, th), baseline = cv2.getTextSize(text, font, font_scale, 1) + x1, y1 = x, y - th - pad_y + x2, y2 = x + tw + pad_x * 2, y + baseline + pad_y + cv2.rectangle(img, (x1, y1), (x2, y2), bg, -1) + cv2.rectangle(img, (x1, y1), (x2, y2), C_BORDER, 1) + cv2.putText(img, text, (x + pad_x, y), font, font_scale, fg, 1, cv2.LINE_AA) + + +def draw_elegant_counting_line(img, line_x, h, pulse_remaining=0): + strength = pulse_remaining / max(LINE_PULSE_FRAMES, 1) + glow_alpha = 0.12 + 0.18 * strength + for offset in (14, 9, 5): + color = tuple(int(c * glow_alpha) for c in C_LINE_GLOW) + cv2.line(img, (line_x - offset, 0), (line_x - offset, h), color, 1, cv2.LINE_AA) + cv2.line(img, (line_x + offset, 0), (line_x + offset, h), color, 1, cv2.LINE_AA) + dash_len, gap = 18, 12 + y = 0 + while y < h: + y_end = min(y + dash_len, h) + cv2.line(img, (line_x, y), (line_x, y_end), C_LINE_CORE, 2, cv2.LINE_AA) + y += dash_len + gap + cv2.putText( + img, + "COUNT LINE", + (line_x - 46, 24), + cv2.FONT_HERSHEY_SIMPLEX, + 0.42, + C_LINE_CORE, + 1, + cv2.LINE_AA, + ) + + +def draw_hero_count(img, line_x, h, count, pulse_remaining=0): + text = str(count) + font = cv2.FONT_HERSHEY_SIMPLEX + boost = 0.35 * (pulse_remaining / max(COUNT_PULSE_FRAMES, 1)) + font_scale, thickness = 1.6 + boost, 3 + (tw, th), _ = cv2.getTextSize(text, font, font_scale, thickness) + pad = 14 + tx, ty = line_x - tw // 2, h // 2 + th // 2 + overlay_rect( + img, tx - pad, ty - th - pad, tx + tw + pad, ty + pad // 2, C_PANEL, alpha=0.78 + ) + cv2.rectangle( + img, (tx - pad, ty - th - pad), (tx + tw + pad, ty + pad // 2), C_LINE_CORE, 2 + ) + cv2.putText(img, text, (tx, ty), font, font_scale, C_GREEN, thickness, cv2.LINE_AA) + + +def draw_hud( + img, w, batch_num, batch_count, total_ayam, elapsed_sec, rate, camera_id, clock +): + bar_h = 52 + overlay_rect(img, 0, 0, w, bar_h, C_PANEL, alpha=0.72) + cv2.line(img, (0, bar_h), (w, bar_h), C_BORDER, 1) + cv2.putText( + img, "BATCH", (16, 20), cv2.FONT_HERSHEY_SIMPLEX, 0.45, C_MUTED, 1, cv2.LINE_AA + ) + batch_label = str(batch_num) if batch_num else "\u2014" + cv2.putText( + img, + batch_label, + (16, 44), + cv2.FONT_HERSHEY_SIMPLEX, + 0.9, + C_ACCENT, + 2, + cv2.LINE_AA, + ) + cv2.putText( + img, "COUNT", (100, 20), cv2.FONT_HERSHEY_SIMPLEX, 0.45, C_MUTED, 1, cv2.LINE_AA + ) + cv2.putText( + img, + str(batch_count), + (100, 44), + cv2.FONT_HERSHEY_SIMPLEX, + 0.9, + C_GREEN, + 2, + cv2.LINE_AA, + ) + cv2.putText( + img, "TOTAL", (190, 20), cv2.FONT_HERSHEY_SIMPLEX, 0.45, C_MUTED, 1, cv2.LINE_AA + ) + cv2.putText( + img, + str(total_ayam), + (190, 44), + cv2.FONT_HERSHEY_SIMPLEX, + 0.7, + C_TEXT, + 1, + cv2.LINE_AA, + ) + cv2.putText( + img, + "UPTIME", + (280, 20), + cv2.FONT_HERSHEY_SIMPLEX, + 0.45, + C_MUTED, + 1, + cv2.LINE_AA, + ) + cv2.putText( + img, + f"{elapsed_sec / 3600:.1f}h", + (280, 44), + cv2.FONT_HERSHEY_SIMPLEX, + 0.7, + C_TEXT, + 1, + cv2.LINE_AA, + ) + cv2.putText( + img, "RATE", (380, 20), cv2.FONT_HERSHEY_SIMPLEX, 0.45, C_MUTED, 1, cv2.LINE_AA + ) + cv2.putText( + img, + f"{rate:.1f}/min", + (380, 44), + cv2.FONT_HERSHEY_SIMPLEX, + 0.7, + C_ACCENT, + 1, + cv2.LINE_AA, + ) + # cv2.putText(img, clock, (w - 180, 36), cv2.FONT_HERSHEY_SIMPLEX, 0.55, C_TEXT, 1, cv2.LINE_AA) + cv2.putText( + img, + f"CAM {camera_id}", + (w - 180, 20), + cv2.FONT_HERSHEY_SIMPLEX, + 0.45, + C_MUTED, + 1, + cv2.LINE_AA, + ) + + +def draw_footer(img, w, h, frame_idx, live_tag): + bar_h = 28 + overlay_rect(img, 0, h - bar_h, w, h, C_PANEL, alpha=0.55) + cv2.putText( + img, + f"{live_tag} | Frame {frame_idx}", + (12, h - 9), + cv2.FONT_HERSHEY_SIMPLEX, + 0.45, + C_MUTED, + 1, + cv2.LINE_AA, + ) + + +def draw_skeleton_bold(img, kpts): + for (a, b), color in zip(SKELETON, SK_COLORS): + if a < len(kpts) and b < len(kpts): + xa, ya = int(kpts[a][0]), int(kpts[a][1]) + xb, yb = int(kpts[b][0]), int(kpts[b][1]) + if xa > 0 and ya > 0 and xb > 0 and yb > 0: + cv2.line(img, (xa, ya), (xb, yb), color, 3, cv2.LINE_AA) + for kp in kpts: + x, y = int(kp[0]), int(kp[1]) + if x > 0 and y > 0: + cv2.circle(img, (x, y), 6, (255, 255, 255), -1, cv2.LINE_AA) + cv2.circle(img, (x, y), 6, (40, 40, 40), 2, cv2.LINE_AA) + + +def draw_popups(img, popups, frame_idx): + alive = [] + for pop in popups: + age = frame_idx - pop["born"] + if age > POPUP_LIFETIME: + continue + alive.append(pop) + fade = 1.0 - age / POPUP_LIFETIME + y = pop["y"] - int(age * 1.8) + color = (int(C_GREEN[0] * fade), int(C_GREEN[1] * fade), int(C_GREEN[2] * fade)) + cv2.putText( + img, + pop["text"], + (pop["x"], y), + cv2.FONT_HERSHEY_SIMPLEX, + 0.7, + color, + 2, + cv2.LINE_AA, + ) + return alive + + +def draw_batch_banner(img, w, batch_num, pulse_remaining): + if pulse_remaining <= 0: + return + text = f"NEW BATCH {batch_num}" + font = cv2.FONT_HERSHEY_SIMPLEX + (tw, th), _ = cv2.getTextSize(text, font, 0.8, 2) + x1, y1 = w // 2 - tw // 2 - 16, 62 + x2, y2 = w // 2 + tw // 2 + 16, 62 + th + 20 + overlay_rect(img, x1, y1, x2, y2, C_PANEL, alpha=0.7) + cv2.rectangle(img, (x1, y1), (x2, y2), C_ACCENT, 2) + cv2.putText( + img, text, (w // 2 - tw // 2, 62 + th + 4), font, 0.8, C_ACCENT, 2, cv2.LINE_AA + ) + + +def connect_stream(source, warmup=WARMUP_FRAMES): + attempts = 0 + while not shutdown_requested: + cap = open_capture(source) + if not cap.isOpened(): + attempts += 1 + if MAX_RECONNECT_ATTEMPTS and attempts >= MAX_RECONNECT_ATTEMPTS: + raise RuntimeError( + f"Cannot open source after {attempts} attempts: {source}" + ) + print(f"Cannot open source, retry in {RECONNECT_DELAY_SEC}s...") + time.sleep(RECONNECT_DELAY_SEC) + continue + if warmup > 0 and source.lower().startswith(("rtsp://", "http://")): + warmup_stream(cap, warmup) + w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + fps = cap.get(cv2.CAP_PROP_FPS) + if not fps or fps <= 1: + fps = OUTPUT_FPS + return cap, w, h, fps + return None, 0, 0, OUTPUT_FPS + + +# ============================================================================= +# Main loop +# ============================================================================= + + +def run(): + global shutdown_requested + + store = BatchStore( + db_path=DB_PATH, + state_file=STATE_FILE, + camera_name=CAMERA_NAME, + object_label=OBJECT_LABEL, + cutoff_time=DAILY_CUTOFF_TIME, + batch_timeout=BATCH_TIMEOUT_SECONDS, + ignore_batch_label_timeout=IGNORE_BATCH_LABEL_TIMEOUT, + min_object_per_batch=MIN_OBJECT_PER_BATCH, + min_duration_per_batch=MIN_DURATION_PER_BATCH, + logger=lambda msg: print(f"[{now_str()}] {msg}"), + ) + store.start_cutoff_watcher() + + cross_logger = None + if EXPORT_CSV: + cross_logger = CsvLogger( + CROSS_CSV, ["batch", "frame", "timestamp", "chicken_id"] + ) + + model = RKNNYOLO( + model_path=MODEL_PATH, + core_mask=CORE_MASK, + imgsz=IMGSZ, + conf=CONF, + num_classes=NUM_CLASSES, + score_sigmoid=SCORE_SIGMOID, + ) + + CLASS_IDS = { + os.getenv("CLASS_AYAM", "ayam"): 0, + os.getenv("CLASS_TALENAN", "talenan"): 1, + } + ayam_cls = CLASS_IDS[CLASS_AYAM] + talenan_cls = CLASS_IDS[CLASS_TALENAN] + + ayam_tracker = ByteTracker( + track_high_thresh=TRACK_HIGH_THRESH, + track_low_thresh=TRACK_LOW_THRESH, + match_thresh=TRACK_MATCH_THRESH, + track_buffer=TRACK_BUFFER, + min_hits=TRACK_MIN_HITS, + ) + talenan_tracker = ByteTracker( + track_high_thresh=TRACK_HIGH_THRESH, + track_low_thresh=TRACK_LOW_THRESH, + match_thresh=TRACK_MATCH_THRESH, + track_buffer=TRACK_BUFFER, + min_hits=TRACK_MIN_HITS, + ) + + ayam_line_crossed = set() + talenan_line_crossed = set() + + ayam_cross_flash = {} + talenan_cross_flash = {} + line_pulse = count_pulse = batch_pulse = 0 + popups = [] + + session_start = time.time() + frame_idx = 0 + video_writer = None + + cap, w, h, fps = connect_stream(SOURCE) + if cap is None: + store.shutdown() + model.release() + return + + line_x = resolve_line_x(w) + print( + f"RKNN+ByteTrack counter | {w}x{h} @ {fps}fps | line x={line_x} | cross={CROSS_DIRECTION}" + ) + print(f"Model: {MODEL_PATH} | imgsz={IMGSZ} | core_mask={CORE_MASK}") + print( + f"ByteTrack: high_thresh={TRACK_HIGH_THRESH} low_thresh={TRACK_LOW_THRESH} " + f"match_thresh={TRACK_MATCH_THRESH} buffer={TRACK_BUFFER}" + ) + print(f"DB: {DB_PATH}") + print(f"State: {STATE_FILE}") + + if RECORD_VIDEO: + video_writer = VideoSegmentWriter(OUTPUT_DIR, w, h, fps, VIDEO_SEGMENT_SEC) + + reconnect_count = 0 + + while not shutdown_requested: + ret, frame = cap.read() + if not ret: + if not IS_LIVE: + break + reconnect_count += 1 + print( + f"Stream dropped (attempt {reconnect_count}), reconnecting in {RECONNECT_DELAY_SEC}s..." + ) + cap.release() + time.sleep(RECONNECT_DELAY_SEC) + cap, w, h, fps = connect_stream(SOURCE) + if cap is None: + break + line_x = resolve_line_x(w) + continue + + now = time.time() + elapsed = now - session_start + mono = time.monotonic() + ayam_crossed_frame = batch_closed_frame = batch_started_frame = False + + detections = model(frame) + + if detections: + ayam_boxes_xyxy = [] + ayam_scores = [] + ayam_kpts_list = [] + ayam_cx_list = [] + + talenan_boxes_xyxy = [] + talenan_scores = [] + talenan_kpts_list = [] + talenan_cx_list = [] + + for det in detections: + bbox = det["bbox"] + score = det["score"] + cls_id = det["cls"] + kpts = det["keypoints"] + cx = (bbox[0] + bbox[2]) / 2.0 + + if cls_id == talenan_cls: + talenan_boxes_xyxy.append(bbox) + talenan_scores.append(score) + talenan_kpts_list.append(kpts) + talenan_cx_list.append(cx) + elif cls_id == ayam_cls: + ayam_boxes_xyxy.append(bbox) + ayam_scores.append(score) + ayam_kpts_list.append(kpts) + ayam_cx_list.append(cx) + + ayam_boxes_xyxy = np.array(ayam_boxes_xyxy, dtype=np.float32) + ayam_scores = np.array(ayam_scores, dtype=np.float32) + talenan_boxes_xyxy = np.array(talenan_boxes_xyxy, dtype=np.float32) + talenan_scores = np.array(talenan_scores, dtype=np.float32) + + ayam_track_map, ayam_det_to_track, ayam_lost_map = ayam_tracker.update( + ayam_boxes_xyxy, ayam_scores + ) + talenan_track_map, talenan_det_to_track, talenan_lost_map = ( + talenan_tracker.update(talenan_boxes_xyxy, talenan_scores) + ) + + # Process talenan crossings + for di in range(len(talenan_boxes_xyxy)): + tid = talenan_det_to_track.get(di) + if tid is None: + continue + cx = talenan_cx_list[di] + bbox = talenan_boxes_xyxy[di] + + if tid in talenan_tracked: + prev_cx = talenan_tracked[tid][0] + if ( + crossed_line(prev_cx, cx, line_x) + and tid not in talenan_line_crossed + ): + talenan_line_crossed.add(tid) + if store.record_talenan_crossing(tid): + batch_closed_frame = True + talenan_cross_flash[tid] = CROSS_FLASH_FRAMES + popups.append( + { + "x": int(cx) - 20, + "y": int((bbox[1] + bbox[3]) / 2), + "born": frame_idx, + "text": "BATCH CLOSED", + } + ) + talenan_tracked[tid] = (cx, mono) + + # Process ayam crossings (including lost tracks for line-cross continuity) + for di in range(len(ayam_boxes_xyxy)): + tid = ayam_det_to_track.get(di) + if tid is None: + continue + cx = ayam_cx_list[di] + + if tid in ayam_tracked: + prev_cx = ayam_tracked[tid][0] + if ( + crossed_line(prev_cx, cx, line_x) + and tid not in ayam_line_crossed + ): + ayam_line_crossed.add(tid) + _, started_new = store.record_ayam_crossing(tid) + if cross_logger: + cross_logger.write_row( + [ + store.current_batch_number, + frame_idx, + datetime.now().isoformat(), + tid, + ] + ) + ayam_crossed_frame = True + if started_new: + batch_started_frame = True + ayam_cross_flash[tid] = CROSS_FLASH_FRAMES + popups.append( + { + "x": int(cx) - 12, + "y": int( + (ayam_boxes_xyxy[di][1] + ayam_boxes_xyxy[di][3]) + / 2 + ), + "born": frame_idx, + "text": "+1", + } + ) + ayam_tracked[tid] = (cx, mono) + + # Also track lost tracks for line-crossing continuity + for tid, cx in ayam_lost_map.items(): + if tid not in ayam_tracked: + ayam_tracked[tid] = (cx, mono) + + # Draw talenan + for di in range(len(talenan_boxes_xyxy)): + tid = talenan_det_to_track.get(di) + if tid is None: + continue + bbox = talenan_boxes_xyxy[di] + x1, y1, x2, y2 = int(bbox[0]), int(bbox[1]), int(bbox[2]), int(bbox[3]) + flash = talenan_cross_flash.get(tid, 0) + color = C_GREEN if flash > 0 else C_TALENAN_BOX + cv2.rectangle(frame, (x1, y1), (x2, y2), color, 3 if flash > 0 else 2) + draw_pill(frame, f"TALENAN {tid}", x1, y1 - 4, color) + + # Draw ayam + for di in range(len(ayam_boxes_xyxy)): + tid = ayam_det_to_track.get(di) + if tid is None: + continue + bbox = ayam_boxes_xyxy[di] + x1, y1, x2, y2 = int(bbox[0]), int(bbox[1]), int(bbox[2]), int(bbox[3]) + flash = ayam_cross_flash.get(tid, 0) + color = C_GREEN if flash > 0 else C_AYAM_BOX + cv2.rectangle(frame, (x1, y1), (x2, y2), color, 3 if flash > 0 else 2) + draw_pill(frame, f"ID {tid}", x1, y1 - 4, color) + kpts = ayam_kpts_list[di] if di < len(ayam_kpts_list) else None + if kpts is not None: + draw_skeleton_bold(frame, kpts) + + if ayam_crossed_frame: + line_pulse = LINE_PULSE_FRAMES + count_pulse = COUNT_PULSE_FRAMES + if batch_closed_frame: + line_pulse = LINE_PULSE_FRAMES + if batch_started_frame: + batch_pulse = BATCH_PULSE_FRAMES + + batch_num = store.current_batch_number or 0 + batch_count = store.current_batch_count + display_total = store.display_total() + rate = (display_total / elapsed * 60) if elapsed > 0 else 0.0 + + draw_elegant_counting_line(frame, line_x, h, line_pulse) + draw_hero_count(frame, line_x, h, batch_count, count_pulse) + draw_hud( + frame, + w, + batch_num, + batch_count, + display_total, + elapsed, + rate, + CAMERA_NAME, + now_str(), + ) + draw_batch_banner(frame, w, batch_num, batch_pulse) + draw_footer( + frame, w, h, frame_idx, "LIVE-RKNN-BT" if IS_LIVE else "FILE-RKNN-BT" + ) + popups = draw_popups(frame, popups, frame_idx) + + for flash_store in (ayam_cross_flash, talenan_cross_flash): + for tid in list(flash_store): + flash_store[tid] -= 1 + if flash_store[tid] <= 0: + del flash_store[tid] + line_pulse = max(0, line_pulse - 1) + count_pulse = max(0, count_pulse - 1) + batch_pulse = max(0, batch_pulse - 1) + + if video_writer is not None: + video_writer.write(frame) + + if LIVE_STREAM_ENABLED and frame_idx % LIVE_STREAM_EVERY_N == 0: + try: + Path(LIVE_STREAM_FRAME_PATH).parent.mkdir(parents=True, exist_ok=True) + _, jpeg = cv2.imencode( + ".jpg", frame, [cv2.IMWRITE_JPEG_QUALITY, LIVE_STREAM_QUALITY] + ) + with open(LIVE_STREAM_FRAME_PATH, "wb") as f: + f.write(jpeg.tobytes()) + except Exception: + pass + + frame_idx += 1 + if frame_idx % FLUSH_EVERY_N_FRAMES == 0: + print( + f"[{now_str()}] Frame {frame_idx} | Batch {batch_num}: {batch_count} " + f"| Total: {display_total} | Uptime {elapsed / 3600:.2f}h" + ) + prune_stale_tracks(ayam_tracked, mono) + prune_stale_tracks(talenan_tracked, mono) + + cap.release() + if video_writer is not None: + video_writer.release() + if cross_logger: + cross_logger.close() + model.release() + store.shutdown() + + print("\n=== Batch Summary (SQLite) ===") + print(f"Database: {DB_PATH}") + + +ayam_tracked = {} +talenan_tracked = {} + + +if __name__ == "__main__": + run() diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..5f705d7 --- /dev/null +++ b/go.mod @@ -0,0 +1,24 @@ +module github.com/proitlab/bytetrack-counter-go + +go 1.23 + +require ( + gocv.io/x/gocv v0.37.0 + modernc.org/sqlite v1.33.1 +) + +require ( + github.com/dustin/go-humanize v1.0.1 // indirect + github.com/google/uuid v1.6.0 // indirect + github.com/hashicorp/golang-lru/v2 v2.0.7 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect + github.com/ncruces/go-strftime v0.1.9 // indirect + github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect + golang.org/x/sys v0.22.0 // indirect + modernc.org/gc/v3 v3.0.0-20240107210532-573471604cb6 // indirect + modernc.org/libc v1.55.3 // indirect + modernc.org/mathutil v1.6.0 // indirect + modernc.org/memory v1.8.0 // indirect + modernc.org/strutil v1.2.0 // indirect + modernc.org/token v1.1.0 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..152a91a --- /dev/null +++ b/go.sum @@ -0,0 +1,51 @@ +github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= +github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= +github.com/google/pprof v0.0.0-20240409012703-83162a5b38cd h1:gbpYu9NMq8jhDVbvlGkMFWCjLFlqqEZjEmObmhUy6Vo= +github.com/google/pprof v0.0.0-20240409012703-83162a5b38cd/go.mod h1:kf6iHlnVGwgKolg33glAes7Yg/8iWP8ukqeldJSO7jw= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= +github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4= +github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +gocv.io/x/gocv v0.37.0 h1:sISHvnApErjoJodz1Dxb8UAkFdITOB3vXGslbVu6Knk= +gocv.io/x/gocv v0.37.0/go.mod h1:lmS802zoQmnNvXETpmGriBqWrENPei2GxYx5KUxJsMA= +golang.org/x/mod v0.16.0 h1:QX4fJ0Rr5cPQCF7O9lh9Se4pmwfwskqZfq5moyldzic= +golang.org/x/mod v0.16.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.22.0 h1:RI27ohtqKCnwULzJLqkv897zojh5/DwS/ENaMzUOaWI= +golang.org/x/sys v0.22.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= +golang.org/x/tools v0.19.0 h1:tfGCXNR1OsFG+sVdLAitlpjAvD/I6dHDKnYrpEZUHkw= +golang.org/x/tools v0.19.0/go.mod h1:qoJWxmGSIBmAeriMx19ogtrEPrGtDbPK634QFIcLAhc= +modernc.org/cc/v4 v4.21.4 h1:3Be/Rdo1fpr8GrQ7IVw9OHtplU4gWbb+wNgeoBMmGLQ= +modernc.org/cc/v4 v4.21.4/go.mod h1:HM7VJTZbUCR3rV8EYBi9wxnJ0ZBRiGE5OeGXNA0IsLQ= +modernc.org/ccgo/v4 v4.19.2 h1:lwQZgvboKD0jBwdaeVCTouxhxAyN6iawF3STraAal8Y= +modernc.org/ccgo/v4 v4.19.2/go.mod h1:ysS3mxiMV38XGRTTcgo0DQTeTmAO4oCmJl1nX9VFI3s= +modernc.org/fileutil v1.3.0 h1:gQ5SIzK3H9kdfai/5x41oQiKValumqNTDXMvKo62HvE= +modernc.org/fileutil v1.3.0/go.mod h1:XatxS8fZi3pS8/hKG2GH/ArUogfxjpEKs3Ku3aK4JyQ= +modernc.org/gc/v2 v2.4.1 h1:9cNzOqPyMJBvrUipmynX0ZohMhcxPtMccYgGOJdOiBw= +modernc.org/gc/v2 v2.4.1/go.mod h1:wzN5dK1AzVGoH6XOzc3YZ+ey/jPgYHLuVckd62P0GYU= +modernc.org/gc/v3 v3.0.0-20240107210532-573471604cb6 h1:5D53IMaUuA5InSeMu9eJtlQXS2NxAhyWQvkKEgXZhHI= +modernc.org/gc/v3 v3.0.0-20240107210532-573471604cb6/go.mod h1:Qz0X07sNOR1jWYCrJMEnbW/X55x206Q7Vt4mz6/wHp4= +modernc.org/libc v1.55.3 h1:AzcW1mhlPNrRtjS5sS+eW2ISCgSOLLNyFzRh/V3Qj/U= +modernc.org/libc v1.55.3/go.mod h1:qFXepLhz+JjFThQ4kzwzOjA/y/artDeg+pcYnY+Q83w= +modernc.org/mathutil v1.6.0 h1:fRe9+AmYlaej+64JsEEhoWuAYBkOtQiMEU7n/XgfYi4= +modernc.org/mathutil v1.6.0/go.mod h1:Ui5Q9q1TR2gFm0AQRqQUaBWFLAhQpCwNcuhBOSedWPo= +modernc.org/memory v1.8.0 h1:IqGTL6eFMaDZZhEWwcREgeMXYwmW83LYW8cROZYkg+E= +modernc.org/memory v1.8.0/go.mod h1:XPZ936zp5OMKGWPqbD3JShgd/ZoQ7899TUuQqxY+peU= +modernc.org/opt v0.1.3 h1:3XOZf2yznlhC+ibLltsDGzABUGVx8J6pnFMS3E4dcq4= +modernc.org/opt v0.1.3/go.mod h1:WdSiB5evDcignE70guQKxYUl14mgWtbClRi5wmkkTX0= +modernc.org/sortutil v1.2.0 h1:jQiD3PfS2REGJNzNCMMaLSp/wdMNieTbKX920Cqdgqc= +modernc.org/sortutil v1.2.0/go.mod h1:TKU2s7kJMf1AE84OoiGppNHJwvB753OYfNl2WRb++Ss= +modernc.org/sqlite v1.33.1 h1:trb6Z3YYoeM9eDL1O8do81kP+0ejv+YzgyFo+Gwy0nM= +modernc.org/sqlite v1.33.1/go.mod h1:pXV2xHxhzXZsgT/RtTFAPY6JJDEvOTcTdwADQCCWD4k= +modernc.org/strutil v1.2.0 h1:agBi9dp1I+eOnxXeiZawM8F4LawKv4NzGWSaLfyeNZA= +modernc.org/strutil v1.2.0/go.mod h1:/mdcBmfOibveCTBxUl5B5l6W+TTH1FXPLHZE6bTosX0= +modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y= +modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM= diff --git a/internal/iou/iou.go b/internal/iou/iou.go new file mode 100644 index 0000000..341d268 --- /dev/null +++ b/internal/iou/iou.go @@ -0,0 +1,110 @@ +package iou + +import ( + "math" + "sort" +) + +func NMS(boxes [][4]float64, scores []float64, iouThr float64) []int { + n := len(scores) + idx := make([]int, n) + for i := range idx { + idx[i] = i + } + sort.Slice(idx, func(i, j int) bool { + return scores[idx[i]] > scores[idx[j]] + }) + + keep := make([]int, 0, n) + for len(idx) > 0 { + cur := idx[0] + keep = append(keep, cur) + if len(idx) == 1 { + break + } + rest := idx[1:] + newOrder := make([]int, 0, len(rest)) + bCur := boxes[cur] + areaCur := (bCur[2] - bCur[0]) * (bCur[3] - bCur[1]) + for _, ri := range rest { + bRi := boxes[ri] + xx1 := math.Max(bCur[0], bRi[0]) + yy1 := math.Max(bCur[1], bRi[1]) + xx2 := math.Min(bCur[2], bRi[2]) + yy2 := math.Min(bCur[3], bRi[3]) + w := math.Max(0, xx2-xx1) + h := math.Max(0, yy2-yy1) + inter := w * h + areaRi := (bRi[2] - bRi[0]) * (bRi[3] - bRi[1]) + ov := inter / (areaCur + areaRi - inter + 1e-16) + if ov < iouThr { + newOrder = append(newOrder, ri) + } + } + idx = newOrder + } + return keep +} + +func PairwiseIoU(boxesA, boxesB [][4]float64) [][]float64 { + n, m := len(boxesA), len(boxesB) + mat := make([][]float64, n) + for i := range mat { + mat[i] = make([]float64, m) + } + if n == 0 || m == 0 { + return mat + } + for i := 0; i < n; i++ { + areaA := (boxesA[i][2] - boxesA[i][0]) * (boxesA[i][3] - boxesA[i][1]) + for j := 0; j < m; j++ { + xx1 := math.Max(boxesA[i][0], boxesB[j][0]) + yy1 := math.Max(boxesA[i][1], boxesB[j][1]) + xx2 := math.Min(boxesA[i][2], boxesB[j][2]) + yy2 := math.Min(boxesA[i][3], boxesB[j][3]) + iw := math.Max(0, xx2-xx1) + ih := math.Max(0, yy2-yy1) + inter := iw * ih + areaB := (boxesB[j][2] - boxesB[j][0]) * (boxesB[j][3] - boxesB[j][1]) + mat[i][j] = inter / (areaA + areaB - inter + 1e-16) + } + } + return mat +} + +func GreedyMatchRowCol(costMatrix [][]float64, threshold float64) []struct{ Row, Col int } { + if len(costMatrix) == 0 || len(costMatrix[0]) == 0 { + return nil + } + n, m := len(costMatrix), len(costMatrix[0]) + + type entry struct { + cost float64 + row, col int + } + entries := make([]entry, 0, n*m) + for i := 0; i < n; i++ { + for j := 0; j < m; j++ { + entries = append(entries, entry{costMatrix[i][j], i, j}) + } + } + sort.Slice(entries, func(a, b int) bool { + return entries[a].cost < entries[b].cost + }) + + rowUsed := make(map[int]bool) + colUsed := make(map[int]bool) + var pairs []struct{ Row, Col int } + for _, e := range entries { + if e.cost >= threshold { + break + } + if rowUsed[e.row] || colUsed[e.col] { + continue + } + rowUsed[e.row] = true + colUsed[e.col] = true + pairs = append(pairs, struct{ Row, Col int }{e.row, e.col}) + } + return pairs +} diff --git a/internal/nms/nms.go b/internal/nms/nms.go new file mode 100644 index 0000000..e3eeb28 --- /dev/null +++ b/internal/nms/nms.go @@ -0,0 +1,45 @@ +package nms + +import ( + "math" + "sort" +) + +func Apply(boxes [][4]float64, scores []float64, iouThr float64) []int { + n := len(scores) + order := make([]int, n) + for i := range order { + order[i] = i + } + sort.Slice(order, func(i, j int) bool { + return scores[order[i]] > scores[order[j]] + }) + + keep := make([]int, 0) + for len(order) > 0 { + idx := order[0] + keep = append(keep, idx) + if len(order) == 1 { + break + } + rest := order[1:] + newOrder := make([]int, 0, len(rest)) + 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 +} diff --git a/pkg/batchstore/store.go b/pkg/batchstore/store.go new file mode 100644 index 0000000..5b967cc --- /dev/null +++ b/pkg/batchstore/store.go @@ -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() + } +} diff --git a/pkg/bytetrack/tracker.go b/pkg/bytetrack/tracker.go new file mode 100644 index 0000000..26e3443 --- /dev/null +++ b/pkg/bytetrack/tracker.go @@ -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 +} diff --git a/pkg/config/config.go b/pkg/config/config.go new file mode 100644 index 0000000..90b2506 --- /dev/null +++ b/pkg/config/config.go @@ -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" +} diff --git a/pkg/drawing/drawing.go b/pkg/drawing/drawing.go new file mode 100644 index 0000000..49cf09c --- /dev/null +++ b/pkg/drawing/drawing.go @@ -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)))) +} diff --git a/pkg/kalman/kalman.go b/pkg/kalman/kalman.go new file mode 100644 index 0000000..ac491b5 --- /dev/null +++ b/pkg/kalman/kalman.go @@ -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 +} diff --git a/pkg/rknn/rknn.go b/pkg/rknn/rknn.go new file mode 100644 index 0000000..4aacf06 --- /dev/null +++ b/pkg/rknn/rknn.go @@ -0,0 +1,145 @@ +package rknn + +/* +#cgo linux LDFLAGS: -lrknnrt +#cgo linux,arm64 LDFLAGS: -lrknnrt + +#include +#include + +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) +} diff --git a/pkg/yolo/yolo.go b/pkg/yolo/yolo.go new file mode 100644 index 0000000..dcb187e --- /dev/null +++ b/pkg/yolo/yolo.go @@ -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() +}