Files
bytetrack-counter-go/cmd/counter/main.go
T
2026-06-30 18:49:39 +07:00

546 lines
14 KiB
Go

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)
}
}
}