Files
zenai-kpc-python/counter_live_rknn.py

1892 lines
68 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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 contextlib
import numpy as np
import cv2
import csv
import json
import os
import shutil
import signal
import socket
import sys
import threading
import time
from collections import deque
from datetime import datetime
from pathlib import Path
from dotenv import load_dotenv
load_dotenv()
from rknnlite.api import RKNNLite
from counter_store import CounterStore
from video_session_writer import VideoSessionWriter
# --- 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_counter.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", "object")
CLASS_OBJECT = os.getenv("CLASS_OBJECT", "object")
LINE_Y1 = int(os.getenv("LINE_Y1")) if os.getenv("LINE_Y1") else None
LINE_Y1_FRAC = float(os.getenv("LINE_Y1_FRAC", "0.33"))
LINE_Y2 = int(os.getenv("LINE_Y2")) if os.getenv("LINE_Y2") else None
LINE_Y2_FRAC = float(os.getenv("LINE_Y2_FRAC", "0.66"))
IMGSZ = int(os.getenv("IMGSZ", "320"))
HALF = os.getenv("HALF", "false").lower() == "true"
CONF = float(os.getenv("CONF", "0.3"))
NMS_IOU = float(os.getenv("NMS_IOU", "0.45"))
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"))
# Dedup guard against ID-switch double counts: ignore a second crossing in the
# same direction that happens within DEDUP_FRAMES and DEDUP_PX (x-distance) of a
# recent crossing.
DEDUP_FRAMES = int(os.getenv("DEDUP_FRAMES", "15"))
DEDUP_PX = float(os.getenv("DEDUP_PX", "60"))
# Trajectory inheritance across ID switches: when a new track appears, inherit the
# last position of a recently-seen nearby track so crossings are not missed when
# the ID changes right at the line.
INHERIT_SEC = float(os.getenv("INHERIT_SEC", "1.0"))
INHERIT_PX = float(os.getenv("INHERIT_PX", "60"))
DAILY_CUTOFF_TIME = os.getenv("DAILY_CUTOFF_TIME", "20:00")
EXPORT_CSV = os.getenv("EXPORT_CSV", "true").lower() == "true"
CROSS_CSV = os.getenv("CROSS_CSV", f"{OUTPUT_DIR}/crossings.csv")
# Save an annotated frame snapshot each time an object crosses a line and the
# counter increases.
SAVE_CROSS_SNAPSHOT = os.getenv("SAVE_CROSS_SNAPSHOT", "false").lower() == "true"
# Also save one snapshot the first time each object is detected (before it crosses),
# named with the same track id so it can be correlated with the crossing snapshot.
SAVE_DETECT_SNAPSHOT = os.getenv("SAVE_DETECT_SNAPSHOT", "false").lower() == "true"
CROSS_SNAPSHOT_DIR = os.getenv("CROSS_SNAPSHOT_DIR", f"{OUTPUT_DIR}/snapshots")
CROSS_SNAPSHOT_QUALITY = int(os.getenv("CROSS_SNAPSHOT_QUALITY", "85"))
# Retention: delete oldest snapshots when either limit is exceeded (0 = disabled).
CROSS_SNAPSHOT_MAX_FILES = int(os.getenv("CROSS_SNAPSHOT_MAX_FILES", "1000"))
CROSS_SNAPSHOT_MAX_AGE_DAYS = float(os.getenv("CROSS_SNAPSHOT_MAX_AGE_DAYS", "7"))
# Run the cleanup sweep at most every N seconds to limit filesystem scans.
CROSS_SNAPSHOT_CLEANUP_SEC = int(os.getenv("CROSS_SNAPSHOT_CLEANUP_SEC", "3600"))
RATE_WINDOW_SEC = int(os.getenv("RATE_WINDOW_SEC", "60"))
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"))
# Max stream-error prints per outage (Cannot open / Stream dropped). After this,
# further errors are silent until the stream recovers, then the budget resets.
# 0 = unlimited (old behavior).
STREAM_ERROR_LOG_MAX = int(os.getenv("STREAM_ERROR_LOG_MAX", "3"))
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")) # legacy segment mode only
OUTPUT_FPS = int(os.getenv("OUTPUT_FPS", "15"))
# ByteTrack-style session recording: encode in SHM, move to VIDEO_OUTPUT_DIR on session end.
SHM_DIR = os.getenv("SHM_DIR", "/dev/shm/zenai-kpc-counter")
VIDEO_OUTPUT_DIR = os.getenv("VIDEO_OUTPUT_DIR", OUTPUT_DIR)
RECORD_END_DELAY = float(os.getenv("RECORD_END_DELAY", "1"))
RECORD_DISCARD_EMPTY = os.getenv("RECORD_DISCARD_EMPTY", "true").lower() == "true"
RECORD_RETENTION_DAYS = int(os.getenv("RECORD_RETENTION_DAYS", "3"))
_crf_raw = os.getenv("RECORD_CRF", "").strip()
RECORD_CRF = int(_crf_raw) if _crf_raw else -1
RECORD_PRESET = os.getenv("RECORD_PRESET", "").strip()
# segment = old hourly dump to OUTPUT_DIR; session = SHM + webhook/control lifecycle
RECORD_MODE = os.getenv("RECORD_MODE", "session").strip().lower()
LIVE_STREAM_ENABLED = os.getenv("LIVE_STREAM_ENABLED", "false").lower() == "true"
LIVE_STREAM_FRAME_PATH = os.getenv(
"LIVE_STREAM_FRAME_PATH", f"{SHM_DIR}/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://"))
MOTION_DETECTION_ENABLED = os.getenv("MOTION_DETECTION_ENABLED", "false").lower() == "true"
MOTION_THRESHOLD = float(os.getenv("MOTION_THRESHOLD", "5.0"))
# Per-pixel intensity change (0-255) for a pixel to count as "moved".
MOTION_PIXEL_DELTA = int(os.getenv("MOTION_PIXEL_DELTA", "25"))
# Fraction of frame pixels that must change (0-1) to trigger inference. Small,
# so an object entering the edge of the frame is detected immediately.
MOTION_MIN_AREA_FRAC = float(os.getenv("MOTION_MIN_AREA_FRAC", "0.002"))
# Always run inference at least every N frames even if no motion (heartbeat), so a
# slow/stationary object is never missed for long.
MOTION_HEARTBEAT_FRAMES = int(os.getenv("MOTION_HEARTBEAT_FRAMES", "15"))
# --- Runtime control (toggle counting on/off on the fly) ---
# When enabled, the process watches a small JSON control file and honors its
# "counting" flag. Set false to always count (ignore the control file).
CONTROL_ENABLED = os.getenv("CONTROL_ENABLED", "false").lower() == "true"
CONTROL_FILE = os.getenv("CONTROL_FILE", f"{OUTPUT_DIR}/control.json")
# Whether counting is active on startup when no control file exists yet.
CONTROL_DEFAULT_COUNTING = os.getenv("CONTROL_DEFAULT_COUNTING", "true").lower() == "true"
# Re-read the control file at most every N seconds.
CONTROL_POLL_SEC = float(os.getenv("CONTROL_POLL_SEC", "1.0"))
# Optional TCP control socket. When enabled, the counter listens for line-based
# commands so counting can be toggled over the network (in addition to the file).
CONTROL_SOCKET_ENABLED = os.getenv("CONTROL_SOCKET_ENABLED", "false").lower() == "true"
CONTROL_SOCKET_HOST = os.getenv("CONTROL_SOCKET_HOST", "127.0.0.1")
CONTROL_SOCKET_PORT = int(os.getenv("CONTROL_SOCKET_PORT", "5090"))
# --- Status webhook gate (direction-locked when IN or OUT) ---
# When enabled, counting runs only while status is IN or OUT, and only that
# direction is counted (IN → in crossings, OUT → out crossings). OFF pauses all.
STATUS_WEBHOOK_ENABLED = os.getenv("STATUS_WEBHOOK_ENABLED", "false").lower() == "true"
STATUS_WEBHOOK_FILE = os.getenv(
"STATUS_WEBHOOK_FILE",
"/opt/zenai-kpc-python/status_webhook_state.json",
)
_active_raw = os.getenv("STATUS_WEBHOOK_ACTIVE_VALUES", "IN,OUT")
STATUS_WEBHOOK_ACTIVE_VALUES = frozenset(
part.strip().upper() for part in _active_raw.split(",") if part.strip()
)
STATUS_WEBHOOK_INACTIVE_VALUE = os.getenv("STATUS_WEBHOOK_INACTIVE_VALUE", "OFF").strip().upper()
CROSS_FLASH_FRAMES = 12
POPUP_LIFETIME = 20
LINE_PULSE_FRAMES = 12
COUNT_PULSE_FRAMES = 15
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_OBJECT_BOX = (0, 165, 255)
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])
def get_cy(self):
return float(self.kf.x[1, 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
remain_orig_idx = np.where(remain)[0]
dets = boxes_xyxy[remain]
det_scores = scores[remain]
is_high = det_scores > self.high_thresh
is_low = ~is_high
else:
remain_orig_idx = np.zeros(0, dtype=np.int64)
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])
orig_idx = int(remain_orig_idx[det_global])
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[orig_idx] = track_pool[ti].track_id
tracked_map[track_pool[ti].track_id] = (track_pool[ti].get_cx(), track_pool[ti].get_cy())
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=1.0 - self.match_thresh
)
for di, uti in matches2:
det_global = int(low_idx[di])
pool_idx = unmatched_tracks[uti]
orig_idx = int(remain_orig_idx[det_global])
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[orig_idx] = track_pool[pool_idx].track_id
tracked_map[track_pool[pool_idx].track_id] = (
track_pool[pool_idx].get_cx(),
track_pool[pool_idx].get_cy(),
)
# --- 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(), trk.get_cy()))
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(), trk.get_cy())
# --- 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:
orig_idx = int(remain_orig_idx[int(dg)])
if orig_idx not in matched_det_ids:
trk = KalmanBoxTracker(dets[int(dg)])
self.tracked_tracks.append(trk)
det_to_track[orig_idx] = trk.track_id
tracked_map[trk.track_id] = (trk.get_cx(), trk.get_cy())
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_y1(frame_height):
if LINE_Y1 is not None:
return LINE_Y1
return int(frame_height * LINE_Y1_FRAC)
def resolve_line_y2(frame_height):
if LINE_Y2 is not None:
return LINE_Y2
return int(frame_height * LINE_Y2_FRAC)
def crossed_top_down(prev_y, y, line_y):
return prev_y < line_y <= y
def crossed_bottom_up(prev_y, y, line_y):
return prev_y > line_y >= y
def is_duplicate_cross(recent, cx, frame_idx):
"""True if a crossing near cx happened within the dedup window (ID-switch guard)."""
while recent and frame_idx - recent[0][0] > DEDUP_FRAMES:
recent.popleft()
for _, prev_cx in recent:
if abs(prev_cx - cx) <= DEDUP_PX:
return True
return False
def _inherit_prev(tracked, new_tid, cx, cy, mono, max_age, max_px):
"""Find a recently-seen nearby track for ID-switch continuation.
Returns (cx, cy, mono, source_tid) or None.
The predecessor must be close in BOTH x and y (same physical object at the same
spot). Matching on x only would let a new track inherit a far-away y, seeding a
huge prev->cur jump that fabricates a line crossing in the wrong direction."""
best = None
best_dist = max_px
for tid, (tcx, tcy, ts) in tracked.items():
if tid == new_tid:
continue
if mono - ts > max_age:
continue
if abs(tcy - cy) > max_px:
continue
dist = abs(tcx - cx)
if dist <= best_dist:
best_dist = dist
best = (tcx, tcy, ts, tid)
return best
def _default_cross_state():
"""Per-track line-crossing state for sequence counting.
seen_out / seen_in: ever registered a pass of that line (arm or count).
counted: already contributed one IN or OUT (at most one per physical object).
"""
return {"seen_out": False, "seen_in": False, "counted": False}
def _copy_cross_state(src):
return {
"seen_out": bool(src.get("seen_out", False)),
"seen_in": bool(src.get("seen_in", False)),
"counted": bool(src.get("counted", False)),
}
def now_str():
return datetime.now().strftime("%Y-%m-%d %H:%M:%S")
def read_counting_flag(default=True):
"""Read the 'counting' flag from the control file. Returns default on any error."""
try:
with open(CONTROL_FILE, "r", encoding="utf-8") as f:
data = json.load(f)
return bool(data.get("counting", default))
except FileNotFoundError:
return default
except Exception:
return default
def read_status_webhook():
"""Return normalized status from webhook file, or inactive value on error."""
try:
with open(STATUS_WEBHOOK_FILE, "r", encoding="utf-8") as f:
data = json.load(f)
status = str(data.get("status", "")).strip().upper()
return status or STATUS_WEBHOOK_INACTIVE_VALUE
except Exception:
return STATUS_WEBHOOK_INACTIVE_VALUE
def status_webhook_allows_counting(status):
return status in STATUS_WEBHOOK_ACTIVE_VALUES
def write_control_file(counting):
"""Create/update the control file atomically (used to seed defaults)."""
try:
Path(CONTROL_FILE).parent.mkdir(parents=True, exist_ok=True)
tmp = f"{CONTROL_FILE}.tmp"
with open(tmp, "w", encoding="utf-8") as f:
json.dump({"counting": bool(counting)}, f)
os.replace(tmp, CONTROL_FILE)
except Exception as exc:
print(f"[{now_str()}] Failed to write control file: {exc}")
def start_control_socket():
"""Start a TCP server for runtime control. Commands (newline-terminated):
START | RESUME | ON -> counting on
STOP | PAUSE | OFF -> counting off
TOGGLE -> flip
STATUS | GET -> report current state
It writes the shared control file, so the main loop's file-poll applies it.
Returns the server socket (call .close() to stop)."""
srv = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
srv.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
srv.bind((CONTROL_SOCKET_HOST, CONTROL_SOCKET_PORT))
srv.listen(5)
def handle(conn, addr):
with conn:
conn.settimeout(30)
try:
buf = b""
while not shutdown_requested:
try:
chunk = conn.recv(256)
except socket.timeout:
break
if not chunk:
break
buf += chunk
while b"\n" in buf:
line, buf = buf.split(b"\n", 1)
cmd = line.decode("utf-8", "ignore").strip().upper()
if not cmd:
continue
current = read_counting_flag(CONTROL_DEFAULT_COUNTING)
if cmd in ("START", "RESUME", "ON"):
write_control_file(True)
resp = "OK counting=on"
elif cmd in ("STOP", "PAUSE", "OFF"):
write_control_file(False)
resp = "OK counting=off"
elif cmd == "TOGGLE":
write_control_file(not current)
resp = f"OK counting={'off' if current else 'on'}"
elif cmd in ("STATUS", "GET"):
resp = f"OK counting={'on' if current else 'off'}"
else:
resp = "ERR unknown command"
conn.sendall((resp + "\n").encode("utf-8"))
except Exception:
pass
def loop():
print(f"Control socket listening on {CONTROL_SOCKET_HOST}:{CONTROL_SOCKET_PORT}")
while not shutdown_requested:
try:
conn, addr = srv.accept()
except OSError:
break
t = threading.Thread(target=handle, args=(conn, addr), daemon=True)
t.start()
threading.Thread(target=loop, daemon=True).start()
return srv
@contextlib.contextmanager
def _quiet_opencv_stderr():
"""Silence OpenCV/FFmpeg stderr (DESCRIBE failed, cap.cpp WARN, etc.)."""
saved_fd = None
devnull_fd = None
try:
stderr_fd = sys.stderr.fileno()
saved_fd = os.dup(stderr_fd)
devnull_fd = os.open(os.devnull, os.O_WRONLY)
os.dup2(devnull_fd, stderr_fd)
except (AttributeError, OSError, ValueError):
yield
return
try:
yield
finally:
os.dup2(saved_fd, stderr_fd)
os.close(saved_fd)
os.close(devnull_fd)
def open_capture(source):
if source.lower().startswith(("rtsp://", "http://")):
os.environ["OPENCV_FFMPEG_CAPTURE_OPTIONS"] = RTSP_FFMPEG_OPTIONS
# When STREAM_ERROR_LOG_MAX > 0, native OpenCV/FFmpeg stderr is suppressed;
# stream_error_log emits at most that many user-facing messages per outage.
if STREAM_ERROR_LOG_MAX > 0:
with _quiet_opencv_stderr():
cap = cv2.VideoCapture(source, cv2.CAP_FFMPEG)
else:
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 cleanup_snapshots(snapshot_dir, max_files, max_age_days):
"""Delete oldest / expired crossing snapshots to bound disk usage."""
d = Path(snapshot_dir)
if not d.is_dir():
return
files = sorted(d.rglob("*.jpg"), key=lambda p: p.stat().st_mtime)
if max_age_days > 0:
cutoff = time.time() - max_age_days * 86400
for p in list(files):
if p.stat().st_mtime < cutoff:
p.unlink(missing_ok=True)
files.remove(p)
if max_files > 0 and len(files) > max_files:
for p in files[: len(files) - max_files]:
p.unlink(missing_ok=True)
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_y, w, pulse_remaining=0, label="LINE"):
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, (0, line_y - offset), (w, line_y - offset), color, 1, cv2.LINE_AA)
cv2.line(img, (0, line_y + offset), (w, line_y + offset), color, 1, cv2.LINE_AA)
dash_len, gap = 18, 12
x = 0
while x < w:
x_end = min(x + dash_len, w)
cv2.line(img, (x, line_y), (x_end, line_y), C_LINE_CORE, 2, cv2.LINE_AA)
x += dash_len + gap
cv2.putText(
img,
label,
(14, line_y - 8),
cv2.FONT_HERSHEY_SIMPLEX,
0.42,
C_LINE_CORE,
1,
cv2.LINE_AA,
)
def draw_line_count(img, w, line_y, count, label, color, pulse_remaining=0, above=True):
text = str(count)
font = cv2.FONT_HERSHEY_SIMPLEX
boost = 0.35 * (pulse_remaining / max(COUNT_PULSE_FRAMES, 1))
font_scale, thickness = 1.2 + boost, 3
(tw, th), _ = cv2.getTextSize(text, font, font_scale, thickness)
(lw, lh), _ = cv2.getTextSize(label, font, 0.45, 1)
pad = 12
box_w = max(tw, lw) + pad * 2
box_h = th + lh + pad * 2 + 6
bx2 = w - 16
bx1 = bx2 - box_w
if above:
by2 = line_y - 10
by1 = by2 - box_h
else:
by1 = line_y + 10
by2 = by1 + box_h
overlay_rect(img, bx1, by1, bx2, by2, C_PANEL, alpha=0.78)
cv2.rectangle(img, (bx1, by1), (bx2, by2), color, 2)
tx = bx1 + (box_w - tw) // 2
ty = by1 + pad + th
cv2.putText(img, text, (tx, ty), font, font_scale, color, thickness, cv2.LINE_AA)
slx = bx1 + (box_w - lw) // 2
sly = ty + lh + 6
cv2.putText(img, label, (slx, sly), font, 0.45, C_MUTED, 1, cv2.LINE_AA)
def draw_hud(img, w, total, total_in, total_out, elapsed_sec, rate):
bar_h = 40
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, "TOTAL", (14, 14), cv2.FONT_HERSHEY_SIMPLEX, 0.32, C_MUTED, 1, cv2.LINE_AA
)
cv2.putText(
img,
str(total),
(14, 32),
cv2.FONT_HERSHEY_SIMPLEX,
0.55,
C_GREEN,
1,
cv2.LINE_AA,
)
cv2.putText(
img, "IN", (100, 14), cv2.FONT_HERSHEY_SIMPLEX, 0.32, C_MUTED, 1, cv2.LINE_AA
)
cv2.putText(
img,
str(total_in),
(100, 32),
cv2.FONT_HERSHEY_SIMPLEX,
0.55,
C_GREEN,
1,
cv2.LINE_AA,
)
cv2.putText(
img, "OUT", (170, 14), cv2.FONT_HERSHEY_SIMPLEX, 0.32, C_MUTED, 1, cv2.LINE_AA
)
cv2.putText(
img,
str(total_out),
(170, 32),
cv2.FONT_HERSHEY_SIMPLEX,
0.55,
C_TEXT,
1,
cv2.LINE_AA,
)
cv2.putText(
img,
"UPTIME",
(250, 14),
cv2.FONT_HERSHEY_SIMPLEX,
0.32,
C_MUTED,
1,
cv2.LINE_AA,
)
cv2.putText(
img,
f"{elapsed_sec / 3600:.1f}h",
(250, 32),
cv2.FONT_HERSHEY_SIMPLEX,
0.45,
C_TEXT,
1,
cv2.LINE_AA,
)
cv2.putText(
img, "RATE", (340, 14), cv2.FONT_HERSHEY_SIMPLEX, 0.32, C_MUTED, 1, cv2.LINE_AA
)
cv2.putText(
img,
f"{rate:.1f}/min",
(340, 32),
cv2.FONT_HERSHEY_SIMPLEX,
0.45,
C_ACCENT,
1,
cv2.LINE_AA,
)
def draw_footer(img, w, h, frame_idx, live_tag, inf_ms=0.0, model_name=""):
bar_h = 28
overlay_rect(img, 0, h - bar_h, w, h, C_PANEL, alpha=0.55)
cv2.putText(
img,
f"{live_tag} | {model_name} | Frame {frame_idx} | Inf {inf_ms:.1f}ms",
(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
class StreamErrorLog:
"""Rate-limit stream error prints to STREAM_ERROR_LOG_MAX per outage.
After the cap, messages are suppressed until note_up() when the stream
is healthy again; the next drop starts a fresh budget.
"""
def __init__(self, max_logs=STREAM_ERROR_LOG_MAX):
self.max_logs = int(max_logs)
self._emitted = 0
self._suppressed = 0
def emit(self, msg: str) -> None:
if self.max_logs <= 0:
print(msg)
return
if self._emitted < self.max_logs:
self._emitted += 1
print(msg)
if self._emitted >= self.max_logs:
print(
f"[{now_str()}] Stream errors capped at {self.max_logs} "
f"this outage — further messages suppressed until stream recovers"
)
else:
self._suppressed += 1
def note_up(self) -> None:
if self._suppressed > 0:
print(
f"[{now_str()}] Stream recovered "
f"(suppressed {self._suppressed} error log(s) during outage)"
)
self._emitted = 0
self._suppressed = 0
stream_error_log = StreamErrorLog()
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}"
)
stream_error_log.emit(
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
stream_error_log.note_up()
return cap, w, h, fps
return None, 0, 0, OUTPUT_FPS
# =============================================================================
# Main loop
# =============================================================================
def run():
global shutdown_requested
store = CounterStore(
db_path=DB_PATH,
state_file=STATE_FILE,
camera_name=CAMERA_NAME,
object_label=OBJECT_LABEL,
cutoff_time=DAILY_CUTOFF_TIME,
logger=lambda msg: print(f"[{now_str()}] {msg}"),
)
store.start_cutoff_watcher()
cross_logger = None
if EXPORT_CSV:
cross_logger = CsvLogger(
CROSS_CSV, ["counting_date", "frame", "direction", "object_id"]
)
model = RKNNYOLO(
model_path=MODEL_PATH,
core_mask=CORE_MASK,
imgsz=IMGSZ,
conf=CONF,
iou=NMS_IOU,
num_classes=NUM_CLASSES,
score_sigmoid=SCORE_SIGMOID,
)
object_cls = int(os.getenv("OBJECT_CLASS_ID", "0"))
object_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,
)
# Per-track sequence state: OUT→IN / IN-only → IN; IN→OUT / OUT-only → OUT.
# Keys survive ID switches via inheritance (see _inherit_prev).
object_cross_state = {}
recent_cross_in = deque()
recent_cross_out = deque()
detect_snapshot_ids = set()
object_cross_flash1 = {}
object_cross_flash2 = {}
line_pulse = count_in_pulse = count_out_pulse = 0
popups = []
session_start = time.time()
frame_idx = 0
inf_ms = 0.0
prev_gray = None
frames_since_infer = 0
video_writer = None
crossing_times = deque()
# Overlay shows daily store totals (resume-safe, resets on counting-day change).
counter_in = store.current_count_in
counter_out = store.current_count_out
last_snapshot_cleanup = 0.0
counting_active = True
status_nyala_active = True
status_webhook_value = STATUS_WEBHOOK_INACTIVE_VALUE
last_control_poll = 0.0
control_socket = None
if CONTROL_ENABLED:
if not Path(CONTROL_FILE).exists():
write_control_file(CONTROL_DEFAULT_COUNTING)
counting_active = read_counting_flag(CONTROL_DEFAULT_COUNTING)
print(
f"Runtime control enabled | file={CONTROL_FILE} | "
f"counting={'ON' if counting_active else 'OFF'}"
)
if CONTROL_SOCKET_ENABLED:
try:
control_socket = start_control_socket()
except Exception as exc:
print(f"[{now_str()}] Failed to start control socket: {exc}")
if STATUS_WEBHOOK_ENABLED:
status_webhook_value = read_status_webhook()
status_nyala_active = status_webhook_allows_counting(status_webhook_value)
print(
f"Status webhook enabled | file={STATUS_WEBHOOK_FILE} | "
f"status={status_webhook_value}"
)
cap, w, h, fps = connect_stream(SOURCE)
if cap is None:
store.shutdown()
model.release()
return
line_y1 = resolve_line_y1(h)
line_y2 = resolve_line_y2(h)
print(
f"RKNN+ByteTrack counter | {w}x{h} @ {fps}fps | "
f"line1 y={line_y1} (IN ↓) line2 y={line_y2} (OUT ↑) | "
f"seq: OUT→IN/IN-only→in, IN→OUT/OUT-only→out"
)
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:
if RECORD_MODE == "segment":
video_writer = VideoSegmentWriter(OUTPUT_DIR, w, h, fps, VIDEO_SEGMENT_SEC)
print(
f"Recording (segment): {OUTPUT_DIR} every {VIDEO_SEGMENT_SEC}s"
)
else:
video_writer = VideoSessionWriter(
shm_dir=SHM_DIR,
video_output_dir=VIDEO_OUTPUT_DIR,
w=w,
h=h,
fps=fps if fps and fps > 1 else OUTPUT_FPS,
end_delay_sec=RECORD_END_DELAY,
retention_days=RECORD_RETENTION_DAYS,
crf=RECORD_CRF,
preset=RECORD_PRESET,
discard_empty=RECORD_DISCARD_EMPTY,
)
print(
f"Recording (session): SHM={SHM_DIR} → out={VIDEO_OUTPUT_DIR} | "
f"end_delay={RECORD_END_DELAY}s discard_empty={RECORD_DISCARD_EMPTY} | "
f"ffmpeg={'yes' if shutil.which('ffmpeg') else 'NO'}"
)
# If webhook already active at startup, begin a session immediately.
if STATUS_WEBHOOK_ENABLED and status_nyala_active:
video_writer.note_session_start(
status_webhook_value, store.get_counting_date()
)
elif not STATUS_WEBHOOK_ENABLED and counting_active:
video_writer.note_session_start("SESSION", store.get_counting_date())
reconnect_count = 0
while not shutdown_requested:
ret, frame = cap.read()
if not ret:
if not IS_LIVE:
break
reconnect_count += 1
stream_error_log.emit(
f"Stream dropped (attempt {reconnect_count}), "
f"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_y1 = resolve_line_y1(h)
line_y2 = resolve_line_y2(h)
continue
now = time.time()
elapsed = now - session_start
mono = time.monotonic()
object_crossed_frame = False
cross_events_frame = []
detect_events_frame = []
if (CONTROL_ENABLED or STATUS_WEBHOOK_ENABLED) and (
now - last_control_poll
) >= CONTROL_POLL_SEC:
last_control_poll = now
if CONTROL_ENABLED:
new_flag = read_counting_flag(CONTROL_DEFAULT_COUNTING)
if new_flag != counting_active:
counting_active = new_flag
print(
f"[{now_str()}] Counting "
f"{'RESUMED' if counting_active else 'PAUSED'} via control file"
)
if (
isinstance(video_writer, VideoSessionWriter)
and not STATUS_WEBHOOK_ENABLED
):
if counting_active:
video_writer.note_session_start(
"SESSION", store.get_counting_date()
)
else:
video_writer.note_session_end()
if STATUS_WEBHOOK_ENABLED:
new_status = read_status_webhook()
if new_status != status_webhook_value:
prev_status = status_webhook_value
status_webhook_value = new_status
status_nyala_active = status_webhook_allows_counting(
status_webhook_value
)
print(f"[{now_str()}] Status webhook {status_webhook_value}")
if isinstance(video_writer, VideoSessionWriter):
was_active = status_webhook_allows_counting(prev_status)
if status_nyala_active and not was_active:
video_writer.note_session_start(
status_webhook_value, store.get_counting_date()
)
elif was_active and not status_nyala_active:
video_writer.note_session_end()
# Raw frames only (before overlay). Session writer encodes here every frame.
if isinstance(video_writer, VideoSessionWriter):
video_writer.feed_frame(frame)
counting_allowed = counting_active and status_nyala_active
skip_inference = False
if not counting_allowed:
skip_inference = True
elif MOTION_DETECTION_ENABLED:
gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
if prev_gray is not None:
diff = cv2.absdiff(gray, prev_gray)
moved = int(np.count_nonzero(diff > MOTION_PIXEL_DELTA))
moved_frac = moved / diff.size
skip_inference = moved_frac < MOTION_MIN_AREA_FRAC
if frames_since_infer >= MOTION_HEARTBEAT_FRAMES:
skip_inference = False
prev_gray = gray
detections = []
if not skip_inference:
frames_since_infer = 0
inf_start = time.time()
detections = model(frame)
inf_ms = inf_ms * 0.9 + (time.time() - inf_start) * 1000 * 0.1
else:
frames_since_infer += 1
object_boxes_xyxy = []
object_scores = []
object_kpts_list = []
object_cx_list = []
object_cy_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
cy = (bbox[1] + bbox[3]) / 2.0
if cls_id == object_cls:
object_boxes_xyxy.append(bbox)
object_scores.append(score)
object_kpts_list.append(kpts)
object_cx_list.append(cx)
object_cy_list.append(cy)
object_boxes_xyxy = np.array(object_boxes_xyxy, dtype=np.float32).reshape(-1, 4)
object_scores = np.array(object_scores, dtype=np.float32)
object_track_map, object_det_to_track, object_lost_map = object_tracker.update(
object_boxes_xyxy, object_scores
)
if os.getenv("DEBUG_TRACKING", "").lower() == "true":
if len(object_boxes_xyxy) > 0:
scores_str = (
f" scores: {object_scores.round(3).tolist()}"
if len(object_scores) > 0
else ""
)
tracks_str = (
f" det->track: {dict(object_det_to_track)}"
if object_det_to_track
else ""
)
crossing_str = ""
if object_cross_state:
line1_ids = sorted(
tid for tid, st in object_cross_state.items() if st.get("seen_in")
)
line2_ids = sorted(
tid for tid, st in object_cross_state.items() if st.get("seen_out")
)
if line1_ids:
crossing_str += f" line1_crossed: {line1_ids}"
if line2_ids:
crossing_str += f" line2_crossed: {line2_ids}"
print(
f"[DEBUG F{frame_idx}] dets={len(object_boxes_xyxy)} "
f"tracks={len(object_track_map)} "
f"line_y1={line_y1} line_y2={line_y2}{scores_str}"
f"{tracks_str}{crossing_str}"
)
# --- crossing detection only on DETECTED objects this frame ---
# (tracker still updates every frame to preserve IDs, but we do NOT
# count on Kalman-predicted/coasting tracks to avoid double counts)
for di in range(len(object_boxes_xyxy)):
tid = object_det_to_track.get(di)
if tid is None:
continue
cx = object_cx_list[di]
cy = object_cy_list[di]
if tid not in detect_snapshot_ids:
detect_snapshot_ids.add(tid)
detect_events_frame.append(tid)
inherited_this_frame = False
if tid not in object_tracked:
inherited = _inherit_prev(
object_tracked, tid, cx, cy, mono, INHERIT_SEC, INHERIT_PX
)
if inherited is not None:
object_tracked[tid] = inherited[:3]
src_tid = inherited[3]
if src_tid in object_cross_state:
object_cross_state[tid] = _copy_cross_state(
object_cross_state[src_tid]
)
inherited_this_frame = True
if os.getenv("DEBUG_TRACKING", "").lower() == "true":
src_st = object_cross_state.get(tid, {})
print(
f"[DEBUG F{frame_idx}] INHERIT prev for new tid={tid} "
f"from tid={src_tid} ({inherited[0]:.1f},{inherited[1]:.1f})"
f" seen_in={src_st.get('seen_in', False)}"
f" seen_out={src_st.get('seen_out', False)}"
f" counted={src_st.get('counted', False)}"
)
if tid not in object_cross_state:
object_cross_state[tid] = _default_cross_state()
st = object_cross_state[tid]
if tid in object_tracked:
prev_cy = object_tracked[tid][1]
# Line 1 = IN, line 2 = OUT.
# Arm (no count): OUT top→down, IN bottom→up.
# Count: IN top→down → IN (OUT→IN or IN-only)
# OUT bottom→up → OUT (IN→OUT or OUT-only)
out_down = (
crossed_top_down(prev_cy, cy, line_y2) and not st["seen_out"]
)
in_up = (
crossed_bottom_up(prev_cy, cy, line_y1) and not st["seen_in"]
)
in_down = (
crossed_top_down(prev_cy, cy, line_y1) and not st["seen_in"]
)
out_up = (
crossed_bottom_up(prev_cy, cy, line_y2) and not st["seen_out"]
)
if out_down or in_up or in_down or out_up:
if out_down:
st["seen_out"] = True
if in_up:
st["seen_in"] = True
if in_down:
st["seen_in"] = True
if out_up:
st["seen_out"] = True
# Prefer decisive motion if both count triggers fire in one jump.
count_in = in_down and not st["counted"]
count_out = out_up and not st["counted"]
if STATUS_WEBHOOK_ENABLED:
# Webhook IN → only IN counts; OUT → only OUT counts.
if status_webhook_value == "IN":
count_out = False
elif status_webhook_value == "OUT":
count_in = False
if count_in and count_out:
if cy >= prev_cy:
count_out = False
else:
count_in = False
direction = "in" if count_in else ("out" if count_out else None)
if direction is None:
# Arm-only pass (OUT↓ or IN↑). No count yet.
if os.getenv("DEBUG_TRACKING", "").lower() == "true":
arm = []
if out_down:
arm.append("out_down")
if in_up:
arm.append("in_up")
print(
f"[DEBUG F{frame_idx}] CROSS ARM: tid={tid} "
f"prev_cy={prev_cy:.1f} -> cy={cy:.1f} "
f"arm={'+'.join(arm)} "
f"seen_in={st['seen_in']} seen_out={st['seen_out']}"
)
object_tracked[tid] = (cx, cy, mono)
continue
recent = recent_cross_in if direction == "in" else recent_cross_out
if is_duplicate_cross(recent, cx, frame_idx):
st["counted"] = True
if os.getenv("DEBUG_TRACKING", "").lower() == "true":
print(
f"[DEBUG F{frame_idx}] DUP CROSS IGNORED: tid={tid} "
f"cx={cx:.1f} dir={direction}"
)
object_tracked[tid] = (cx, cy, mono)
continue
if os.getenv("DEBUG_TRACKING", "").lower() == "true":
line_n = 1 if direction == "in" else 2
line_y = line_y1 if direction == "in" else line_y2
print(
f"[DEBUG F{frame_idx}] CROSS DETECTED: tid={tid} "
f"prev_cy={prev_cy:.1f} -> cy={cy:.1f} "
f"line={line_n} y={line_y} dir={direction}"
)
recent.append((frame_idx, cx))
st["counted"] = True
_, _, counted = store.record_object_crossing(tid, direction)
counter_in = store.current_count_in
counter_out = store.current_count_out
if not counted:
# Dedup in the store rejected this event; skip overlay/CSV.
object_tracked[tid] = (cx, cy, mono)
continue
if direction == "in":
count_in_pulse = COUNT_PULSE_FRAMES
else:
count_out_pulse = COUNT_PULSE_FRAMES
if cross_logger:
cross_logger.write_row(
[
store.get_counting_date(),
frame_idx,
direction,
tid,
]
)
object_crossed_frame = True
cross_events_frame.append((tid, direction))
crossing_times.append(mono)
if isinstance(video_writer, VideoSessionWriter):
video_writer.note_crossing()
object_cross_flash1[tid] = CROSS_FLASH_FRAMES
object_cross_flash2[tid] = CROSS_FLASH_FRAMES
popups.append(
{
"x": int(cx) - 12,
"y": int(cy),
"born": frame_idx,
"text": "+1",
}
)
object_tracked[tid] = (cx, cy, mono)
for tid, (cx, cy) in object_lost_map.items():
if tid not in object_tracked:
object_tracked[tid] = (cx, cy, mono)
for di in range(len(object_boxes_xyxy)):
tid = object_det_to_track.get(di)
if tid is None:
continue
bbox = object_boxes_xyxy[di]
x1, y1, x2, y2 = int(bbox[0]), int(bbox[1]), int(bbox[2]), int(bbox[3])
flash = max(
object_cross_flash1.get(tid, 0),
object_cross_flash2.get(tid, 0),
)
color = C_GREEN if flash > 0 else C_OBJECT_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 = object_kpts_list[di] if di < len(object_kpts_list) else None
if kpts is not None:
draw_skeleton_bold(frame, kpts)
if object_crossed_frame:
line_pulse = LINE_PULSE_FRAMES
# Keep overlay aligned with daily store (also resets after cutoff).
counter_in = store.current_count_in
counter_out = store.current_count_out
while crossing_times and mono - crossing_times[0] > RATE_WINDOW_SEC:
crossing_times.popleft()
rate = (len(crossing_times) / RATE_WINDOW_SEC * 60) if crossing_times else 0.0
draw_elegant_counting_line(frame, line_y1, w, line_pulse, label="LINE IN")
draw_elegant_counting_line(frame, line_y2, w, line_pulse, label="LINE OUT")
draw_line_count(frame, w, line_y1, counter_in, "IN", C_GREEN, count_in_pulse, above=True)
draw_line_count(frame, w, line_y2, counter_out, "OUT", C_OBJECT_BOX, count_out_pulse, above=False)
draw_hud(
frame,
w,
counter_in + counter_out,
counter_in,
counter_out,
elapsed,
rate,
)
draw_footer(
frame,
w,
h,
frame_idx,
"LIVE" if IS_LIVE else "FILE",
inf_ms,
Path(MODEL_PATH).name,
)
popups = draw_popups(frame, popups, frame_idx)
if (CONTROL_ENABLED or STATUS_WEBHOOK_ENABLED) and not counting_allowed:
if STATUS_WEBHOOK_ENABLED and not status_nyala_active:
badge = status_webhook_value
else:
badge = "COUNTING PAUSED"
(bw, bh), _ = cv2.getTextSize(badge, cv2.FONT_HERSHEY_SIMPLEX, 0.6, 2)
bx = w // 2 - bw // 2
overlay_rect(frame, bx - 14, 48, bx + bw + 14, 48 + bh + 18, C_PANEL, alpha=0.75)
cv2.rectangle(frame, (bx - 14, 48), (bx + bw + 14, 48 + bh + 18), C_ACCENT, 2)
cv2.putText(frame, badge, (bx, 48 + bh + 6), cv2.FONT_HERSHEY_SIMPLEX, 0.6, C_ACCENT, 2, cv2.LINE_AA)
for flash_store in (
object_cross_flash1, object_cross_flash2,
):
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_in_pulse = max(0, count_in_pulse - 1)
count_out_pulse = max(0, count_out_pulse - 1)
if isinstance(video_writer, VideoSegmentWriter):
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)
ok, jpeg = cv2.imencode(
".jpg", frame, [cv2.IMWRITE_JPEG_QUALITY, LIVE_STREAM_QUALITY]
)
if ok:
tmp_path = f"{LIVE_STREAM_FRAME_PATH}.tmp"
with open(tmp_path, "wb") as f:
f.write(jpeg.tobytes())
os.replace(tmp_path, LIVE_STREAM_FRAME_PATH)
except Exception:
pass
if (SAVE_DETECT_SNAPSHOT and detect_events_frame) or (
SAVE_CROSS_SNAPSHOT and cross_events_frame
):
try:
ts = datetime.now().strftime("%Y%m%d_%H%M%S_%f")[:-3]
if SAVE_DETECT_SNAPSHOT and detect_events_frame:
detect_dir = Path(CROSS_SNAPSHOT_DIR) / "detect"
detect_dir.mkdir(parents=True, exist_ok=True)
for tid in detect_events_frame:
fname = f"{ts}_detect_id{tid}_f{frame_idx}.jpg"
cv2.imwrite(
str(detect_dir / fname),
frame,
[cv2.IMWRITE_JPEG_QUALITY, CROSS_SNAPSHOT_QUALITY],
)
if SAVE_CROSS_SNAPSHOT and cross_events_frame:
cross_dir = Path(CROSS_SNAPSHOT_DIR) / "cross"
cross_dir.mkdir(parents=True, exist_ok=True)
for tid, direction in cross_events_frame:
fname = f"{ts}_{direction}_id{tid}_f{frame_idx}.jpg"
cv2.imwrite(
str(cross_dir / fname),
frame,
[cv2.IMWRITE_JPEG_QUALITY, CROSS_SNAPSHOT_QUALITY],
)
if now - last_snapshot_cleanup >= CROSS_SNAPSHOT_CLEANUP_SEC:
cleanup_snapshots(
CROSS_SNAPSHOT_DIR,
CROSS_SNAPSHOT_MAX_FILES,
CROSS_SNAPSHOT_MAX_AGE_DAYS,
)
last_snapshot_cleanup = now
except Exception as exc:
print(f"[{now_str()}] Failed to save snapshot: {exc}")
frame_idx += 1
prune_stale_tracks(object_tracked, mono)
cap.release()
if isinstance(video_writer, VideoSessionWriter):
video_writer.shutdown()
elif video_writer is not None:
video_writer.release()
if cross_logger:
cross_logger.close()
if control_socket is not None:
try:
control_socket.close()
except Exception:
pass
model.release()
store.shutdown()
print("\n=== Daily Counter Summary (SQLite) ===")
print(f"Database: {DB_PATH}")
object_tracked = {}
if __name__ == "__main__":
run()