First commit

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

No files matched your search

+467
View File
@@ -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()
}
}
+195
View File
@@ -0,0 +1,195 @@
package bytetrack
import (
"github.com/anomalyco/bytetrack-counter-go/internal/iou"
"github.com/anomalyco/bytetrack-counter-go/pkg/kalman"
)
type ByteTracker struct {
HighThresh float64
LowThresh float64
MatchThresh float64
Buffer int
MinHits int
tracked []*kalman.Tracker
lost []*kalman.Tracker
removed []*kalman.Tracker
frameID int
}
func New(highThresh, lowThresh, matchThresh float64, buffer, minHits int) *ByteTracker {
return &ByteTracker{
HighThresh: highThresh,
LowThresh: lowThresh,
MatchThresh: matchThresh,
Buffer: buffer,
MinHits: minHits,
}
}
type UpdateResult struct {
Tracked map[int]float64
DetToTrack map[int]int
Lost map[int]float64
}
func (bt *ByteTracker) Update(boxes [][4]float64, scores []float64) UpdateResult {
bt.frameID++
result := UpdateResult{
Tracked: make(map[int]float64),
DetToTrack: make(map[int]int),
Lost: make(map[int]float64),
}
n := len(boxes)
var dets [][4]float64
var filtToGlobal []int
var highFiltIdx []int
var lowFiltIdx []int
for i := 0; i < n; i++ {
if scores[i] > bt.LowThresh {
filtToGlobal = append(filtToGlobal, i)
dets = append(dets, boxes[i])
if scores[i] > bt.HighThresh {
highFiltIdx = append(highFiltIdx, len(dets)-1)
} else {
lowFiltIdx = append(lowFiltIdx, len(dets)-1)
}
}
}
pool := make([]*kalman.Tracker, 0, len(bt.tracked)+len(bt.lost))
pool = append(pool, bt.tracked...)
pool = append(pool, bt.lost...)
numPool := len(pool)
var poolBoxes [][4]float64
if numPool > 0 {
poolBoxes = make([][4]float64, numPool)
for i, trk := range pool {
trk.Predict()
poolBoxes[i] = trk.GetState()
}
}
matchedPool := make(map[int]bool)
if numPool > 0 && len(highFiltIdx) > 0 {
highDets := make([][4]float64, len(highFiltIdx))
for i, hi := range highFiltIdx {
highDets[i] = dets[hi]
}
iouMat := iou.PairwiseIoU(highDets, poolBoxes)
cost := make([][]float64, len(iouMat))
for i := range cost {
cost[i] = make([]float64, len(iouMat[i]))
for j := range cost[i] {
cost[i][j] = 1.0 - iouMat[i][j]
}
}
matches := iou.GreedyMatchRowCol(cost, 1.0-bt.MatchThresh)
for _, m := range matches {
filtIdx := highFiltIdx[m.Row]
globalIdx := filtToGlobal[filtIdx]
ti := m.Col
pool[ti].Update(dets[filtIdx])
if pool[ti].HitStreak < 1 {
pool[ti].HitStreak = 1
}
matchedPool[ti] = true
result.DetToTrack[globalIdx] = pool[ti].ID
result.Tracked[pool[ti].ID] = pool[ti].GetCX()
}
}
var unmatchedPool []int
for i := 0; i < numPool; i++ {
if !matchedPool[i] {
unmatchedPool = append(unmatchedPool, i)
}
}
if len(lowFiltIdx) > 0 && len(unmatchedPool) > 0 {
lowDets := make([][4]float64, len(lowFiltIdx))
for i, li := range lowFiltIdx {
lowDets[i] = dets[li]
}
unmatchedBoxes := make([][4]float64, len(unmatchedPool))
for i, pi := range unmatchedPool {
unmatchedBoxes[i] = poolBoxes[pi]
}
iouMat2 := iou.PairwiseIoU(lowDets, unmatchedBoxes)
cost2 := make([][]float64, len(iouMat2))
for i := range cost2 {
cost2[i] = make([]float64, len(iouMat2[i]))
for j := range cost2[i] {
cost2[i][j] = 1.0 - iouMat2[i][j]
}
}
matches2 := iou.GreedyMatchRowCol(cost2, 0.5)
for _, m := range matches2 {
filtIdx := lowFiltIdx[m.Row]
globalIdx := filtToGlobal[filtIdx]
ti := unmatchedPool[m.Col]
pool[ti].Update(dets[filtIdx])
if pool[ti].HitStreak < 1 {
pool[ti].HitStreak = 1
}
matchedPool[ti] = true
result.DetToTrack[globalIdx] = pool[ti].ID
result.Tracked[pool[ti].ID] = pool[ti].GetCX()
}
}
for i, trk := range pool {
if !matchedPool[i] {
trk.HitStreak = 0
}
}
newTracked := make([]*kalman.Tracker, 0)
newLost := make([]*kalman.Tracker, 0)
for _, trk := range pool {
if trk.TimeSinceUpd > bt.Buffer {
bt.removed = append(bt.removed, trk)
} else if trk.TimeSinceUpd > 0 {
newLost = append(newLost, trk)
} else {
newTracked = append(newTracked, trk)
}
}
bt.tracked = newTracked
bt.lost = newLost
for _, trk := range bt.tracked {
if trk.HitStreak >= bt.MinHits || trk.Hits >= bt.MinHits {
result.Tracked[trk.ID] = trk.GetCX()
}
}
for _, trk := range bt.lost {
if trk.HitStreak >= bt.MinHits || trk.Hits >= bt.MinHits {
result.Tracked[trk.ID] = trk.GetCX()
result.Lost[trk.ID] = trk.GetCX()
}
}
matchedDets := make(map[int]bool)
for idx := range result.DetToTrack {
matchedDets[idx] = true
}
for _, fi := range highFiltIdx {
globalIdx := filtToGlobal[fi]
if !matchedDets[globalIdx] {
trk := kalman.NewTracker(dets[fi])
bt.tracked = append(bt.tracked, trk)
result.DetToTrack[globalIdx] = trk.ID
result.Tracked[trk.ID] = trk.GetCX()
}
}
return result
}
+173
View File
@@ -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"
}
+241
View File
@@ -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))))
}
+396
View File
@@ -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
}
+145
View File
@@ -0,0 +1,145 @@
package rknn
/*
#cgo linux LDFLAGS: -lrknnrt
#cgo linux,arm64 LDFLAGS: -lrknnrt
#include <stdint.h>
#include <stdlib.h>
typedef void* rknn_context;
typedef enum {
RKNN_QUERY_IN_OUT_NUM = 0,
RKNN_QUERY_INPUT_ATTR = 1,
RKNN_QUERY_OUTPUT_ATTR = 2,
RKNN_QUERY_SDK_VERSION = 5,
} rknn_query_cmd;
typedef enum {
RKNN_TENSOR_UINT8 = 0,
RKNN_TENSOR_FLOAT32 = 2,
} rknn_tensor_type;
typedef enum {
RKNN_TENSOR_NCHW = 0,
RKNN_TENSOR_NHWC = 1,
} rknn_tensor_format;
typedef struct {
uint32_t index;
int32_t type;
int32_t fmt;
void* buf;
uint32_t size;
uint8_t want_float;
uint8_t is_prealloc;
} rknn_output;
typedef struct {
uint32_t n_input;
uint32_t n_output;
} rknn_input_output_num;
typedef struct {
void* buf;
uint32_t size;
uint8_t pass_through;
uint32_t type;
uint32_t fmt;
} rknn_input;
extern int rknn_init(rknn_context* ctx, void* model, uint32_t size, uint32_t flag, rknn_input_output_num* io_num);
extern int rknn_query(rknn_context ctx, rknn_query_cmd cmd, void* info, uint32_t size);
extern int rknn_inputs_set(rknn_context ctx, uint32_t n_inputs, rknn_input inputs[]);
extern int rknn_run(rknn_context ctx, void* ext);
extern int rknn_outputs_get(rknn_context ctx, uint32_t n_outputs, rknn_output outputs[], void* ext);
extern int rknn_outputs_release(rknn_context ctx, uint32_t n_outputs, rknn_output outputs[]);
extern int rknn_destroy(rknn_context ctx);
*/
import "C"
import (
"fmt"
"os"
"unsafe"
)
type Context struct {
ctx C.rknn_context
nInput uint32
nOutput uint32
}
func LoadModel(modelPath string) (*Context, error) {
data, err := os.ReadFile(modelPath)
if err != nil {
return nil, fmt.Errorf("rknn: read model file: %w", err)
}
var ctx C.rknn_context
var ioNum C.rknn_input_output_num
ret := C.rknn_init(
&ctx,
unsafe.Pointer(&data[0]),
C.uint32_t(len(data)),
0,
&ioNum,
)
if ret != 0 {
return nil, fmt.Errorf("rknn_init: error %d", int(ret))
}
return &Context{
ctx: ctx,
nInput: uint32(ioNum.n_input),
nOutput: uint32(ioNum.n_output),
}, nil
}
func (c *Context) InferenceRGB(inputData []uint8, height, width int) ([][]float32, error) {
rknnIn := C.rknn_input{
buf: unsafe.Pointer(&inputData[0]),
size: C.uint32_t(len(inputData)),
pass_through: 0,
_type: C.RKNN_TENSOR_UINT8,
fmt: C.RKNN_TENSOR_NHWC,
}
ret := C.rknn_inputs_set(c.ctx, 1, &rknnIn)
if ret < 0 {
return nil, fmt.Errorf("rknn_inputs_set: error %d", int(ret))
}
ret = C.rknn_run(c.ctx, nil)
if ret < 0 {
return nil, fmt.Errorf("rknn_run: error %d", int(ret))
}
outputs := make([]C.rknn_output, c.nOutput)
for i := range outputs {
outputs[i].want_float = 1
}
ret = C.rknn_outputs_get(c.ctx, C.uint32_t(c.nOutput), &outputs[0], nil)
if ret < 0 {
return nil, fmt.Errorf("rknn_outputs_get: error %d", int(ret))
}
results := make([][]float32, c.nOutput)
for i := range outputs {
nVals := int(outputs[i].size) / 4
vals := make([]float32, nVals)
src := unsafe.Slice((*float32)(outputs[i].buf), nVals)
copy(vals, src)
results[i] = vals
}
C.rknn_outputs_release(c.ctx, C.uint32_t(c.nOutput), &outputs[0])
return results, nil
}
func (c *Context) Release() {
C.rknn_destroy(c.ctx)
}
+236
View File
@@ -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()
}