From 8285400254b8abcc3321310683a85651df2041bc Mon Sep 17 00:00:00 2001
From: asus
Date: Wed, 5 Aug 2026 15:56:11 +0700
Subject: [PATCH] update from asus 106
---
algoritma-batch/batch_video_cropper.py | 485 ++++++++++++++
algoritma-batch/cfg/tracker.yaml | 26 +
algoritma-batch/src/__init__.py | 1 +
algoritma-batch/src/batch.py | 390 +++++++++++
algoritma-batch/src/counting.py | 237 +++++++
algoritma-batch/src/dashboard.py | 251 +++++++
algoritma-batch/src/detection.py | 82 +++
algoritma-batch/src/interfaces.py | 98 +++
algoritma-batch/src/stabilizer.py | 126 ++++
algoritma-batch/src/tracking.py | 83 +++
algoritma-batch/src/truck_roi.py | 151 +++++
algoritma-batch/zones.json | 43 ++
backend/api/batches.py | 127 +++-
backend/api/models.py | 2 +
backend/autolabel.py | 133 ++--
backend/dataset.py | 48 +-
backend/hardware.py | 6 +-
backend/labeling.py | 15 +-
backend/projects.py | 5 +
backend/review.py | 50 ++
backend/training.py | 13 +-
frontend/src/App.jsx | 6 +
frontend/src/api.js | 24 +
frontend/src/components/Sidebar.jsx | 6 +-
frontend/src/pages/LibraryPage.jsx | 766 ++++++++++++++--------
frontend/src/pages/ModelsPage.jsx | 165 ++---
frontend/src/pages/Sam3PlaygroundPage.jsx | 328 +++++++++
scripts/restart_app.sh | 7 +-
28 files changed, 3215 insertions(+), 459 deletions(-)
create mode 100644 algoritma-batch/batch_video_cropper.py
create mode 100644 algoritma-batch/cfg/tracker.yaml
create mode 100644 algoritma-batch/src/__init__.py
create mode 100644 algoritma-batch/src/batch.py
create mode 100644 algoritma-batch/src/counting.py
create mode 100644 algoritma-batch/src/dashboard.py
create mode 100644 algoritma-batch/src/detection.py
create mode 100644 algoritma-batch/src/interfaces.py
create mode 100644 algoritma-batch/src/stabilizer.py
create mode 100644 algoritma-batch/src/tracking.py
create mode 100644 algoritma-batch/src/truck_roi.py
create mode 100644 algoritma-batch/zones.json
create mode 100644 frontend/src/pages/Sam3PlaygroundPage.jsx
diff --git a/algoritma-batch/batch_video_cropper.py b/algoritma-batch/batch_video_cropper.py
new file mode 100644
index 0000000..4b97065
--- /dev/null
+++ b/algoritma-batch/batch_video_cropper.py
@@ -0,0 +1,485 @@
+"""
+Batch Video Cropper — Rekam Video RTSP per Sesi Batch Truk
+
+Program ini membaca livestream CCTV (RTSP), menjalankan algoritma penentuan batch
+menggunakan YOLO + State Machine, dan menyimpan potongan video per-batch ke folder
+yang terorganisir berdasarkan tanggal.
+
+Struktur Output:
+ ~/reTraining/data/archive/
+ ├── 2026-08-05/
+ │ ├── batch_1_09-15-30.mp4
+ │ ├── batch_2_10-22-45.mp4
+ │ └── batch_3_14-08-12.mp4
+ └── 2026-08-06/
+ └── batch_1_07-30-00.mp4
+
+Menjalankan:
+ cd ~/reTraining/algoritma-batch
+ python3 batch_video_cropper.py
+"""
+
+import os
+import cv2
+import numpy as np
+import time
+import json
+import threading
+from datetime import datetime, timedelta
+from shapely.geometry import Point, Polygon
+from ultralytics import YOLO
+
+# Import modules
+from src.tracking import ByteTrackTracker
+from src.stabilizer import BboxStabilizer
+from src.truck_roi import TruckROI
+from src.counting import LineCrossCounter
+from src.batch import BatchLifecycleManager, BatchRecord
+
+# =====================================================================
+# 1. KONFIGURASI
+# =====================================================================
+
+import platform
+IS_WINDOWS = platform.system() == "Windows"
+
+# --- Path Konfigurasi ---
+BASE_DIR = os.path.dirname(os.path.abspath(__file__))
+ZONES_JSON = os.path.join(BASE_DIR, "zones.json")
+
+if IS_WINDOWS:
+ MODEL_PATH = os.path.join(BASE_DIR, "v1-best.pt")
+ ARCHIVE_BASE = os.path.join(BASE_DIR, "archive_output")
+else:
+ MODEL_PATH = os.path.join(BASE_DIR, "v1-best.pt")
+ ARCHIVE_BASE = os.path.expanduser("~/reTraining/data/archive")
+
+# --- Sumber Video RTSP ---
+RTSP_URL = "rtsp://frigate:zenai@192.168.192.209:8554/camera_stream_640"
+
+# --- Batas Pergantian Hari (Cutoff) ---
+DAILY_CUTOFF_TIME = "20:00"
+
+# --- Parameter State Machine ---
+SACK_IDLE_TIMEOUT = 5.0 # Jeda aktivitas sebelum masuk WAITING_FOR_ACTIVITY
+MIN_BATCH_DURATION = 2.0 # Durasi minimal batch sebelum boleh masuk WAITING
+TOLERANCE_LOW_COUNT = 0.0 # Instan (0s) — batch langsung berakhir saat area truk kosong
+TOLERANCE_MED_COUNT = 0.0
+TOLERANCE_HIGH_COUNT = 0.0
+
+# --- Video Recording ---
+VIDEO_FPS = 10.0 # FPS output video (10 fps sudah cukup untuk rekaman arsip)
+VIDEO_CODEC = "mp4v" # Codec untuk .mp4
+
+
+# =====================================================================
+# 2. THREADED RTSP READER (Menghindari Lag Buffer)
+# =====================================================================
+class RTSPStreamReader:
+ """Threaded RTSP reader yang selalu mengambil frame terbaru."""
+
+ def __init__(self, source_url):
+ self.source_url = source_url
+ self.cap = cv2.VideoCapture(source_url)
+ self.frame = None
+ self.ret = False
+ self.running = True
+ self.lock = threading.Lock()
+ self.new_frame_event = threading.Event()
+ self.thread = threading.Thread(target=self._update, daemon=True)
+ self.thread.start()
+
+ def _update(self):
+ while self.running:
+ if not self.cap.isOpened():
+ print("[RTSP] Stream terputus, mencoba reconnect dalam 5 detik...")
+ time.sleep(5)
+ self.cap = cv2.VideoCapture(self.source_url)
+ continue
+ ret, frame = self.cap.read()
+ if not ret:
+ time.sleep(0.01)
+ continue
+ with self.lock:
+ self.ret = ret
+ self.frame = frame
+ self.new_frame_event.set()
+ time.sleep(0.001)
+
+ def read(self):
+ if self.new_frame_event.wait(timeout=2.0):
+ self.new_frame_event.clear()
+ with self.lock:
+ if self.frame is None:
+ return False, None
+ return self.ret, self.frame.copy()
+ else:
+ with self.lock:
+ if self.frame is None:
+ return False, None
+ return self.ret, self.frame.copy()
+
+ def isOpened(self):
+ return self.cap.isOpened()
+
+ def release(self):
+ self.running = False
+ if self.cap.isOpened():
+ self.cap.release()
+
+
+# =====================================================================
+# 3. FUNGSI UTILITAS TANGGAL & FOLDER
+# =====================================================================
+def get_counting_date(dt=None):
+ """Menentukan tanggal kerja berdasarkan cutoff harian."""
+ if dt is None:
+ dt = datetime.now()
+ try:
+ cutoff = datetime.strptime(DAILY_CUTOFF_TIME, "%H:%M").time()
+ except Exception:
+ cutoff = datetime.strptime("20:00", "%H:%M").time()
+
+ if cutoff.hour == 0 and cutoff.minute == 0:
+ return dt.date().isoformat()
+
+ if dt.time() < cutoff:
+ return dt.date().isoformat()
+ return (dt.date() + timedelta(days=1)).isoformat()
+
+
+def ensure_date_folder(counting_date):
+ """Membuat folder tanggal di archive jika belum ada. Mengembalikan path folder."""
+ folder = os.path.join(ARCHIVE_BASE, counting_date)
+ os.makedirs(folder, exist_ok=True)
+ return folder
+
+
+# =====================================================================
+# 4. BATCH VIDEO RECORDER (Mengelola VideoWriter per Batch)
+# =====================================================================
+class BatchVideoRecorder:
+ """Mengelola pembukaan dan penutupan file video per sesi batch."""
+
+ def __init__(self, archive_base, video_fps=10.0, codec="mp4v"):
+ self.archive_base = archive_base
+ self.video_fps = video_fps
+ self.codec = codec
+ self.writer = None
+ self.current_path = None
+ self.frame_count = 0
+
+ def start_recording(self, batch_number, counting_date, frame_width=1280, frame_height=720):
+ """Membuka file video baru untuk batch ini."""
+ self.stop_recording() # Pastikan writer sebelumnya ditutup
+
+ folder = ensure_date_folder(counting_date)
+ timestamp_str = datetime.now().strftime("%H-%M-%S")
+ filename = f"batch_{batch_number}_{timestamp_str}.mp4"
+ self.current_path = os.path.join(folder, filename)
+
+ fourcc = cv2.VideoWriter_fourcc(*self.codec)
+ self.writer = cv2.VideoWriter(
+ self.current_path, fourcc, self.video_fps, (frame_width, frame_height)
+ )
+ self.frame_count = 0
+
+ if self.writer.isOpened():
+ print(f"[RECORD] Mulai merekam video batch #{batch_number} -> {self.current_path}")
+ else:
+ print(f"[RECORD ERROR] Gagal membuka VideoWriter untuk: {self.current_path}")
+ self.writer = None
+
+ def write_frame(self, frame):
+ """Menulis satu frame ke video aktif."""
+ if self.writer is not None and self.writer.isOpened():
+ self.writer.write(frame)
+ self.frame_count += 1
+
+ def stop_recording(self):
+ """Menutup file video yang sedang aktif."""
+ if self.writer is not None:
+ self.writer.release()
+ self.writer = None
+ if self.current_path and self.frame_count > 0:
+ print(f"[RECORD] Video selesai disimpan: {self.current_path} ({self.frame_count} frames)")
+ elif self.current_path and self.frame_count == 0:
+ # Hapus file kosong
+ try:
+ os.remove(self.current_path)
+ print(f"[RECORD] File video kosong dihapus: {self.current_path}")
+ except Exception:
+ pass
+ self.current_path = None
+ self.frame_count = 0
+
+ @property
+ def is_recording(self):
+ return self.writer is not None and self.writer.isOpened()
+
+
+# =====================================================================
+# 5. FUNGSI UTAMA — LOOP UTAMA DETEKSI & PEREKAMAN
+# =====================================================================
+def run_batch_video_cropper():
+ print("=" * 60)
+ print(" BATCH VIDEO CROPPER — RTSP → Per-Batch MP4 Recorder")
+ print("=" * 60)
+ print(f" Model : {MODEL_PATH}")
+ print(f" RTSP : {RTSP_URL}")
+ print(f" Archive : {ARCHIVE_BASE}")
+ print(f" Cutoff : {DAILY_CUTOFF_TIME}")
+ print(f" Tolerance : Instan (0s)")
+ print("=" * 60)
+
+ # Pastikan folder archive ada
+ os.makedirs(ARCHIVE_BASE, exist_ok=True)
+
+ # --- Load Model YOLO ---
+ print("[INFO] Memuat model YOLO...")
+ model = YOLO(MODEL_PATH)
+
+ # Warm-up model
+ print("[INFO] Warm-up model...")
+ dummy = np.zeros((720, 1280, 3), dtype=np.uint8)
+ device = "cuda" if os.path.exists("/usr/local/cuda") else "cpu"
+ try:
+ import torch
+ if torch.cuda.is_available():
+ device = "cuda"
+ except ImportError:
+ pass
+ _ = model(dummy, imgsz=640, device=device, verbose=False)
+ print(f"[INFO] Model siap. Device: {device}")
+
+ # --- Setup Components ---
+ tracker = ByteTrackTracker(model, conf=0.55)
+ stabilizer = BboxStabilizer(
+ ema_alpha=0.35,
+ max_hold_frames=10,
+ max_height_ratio=1.5,
+ min_height_ratio=0.70,
+ )
+
+ # Skala koordinat (kalibrasi 1920x1080 -> 1280x720)
+ scale_x = 1280.0 / 1920.0
+ scale_y = 720.0 / 1080.0
+
+ # Detection polygon
+ detection_poly_pts = [
+ [int(574 * scale_x), int(50 * scale_y)],
+ [int(586 * scale_x), int(1077 * scale_y)],
+ [int(1418 * scale_x), int(1076 * scale_y)],
+ [int(1397 * scale_x), int(50 * scale_y)],
+ ]
+ detection_polygon = Polygon(detection_poly_pts)
+
+ # Truck polygon (untuk monitoring kehadiran karung)
+ truck_poly_pts = [
+ [int(600 * scale_x), int(385 * scale_y)],
+ [int(609 * scale_x), int(1076 * scale_y)],
+ [int(1404 * scale_x), int(1078 * scale_y)],
+ [int(1381 * scale_x), int(343 * scale_y)],
+ ]
+ truck_polygon = Polygon(truck_poly_pts)
+
+ # Line crossing
+ static_line_y = int(330 * scale_y)
+ static_line_x_start = int(577 * scale_x)
+ static_line_x_end = int(1401 * scale_x)
+
+ static_roi = TruckROI(
+ x1=int(600 * scale_x),
+ y1=int(343 * scale_y),
+ x2=int(1404 * scale_x),
+ y2=int(1078 * scale_y),
+ line_y=static_line_y,
+ confidence=1.0,
+ )
+
+ counter = LineCrossCounter(
+ line_y=static_line_y,
+ line_x_start=static_line_x_start,
+ line_x_end=static_line_x_end,
+ margin=20,
+ dedup_radius=60.0,
+ )
+
+ batch_mgr = BatchLifecycleManager(
+ stabilize_seconds=0.0,
+ stabilize_threshold_px=9999.0,
+ sack_idle_timeout=SACK_IDLE_TIMEOUT,
+ min_batch_duration=MIN_BATCH_DURATION,
+ truck_gone_tolerance=TOLERANCE_LOW_COUNT,
+ )
+
+ recorder = BatchVideoRecorder(
+ archive_base=ARCHIVE_BASE,
+ video_fps=VIDEO_FPS,
+ codec=VIDEO_CODEC,
+ )
+
+ # --- Buka RTSP Stream ---
+ print(f"[INFO] Membuka RTSP stream: {RTSP_URL}")
+ cap = RTSPStreamReader(RTSP_URL)
+ if not cap.isOpened():
+ print(f"[ERROR] Gagal membuka RTSP stream: {RTSP_URL}")
+ return
+
+ # Tracking state
+ batch_counter = 0
+ frame_idx = 0
+ last_fps_time = time.time()
+ fps_counter = 0
+
+ print("\n[INFO] Memulai loop utama... Tekan Ctrl+C untuk berhenti.\n")
+
+ try:
+ while True:
+ ret, frame = cap.read()
+ if not ret or frame is None:
+ time.sleep(0.01)
+ continue
+
+ # Resize ke 1280x720 (sesuai kalibrasi koordinat zona)
+ frame = cv2.resize(frame, (1280, 720))
+ timestamp = time.time()
+ frame_idx += 1
+
+ # FPS counter
+ fps_counter += 1
+ if fps_counter % 100 == 0:
+ elapsed = time.time() - last_fps_time
+ fps = 100.0 / elapsed if elapsed > 0 else 0
+ last_fps_time = time.time()
+ state_str = batch_mgr.state
+ rec_str = "REC" if recorder.is_recording else "---"
+ print(f"[FPS] {fps:.1f} fps | State: {state_str} | {rec_str} | Frames: {frame_idx}")
+
+ # Simpan state sebelum update
+ prev_active = batch_mgr.is_active
+ prev_state = batch_mgr.state
+
+ # --- 1. YOLO Tracking ---
+ raw_tracked_all = tracker.update(frame, [])
+ raw_tracked_sacks = [d for d in raw_tracked_all if d.class_name == "sack"]
+
+ # --- 2. Stabilizer ---
+ stable = stabilizer.update(raw_tracked_sacks)
+
+ # --- 3. Filter Detection Polygon ---
+ stable = [
+ d for d in stable
+ if detection_polygon.contains(
+ Point((d.bbox[0] + d.bbox[2]) / 2.0, (d.bbox[1] + d.bbox[3]) / 2.0)
+ )
+ ]
+
+ # --- 4. Hitung karung di 70% area bawah truk ---
+ min_ty, max_ty = truck_polygon.bounds[1], truck_polygon.bounds[3]
+ truck_height = max_ty - min_ty
+ truck_cutoff_y = min_ty + 0.30 * truck_height
+
+ sacks_in_truck_area = 0
+ for d in stable:
+ cx = (d.bbox[0] + d.bbox[2]) / 2.0
+ cy = (d.bbox[1] + d.bbox[3]) / 2.0
+ if truck_polygon.contains(Point(cx, cy)) and cy >= truck_cutoff_y:
+ sacks_in_truck_area += 1
+
+ # --- 5. Line Crossing ---
+ tracked_in_roi = [
+ d for d in stable
+ if static_roi.contains_x((d.bbox[0] + d.bbox[2]) / 2.0)
+ ]
+ events = counter.update(tracked_in_roi)
+ has_crossing = len(events) > 0
+
+ # =============================================================
+ # LOGIKA ALGORITMA PENENTUAN BATCH (STATE MACHINE)
+ # =============================================================
+
+ # A. Mulai Batch
+ if batch_mgr.state in ("IDLE", "TRUCK_STABILIZING"):
+ batch_mgr.update_truck(has_crossing, (0.0, 0.0), timestamp)
+
+ # B. Monitoring Batch Aktif
+ if batch_mgr.state in ("COUNTING_SACKS", "WAITING_FOR_ACTIVITY"):
+ current_count = counter.loading_count
+ if current_count < 20:
+ batch_mgr._truck_gone_tolerance = TOLERANCE_LOW_COUNT
+ elif current_count >= 40:
+ batch_mgr._truck_gone_tolerance = TOLERANCE_HIGH_COUNT
+ else:
+ batch_mgr._truck_gone_tolerance = TOLERANCE_MED_COUNT
+
+ batch_mgr.update_sacks(
+ has_crossing_event=has_crossing,
+ sacks_in_area_count=sacks_in_truck_area,
+ timestamp=timestamp,
+ loading_count=counter.loading_count,
+ unloading_count=counter.unloading_count,
+ )
+
+ if batch_mgr.state == "WAITING_FOR_ACTIVITY":
+ batch_mgr.update_truck(sacks_in_truck_area > 0, None, timestamp)
+
+ for ev in events:
+ now_str = datetime.now().strftime("%H:%M:%S")
+ print(f"[{now_str}] [KARUNG] #{ev['track_id']} melintasi garis. Total: {counter.loading_count}")
+
+ # =============================================================
+ # TRANSISI BATCH — MULAI/SELESAI REKAMAN VIDEO
+ # =============================================================
+
+ # C. Batch baru saja dimulai
+ if batch_mgr.is_active and not prev_active:
+ counting_date = get_counting_date()
+ batch_counter += 1
+ now_str = datetime.now().strftime("%H:%M:%S")
+ print(f"\n>>> [{now_str}] BATCH #{batch_counter} DIMULAI (tanggal: {counting_date}) <<<")
+ recorder.start_recording(batch_counter, counting_date, 1280, 720)
+
+ # D. Batch baru saja selesai
+ elif not batch_mgr.is_active and prev_active:
+ final_count = counter.loading_count
+ now_str = datetime.now().strftime("%H:%M:%S")
+ print(f"\n>>> [{now_str}] BATCH #{batch_counter} SELESAI. Total karung: {final_count} <<<")
+ recorder.stop_recording()
+
+ # Reset counter dan stabilizer untuk batch berikutnya
+ counter.reset()
+ stabilizer.reset()
+
+ # E. Log transisi status
+ if batch_mgr.state != prev_state:
+ now_str = datetime.now().strftime("%H:%M:%S")
+ print(f"[{now_str}] [STATE] {prev_state} -> {batch_mgr.state}")
+
+ # =============================================================
+ # TULIS FRAME KE VIDEO (jika batch aktif)
+ # =============================================================
+ if batch_mgr.is_active and recorder.is_recording:
+ recorder.write_frame(frame)
+
+ except KeyboardInterrupt:
+ print("\n\n[INFO] Program dihentikan oleh pengguna (Ctrl+C).")
+
+ finally:
+ # Tutup rekaman yang masih terbuka
+ if recorder.is_recording:
+ print("[INFO] Menyimpan rekaman batch terakhir...")
+ recorder.stop_recording()
+
+ cap.release()
+
+ print("\n" + "=" * 60)
+ print(" BATCH VIDEO CROPPER SELESAI")
+ print(f" Total batch terekam: {batch_counter}")
+ print(f" Total frame diproses: {frame_idx}")
+ print(f" Folder output: {ARCHIVE_BASE}")
+ print("=" * 60)
+
+
+if __name__ == "__main__":
+ run_batch_video_cropper()
diff --git a/algoritma-batch/cfg/tracker.yaml b/algoritma-batch/cfg/tracker.yaml
new file mode 100644
index 0000000..d654afd
--- /dev/null
+++ b/algoritma-batch/cfg/tracker.yaml
@@ -0,0 +1,26 @@
+# Custom FastTrack config tuned for sack counting:
+# - track_buffer=60: hold lost tracks for 60 frames (~2.4s at 25fps)
+# to survive worker occlusion
+# - new_track_thresh=0.3: harder to spawn duplicate IDs
+# - track_low_thresh=0.05: recover faint detections behind workers
+# - active_occ_to_lost_thresh=15: tolerate 15 occluded frames
+# - occ_reappear_window=60: re-find tracks after long occlusion
+# - enlarge_bbox_occ=1.15: widen search region during occlusion
+
+tracker_type: bytetrack
+track_high_thresh: 0.20
+track_low_thresh: 0.05
+new_track_thresh: 0.30
+track_buffer: 60
+match_thresh: 0.85
+fuse_score: true
+
+# Occlusion handling (FastTrack-specific)
+reset_velocity_offset_occ: 5
+reset_pos_offset_occ: 3
+enlarge_bbox_occ: 1.15
+dampen_motion_occ: 0.4
+active_occ_to_lost_thresh: 15
+occ_cover_thresh: 0.6
+occ_reappear_window: 60
+init_iou_suppress: 0.65
diff --git a/algoritma-batch/src/__init__.py b/algoritma-batch/src/__init__.py
new file mode 100644
index 0000000..25e317d
--- /dev/null
+++ b/algoritma-batch/src/__init__.py
@@ -0,0 +1 @@
+"""Package marker."""
diff --git a/algoritma-batch/src/batch.py b/algoritma-batch/src/batch.py
new file mode 100644
index 0000000..1dba93f
--- /dev/null
+++ b/algoritma-batch/src/batch.py
@@ -0,0 +1,390 @@
+"""Batch lifecycle manager — 4-state machine for truck+sack sessions.
+
+State machine:
+ IDLE ──truck detected──▶ TRUCK_STABILIZING ──stable 5s──▶ COUNTING_SACKS
+ ▲ │ truck gone │ ▲
+ │ └──────▶ IDLE │ │
+ │ │ │
+ │ 10s no sack activity │ │ sacks resume
+ │ ▼ │
+ │ WAITING_FOR_ACTIVITY
+ │ (batch OPEN)
+ │ │
+ └──────────────── truck leaves ─────────────────────────────┘
+ (batch finalized)
+"""
+
+from __future__ import annotations
+
+import time
+import math
+from dataclasses import dataclass, field
+from enum import Enum, auto
+
+
+class BatchState(Enum):
+ IDLE = auto()
+ TRUCK_STABILIZING = auto()
+ COUNTING_SACKS = auto()
+ WAITING_FOR_ACTIVITY = auto() # Paused: no sacks, but truck still here
+
+
+@dataclass
+class BatchRecord:
+ """Completed batch summary."""
+
+ batch_id: int
+ start_time: float
+ end_time: float
+ loading_count: int
+ unloading_count: int
+
+ @property
+ def net_count(self) -> int:
+ return self.loading_count - self.unloading_count
+
+ @property
+ def duration_seconds(self) -> float:
+ return self.end_time - self.start_time
+
+
+class BatchLifecycleManager:
+ """Manages batch transitions based on truck stability and sack activity.
+
+ State flow:
+ - IDLE: waiting for truck to appear in ROI polygon
+ - TRUCK_STABILIZING: truck seen, tracking centroid stability
+ - COUNTING_SACKS: actively counting sacks crossing line
+ - WAITING_FOR_ACTIVITY: sacks idle, but truck still present — batch stays open
+ """
+
+ def __init__(
+ self,
+ stabilize_seconds: float = 5.0,
+ stabilize_threshold_px: float = 15.0,
+ sack_idle_timeout: float = 10.0,
+ min_batch_duration: float = 30.0,
+ truck_gone_tolerance: float = 3.0,
+ timeout_seconds: float = 30.0, # kept for backward compat (unused)
+ ) -> None:
+ # Tunable parameters
+ self._stabilize_seconds = stabilize_seconds
+ self._stabilize_threshold_px = stabilize_threshold_px
+ self._sack_idle_timeout = sack_idle_timeout
+ self._min_batch_duration = min_batch_duration
+ self._truck_gone_tolerance = truck_gone_tolerance
+
+ # Internal state
+ self._state = BatchState.IDLE
+ self._batch_counter = 0
+ self._current_batch_id: int | None = None
+ self._batch_start_time = 0.0
+ self._history: list[BatchRecord] = []
+
+ # Truck stabilization tracking
+ self._truck_first_seen_time = 0.0
+ self._truck_last_centroid: tuple[float, float] | None = None
+ self._truck_stable_since = 0.0
+ self._truck_is_stable = False
+ self._truck_last_seen = 0.0 # timestamp when truck was last detected
+
+ # Sack activity tracking (for pause condition)
+ self._last_sack_crossing_time = 0.0
+ self._last_sack_seen_in_area_time = 0.0
+
+ # Waiting state tracking
+ self._waiting_since = 0.0
+
+ # Callbacks
+ self._on_batch_start: list = []
+ self._on_batch_end: list = []
+
+ # -- Public API: Register callbacks --
+
+ def on_batch_start(self, callback) -> None:
+ """Register callback: fn(batch_id, timestamp)."""
+ self._on_batch_start.append(callback)
+
+ def on_batch_end(self, callback) -> None:
+ """Register callback: fn(BatchRecord)."""
+ self._on_batch_end.append(callback)
+
+ # -- Public API: State update methods --
+
+ def update_truck(
+ self,
+ truck_detected: bool,
+ truck_centroid: tuple[float, float] | None,
+ timestamp: float,
+ ) -> None:
+ """Called during IDLE, TRUCK_STABILIZING, and WAITING_FOR_ACTIVITY states.
+
+ Args:
+ truck_detected: whether a truck is detected in the ROI polygon
+ truck_centroid: (cx, cy) of the truck bounding box, or None
+ timestamp: current time.time()
+ """
+ if self._state == BatchState.IDLE:
+ if truck_detected and truck_centroid is not None:
+ # Transition to STABILIZING
+ self._state = BatchState.TRUCK_STABILIZING
+ self._truck_first_seen_time = timestamp
+ self._truck_last_centroid = truck_centroid
+ self._truck_stable_since = timestamp
+ self._truck_is_stable = False
+ self._truck_last_seen = timestamp
+ print(f"[BATCH] Truk terdeteksi di area. Memantau stabilitas...")
+
+ elif self._state == BatchState.TRUCK_STABILIZING:
+ if truck_detected:
+ self._truck_last_seen = timestamp
+
+ # Check if truck has been gone for too long (tolerance)
+ time_since_last_seen = timestamp - self._truck_last_seen
+ if not truck_detected and time_since_last_seen >= self._truck_gone_tolerance:
+ print(f"[BATCH] Truk hilang selama {time_since_last_seen:.1f}s. Kembali ke IDLE.")
+ self._state = BatchState.IDLE
+ self._truck_last_centroid = None
+ self._truck_is_stable = False
+ return
+
+ if truck_centroid is not None and self._truck_last_centroid is not None:
+ # Calculate centroid displacement
+ dx = truck_centroid[0] - self._truck_last_centroid[0]
+ dy = truck_centroid[1] - self._truck_last_centroid[1]
+ displacement = math.sqrt(dx * dx + dy * dy)
+
+ if displacement > self._stabilize_threshold_px:
+ # Truck moved too much -> reset stability timer
+ self._truck_stable_since = timestamp
+ self._truck_is_stable = False
+
+ self._truck_last_centroid = truck_centroid
+
+ # Check if stable long enough
+ stable_duration = timestamp - self._truck_stable_since
+ if stable_duration >= self._stabilize_seconds:
+ if not self._truck_is_stable:
+ self._truck_is_stable = True
+ print(f"[BATCH] Truk stabil selama {stable_duration:.1f}s. Memulai counting...")
+ self._start_batch(timestamp)
+
+ elif self._state == BatchState.WAITING_FOR_ACTIVITY:
+ if truck_detected:
+ self._truck_last_seen = timestamp
+
+ # Check if truck has been gone for tolerance period
+ time_since_last_seen = timestamp - self._truck_last_seen
+ if not truck_detected and time_since_last_seen >= self._truck_gone_tolerance:
+ # Truck has truly left! NOW we finalize the batch.
+ wait_duration = timestamp - self._waiting_since
+ print(
+ f"[BATCH] Truk pergi setelah menunggu {wait_duration:.0f}s. "
+ f"Batch selesai."
+ )
+ self._end_batch(timestamp, self._pending_loading, self._pending_unloading)
+
+ def update_sacks(
+ self,
+ has_crossing_event: bool,
+ sacks_in_area_count: int,
+ timestamp: float,
+ loading_count: int = 0,
+ unloading_count: int = 0,
+ ) -> None:
+ """Called during COUNTING_SACKS and WAITING_FOR_ACTIVITY states.
+
+ Args:
+ has_crossing_event: True if a sack crossed the counting line this frame
+ sacks_in_area_count: number of sacks currently detected in truck area
+ timestamp: current time.time()
+ loading_count: current cumulative loading count
+ unloading_count: current cumulative unloading count
+ """
+ # WAITING_FOR_ACTIVITY: if sacks appear again, resume counting in the SAME batch
+ if self._state == BatchState.WAITING_FOR_ACTIVITY:
+ if has_crossing_event or sacks_in_area_count > 0:
+ wait_duration = timestamp - self._waiting_since
+ print(
+ f"[BATCH] Aktivitas karung terdeteksi setelah {wait_duration:.0f}s menunggu. "
+ f"Melanjutkan counting batch #{self._current_batch_id}..."
+ )
+ self._state = BatchState.COUNTING_SACKS
+ self._last_sack_crossing_time = timestamp
+ self._last_sack_seen_in_area_time = timestamp
+ # Fall through to counting logic below
+ else:
+ return
+
+ if self._state != BatchState.COUNTING_SACKS:
+ return
+
+ # Update activity timers
+ if has_crossing_event:
+ self._last_sack_crossing_time = timestamp
+
+ if sacks_in_area_count > 0:
+ self._last_sack_seen_in_area_time = timestamp
+
+ # Store latest counts for when batch eventually ends
+ self._pending_loading = loading_count
+ self._pending_unloading = unloading_count
+
+ # Check pause condition: no sack activity for timeout period
+ batch_duration = timestamp - self._batch_start_time
+ time_since_last_crossing = timestamp - self._last_sack_crossing_time
+ time_since_last_sack_seen = timestamp - self._last_sack_seen_in_area_time
+
+ if (
+ batch_duration >= self._min_batch_duration
+ and time_since_last_crossing >= self._sack_idle_timeout
+ and time_since_last_sack_seen >= self._sack_idle_timeout
+ ):
+ print(
+ f"[BATCH] Tidak ada aktivitas karung selama {self._sack_idle_timeout}s. "
+ f"Menunggu truk pergi atau palet selanjutnya..."
+ )
+ self._state = BatchState.WAITING_FOR_ACTIVITY
+ self._waiting_since = timestamp
+
+ # -- Public API: Backward-compatible update (legacy) --
+
+ def update(
+ self,
+ truck_detected: bool,
+ timestamp: float,
+ loading_count: int = 0,
+ unloading_count: int = 0,
+ ) -> None:
+ """Legacy update method — kept for backward compatibility."""
+ if self._state in (BatchState.IDLE, BatchState.TRUCK_STABILIZING):
+ self.update_truck(truck_detected, None, timestamp)
+ elif self._state in (BatchState.COUNTING_SACKS, BatchState.WAITING_FOR_ACTIVITY):
+ self.update_sacks(
+ has_crossing_event=False,
+ sacks_in_area_count=1 if truck_detected else 0,
+ timestamp=timestamp,
+ loading_count=loading_count,
+ unloading_count=unloading_count,
+ )
+
+ # -- Properties --
+
+ @property
+ def state(self) -> str:
+ """Return current state as human-readable string."""
+ return self._state.name
+
+ @property
+ def current_batch_id(self) -> int | None:
+ return self._current_batch_id
+
+ @property
+ def is_active(self) -> bool:
+ """True during COUNTING or WAITING (batch is still open)."""
+ return self._state in (BatchState.COUNTING_SACKS, BatchState.WAITING_FOR_ACTIVITY)
+
+ @property
+ def is_counting(self) -> bool:
+ """True only during active sack counting."""
+ return self._state == BatchState.COUNTING_SACKS
+
+ @property
+ def is_waiting(self) -> bool:
+ """True when paused waiting for next pallet or truck departure."""
+ return self._state == BatchState.WAITING_FOR_ACTIVITY
+
+ @property
+ def is_stabilizing(self) -> bool:
+ return self._state == BatchState.TRUCK_STABILIZING
+
+ @property
+ def history(self) -> list[BatchRecord]:
+ return list(self._history)
+
+ @property
+ def batch_duration(self) -> float:
+ """Duration of current batch in seconds (0 if not active)."""
+ if not self.is_active:
+ return 0.0
+ return time.time() - self._batch_start_time
+
+ @property
+ def time_since_last_sack_activity(self) -> float:
+ """Seconds since last sack crossed line or seen in area."""
+ if not self.is_active:
+ return 0.0
+ now = time.time()
+ last_activity = max(self._last_sack_crossing_time, self._last_sack_seen_in_area_time)
+ return now - last_activity if last_activity > 0 else 0.0
+
+ @property
+ def waiting_duration(self) -> float:
+ """How long we've been in WAITING_FOR_ACTIVITY state."""
+ if self._state != BatchState.WAITING_FOR_ACTIVITY:
+ return 0.0
+ return time.time() - self._waiting_since
+
+ @property
+ def stabilize_progress(self) -> float:
+ """Progress of truck stabilization (0.0 to 1.0)."""
+ if self._state != BatchState.TRUCK_STABILIZING:
+ return 0.0
+ if self._stabilize_seconds <= 0.0:
+ return 1.0
+ elapsed = time.time() - self._truck_stable_since
+ return min(1.0, elapsed / self._stabilize_seconds)
+
+ def resume_batch(
+ self,
+ batch_id: int,
+ start_time: float,
+ loading_count: int,
+ unloading_count: int,
+ ) -> None:
+ """Resume a previously finalized batch."""
+ self._current_batch_id = batch_id
+ self._batch_counter = max(self._batch_counter, batch_id)
+ self._batch_start_time = start_time
+ self._pending_loading = loading_count
+ self._pending_unloading = unloading_count
+ self._state = BatchState.COUNTING_SACKS
+
+ # Pop from history if it was just completed
+ if self._history and self._history[-1].batch_id == batch_id:
+ self._history.pop()
+
+ # -- Private methods --
+
+ def _start_batch(self, timestamp: float) -> None:
+ self._batch_counter += 1
+ self._current_batch_id = self._batch_counter
+ self._batch_start_time = timestamp
+ self._last_sack_crossing_time = timestamp # Grace period
+ self._last_sack_seen_in_area_time = timestamp # Grace period
+ self._pending_loading = 0
+ self._pending_unloading = 0
+ self._state = BatchState.COUNTING_SACKS
+ for cb in self._on_batch_start:
+ cb(self._current_batch_id, timestamp)
+
+ def _end_batch(
+ self,
+ timestamp: float,
+ loading_count: int,
+ unloading_count: int,
+ ) -> None:
+ record = BatchRecord(
+ batch_id=self._current_batch_id or 0,
+ start_time=self._batch_start_time,
+ end_time=timestamp,
+ loading_count=loading_count,
+ unloading_count=unloading_count,
+ )
+ self._history.append(record)
+ self._state = BatchState.IDLE
+ self._current_batch_id = None
+ self._truck_last_centroid = None
+ self._truck_is_stable = False
+ for cb in self._on_batch_end:
+ cb(record)
diff --git a/algoritma-batch/src/counting.py b/algoritma-batch/src/counting.py
new file mode 100644
index 0000000..ff6ac33
--- /dev/null
+++ b/algoritma-batch/src/counting.py
@@ -0,0 +1,237 @@
+"""Line-crossing counter — hybrid zone-based state tracking.
+
+Counting logic (Low-FPS robust):
+ Uses y1 (top edge) of the stabilized sack bounding box.
+
+ Each track_id goes through states:
+ UNKNOWN → ABOVE → COUNTED (when seen below line)
+ UNKNOWN → BELOW (ghost/appeared below line first → never counted)
+
+ Loading: track had state ABOVE, now detected BELOW the zone
+ Unloading: track had state BELOW, now detected ABOVE the zone (if needed)
+
+ 3-Layer deduplication:
+ Layer 1: State guard — must have been ABOVE before counting
+ Layer 2: Spatial dedup radius — same position can't trigger twice
+ Layer 3: Track ID — one track_id can only be counted once per direction
+
+ This approach is immune to low FPS because it doesn't require
+ detecting the exact frame of crossing. It only needs the track
+ to have been seen ABOVE the line at ANY point in its lifetime.
+"""
+
+from __future__ import annotations
+
+import time
+
+from src.interfaces import Detection
+
+
+class LineCrossCounter:
+ """Counts sacks crossing a horizontal zone using y1 (top edge).
+
+ The zone is a band [line_y - margin, line_y + margin].
+ A sack is "above" if y1 < line_y - margin,
+ "below" if y1 > line_y + margin.
+ While y1 is inside the band, state is held (no trigger).
+
+ Loading = track was ever "above", now "below" (entered truck)
+ Unloading = track was ever "below", now "above" (left truck)
+ """
+
+ def __init__(
+ self,
+ line_y: int,
+ line_x_start: int,
+ line_x_end: int,
+ margin: int = 20,
+ dedup_radius: float = 60.0,
+ ) -> None:
+ self._line_y = line_y
+ self._line_x_start = line_x_start
+ self._line_x_end = line_x_end
+ self._margin = margin
+ self._dedup_radius = dedup_radius
+
+ self._loading_count = 0
+ self._unloading_count = 0
+
+ # track_id -> zone state for y1: "above" | "below" | None
+ self._state: dict[int, str | None] = {}
+ # track_id -> whether this track has EVER been in each zone
+ self._has_been_above: dict[int, bool] = {}
+ self._has_been_below: dict[int, bool] = {}
+ # track_id -> set of directions already counted
+ self._counted: dict[int, set[str]] = {}
+ # track_id -> initial coordinates (cx, y1) when first tracked
+ self._entry_points: dict[int, tuple[float, float]] = {}
+ # list of active deduplication circles
+ self._dedup_circles: list[dict] = []
+
+ @property
+ def entry_points(self) -> dict[int, tuple[float, float]]:
+ return self._entry_points
+
+ @property
+ def counted_tracks(self) -> dict[int, set[str]]:
+ return self._counted
+
+ @property
+ def line_y(self) -> int:
+ return self._line_y
+
+ @line_y.setter
+ def line_y(self, value: int) -> None:
+ self._line_y = value
+
+ @property
+ def line_x_start(self) -> int:
+ return self._line_x_start
+
+ @line_x_start.setter
+ def line_x_start(self, value: int) -> None:
+ self._line_x_start = value
+
+ @property
+ def line_x_end(self) -> int:
+ return self._line_x_end
+
+ @line_x_end.setter
+ def line_x_end(self, value: int) -> None:
+ self._line_x_end = value
+
+ def update(self, detections: list[Detection]) -> list[dict]:
+ """Process detections, return list of crossing events.
+
+ Hybrid approach:
+ - Tracks zone state per frame (above/below/in-band)
+ - BUT uses accumulated history (has_been_above) for counting decision
+ - A track counts as "loading" when:
+ 1. It has been seen ABOVE the line at any previous point
+ 2. Its current y1 is now BELOW the line
+ 3. It hasn't been counted for loading yet
+ 4. It passes spatial dedup check
+ """
+ now_t = time.time()
+ events: list[dict] = []
+ upper = self._line_y - self._margin
+ lower = self._line_y + self._margin
+
+ # Clean up expired dedup circles (older than 3.0 seconds)
+ self._dedup_circles = [c for c in self._dedup_circles if (now_t - c["time"]) <= 3.0]
+
+ for det in detections:
+ if det.track_id is None:
+ continue
+
+ x1, y1, x2, y2 = det.bbox
+ cx = (x1 + x2) / 2.0
+ tid = det.track_id
+
+ if tid not in self._entry_points:
+ self._entry_points[tid] = (cx, y1)
+
+ # Skip if centroid X outside counting bounds
+ if cx < self._line_x_start or cx > self._line_x_end:
+ continue
+
+ counted_dirs = self._counted.setdefault(tid, set())
+
+ # Determine y1 zone state (top edge of sack bbox)
+ if y1 < upper:
+ new_state = "above"
+ elif y1 > lower:
+ new_state = "below"
+ else:
+ new_state = self._state.get(tid) # in band: hold
+
+ prev_state = self._state.get(tid)
+ self._state[tid] = new_state
+
+ # Track zone history — CRITICAL for low-FPS robustness
+ # Once a track has been seen above/below, it stays recorded forever
+ if new_state == "above":
+ self._has_been_above[tid] = True
+ elif new_state == "below":
+ self._has_been_below[tid] = True
+
+ # --- HYBRID COUNTING LOGIC ---
+ # Loading: track was EVER above, NOW below (entered truck from top)
+ # This works even if the track jumped over the line between frames
+ is_loading = (
+ new_state == "below"
+ and self._has_been_above.get(tid, False)
+ and "loading" not in counted_dirs
+ )
+
+ # Unloading: track was EVER below, NOW above (left truck)
+ is_unloading = (
+ new_state == "above"
+ and self._has_been_below.get(tid, False)
+ and "unloading" not in counted_dirs
+ )
+
+ if is_loading or is_unloading:
+ # Check spatial distance against all active dedup circles
+ is_duplicate = False
+ for circle in self._dedup_circles:
+ dist = ((cx - circle["x"]) ** 2 + (y1 - circle["y"]) ** 2) ** 0.5
+ if dist <= self._dedup_radius:
+ is_duplicate = True
+ break
+
+ if is_duplicate:
+ continue
+
+ # Add this coordinate to the active dedup circles
+ self._dedup_circles.append({
+ "x": cx,
+ "y": y1,
+ "time": now_t,
+ "track_id": tid
+ })
+
+ if is_loading:
+ self._loading_count += 1
+ counted_dirs.add("loading")
+ events.append({
+ "track_id": tid,
+ "direction": "loading",
+ "cx": cx,
+ "cy": y1
+ })
+
+ elif is_unloading:
+ self._unloading_count += 1
+ counted_dirs.add("unloading")
+ events.append({
+ "track_id": tid,
+ "direction": "unloading",
+ "cx": cx,
+ "cy": y1
+ })
+
+ return events
+
+ @property
+ def loading_count(self) -> int:
+ return self._loading_count
+
+ @property
+ def unloading_count(self) -> int:
+ return self._unloading_count
+
+ @property
+ def net_count(self) -> int:
+ return self._loading_count - self._unloading_count
+
+ def reset(self) -> None:
+ """Reset all counters (new batch)."""
+ self._loading_count = 0
+ self._unloading_count = 0
+ self._state.clear()
+ self._has_been_above.clear()
+ self._has_been_below.clear()
+ self._counted.clear()
+ self._entry_points.clear()
+ self._dedup_circles.clear()
diff --git a/algoritma-batch/src/dashboard.py b/algoritma-batch/src/dashboard.py
new file mode 100644
index 0000000..152503a
--- /dev/null
+++ b/algoritma-batch/src/dashboard.py
@@ -0,0 +1,251 @@
+"""Dashboard overlay — draws counting info onto the video frame.
+
+Draws: truck ROI, counting zone (band), sack bounding boxes with y1
+marker (the crossing trigger edge), stats panel, batch history,
+and system state indicator.
+"""
+
+from __future__ import annotations
+
+import cv2
+import numpy as np
+
+from src.batch import BatchRecord
+from src.interfaces import Detection
+from src.truck_roi import TruckROI
+
+
+# Colors (BGR)
+GREEN = (0, 200, 0)
+RED = (0, 0, 220)
+CYAN = (220, 200, 0)
+WHITE = (255, 255, 255)
+YELLOW = (0, 230, 255)
+MAGENTA = (255, 0, 255)
+ORANGE = (0, 165, 255)
+GRAY = (140, 140, 140)
+DARK_GREEN = (0, 130, 0)
+LIGHT_BLUE = (255, 200, 100)
+
+
+# State display labels and colors
+STATE_DISPLAY = {
+ "IDLE": ("MENCARI TRUK...", ORANGE),
+ "TRUCK_STABILIZING": ("TRUK TERDETEKSI - STABILISASI", YELLOW),
+ "COUNTING_SACKS": ("MENGHITUNG KARUNG", GREEN),
+ "WAITING_FOR_ACTIVITY": ("MENUNGGU PALET / TRUK PERGI", LIGHT_BLUE),
+}
+
+
+class DashboardOverlay:
+ """Draws detection boxes, ROI, counting line, and stats onto frames."""
+
+ def draw(
+ self,
+ frame: np.ndarray,
+ detections: list[Detection],
+ roi: TruckROI | None,
+ loading_count: int,
+ unloading_count: int,
+ batch_id: int | None,
+ history: list[BatchRecord] | None = None,
+ system_state: str = "IDLE",
+ batch_duration: float = 0.0,
+ idle_timer: float = 0.0,
+ stabilize_progress: float = 0.0,
+ waiting_duration: float = 0.0,
+ ) -> np.ndarray:
+ out = frame.copy()
+ if roi is not None:
+ self._draw_roi(out, roi)
+ self._draw_counting_zone(out, roi)
+ self._draw_detections(out, detections)
+ self._draw_stats(out, loading_count, unloading_count, batch_id, history)
+ self._draw_system_state(
+ out, system_state, batch_duration, idle_timer, stabilize_progress,
+ waiting_duration,
+ )
+ return out
+
+ def _draw_roi(self, frame: np.ndarray, roi: TruckROI) -> None:
+ cv2.rectangle(
+ frame, (roi.x1, roi.y1), (roi.x2, roi.y2), ORANGE, 2,
+ )
+ cv2.putText(
+ frame, f"TRUCK ROI ({roi.confidence:.0%})",
+ (roi.x1, roi.y1 - 8),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.5, ORANGE, 1,
+ )
+ # Draw truck top reference line (dashed via short segments)
+ for x in range(roi.x1, roi.x2, 20):
+ cv2.line(frame, (x, roi.y1), (min(x + 10, roi.x2), roi.y1), GRAY, 1)
+
+ def _draw_counting_zone(
+ self, frame: np.ndarray, roi: TruckROI, margin: int = 20,
+ ) -> None:
+ y = roi.line_y
+ # Draw zone band (semi-transparent)
+ overlay = frame.copy()
+ cv2.rectangle(
+ overlay, (roi.x1, y - margin), (roi.x2, y + margin),
+ MAGENTA, -1,
+ )
+ cv2.addWeighted(overlay, 0.15, frame, 0.85, 0, frame)
+ # Draw center line
+ cv2.line(frame, (roi.x1, y), (roi.x2, y), MAGENTA, 2)
+ cv2.putText(
+ frame,
+ f"COUNT LINE Y={y} (y1 trigger)",
+ (roi.x1, y - margin - 8),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.5, MAGENTA, 1,
+ )
+
+ def _draw_detections(
+ self,
+ frame: np.ndarray,
+ detections: list[Detection],
+ ) -> None:
+ for det in detections:
+ x1, y1, x2, y2 = [int(v) for v in det.bbox]
+ label = "sack"
+ if det.track_id is not None:
+ label += f" #{det.track_id}"
+ label += f" {det.confidence:.0%}"
+
+ cv2.rectangle(frame, (x1, y1), (x2, y2), CYAN, 2)
+ cv2.putText(
+ frame, label, (x1, y1 - 6),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.4, CYAN, 1,
+ )
+ # Mark TOP edge (y1) — the crossing trigger
+ cv2.line(frame, (x1, y1), (x2, y1), GREEN, 3)
+
+ def _draw_stats(
+ self,
+ frame: np.ndarray,
+ loading: int,
+ unloading: int,
+ batch_id: int | None,
+ history: list[BatchRecord] | None,
+ ) -> None:
+ # Panel background
+ cv2.rectangle(frame, (10, 10), (320, 160), (0, 0, 0), -1)
+ cv2.rectangle(frame, (10, 10), (320, 160), WHITE, 1)
+
+ batch_text = f"Batch #{batch_id}" if batch_id else "IDLE"
+ net = loading - unloading
+ y0 = 35
+
+ cv2.putText(
+ frame, batch_text, (20, y0),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.7, YELLOW, 2,
+ )
+ cv2.putText(
+ frame, f"Loading: {loading}", (20, y0 + 30),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.6, GREEN, 2,
+ )
+ cv2.putText(
+ frame, f"Unloading: {unloading}", (20, y0 + 60),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.6, RED, 2,
+ )
+ cv2.putText(
+ frame, f"Net: {net}", (20, y0 + 90),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.6, WHITE, 2,
+ )
+
+ # History (last 3 batches)
+ if history:
+ y_h = 180
+ cv2.putText(
+ frame, "HISTORY", (20, y_h),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.5, YELLOW, 1,
+ )
+ for rec in history[-3:]:
+ y_h += 22
+ txt = (
+ f"B#{rec.batch_id}: "
+ f"L={rec.loading_count} "
+ f"U={rec.unloading_count} "
+ f"Net={rec.net_count}"
+ )
+ cv2.putText(
+ frame, txt, (20, y_h),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.4, WHITE, 1,
+ )
+
+ def _draw_system_state(
+ self,
+ frame: np.ndarray,
+ system_state: str,
+ batch_duration: float,
+ idle_timer: float,
+ stabilize_progress: float,
+ waiting_duration: float = 0.0,
+ ) -> None:
+ """Draw system state indicator bar at bottom of frame."""
+ h, w = frame.shape[:2]
+
+ # Get display info for current state
+ label, color = STATE_DISPLAY.get(system_state, ("UNKNOWN", GRAY))
+
+ # Draw state bar background
+ bar_h = 36
+ bar_y = h - bar_h
+ overlay = frame.copy()
+ cv2.rectangle(overlay, (0, bar_y), (w, h), (0, 0, 0), -1)
+ cv2.addWeighted(overlay, 0.7, frame, 0.3, 0, frame)
+
+ # Draw colored indicator dot
+ cv2.circle(frame, (20, bar_y + bar_h // 2), 8, color, -1)
+ cv2.circle(frame, (20, bar_y + bar_h // 2), 8, WHITE, 1)
+
+ # Draw state label
+ cv2.putText(
+ frame, label, (36, bar_y + bar_h // 2 + 5),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.55, color, 2,
+ )
+
+ # Draw additional info based on state
+ if system_state == "TRUCK_STABILIZING":
+ # Draw stabilization progress bar
+ prog_x = 340
+ prog_w = 150
+ prog_h = 14
+ prog_y = bar_y + (bar_h - prog_h) // 2
+
+ cv2.rectangle(frame, (prog_x, prog_y), (prog_x + prog_w, prog_y + prog_h), GRAY, 1)
+ fill_w = int(prog_w * stabilize_progress)
+ if fill_w > 0:
+ cv2.rectangle(frame, (prog_x, prog_y), (prog_x + fill_w, prog_y + prog_h), YELLOW, -1)
+
+ pct_text = f"{stabilize_progress * 100:.0f}%"
+ cv2.putText(
+ frame, pct_text, (prog_x + prog_w + 8, prog_y + prog_h - 2),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.4, YELLOW, 1,
+ )
+
+ elif system_state == "COUNTING_SACKS":
+ # Draw batch duration and idle timer
+ info_x = 340
+ dur_text = f"Durasi: {batch_duration:.0f}s"
+ cv2.putText(
+ frame, dur_text, (info_x, bar_y + 15),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.4, WHITE, 1,
+ )
+
+ if idle_timer > 0:
+ idle_color = RED if idle_timer > 7.0 else (YELLOW if idle_timer > 4.0 else WHITE)
+ idle_text = f"Idle: {idle_timer:.1f}s / 10s"
+ cv2.putText(
+ frame, idle_text, (info_x, bar_y + 30),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.4, idle_color, 1,
+ )
+
+ elif system_state == "WAITING_FOR_ACTIVITY":
+ # Show waiting duration and batch info
+ info_x = 360
+ wait_text = f"Menunggu: {waiting_duration:.0f}s | Batch masih terbuka"
+ cv2.putText(
+ frame, wait_text, (info_x, bar_y + bar_h // 2 + 5),
+ cv2.FONT_HERSHEY_SIMPLEX, 0.4, LIGHT_BLUE, 1,
+ )
diff --git a/algoritma-batch/src/detection.py b/algoritma-batch/src/detection.py
new file mode 100644
index 0000000..ade6a95
--- /dev/null
+++ b/algoritma-batch/src/detection.py
@@ -0,0 +1,82 @@
+"""YOLO-based detectors for sacks and trucks.
+
+Each detector is a single-responsibility unit (S). New model types can be
+added as new classes without touching these (O).
+"""
+
+from __future__ import annotations
+
+import numpy as np
+from ultralytics import YOLO
+
+from src.interfaces import Detection
+
+
+class SackDetector:
+ """Detects sacks (and persons) using a YOLO segmentation model."""
+
+ def __init__(self, model_path: str, conf: float = 0.40) -> None:
+ self._model = YOLO(model_path)
+ self._conf = conf
+
+ def detect(self, frame: np.ndarray) -> list[Detection]:
+ results = self._model.predict(
+ frame, conf=self._conf, verbose=False
+ )
+ return self._parse(results[0])
+
+ def _parse(self, result) -> list[Detection]:
+ detections: list[Detection] = []
+ masks = result.masks
+ for i, box in enumerate(result.boxes):
+ cls_id = int(box.cls[0])
+ name = self._model.names[cls_id]
+ if name != "sack":
+ continue
+ x1, y1, x2, y2 = box.xyxy[0].tolist()
+ mask = None
+ if masks is not None and i < len(masks):
+ mask = masks[i].data.cpu().numpy().squeeze()
+ detections.append(
+ Detection(
+ bbox=(x1, y1, x2, y2),
+ confidence=float(box.conf[0]),
+ class_id=cls_id,
+ class_name=name,
+ mask=mask,
+ )
+ )
+ return detections
+
+
+class TruckDetector:
+ """Detects trucks using a YOLO detection model."""
+
+ def __init__(self, model_path_or_model: str | YOLO, conf: float = 0.50) -> None:
+ if isinstance(model_path_or_model, str):
+ self._model = YOLO(model_path_or_model)
+ else:
+ self._model = model_path_or_model
+ self._conf = conf
+
+ def detect(self, frame: np.ndarray) -> list[Detection]:
+ results = self._model.predict(
+ frame, conf=self._conf, verbose=False
+ )
+ return self._parse(results[0])
+
+ def _parse(self, result) -> list[Detection]:
+ detections: list[Detection] = []
+ for box in result.boxes:
+ cls_id = int(box.cls[0])
+ name = self._model.names[cls_id]
+ x1, y1, x2, y2 = box.xyxy[0].tolist()
+ detections.append(
+ Detection(
+ bbox=(x1, y1, x2, y2),
+ confidence=float(box.conf[0]),
+ class_id=cls_id,
+ class_name=name,
+ )
+ )
+ return detections
diff --git a/algoritma-batch/src/interfaces.py b/algoritma-batch/src/interfaces.py
new file mode 100644
index 0000000..7a568dd
--- /dev/null
+++ b/algoritma-batch/src/interfaces.py
@@ -0,0 +1,98 @@
+"""Abstract interfaces — all components code against these, never concretions.
+
+Keeps Interface Segregation (I) and Dependency Inversion (D) satisfied.
+Each protocol is tiny and single-purpose (S).
+"""
+
+from __future__ import annotations
+
+from dataclasses import dataclass, field
+from typing import Protocol, runtime_checkable
+
+import numpy as np
+
+
+# ── Data transfer objects ────────────────────────────────────────────────
+
+
+@dataclass
+class Detection:
+ """Single object detection."""
+
+ bbox: tuple[float, float, float, float] # x1, y1, x2, y2
+ confidence: float
+ class_id: int
+ class_name: str
+ track_id: int | None = None
+ mask: np.ndarray | None = None # segmentation mask (optional)
+
+
+@dataclass
+class FrameResult:
+ """All detections for one frame."""
+
+ detections: list[Detection] = field(default_factory=list)
+ frame_index: int = 0
+ timestamp: float = 0.0
+
+
+# ── Protocols ────────────────────────────────────────────────────────────
+
+
+@runtime_checkable
+class StreamSource(Protocol):
+ """Reads frames from a video source."""
+
+ def open(self) -> bool: ...
+ def read(self) -> tuple[bool, np.ndarray | None]: ...
+ def release(self) -> None: ...
+ @property
+ def fps(self) -> float: ...
+ @property
+ def frame_size(self) -> tuple[int, int]: ...
+
+
+@runtime_checkable
+class Detector(Protocol):
+ """Runs inference on a frame and returns detections."""
+
+ def detect(self, frame: np.ndarray) -> list[Detection]: ...
+
+
+@runtime_checkable
+class Tracker(Protocol):
+ """Assigns persistent IDs to detections across frames."""
+
+ def update(
+ self, frame: np.ndarray, detections: list[Detection]
+ ) -> list[Detection]: ...
+
+ def reset(self) -> None: ...
+
+
+@runtime_checkable
+class Counter(Protocol):
+ """Counts objects crossing a virtual boundary."""
+
+ def update(self, detections: list[Detection]) -> None: ...
+
+ @property
+ def loading_count(self) -> int: ...
+
+ @property
+ def unloading_count(self) -> int: ...
+
+ def reset(self) -> None: ...
+
+
+@runtime_checkable
+class BatchManager(Protocol):
+ """Manages batch lifecycle based on truck presence."""
+
+ def update(self, truck_detected: bool, timestamp: float) -> None: ...
+
+ @property
+ def current_batch_id(self) -> int | None: ...
+
+ @property
+ def is_active(self) -> bool: ...
diff --git a/algoritma-batch/src/stabilizer.py b/algoritma-batch/src/stabilizer.py
new file mode 100644
index 0000000..348d2a1
--- /dev/null
+++ b/algoritma-batch/src/stabilizer.py
@@ -0,0 +1,126 @@
+"""Bounding-box stabilizer — EMA smoothing + dual height clamping + dropout hold.
+
+Addresses four occlusion/flickering problems:
+1. Bbox jitter: raw detections jump 10-50px between frames.
+ Fix: EMA (exponential moving average) on bbox coordinates.
+2. Bbox loss at line: worker's head/body blocks sack for 5-10 frames.
+ Fix: hold last known smoothed bbox for `max_hold_frames` (10 frames).
+3. Height expansion spike: worker body merges into sack bbox.
+ Fix: clamp height expansion (`raw_h > smooth_h * max_h_ratio`).
+4. Height shrinkage collapse: worker head/shoulder covers bottom of sack.
+ Fix: clamp height shrinkage (`raw_h < smooth_h * min_h_ratio`).
+
+All rules are applied per track ID to maintain smooth trajectories.
+"""
+
+from __future__ import annotations
+
+from src.interfaces import Detection
+
+
+class BboxStabilizer:
+ """Smooths and holds bounding boxes per track ID against worker occlusion."""
+
+ def __init__(
+ self,
+ ema_alpha: float = 0.35,
+ max_hold_frames: int = 10,
+ max_height_ratio: float = 1.5,
+ min_height_ratio: float = 0.70,
+ ) -> None:
+ self._alpha = ema_alpha
+ self._max_hold = max_hold_frames
+ self._max_h_ratio = max_height_ratio
+ self._min_h_ratio = min_height_ratio
+
+ # tid -> (smoothed_x1, smoothed_y1, smoothed_x2, smoothed_y2)
+ self._smooth: dict[int, tuple[float, float, float, float]] = {}
+ # tid -> frames since last real detection
+ self._age: dict[int, int] = {}
+ # tid -> last confidence and class info
+ self._meta: dict[int, tuple[float, int, str]] = {}
+
+ def update(self, detections: list[Detection]) -> list[Detection]:
+ """Smooth incoming detections + inject held tracks during occlusion."""
+ seen_tids: set[int] = set()
+ result: list[Detection] = []
+
+ # 1. Process real detections — apply EMA & dual height clamping
+ for det in detections:
+ tid = det.track_id
+ if tid is None:
+ result.append(det)
+ continue
+
+ seen_tids.add(tid)
+ self._age[tid] = 0
+ self._meta[tid] = (det.confidence, det.class_id, det.class_name)
+
+ x1, y1, x2, y2 = det.bbox
+ raw_h = y2 - y1
+
+ if tid in self._smooth:
+ sx1, sy1, sx2, sy2 = self._smooth[tid]
+ smooth_h = sy2 - sy1
+
+ if smooth_h > 0:
+ # Height expansion clamp (worker body merged)
+ if raw_h > smooth_h * self._max_h_ratio:
+ y2 = y1 + smooth_h * self._max_h_ratio
+ # Height shrinkage clamp (worker head/shoulder blocking bottom)
+ elif raw_h < smooth_h * self._min_h_ratio:
+ y2 = y1 + smooth_h * self._min_h_ratio
+
+ a = self._alpha
+ sx1 = a * x1 + (1 - a) * sx1
+ sy1 = a * y1 + (1 - a) * sy1
+ sx2 = a * x2 + (1 - a) * sx2
+ sy2 = a * y2 + (1 - a) * sy2
+ else:
+ sx1, sy1, sx2, sy2 = x1, y1, x2, y2
+
+ self._smooth[tid] = (sx1, sy1, sx2, sy2)
+
+ result.append(Detection(
+ bbox=(sx1, sy1, sx2, sy2),
+ confidence=det.confidence,
+ class_id=det.class_id,
+ class_name=det.class_name,
+ track_id=tid,
+ mask=det.mask,
+ ))
+
+ # 2. Hold tracks missing this frame (occlusion tolerance)
+ expired: list[int] = []
+ for tid in list(self._age.keys()):
+ if tid in seen_tids:
+ continue
+ self._age[tid] += 1
+ if self._age[tid] > self._max_hold:
+ expired.append(tid)
+ continue
+
+ # Inject held bbox from last smoothed position
+ sx1, sy1, sx2, sy2 = self._smooth[tid]
+ conf, cls_id, cls_name = self._meta[tid]
+ result.append(Detection(
+ bbox=(sx1, sy1, sx2, sy2),
+ confidence=conf * 0.85, # gentle decay during occlusion hold
+ class_id=cls_id,
+ class_name=cls_name,
+ track_id=tid,
+ ))
+
+ # 3. Clean up expired tracks
+ for tid in expired:
+ del self._smooth[tid]
+ del self._age[tid]
+ del self._meta[tid]
+
+ return result
+
+ def reset(self) -> None:
+ """Clear all state (new batch)."""
+ self._smooth.clear()
+ self._age.clear()
+ self._meta.clear()
diff --git a/algoritma-batch/src/tracking.py b/algoritma-batch/src/tracking.py
new file mode 100644
index 0000000..30ead0e
--- /dev/null
+++ b/algoritma-batch/src/tracking.py
@@ -0,0 +1,83 @@
+"""FastTrack wrapper — occlusion-aware tracker with custom tuning.
+
+Uses Ultralytics FastTrack which handles:
+ - Kalman rollback on occlusion onset (restores pre-occlusion velocity)
+ - Enlarged search region during occlusion
+ - Re-identification of occluded tracks after reappearance
+
+Our custom cfg/tracker.yaml tunes:
+ - track_buffer=60 (hold lost tracks ~2.4s to survive worker occlusion)
+ - new_track_thresh=0.3 (prevent duplicate IDs from spawning)
+ - active_occ_to_lost_thresh=15 (tolerate 15 occluded frames)
+"""
+
+from __future__ import annotations
+
+import os
+
+import numpy as np
+from ultralytics import YOLO
+
+from src.interfaces import Detection
+
+_TRACKER_CFG = os.path.join(
+ os.path.dirname(os.path.dirname(__file__)), "cfg", "tracker.yaml"
+)
+
+
+class ByteTrackTracker:
+ """Tracks sacks across frames using FastTrack (occlusion-aware)."""
+
+ def __init__(self, model_path_or_model: str | YOLO, conf: float = 0.35) -> None:
+ if isinstance(model_path_or_model, str):
+ self._model = YOLO(model_path_or_model)
+ self._model_path = model_path_or_model
+ else:
+ self._model = model_path_or_model
+ self._model_path = model_path_or_model.ckpt_path if hasattr(model_path_or_model, 'ckpt_path') else ""
+ self._conf = conf
+ self._tracker_cfg = _TRACKER_CFG
+
+ def update(
+ self, frame: np.ndarray, detections: list[Detection]
+ ) -> list[Detection]:
+ """Run tracking on the frame, return detections with track IDs."""
+ results = self._model.track(
+ frame,
+ conf=self._conf,
+ persist=True,
+ tracker=self._tracker_cfg,
+ verbose=False,
+ )
+ return self._parse(results[0])
+
+ def _parse(self, result) -> list[Detection]:
+ tracked: list[Detection] = []
+ ids = result.boxes.id
+ for i, box in enumerate(result.boxes):
+ cls_id = int(box.cls[0])
+ name = self._model.names[cls_id]
+ # Retain both 'sack' and 'truck' classes
+ if name not in ("sack", "truck"):
+ continue
+ track_id = int(ids[i]) if ids is not None else None
+ x1, y1, x2, y2 = box.xyxy[0].tolist()
+
+ mask = None
+
+ tracked.append(
+ Detection(
+ bbox=(x1, y1, x2, y2),
+ confidence=float(box.conf[0]),
+ class_id=cls_id,
+ class_name=name,
+ track_id=track_id,
+ mask=mask,
+ )
+ )
+ return tracked
+
+ def reset(self) -> None:
+ """Reset tracker state (new batch / new truck)."""
+ if self._model_path:
+ self._model = YOLO(self._model_path)
diff --git a/algoritma-batch/src/truck_roi.py b/algoritma-batch/src/truck_roi.py
new file mode 100644
index 0000000..3ad91b1
--- /dev/null
+++ b/algoritma-batch/src/truck_roi.py
@@ -0,0 +1,151 @@
+"""Truck ROI tracker — identifies the main truck and provides a stable ROI.
+
+Uses exponential moving average (EMA) to smooth the bounding box across
+frames, preventing jitter from frame-to-frame detection variance.
+
+For y2-based counting (bottom edge of sack bbox), the counting line is
+placed `LINE_OFFSET_PX` pixels relative to the truck bottom edge (`roi.y2`).
+With `offset = +20`, the line sits at `roi.y2 + 20` (~620px), cleanly
+separating sacks on the ground (`y2 > 650`) from loaded sacks (`y2 < 580`).
+"""
+
+from __future__ import annotations
+
+from dataclasses import dataclass
+
+from src.interfaces import Detection
+
+# Offset for y1 counting line relative to truck top edge (px).
+# Positive = below truck top edge (into the truck).
+# Negative = above truck top edge (towards the camera).
+LINE_OFFSET_PX = 0
+
+
+@dataclass
+class TruckROI:
+ """Region of interest derived from the main truck bbox."""
+
+ x1: int
+ y1: int
+ x2: int
+ y2: int
+ line_y: int # counting line Y position (pixels)
+ confidence: float
+
+ @property
+ def width(self) -> int:
+ return self.x2 - self.x1
+
+ @property
+ def height(self) -> int:
+ return self.y2 - self.y1
+
+ def contains_x(self, cx: float) -> bool:
+ """Check if a centroid X falls within the truck X bounds."""
+ return self.x1 <= cx <= self.x2
+
+
+class TruckROITracker:
+ """Tracks the main truck and provides a smoothed ROI + counting line.
+
+ Main truck = largest truck detection whose center X falls in the
+ expected lane (center region of the frame).
+
+ The counting line is placed `line_offset` pixels below the truck
+ bottom edge (`roi.y2`).
+ """
+
+ def __init__(
+ self,
+ frame_width: int,
+ frame_height: int,
+ lane_x_min: float = 0.35,
+ lane_x_max: float = 0.80,
+ ema_alpha: float = 0.15,
+ line_offset: int = LINE_OFFSET_PX,
+ ) -> None:
+ self._fw = frame_width
+ self._fh = frame_height
+ self._lane_x_min = int(lane_x_min * frame_width)
+ self._lane_x_max = int(lane_x_max * frame_width)
+ self._alpha = ema_alpha
+ self._line_offset = line_offset
+
+ # Smoothed bbox (None until first detection)
+ self._sx1: float | None = None
+ self._sy1: float | None = None
+ self._sx2: float | None = None
+ self._sy2: float | None = None
+
+ self._last_roi: TruckROI | None = None
+ self._frames_without_truck = 0
+
+ def update(self, truck_detections: list[Detection]) -> TruckROI | None:
+ """Pick the main truck, smooth its bbox, return ROI."""
+ main = self._pick_main_truck(truck_detections)
+
+ if main is None:
+ self._frames_without_truck += 1
+ if self._frames_without_truck > 5: # Clear ROI if truck is missing for >5 updates (~3 seconds)
+ self.reset()
+ return None
+ return self._last_roi # hold last known ROI briefly
+
+ self._frames_without_truck = 0
+ x1, y1, x2, y2 = main.bbox
+
+ # EMA smoothing
+ if self._sx1 is None:
+ self._sx1, self._sy1 = float(x1), float(y1)
+ self._sx2, self._sy2 = float(x2), float(y2)
+ else:
+ a = self._alpha
+ self._sx1 = a * x1 + (1 - a) * self._sx1
+ self._sy1 = a * y1 + (1 - a) * self._sy1
+ self._sx2 = a * x2 + (1 - a) * self._sx2
+ self._sy2 = a * y2 + (1 - a) * self._sy2
+
+ # Build ROI — line placed at truck top edge
+ roi_x1 = max(0, int(self._sx1))
+ roi_y1 = max(0, int(self._sy1))
+ roi_x2 = min(self._fw, int(self._sx2))
+ roi_y2 = min(self._fh, int(self._sy2))
+ line_y = roi_y1 + self._line_offset
+
+ self._last_roi = TruckROI(
+ x1=roi_x1, y1=roi_y1, x2=roi_x2, y2=roi_y2,
+ line_y=line_y, confidence=main.confidence,
+ )
+ return self._last_roi
+
+ @property
+ def roi(self) -> TruckROI | None:
+ return self._last_roi
+
+ @property
+ def frames_without_truck(self) -> int:
+ return self._frames_without_truck
+
+ def reset(self) -> None:
+ self._sx1 = self._sy1 = self._sx2 = self._sy2 = None
+ self._last_roi = None
+ self._frames_without_truck = 0
+
+ def _pick_main_truck(
+ self, detections: list[Detection]
+ ) -> Detection | None:
+ """Select the largest truck whose center X is in the expected lane."""
+ best: Detection | None = None
+ best_area = 0
+
+ for det in detections:
+ x1, y1, x2, y2 = det.bbox
+ cx = (x1 + x2) / 2
+ if not (self._lane_x_min <= cx <= self._lane_x_max):
+ continue
+ area = (x2 - x1) * (y2 - y1)
+ if area > best_area:
+ best = det
+ best_area = area
+
+ return best
diff --git a/algoritma-batch/zones.json b/algoritma-batch/zones.json
new file mode 100644
index 0000000..edee18c
--- /dev/null
+++ b/algoritma-batch/zones.json
@@ -0,0 +1,43 @@
+{
+ "palet": [
+ [
+ 612,
+ 628
+ ],
+ [
+ 1454,
+ 630
+ ],
+ [
+ 1458,
+ 1074
+ ],
+ [
+ 608,
+ 1071
+ ]
+ ],
+ "truck": [
+ [600, 385],
+ [609, 1076],
+ [1404, 1078],
+ [1381, 343]
+ ],
+ "counting": [
+ [574, 50],
+ [586, 1077],
+ [1418, 1076],
+ [1397, 50]
+ ],
+ "left_limit": 0.27578,
+ "right_limit": 0.72578,
+ "duplicate_circle_radius": 60,
+ "min_valid_area": 15000,
+ "jarak_toleransi_duplikat": 30,
+ "max_reid_transit_distance": 400,
+ "circle_stay_timeout_sec": 10.0,
+ "inference_stride": 2,
+ "confirm_delay_sec": 0.5,
+ "exit_confirm_delay_sec": 6.0,
+ "external_stream_url": "http://192.168.192.96:8888/cam/"
+}
\ No newline at end of file
diff --git a/backend/api/batches.py b/backend/api/batches.py
index f2a3233..9f46ae7 100644
--- a/backend/api/batches.py
+++ b/backend/api/batches.py
@@ -1,13 +1,14 @@
-"""Batch, frame and auto-annotation routes (REQ-020…034)."""
-
+import json
import os
+import shutil
+import tempfile
from typing import Optional
-from fastapi import APIRouter, HTTPException, Response
+from fastapi import APIRouter, File, Form, HTTPException, Response, UploadFile
from fastapi.responses import FileResponse
from pydantic import BaseModel
-from backend import autolabel, dataset, library
+from backend import autolabel, dataset, library, projects
from backend import batches as batch_store
from backend import review as review_store
from backend.api.common import project_or_404, thumbnail
@@ -27,10 +28,12 @@ class AutolabelRequest(BaseModel):
engines: Optional[list[str]] = None
class_ids: Optional[list[int]] = None
engine_classes: Optional[dict[str, list[str]]] = None
+ target_class_names: Optional[list[str]] = None
threshold: float = autolabel.DEFAULT_THRESHOLD
iou_threshold: float = autolabel.DEFAULT_IOU
min_box_frac: float = 0.0
resume: bool = False
+ append: bool = False
@router.post("/api/projects/{project_id}/batches")
@@ -90,11 +93,115 @@ def start_autolabel(batch_id: int, request: AutolabelRequest) -> dict:
try:
engine_list = request.engines if (request.engines and len(request.engines) > 0) else [request.engine]
return autolabel.start(batch_id, request.threshold, request.iou_threshold,
- request.min_box_frac, resume=request.resume,
+ request.min_box_frac, resume=request.resume, append=request.append,
engines=engine_list, class_ids=request.class_ids,
- engine_classes=request.engine_classes)
+ engine_classes=request.engine_classes,
+ target_class_names=request.target_class_names)
except batch_store.BatchError as exc:
raise HTTPException(400, str(exc))
+@router.post("/api/batches/inspect-model")
+async def inspect_model(file: UploadFile = File(...)) -> dict:
+ if not (file.filename or "").endswith(".pt"):
+ raise HTTPException(400, "Model must be a .pt file")
+ with tempfile.NamedTemporaryFile(suffix=".pt", delete=False) as staged:
+ shutil.copyfileobj(file.file, staged)
+ staged_path = staged.name
+ await file.close()
+ try:
+ classes = projects.read_model_classes(staged_path)
+ return {"filename": file.filename, "classes": classes, "staged_path": staged_path}
+ except Exception as exc:
+ if os.path.exists(staged_path):
+ os.unlink(staged_path)
+ raise HTTPException(400, f"Could not inspect model: {exc}")
+
+
+@router.post("/api/batches/{batch_id}/autolabel-with-model")
+async def autolabel_with_model(
+ batch_id: int,
+ file: UploadFile = File(...),
+ threshold: float = Form(0.35),
+ iou_threshold: float = Form(0.8),
+ selected_classes: str = Form("[]"),
+ append: bool = Form(True),
+) -> dict:
+ if not (file.filename or "").endswith(".pt"):
+ raise HTTPException(400, "Model must be a .pt file")
+ with tempfile.NamedTemporaryFile(suffix=".pt", delete=False) as staged:
+ shutil.copyfileobj(file.file, staged)
+ staged_path = staged.name
+ await file.close()
+
+ try:
+ target_classes = json.loads(selected_classes) if selected_classes else None
+ return autolabel.start(
+ batch_id,
+ threshold=threshold,
+ iou_threshold=iou_threshold,
+ append=append,
+ custom_model_path=staged_path,
+ target_class_names=target_classes,
+ )
+ except Exception as exc:
+ if os.path.exists(staged_path):
+ os.unlink(staged_path)
+ raise HTTPException(400, f"Auto-annotation failed to start: {exc}")
+
+
+@router.post("/api/sam3/playground-test")
+async def sam3_playground_test(
+ file: UploadFile = File(...),
+ prompts: str = Form(...),
+ threshold: float = Form(0.35),
+ iou_threshold: float = Form(0.8),
+) -> dict:
+ from PIL import Image
+ from backend import labeling
+ from backend.sam3_engine import get_engine
+
+ try:
+ image = Image.open(file.file).convert("RGB")
+ except Exception as exc:
+ raise HTTPException(400, f"Could not read image: {exc}")
+
+ width, height = image.size
+ prompt_list = [p.strip() for p in prompts.split(",") if p.strip()]
+ if not prompt_list:
+ raise HTTPException(400, "At least one text prompt is required")
+
+ try:
+ engine = get_engine()
+ raw_dets = engine.detect(image, prompt_list, threshold)
+ kept_dets = labeling.deduplicate(raw_dets, iou_threshold=iou_threshold)
+ except Exception as exc:
+ raise HTTPException(500, f"SAM3 inference failed: {exc}")
+
+ results = []
+ for det in kept_dets:
+ norm_box = [
+ det.box[0] / width,
+ det.box[1] / height,
+ det.box[2] / width,
+ det.box[3] / height,
+ ]
+ polys = []
+ if det.mask is not None:
+ raw_polys = review_store.mask_to_polygons(det.mask)
+ polys = [[[float(pt[0]), float(pt[1])] for pt in poly] for poly in raw_polys]
+
+ results.append({
+ "class_id": det.class_id,
+ "prompt": det.class_name,
+ "score": round(float(det.score), 4),
+ "box": [round(v, 5) for v in norm_box],
+ "polygons": polys,
+ })
+
+ return {
+ "width": width,
+ "height": height,
+ "detections": results,
+ }
@@ -150,3 +257,11 @@ def clear_batch_class_annotations(batch_id: int, class_id: int) -> dict:
deleted = review_store.clear_batch_class_annotations(batch_id, class_id)
return {"deleted": deleted}
+
+@router.post("/api/batches/{batch_id}/reset-auto-annotations")
+def reset_batch_auto_annotations(batch_id: int) -> dict:
+ if batch_store.get(batch_id) is None:
+ raise HTTPException(404, "No such batch")
+ deleted = review_store.clear_batch_auto_annotations(batch_id)
+ return {"deleted": deleted}
+
diff --git a/backend/api/models.py b/backend/api/models.py
index 7a81e72..757374c 100644
--- a/backend/api/models.py
+++ b/backend/api/models.py
@@ -19,6 +19,7 @@ class TrainRequest(BaseModel):
imgsz: Optional[int] = None
device: Optional[Union[int, str]] = None
batch_ids: Optional[list] = None
+ class_ids: Optional[list] = None
@router.get("/api/hardware")
@@ -34,6 +35,7 @@ def start_training(project_id: int, request: TrainRequest) -> dict:
project_id, request.epochs,
{"batch": request.batch, "imgsz": request.imgsz, "device": request.device},
batch_ids=request.batch_ids,
+ class_ids=request.class_ids,
)
except training.TrainingError as exc:
raise HTTPException(400, str(exc))
diff --git a/backend/autolabel.py b/backend/autolabel.py
index 30f0e8e..da07011 100644
--- a/backend/autolabel.py
+++ b/backend/autolabel.py
@@ -18,10 +18,12 @@ DEFAULT_IOU = 0.8
def start(batch_id: int, threshold: float = DEFAULT_THRESHOLD,
iou_threshold: float = DEFAULT_IOU, min_box_frac: float = 0.0,
- resume: bool = False, engine: str = "sam3",
+ resume: bool = False, append: bool = False, engine: str = "base_model",
engines: Optional[List[str]] = None,
class_ids: Optional[List[int]] = None,
- engine_classes: Optional[dict[str, List[str]]] = None) -> dict:
+ engine_classes: Optional[dict[str, List[str]]] = None,
+ custom_model_path: Optional[str] = None,
+ target_class_names: Optional[List[str]] = None) -> dict:
batch = batches.get(batch_id)
if batch is None:
raise batches.BatchError("No such batch")
@@ -34,8 +36,9 @@ def start(batch_id: int, threshold: float = DEFAULT_THRESHOLD,
"autolabel",
params={"batch_id": batch_id, "threshold": threshold,
"iou_threshold": iou_threshold, "min_box_frac": min_box_frac,
- "resume": resume, "engine": active_engines[0], "engines": active_engines,
- "class_ids": class_ids, "engine_classes": engine_classes},
+ "resume": resume, "append": append, "engine": active_engines[0], "engines": active_engines,
+ "class_ids": class_ids, "engine_classes": engine_classes,
+ "custom_model_path": custom_model_path, "target_class_names": target_class_names},
project_id=batch["project_id"],
batch_id=batch_id,
message=f"{batch['date_label']}/{batch['batch_label']} ({'+'.join(e.upper() for e in active_engines)})",
@@ -62,16 +65,7 @@ def _run_autolabel(job) -> None:
raise batches.BatchError("The batch disappeared before labeling started")
project = projects.get(batch["project_id"])
- raw_active = job.params.get("engines") or [job.params.get("engine", "sam3")]
- expanded_engines = []
- for eng in raw_active:
- if eng == "both":
- expanded_engines.extend(["base_model", "secondary_model"])
- elif eng == "sam3+model1":
- expanded_engines.extend(["sam3", "base_model"])
- else:
- expanded_engines.append(eng)
- expanded_engines = list(dict.fromkeys(expanded_engines))
+ selected_engine = job.params.get("engine", "base_model")
frames = batches.frames(batch["id"])
batches.set_status(batch["id"], "labeling")
@@ -86,42 +80,34 @@ def _run_autolabel(job) -> None:
attempted = 0
failures = []
- from ultralytics import YOLO
- m1_path = projects.training_start_point(project)
- with db.cursor() as cur:
- cur.execute("SELECT weights_path FROM model_versions WHERE project_id = ? ORDER BY version DESC LIMIT 1", (project["id"],))
- row = cur.fetchone()
- if row and os.path.isfile(row[0]):
- m1_path = row[0]
-
- yolo_models = {}
- if "base_model" in expanded_engines or "yolo" in expanded_engines:
- job.log(f"Loading Base/Trained Model: {os.path.basename(m1_path)}...")
- yolo_models["base_model"] = YOLO(m1_path)
-
- if "secondary_model" in expanded_engines:
- m2_path = project["secondary_model_path"] if (project.get("secondary_model_path") and os.path.isfile(project["secondary_model_path"])) else m1_path
- label_name = project.get("secondary_model_name") or os.path.basename(m2_path)
- job.log(f"Loading Secondary Model: {label_name}...")
- yolo_models["secondary_model"] = YOLO(m2_path)
-
- allowed_class_ids = set(job.params["class_ids"]) if job.params.get("class_ids") is not None else None
- engine_classes = job.params.get("engine_classes") or {}
+ conf = job.params.get("threshold", DEFAULT_THRESHOLD)
+ iou_thresh = job.params.get("iou_threshold", DEFAULT_IOU)
+ yolo_model = None
sam3_target_classes = []
- if "sam3" in expanded_engines:
- sam3_classes = engine_classes.get("sam3")
- if sam3_classes is not None:
- allowed_set = {c.strip().lower() for c in sam3_classes}
- sam3_target_classes = [
- c for c in project["classes"]
- if c["name"].strip().lower() in allowed_set or c["prompt"].strip().lower() in allowed_set
- ]
+
+ custom_path = job.params.get("custom_model_path")
+ target_class_names = job.params.get("target_class_names")
+
+ if selected_engine == "sam3" and not custom_path:
+ allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
+ if allowed_classes_set:
+ sam3_target_classes = [c for c in project["classes"] if c["name"].strip().lower() in allowed_classes_set or c["prompt"].strip().lower() in allowed_classes_set]
+ # Add any new target class names that aren't in project classes yet
+ existing_names = {c["name"].strip().lower() for c in project["classes"]}
+ for name in target_class_names:
+ if name.strip().lower() not in existing_names:
+ try:
+ updated_proj = projects.add_class(project["id"], name=name.strip(), prompt=name.strip())
+ project["classes"] = updated_proj["classes"]
+ for new_c in project["classes"]:
+ if new_c["name"].strip().lower() == name.strip().lower() and new_c not in sam3_target_classes:
+ sam3_target_classes.append(new_c)
+ except Exception:
+ pass
else:
- sam3_target_classes = [
- c for c in project["classes"]
- if (allowed_class_ids is None or c["class_id"] in allowed_class_ids)
- ]
+ sam3_target_classes = [c for c in project["classes"]]
+
prompts = [c["prompt"] for c in sam3_target_classes]
if prompts:
from backend.sam3_engine import engine_is_loaded, get_engine
@@ -130,13 +116,27 @@ def _run_autolabel(job) -> None:
engine = get_engine()
job.log(f"SAM3 ready on {engine.device}; prompts: {', '.join(prompts)}")
else:
- job.log("SAM3 selected but 0 prompts match class filter.")
+ job.log("SAM3 selected but 0 prompts match project classes.")
+ else:
+ from ultralytics import YOLO
+ if custom_path and os.path.isfile(custom_path):
+ m_path = custom_path
+ job.log(f"Loading Custom Model: {os.path.basename(m_path)}...")
+ else:
+ m_path = projects.training_start_point(project)
+ with db.cursor() as cur:
+ cur.execute("SELECT weights_path FROM model_versions WHERE project_id = ? ORDER BY version DESC LIMIT 1", (project["id"],))
+ row = cur.fetchone()
+ if row and os.path.isfile(row[0]):
+ m_path = row[0]
+ job.log(f"Loading Base Model: {os.path.basename(m_path)}...")
+
+ yolo_model = YOLO(m_path)
name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
- conf = job.params.get("threshold", DEFAULT_THRESHOLD)
- iou_thresh = job.params.get("iou_threshold", DEFAULT_IOU)
+ allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
- job.log(f"Starting multi-engine auto-labeling ({', '.join(expanded_engines)})...")
+ job.log(f"Starting auto-labeling with {selected_engine}...")
for index, frame in enumerate(frames):
if job.cancelled:
@@ -151,29 +151,21 @@ def _run_autolabel(job) -> None:
frame_file = os.path.join(directory, frame["filename"])
all_raw_detections = []
- for eng_key, y_model in yolo_models.items():
- allowed_for_eng = engine_classes.get(eng_key)
- if allowed_for_eng is not None and len(allowed_for_eng) == 0:
- continue
- results = y_model.predict(frame_file, conf=conf, verbose=False)
+ if yolo_model is not None:
+ results = yolo_model.predict(frame_file, conf=conf, verbose=False)
if results and len(results) > 0:
model_names = results[0].names
for box in results[0].boxes:
cls_idx = int(box.cls[0].item())
cls_name = str(model_names.get(cls_idx, cls_idx)).strip().lower()
- if allowed_for_eng is not None and cls_name not in [c.strip().lower() for c in allowed_for_eng]:
+
+ if allowed_classes_set is not None and cls_name not in allowed_classes_set:
continue
- if cls_name not in name_to_class_id:
- try:
- updated_proj = projects.add_class(project["id"], {"name": cls_name, "prompt": cls_name})
- project["classes"] = updated_proj["classes"]
- name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
- except Exception:
- pass
- if cls_name in name_to_class_id:
- target_class_id = name_to_class_id[cls_name]
- else:
+
+ target_class_id = name_to_class_id.get(cls_name)
+ if target_class_id is None:
continue
+
score = float(box.conf[0].item())
xyxyn = box.xyxyn[0].tolist()
all_raw_detections.append(labeling.Detection(
@@ -184,7 +176,7 @@ def _run_autolabel(job) -> None:
mask=None
))
- if "sam3" in expanded_engines and sam3_target_classes:
+ elif selected_engine == "sam3" and sam3_target_classes:
prompts = [c["prompt"] for c in sam3_target_classes]
res = labeling.label_image(
frame_file, frame["filename"], prompts, conf,
@@ -208,7 +200,10 @@ def _run_autolabel(job) -> None:
for geometry in _geometries(det, frame["width"], frame["height"], project["label_type"]):
items.append({"class_id": det.class_id, "geometry": geometry, "score": det.score})
- review.replace_auto(frame["id"], items)
+ if job.params.get("append"):
+ review.append_auto(frame["id"], items)
+ else:
+ review.replace_auto(frame["id"], items)
written += len(items)
job.progress(index + 1, len(frames), f"{frame['filename']}: {len(items)} shape(s)")
except Exception as exc:
diff --git a/backend/dataset.py b/backend/dataset.py
index 3c72a3b..85441b0 100644
--- a/backend/dataset.py
+++ b/backend/dataset.py
@@ -14,6 +14,7 @@ Label files are plain YOLO:
import os
import shutil
import time
+from typing import List, Optional
from backend import batches, config, db, jobs, projects, review
@@ -74,12 +75,55 @@ def _next_split(cur, project_id: int, val_every: int) -> str:
return "val" if position % val_every == val_every - 1 else "train"
-def write_data_yaml(project: dict, batch_ids: list = None) -> str:
+def sync_labels(project_id: int, selected_class_ids: Optional[List[int]] = None) -> dict:
+ """Re-sync label files on disk for all merged frames in the project dataset."""
+ project = projects.get(project_id)
+ root = dataset_dir(project["slug"])
+ with db.cursor() as cur:
+ cur.execute(
+ "SELECT d.frame_id, d.label_rel FROM dataset_items d WHERE d.project_id = ?",
+ (project_id,),
+ )
+ items = cur.fetchall()
+
+ class_map = None
+ if selected_class_ids is not None and len(selected_class_ids) > 0:
+ class_map = {cid: idx for idx, cid in enumerate(sorted(selected_class_ids))}
+
+ synced_files = 0
+ total_lines = 0
+ for frame_id, label_rel in items:
+ annotations = review.listing(frame_id)
+ if class_map is not None:
+ annotations = [a for a in annotations if a["class_id"] in class_map]
+
+ lines = []
+ for item in annotations:
+ mapped_cid = class_map[item["class_id"]] if class_map is not None else item["class_id"]
+ lines.append(_label_line(mapped_cid, item["geometry"], project["label_type"]))
+
+ path = os.path.join(root, label_rel)
+ os.makedirs(os.path.dirname(path), exist_ok=True)
+ with open(path, "w", encoding="utf-8") as f:
+ f.write("\n".join(lines) + ("\n" if lines else ""))
+ synced_files += 1
+ total_lines += len(lines)
+
+ return {"synced_files": synced_files, "total_lines": total_lines}
+
+
+def write_data_yaml(project: dict, batch_ids: list = None, selected_class_ids: Optional[List[int]] = None) -> str:
"""Rebuild data.yaml from the project's classes (REQ-051)."""
+ sync_labels(project["id"], selected_class_ids=selected_class_ids)
root = dataset_dir(project["slug"])
os.makedirs(root, exist_ok=True)
counts = summary(project["id"])["splits"]
- names = ", ".join(f"'{item['name']}'" for item in project["classes"])
+
+ target_classes = project["classes"]
+ if selected_class_ids is not None and len(selected_class_ids) > 0:
+ target_classes = [c for c in project["classes"] if c["class_id"] in selected_class_ids]
+
+ names = ", ".join(f"'{item['name']}'" for item in target_classes)
if batch_ids:
with db.cursor() as cur:
diff --git a/backend/hardware.py b/backend/hardware.py
index 9f4d8d2..df744a7 100644
--- a/backend/hardware.py
+++ b/backend/hardware.py
@@ -44,9 +44,9 @@ def defaults(epochs: int = 50) -> dict:
if info["device"] == "cpu":
settings = {"batch": 4, "imgsz": 512, "device": "cpu", "workers": 2}
note = "No GPU visible — training on CPU will be very slow."
- elif vram < 8:
- settings = {"batch": 8, "imgsz": 640, "device": 0, "workers": 2}
- note = f"{vram} GB of VRAM: small batches, 640 px."
+ elif vram < 6:
+ settings = {"batch": 16, "imgsz": 640, "device": 0, "workers": 4}
+ note = f"{vram} GB of VRAM: batch 16, 640 px."
elif vram <= 16:
settings = {"batch": 32, "imgsz": 640, "device": 0, "workers": 8}
note = f"{vram} GB of VRAM: optimized batch 32, 640 px."
diff --git a/backend/labeling.py b/backend/labeling.py
index cc4fa3c..fbdb43d 100644
--- a/backend/labeling.py
+++ b/backend/labeling.py
@@ -42,11 +42,18 @@ def _iou(box_a: List[float], box_b: List[float]) -> float:
def deduplicate(detections: List[Detection], iou_threshold: float = 0.8) -> List[Detection]:
- """Greedy NMS across all prompts: highest score wins an overlapping region."""
+ """Greedy NMS per class: highest score wins within the SAME class."""
+ by_class: dict[int, List[Detection]] = {}
+ for det in detections:
+ by_class.setdefault(det.class_id, []).append(det)
+
kept: List[Detection] = []
- for det in sorted(detections, key=lambda d: d.score, reverse=True):
- if all(_iou(det.box, k.box) < iou_threshold for k in kept):
- kept.append(det)
+ for cls_dets in by_class.values():
+ cls_kept: List[Detection] = []
+ for det in sorted(cls_dets, key=lambda d: d.score, reverse=True):
+ if all(_iou(det.box, k.box) < iou_threshold for k in cls_kept):
+ cls_kept.append(det)
+ kept.extend(cls_kept)
return kept
diff --git a/backend/projects.py b/backend/projects.py
index 407f8b0..8305640 100644
--- a/backend/projects.py
+++ b/backend/projects.py
@@ -384,6 +384,11 @@ def _kept_prompts(project: dict, names: List[str]) -> List[str]:
def training_start_point(project: dict) -> str:
"""The weights a training run should start from (REQ-060, REQ-004)."""
path = project["base_model_path"]
+ if path and not os.path.isfile(path):
+ if path.startswith("/data/"):
+ alt_path = os.path.join(config.DATA_DIR, path[6:])
+ if os.path.isfile(alt_path):
+ path = alt_path
if path and os.path.isfile(path):
return path
return PRETRAINED[project["label_type"]]
diff --git a/backend/review.py b/backend/review.py
index 2364050..237449a 100644
--- a/backend/review.py
+++ b/backend/review.py
@@ -234,6 +234,38 @@ def replace_auto(frame_id: int, items: List[dict]) -> int:
return len(items)
+def append_auto(frame_id: int, items: List[dict]) -> int:
+ """Add new automatic shapes to this frame without duplicating existing ones."""
+ if not items:
+ return 0
+ existing = listing(frame_id)
+ filtered_items = []
+ for item in items:
+ is_dup = False
+ item_box = to_box(item["geometry"])
+ for ex in existing:
+ if ex["class_id"] == item["class_id"]:
+ ex_box = to_box(ex["geometry"])
+ from backend.labeling import _iou
+ if _iou(item_box, ex_box) >= 0.85:
+ is_dup = True
+ break
+ if not is_dup:
+ filtered_items.append(item)
+
+ if not filtered_items:
+ return 0
+
+ with db.cursor() as cur:
+ cur.executemany(
+ """INSERT INTO annotations (frame_id, class_id, geometry, score, source, created_at)
+ VALUES (?, ?, ?, ?, 'auto', ?)""",
+ [(frame_id, item["class_id"], json.dumps(item["geometry"]),
+ item.get("score", 1.0), time.time()) for item in filtered_items],
+ )
+ return len(filtered_items)
+
+
def frames_with_auto(batch_id: int) -> set:
"""Frame ids that already carry automatic shapes — the resume skip-list
for REQ-035."""
@@ -339,6 +371,24 @@ def clear_batch_class_annotations(batch_id: int, class_id: int) -> int:
return cur.rowcount
+def clear_batch_auto_annotations(batch_id: int) -> int:
+ """Delete all automatic annotations (source = 'auto') for a batch and reset frame statuses."""
+ with db.cursor() as cur:
+ cur.execute(
+ """DELETE FROM annotations
+ WHERE source = 'auto' AND frame_id IN (
+ SELECT id FROM frames WHERE batch_id = ?
+ )""",
+ (batch_id,),
+ )
+ deleted = cur.rowcount
+ cur.execute(
+ "UPDATE frames SET review_status = 'pending' WHERE batch_id = ?",
+ (batch_id,),
+ )
+ return deleted
+
+
def _check_class(project_id: int, class_id: int) -> None:
with db.cursor() as cur:
cur.execute(
diff --git a/backend/training.py b/backend/training.py
index 5d187a4..c5340f8 100644
--- a/backend/training.py
+++ b/backend/training.py
@@ -25,7 +25,7 @@ def models_dir(project_slug: str) -> str:
return os.path.join(config.project_dir(project_slug), "models")
-def start(project_id: int, epochs: int = 50, overrides: Optional[dict] = None, batch_ids: Optional[list] = None) -> dict:
+def start(project_id: int, epochs: int = 50, overrides: Optional[dict] = None, batch_ids: Optional[list] = None, class_ids: Optional[list] = None) -> dict:
project = projects.get(project_id)
if project is None:
raise TrainingError("No such project")
@@ -38,7 +38,7 @@ def start(project_id: int, epochs: int = 50, overrides: Optional[dict] = None, b
settings = hardware.resolve(overrides, epochs)
job = jobs.create(
"train",
- params={"project_id": project_id, "settings": settings, "batch_ids": batch_ids},
+ params={"project_id": project_id, "settings": settings, "batch_ids": batch_ids, "class_ids": class_ids},
project_id=project_id,
message=f"{counts['train']} train / {counts['val']} val",
)
@@ -102,7 +102,8 @@ def _run_train(job) -> None:
project = projects.get(job.params["project_id"])
settings = job.params["settings"]
batch_ids = job.params.get("batch_ids")
- data_yaml = dataset.write_data_yaml(project, batch_ids=batch_ids)
+ class_ids = job.params.get("class_ids")
+ data_yaml = dataset.write_data_yaml(project, batch_ids=batch_ids, selected_class_ids=class_ids)
# SAM3 and a training run must not hold VRAM at the same time (REQ-065).
from backend.sam3_engine import release_engine
@@ -133,6 +134,10 @@ def _run_train(job) -> None:
model.add_callback("on_fit_epoch_end", on_epoch)
job.progress(0, settings["epochs"])
+ import torch
+ if torch.cuda.is_available():
+ torch.backends.cudnn.benchmark = True
+
keep_run_dir = False
try:
model.train(
@@ -142,7 +147,7 @@ def _run_train(job) -> None:
batch=settings["batch"],
device=settings["device"],
workers=settings.get("workers", 8),
- cache=False,
+ cache="ram",
project=os.path.join(out_dir, "runs"),
name="train",
exist_ok=True,
diff --git a/frontend/src/App.jsx b/frontend/src/App.jsx
index 1d1fe60..4abc3cd 100644
--- a/frontend/src/App.jsx
+++ b/frontend/src/App.jsx
@@ -6,6 +6,7 @@ import LibraryPage from './pages/LibraryPage'
import TrimPage from './pages/TrimPage'
import ReviewPage from './pages/ReviewPage'
import ModelsPage from './pages/ModelsPage'
+import Sam3PlaygroundPage from './pages/Sam3PlaygroundPage'
import './roboflow.css'
function parseRoute(hash) {
@@ -14,6 +15,10 @@ function parseRoute(hash) {
const query = new URLSearchParams(queryString || '')
const parts = path.split('/').filter(Boolean)
+ if (parts[0] === 'sam3-playground') {
+ return { name: 'sam3-playground' }
+ }
+
if (parts[0] === 'batches' && parts[1]) {
return { name: 'review', batchId: Number(parts[1]) }
}
@@ -135,6 +140,7 @@ export default function App() {
onProject={(p) => setCurrentProject(p)}
/>
)}
+ {route.name === 'sam3-playground' && }
diff --git a/frontend/src/api.js b/frontend/src/api.js
index 980dd96..85b82a4 100644
--- a/frontend/src/api.js
+++ b/frontend/src/api.js
@@ -60,8 +60,32 @@ export const api = {
startAutolabel: (batchId, body) =>
request(`/batches/${batchId}/autolabel`, { method: 'POST', body: body ?? {} }),
+ inspectModel: (file) => {
+ const form = new FormData()
+ form.append('file', file)
+ return request('/batches/inspect-model', { method: 'POST', form })
+ },
+ autolabelWithModel: (batchId, file, threshold, selectedClasses, iouThreshold = 0.8) => {
+ const form = new FormData()
+ form.append('file', file)
+ form.append('threshold', threshold)
+ form.append('iou_threshold', iouThreshold)
+ form.append('selected_classes', JSON.stringify(selectedClasses))
+ form.append('append', 'true')
+ return request(`/batches/${batchId}/autolabel-with-model`, { method: 'POST', form })
+ },
+ sam3PlaygroundTest: (file, prompts, threshold = 0.35, iouThreshold = 0.8) => {
+ const form = new FormData()
+ form.append('file', file)
+ form.append('prompts', prompts)
+ form.append('threshold', threshold)
+ form.append('iou_threshold', iouThreshold)
+ return request('/sam3/playground-test', { method: 'POST', form })
+ },
clearBatchClassAnnotations: (batchId, classId) =>
request(`/batches/${batchId}/classes/${classId}/annotations`, { method: 'DELETE' }),
+ resetBatchAutoAnnotations: (batchId) =>
+ request(`/batches/${batchId}/reset-auto-annotations`, { method: 'POST' }),
nextPending: (batchId, afterIdx = -1) =>
request(`/batches/${batchId}/next-pending?after_idx=${afterIdx}`),
diff --git a/frontend/src/components/Sidebar.jsx b/frontend/src/components/Sidebar.jsx
index f0fe683..7e840cc 100644
--- a/frontend/src/components/Sidebar.jsx
+++ b/frontend/src/components/Sidebar.jsx
@@ -74,11 +74,15 @@ export default function Sidebar({ route, currentProject, theme, onToggleTheme })
diff --git a/frontend/src/pages/LibraryPage.jsx b/frontend/src/pages/LibraryPage.jsx
index 7411514..788a57d 100644
--- a/frontend/src/pages/LibraryPage.jsx
+++ b/frontend/src/pages/LibraryPage.jsx
@@ -43,40 +43,27 @@ function ActiveJobsBanner({ jobs, onCancel }) {
function BatchList({ project, batches, activeJobs, onChanged, onError }) {
const [busyId, setBusyId] = useState(null)
- const [selectedEngine, setSelectedEngine] = useState({})
- const [engineClassMap, setEngineClassMap] = useState({})
- const [expandedFilterBatchId, setExpandedFilterBatchId] = useState(null)
+ const [appendChoiceBatch, setAppendChoiceBatch] = useState(null)
+ const [appendModalState, setAppendModalState] = useState(null)
+ const [sam3AppendState, setSam3AppendState] = useState(null)
+ const [customPromptInput, setCustomPromptInput] = useState('')
+ const [baseModelModalState, setBaseModelModalState] = useState(null)
- const hasSecondaryModel = Boolean(project?.secondary_model_path)
- const model1Classes = project?.classes?.map((c) => c.name) || []
- const model2Classes = project?.secondary_model_classes?.length > 0
- ? project.secondary_model_classes
- : model1Classes
+ function openBaseModelAutolabelModal(batch) {
+ const projectClasses = project?.classes?.map((c) => c.name) || []
+ setBaseModelModalState({
+ batch,
+ selectedClasses: [...projectClasses],
+ threshold: 0.35,
+ iouThreshold: 0.8,
+ })
+ }
- const defaultEngines = project?.base_model_path ? ['base_model'] : ['sam3']
- const [selectedThreshold, setSelectedThreshold] = useState({})
-
- async function autolabel(batch, resume = false) {
+ async function resetAutoAnnotations(batch) {
+ if (!window.confirm(`Reset auto-annotations for batch "${batch.batch_label}"?\n\nThis will clear all auto-generated shapes and reset frame review statuses back to pending.`)) return
setBusyId(batch.id)
- const engines = selectedEngine[batch.id] || defaultEngines
- const threshold = selectedThreshold[batch.id] ?? 0.35
-
- // Default class filter for each active engine if not customized
- const currentBatchMap = engineClassMap[batch.id] || {}
- const engine_classes = {}
-
- if (engines.includes('sam3')) {
- engine_classes.sam3 = currentBatchMap.sam3 || model1Classes
- }
- if (engines.includes('base_model')) {
- engine_classes.base_model = currentBatchMap.base_model || model1Classes
- }
- if (engines.includes('secondary_model')) {
- engine_classes.secondary_model = currentBatchMap.secondary_model || model2Classes
- }
-
try {
- await api.startAutolabel(batch.id, { resume, engines, engine_classes, threshold })
+ await api.resetBatchAutoAnnotations(batch.id)
onChanged()
} catch (exc) {
onError(exc.message)
@@ -85,35 +72,45 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
}
}
- const toggleEngine = (batchId, engineKey) => {
- const current = selectedEngine[batchId] || defaultEngines
- let next
- if (current.includes(engineKey)) {
- if (current.length === 1) return // keep at least 1 engine selected
- next = current.filter((e) => e !== engineKey)
- } else {
- next = [...current, engineKey]
+ function openFilePickerForYolo(batch) {
+ setAppendChoiceBatch(null)
+ const input = document.createElement('input')
+ input.type = 'file'
+ input.accept = '.pt'
+ input.onchange = async (e) => {
+ const file = e.target.files?.[0]
+ if (!file) return
+ setBusyId(batch.id)
+ try {
+ const info = await api.inspectModel(file)
+ setAppendModalState({
+ batch,
+ file,
+ filename: info.filename,
+ classes: info.classes || [],
+ selectedClasses: info.classes || [],
+ threshold: 0.35,
+ iouThreshold: 0.8,
+ })
+ } catch (err) {
+ onError(err.message)
+ } finally {
+ setBusyId(null)
+ }
}
- setSelectedEngine({ ...selectedEngine, [batchId]: next })
+ input.click()
}
- const toggleEngineClass = (batchId, engineKey, className, defaultClasses) => {
- const batchFilters = engineClassMap[batchId] || {}
- const currentEngClasses = batchFilters[engineKey] || defaultClasses
- let next
- if (currentEngClasses.includes(className)) {
- if (currentEngClasses.length === 1) return // keep at least 1 class
- next = currentEngClasses.filter((c) => c !== className)
- } else {
- next = [...currentEngClasses, className]
- }
- setEngineClassMap({
- ...engineClassMap,
- [batchId]: {
- ...batchFilters,
- [engineKey]: next,
- },
+ function openSam3AppendModal(batch) {
+ setAppendChoiceBatch(null)
+ const projectClasses = project?.classes?.map((c) => c.name) || []
+ setSam3AppendState({
+ batch,
+ selectedClasses: [...projectClasses],
+ threshold: 0.35,
+ iouThreshold: 0.8,
})
+ setCustomPromptInput('')
}
async function deleteBatch(batch) {
@@ -151,20 +148,13 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
Batch Range Frames Reviewed
- Shapes Model & Class Configuration Status
+ Shapes Status
{batches.map((batch) => {
const batchJob = activeJobs?.find((j) => j.batch_id === batch.id)
const isProcessing = Boolean(batchJob || busyId === batch.id)
- const activeEngines = selectedEngine[batch.id] || defaultEngines
- const isFilterExpanded = expandedFilterBatchId === batch.id
- const batchMap = engineClassMap[batch.id] || {}
-
- const sam3Active = batchMap.sam3 || model1Classes
- const model1Active = batchMap.base_model || model1Classes
- const model2Active = batchMap.secondary_model || model2Classes
return (
@@ -180,96 +170,6 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
{batch.frame_count}
{batch.reviewed}/{batch.frame_count}
{batch.annotation_count}
-
-
-
- toggleEngine(batch.id, 'sam3')}
- disabled={isProcessing}
- title="SAM3 Zero-shot Text Prompt Engine"
- >
- 🤖 SAM3
-
- toggleEngine(batch.id, 'base_model')}
- disabled={isProcessing}
- title="Primary Model 1 (Base / Trained)"
- >
- ⚡ Model 1
-
- {hasSecondaryModel && (
- toggleEngine(batch.id, 'secondary_model')}
- disabled={isProcessing}
- title={project?.secondary_model_name || "Secondary Uploaded Model 2"}
- >
- 🎯 Model 2
-
- )}
- setExpandedFilterBatchId(isFilterExpanded ? null : batch.id)}
- disabled={isProcessing}
- >
- ⚙️ {isFilterExpanded ? 'Hide' : 'Per-Engine'}
-
-
-
-
- Classes:
- {model1Classes.map((clsName) => {
- const isModel1On = activeEngines.includes('base_model') && model1Active.includes(clsName)
- const isSam3On = activeEngines.includes('sam3') && sam3Active.includes(clsName)
- const isSecOn = activeEngines.includes('secondary_model') && hasSecondaryModel && model2Active.includes(clsName)
- const isActiveAny = isModel1On || isSam3On || isSecOn
-
- return (
- {
- if (activeEngines.includes('base_model')) toggleEngineClass(batch.id, 'base_model', clsName, model1Classes)
- if (activeEngines.includes('sam3')) toggleEngineClass(batch.id, 'sam3', clsName, model1Classes)
- if (activeEngines.includes('secondary_model')) toggleEngineClass(batch.id, 'secondary_model', clsName, model2Classes)
- }}
- >
- {isActiveAny ? '✓ ' : ''}{clsName}
-
- )
- })}
-
-
-
{batchJob ? `${batchJob.type} (${batchJob.status})` : batch.status}
@@ -278,16 +178,20 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
autolabel(batch)}>
- {batchJob?.type === 'autolabel' ? 'Processing…' : `Auto-annotate (${activeEngines.map(e => e === 'sam3' ? 'SAM3' : e === 'base_model' ? 'Model 1' : 'Model 2').join('+')})`}
+ title="Annotate using project base model with target class selection"
+ onClick={() => openBaseModelAutolabelModal(batch)}>
+ {batchJob?.type === 'autolabel' ? 'Processing…' : 'Auto-annotate'}
+
+ setAppendChoiceBatch(batch)}>
+ + Append
+
+ resetAutoAnnotations(batch)}>
+ 🔄 Reset Auto
- {batch.annotation_count > 0 && (
- autolabel(batch, true)}>
- Resume
-
- )}
navigate(`/projects/${project.id}/review?batch=${batch.id}`)}>
Review ({batch.reviewed}/{batch.frame_count})
@@ -299,122 +203,438 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
-
- {/* Inline Per-Engine Class Filter Panel */}
- {isFilterExpanded && (
-
-
-
-
- 🛠️ Inline Per-Engine Class Filters (Select target classes per detector):
-
-
-
- {/* SAM3 Section */}
- {activeEngines.includes('sam3') && (
-
-
- 🤖 SAM3 Text Prompts
-
-
- {model1Classes.map((clsName) => {
- const isActive = sam3Active.includes(clsName)
- return (
- toggleEngineClass(batch.id, 'sam3', clsName, model1Classes)}
- >
- {isActive ? '✓ ' : ''}{clsName}
-
- )
- })}
-
-
- )}
-
- {/* Model 1 Section */}
- {activeEngines.includes('base_model') && (
-
-
- ⚡ Model 1 (Primary Base) Classes
-
-
- {model1Classes.map((clsName) => {
- const isActive = model1Active.includes(clsName)
- return (
- toggleEngineClass(batch.id, 'base_model', clsName, model1Classes)}
- >
- {isActive ? '✓ ' : ''}{clsName}
-
- )
- })}
-
-
- )}
-
- {/* Model 2 Section */}
- {activeEngines.includes('secondary_model') && hasSecondaryModel && (
-
-
- 🎯 Model 2 (Secondary Engine) Native Classes
-
-
- {model2Classes.map((clsName) => {
- const isActive = model2Active.includes(clsName)
- return (
- toggleEngineClass(batch.id, 'secondary_model', clsName, model2Classes)}
- >
- {isActive ? '✓ ' : ''}{clsName}
-
- )
- })}
-
-
- )}
-
-
-
-
- )}
)
})}
+
+ {/* Choice Modal: SAM3 vs Custom YOLO */}
+ {appendChoiceBatch && (
+
+
+
Select Engine to Append Annotations
+
+ Choose how you want to detect and append new classes to batch {appendChoiceBatch.batch_label} :
+
+
+
+ {/* SAM3 Card */}
+
openSam3AppendModal(appendChoiceBatch)}
+ >
+
+
🤖 SAM3 Zero-Shot (Text Prompts)
+ Select >
+
+
+ Detect objects by typing any text prompt (e.g. "box", "sack", "person") without needing a pre-trained model.
+
+
+
+ {/* YOLO Card */}
+
openFilePickerForYolo(appendChoiceBatch)}
+ >
+
+
⚡ Custom YOLO Model (.pt)
+ Upload >
+
+
+ Upload a custom trained `.pt` model file from your computer and select which classes to append.
+
+
+
+
+
+ setAppendChoiceBatch(null)}>Cancel
+
+
+
+ )}
+
+ {/* SAM3 Append Modal */}
+ {sam3AppendState && (
+
+
+
Append Annotations with SAM3
+
+ Select target text prompts to detect with SAM3:
+
+
+
+
+ Confidence Threshold:
+ {sam3AppendState.threshold}
+
+
setSam3AppendState({ ...sam3AppendState, threshold: parseFloat(e.target.value) })}
+ style={{ width: '100%', cursor: 'pointer' }}
+ />
+
+
+
+
+ NMS IoU Threshold:
+ {sam3AppendState.iouThreshold ?? 0.8}
+
+
setSam3AppendState({ ...sam3AppendState, iouThreshold: parseFloat(e.target.value) })}
+ style={{ width: '100%', cursor: 'pointer' }}
+ />
+
+
+
+
Target Prompts to Detect:
+
+ {sam3AppendState.selectedClasses.map((clsName) => (
+ {
+ const next = sam3AppendState.selectedClasses.filter((c) => c !== clsName)
+ setSam3AppendState({ ...sam3AppendState, selectedClasses: next })
+ }}
+ >
+ ✕ {clsName}
+
+ ))}
+
+
+ {/* Add Custom SAM3 Prompt */}
+
+ setCustomPromptInput(e.target.value)}
+ onKeyDown={(e) => {
+ if (e.key === 'Enter' && customPromptInput.trim()) {
+ const val = customPromptInput.trim().toLowerCase()
+ if (!sam3AppendState.selectedClasses.includes(val)) {
+ setSam3AppendState({
+ ...sam3AppendState,
+ selectedClasses: [...sam3AppendState.selectedClasses, val]
+ })
+ }
+ setCustomPromptInput('')
+ }
+ }}
+ style={{ flex: 1, padding: '6px 10px', fontSize: '0.82rem', background: '#09090b', border: '1px solid rgba(255,255,255,0.15)', borderRadius: 4, color: '#fff' }}
+ />
+ {
+ if (customPromptInput.trim()) {
+ const val = customPromptInput.trim().toLowerCase()
+ if (!sam3AppendState.selectedClasses.includes(val)) {
+ setSam3AppendState({
+ ...sam3AppendState,
+ selectedClasses: [...sam3AppendState.selectedClasses, val]
+ })
+ }
+ setCustomPromptInput('')
+ }
+ }}
+ >
+ + Add Prompt
+
+
+
+
+
+ setSam3AppendState(null)}>Cancel
+ {
+ const { batch, threshold, iouThreshold, selectedClasses } = sam3AppendState
+ setSam3AppendState(null)
+ setBusyId(batch.id)
+ try {
+ await api.startAutolabel(batch.id, {
+ resume: false,
+ append: true,
+ engine: 'sam3',
+ threshold,
+ iou_threshold: iouThreshold ?? 0.8,
+ engine_classes: { sam3: selectedClasses },
+ class_ids: null
+ })
+ onChanged()
+ } catch (err) {
+ onError(err.message)
+ } finally {
+ setBusyId(null)
+ }
+ }}
+ >
+ Start SAM3 Append
+
+
+
+
+ )}
+
+ {/* YOLO Custom Model Append Modal */}
+ {appendModalState && (
+
+
+
Append Annotations with Custom Model
+
+ Model file: {appendModalState.filename}
+
+
+
+
+ Confidence Threshold:
+ {appendModalState.threshold}
+
+
setAppendModalState({ ...appendModalState, threshold: parseFloat(e.target.value) })}
+ style={{ width: '100%', cursor: 'pointer' }}
+ />
+
+
+
+
+ NMS IoU Threshold:
+ {appendModalState.iouThreshold ?? 0.8}
+
+
setAppendModalState({ ...appendModalState, iouThreshold: parseFloat(e.target.value) })}
+ style={{ width: '100%', cursor: 'pointer' }}
+ />
+
+
+
+
+ Select Classes to Append:
+
+ {appendModalState.selectedClasses.length} of {appendModalState.classes.length} selected
+
+
+
+ {appendModalState.classes.length === 0 ? (
+
No embedded class names found in model file. All predictions will be appended.
+ ) : (
+
+ {appendModalState.classes.map((clsName) => {
+ const isChecked = appendModalState.selectedClasses.includes(clsName)
+ return (
+ {
+ const next = isChecked
+ ? appendModalState.selectedClasses.filter((c) => c !== clsName)
+ : [...appendModalState.selectedClasses, clsName]
+ setAppendModalState({ ...appendModalState, selectedClasses: next })
+ }}
+ >
+ {isChecked ? '✓ ' : ''}{clsName}
+
+ )
+ })}
+
+ )}
+
+
+
+ setAppendModalState(null)}>Cancel
+ 0 && appendModalState.selectedClasses.length === 0}
+ onClick={async () => {
+ const { batch, file, threshold, iouThreshold, selectedClasses } = appendModalState
+ setAppendModalState(null)
+ setBusyId(batch.id)
+ try {
+ await api.autolabelWithModel(batch.id, file, threshold, selectedClasses, iouThreshold ?? 0.8)
+ onChanged()
+ } catch (err) {
+ onError(err.message)
+ } finally {
+ setBusyId(null)
+ }
+ }}
+ >
+ Start Append Auto-Annotation
+
+
+
+
+ )}
+
+ {/* Base Model Auto-annotate Modal */}
+ {baseModelModalState && (
+
+
+
Auto-annotate Batch (Base Model)
+
+ Select target classes to detect using the project's base model:
+
+
+
+
+ Confidence Threshold:
+ {baseModelModalState.threshold}
+
+
setBaseModelModalState({ ...baseModelModalState, threshold: parseFloat(e.target.value) })}
+ style={{ width: '100%', cursor: 'pointer' }}
+ />
+
+
+
+
+ NMS IoU Threshold:
+ {baseModelModalState.iouThreshold ?? 0.8}
+
+
setBaseModelModalState({ ...baseModelModalState, iouThreshold: parseFloat(e.target.value) })}
+ style={{ width: '100%', cursor: 'pointer' }}
+ />
+
+
+
+
Target Classes to Detect:
+
+ {project?.classes?.map((cls) => {
+ const isChecked = baseModelModalState.selectedClasses.includes(cls.name)
+ return (
+ {
+ const newSel = isChecked
+ ? baseModelModalState.selectedClasses.filter((n) => n !== cls.name)
+ : [...baseModelModalState.selectedClasses, cls.name]
+ setBaseModelModalState({ ...baseModelModalState, selectedClasses: newSel })
+ }}
+ >
+ {isChecked ? '✓ ' : ''}{cls.name}
+
+ )
+ })}
+
+
+
+
+ setBaseModelModalState(null)}>Cancel
+ {
+ const { batch, threshold, iouThreshold, selectedClasses } = baseModelModalState
+ setBaseModelModalState(null)
+ setBusyId(batch.id)
+ try {
+ await api.startAutolabel(batch.id, {
+ resume: false,
+ append: false,
+ engine: 'base_model',
+ threshold,
+ iou_threshold: iouThreshold ?? 0.8,
+ target_class_names: selectedClasses,
+ })
+ onChanged()
+ } catch (err) {
+ onError(err.message)
+ } finally {
+ setBusyId(null)
+ }
+ }}
+ >
+ Start Auto-Annotation
+
+
+
+
+ )}
)
}
diff --git a/frontend/src/pages/ModelsPage.jsx b/frontend/src/pages/ModelsPage.jsx
index 7242292..2ec8ed4 100644
--- a/frontend/src/pages/ModelsPage.jsx
+++ b/frontend/src/pages/ModelsPage.jsx
@@ -91,6 +91,7 @@ export default function ModelsPage({ projectId, onProject }) {
const [error, setError] = useState('')
const [selectedBatchIds, setSelectedBatchIds] = useState([])
+ const [selectedClassIds, setSelectedClassIds] = useState([])
const load = useCallback(async () => {
const [loadedProject, loadedSummary, modelPayload, hw, jobsPayload] = await Promise.all([
@@ -103,6 +104,7 @@ export default function ModelsPage({ projectId, onProject }) {
setModels(modelPayload.models)
setHardware(hw)
setSelectedBatchIds(loadedSummary.batches.map((b) => b.id))
+ setSelectedClassIds(loadedProject.classes?.map((c) => c.class_id) || [])
const activeJob = jobsPayload.jobs?.find((j) => ['running', 'queued'].includes(j.status))
if (activeJob) setJob(activeJob)
}, [projectId])
@@ -126,6 +128,7 @@ export default function ModelsPage({ projectId, onProject }) {
setJob(await api.startTraining(projectId, {
epochs: Number(epochs),
batch_ids: selectedBatchIds.length > 0 ? selectedBatchIds : null,
+ class_ids: selectedClassIds.length > 0 ? selectedClassIds : null,
}))
} catch (exc) {
setError(exc.message)
@@ -146,6 +149,12 @@ export default function ModelsPage({ projectId, onProject }) {
}
}
+ const toggleClassSelect = (classId) => {
+ setSelectedClassIds((prev) =>
+ prev.includes(classId) ? prev.filter((cId) => cId !== classId) : [...prev, classId]
+ )
+ }
+
if (error && !project) return {error}
if (!project || !summary) return Loading…
@@ -169,103 +178,40 @@ export default function ModelsPage({ projectId, onProject }) {
{error}
}
-
-
Select Engine
-
- Select how you want to train your model. Configure custom fine-tuning parameters or use SAM3 grounding backbone.
-
-
-
-
-
- Selected
- Custom Training (YOLO11)
-
-
- Fine-tune on the merged master dataset using pre-configured hardware batch size and image resolution.
-
-
-
-
-
- Rapid NAS
- Neural Architecture Search
-
-
- Automated model selection optimized for latency and accuracy trade-offs on your specific project dataset.
-
-
-
-
-
-
Uploaded Models for Training & Auto-Annotation
-
Upload up to 2 models to use for fine-tuning baseline or auto-annotating new video batches:
+
Base Model Configuration
+
Upload a base YOLO model checkpoint (.pt) to use for fine-tuning baseline and auto-annotation:
-
-
-
Primary Model 1 (Base Model)
-
- Used as fine-tuning starting point and benchmark baseline.
-
-
- Status: {project.base_model_path ? 'Custom model.pt loaded' : 'Default yolo11n.pt'}
-
-
- Model 1 Classes: {project.classes?.map((c) => c.name).join(', ')}
-
-
{
- const file = e.target.files?.[0]
- if (!file) return
- try {
- await api.uploadBaseModel(projectId, file)
- load()
- } catch (err) {
- setError(err.message)
- }
- }}
- />
-
- Upload Primary Model 1 (.pt)
-
+
+
Base Model
+
+ Used as fine-tuning starting point and baseline benchmark.
+
+
+ Status: {project.base_model_path ? 'Custom model.pt loaded' : 'Default yolo11n.pt'}
-
-
-
Secondary Model 2 (Auto-Annotate Engine)
-
- Used as an auxiliary engine choice for fast auto-annotation.
-
-
- Status: {project.secondary_model_path ? (project.secondary_model_name || 'secondary_model.pt') : 'None uploaded'}
-
-
- Model 2 Native Classes: {project.secondary_model_classes?.length > 0 ? project.secondary_model_classes.join(', ') : 'Extracted on upload'}
-
-
{
- const file = e.target.files?.[0]
- if (!file) return
- try {
- await api.uploadSecondaryModel(projectId, file)
- load()
- } catch (err) {
- setError(err.message)
- }
- }}
- />
-
- Upload Secondary Model 2 (.pt)
-
+
+ Base Model Classes: {project.classes?.map((c) => c.name).join(', ')}
+
{
+ const file = e.target.files?.[0]
+ if (!file) return
+ try {
+ await api.uploadBaseModel(projectId, file)
+ load()
+ } catch (err) {
+ setError(err.message)
+ }
+ }}
+ />
+
+ Upload Base Model (.pt)
+
@@ -291,6 +237,34 @@ export default function ModelsPage({ projectId, onProject }) {
: `${project.base_model_fallback} (no base model uploaded)`}{' '}
on the selected dataset batches ({selectedBatchIds.length}/{summary.batches.length} selected).
+
+
+
Target Classes to Train:
+
+ {project.classes?.map((cls) => {
+ const isChecked = selectedClassIds.includes(cls.class_id)
+ return (
+ toggleClassSelect(cls.class_id)}
+ >
+ {isChecked ? '✓ ' : ''}{cls.name}
+
+ )
+ })}
+
+
+
Epochs
)}
+ disabled={running || summary.splits.train === 0 || selectedBatchIds.length === 0 || selectedClassIds.length === 0}>
{running ? 'Training…' : 'Start training'}
{summary.splits.train === 0 && (
@@ -312,6 +286,9 @@ export default function ModelsPage({ projectId, onProject }) {
{selectedBatchIds.length === 0 && summary.splits.train > 0 && (
Select at least one dataset batch to train.
)}
+ {selectedClassIds.length === 0 && (
+
Select at least one class to train.
+ )}
{job && (
diff --git a/frontend/src/pages/Sam3PlaygroundPage.jsx b/frontend/src/pages/Sam3PlaygroundPage.jsx
new file mode 100644
index 0000000..7067cb7
--- /dev/null
+++ b/frontend/src/pages/Sam3PlaygroundPage.jsx
@@ -0,0 +1,328 @@
+import React, { useState, useRef } from 'react'
+import { api } from '../api'
+
+const COLORS = [
+ '#38bdf8', '#c084fc', '#f43f5e', '#34d399', '#fbbf24',
+ '#818cf8', '#f472b6', '#a3e635', '#2dd4bf', '#fb923c'
+]
+
+function getColor(index) {
+ return COLORS[index % COLORS.length]
+}
+
+export default function Sam3PlaygroundPage() {
+ const [imageFile, setImageFile] = useState(null)
+ const [previewUrl, setPreviewUrl] = useState(null)
+ const [promptsText, setPromptsText] = useState('white plastic bag, forklift, pallet')
+ const [threshold, setThreshold] = useState(0.35)
+ const [iouThreshold, setIouThreshold] = useState(0.8)
+ const [loading, setLoading] = useState(false)
+ const [error, setError] = useState('')
+ const [result, setResult] = useState(null)
+ const fileInputRef = useRef(null)
+
+ const handleFileSelect = (file) => {
+ if (!file) return
+ setImageFile(file)
+ setPreviewUrl(URL.createObjectURL(file))
+ setResult(null)
+ setError('')
+ }
+
+ const handleDrop = (e) => {
+ e.preventDefault()
+ const file = e.dataTransfer?.files?.[0]
+ if (file && file.type.startsWith('image/')) {
+ handleFileSelect(file)
+ }
+ }
+
+ const runTest = async () => {
+ if (!imageFile) {
+ setError('Please select or upload an image first.')
+ return
+ }
+ if (!promptsText.trim()) {
+ setError('Please enter at least one text prompt.')
+ return
+ }
+ setLoading(true)
+ setError('')
+ try {
+ const payload = await api.sam3PlaygroundTest(imageFile, promptsText, threshold, iouThreshold)
+ setResult(payload)
+ } catch (err) {
+ setError(err.message || 'SAM3 test failed.')
+ } finally {
+ setLoading(false)
+ }
+ }
+
+ return (
+
+
+
+
+ 🤖 SAM3 Global Playground
+
+
+ Test SAM3 zero-shot text grounding prompts on any image without affecting projects or database records.
+
+
+
+
+ {error && (
+
+ ⚠️ {error}
+
+ )}
+
+
+ {/* Left Control Panel */}
+
+
Configuration
+
+ {/* Image Selection / Dropzone */}
+
+
+ Select Test Image:
+
+
e.preventDefault()}
+ onDrop={handleDrop}
+ onClick={() => fileInputRef.current?.click()}
+ style={{
+ border: '2px dashed rgba(192, 132, 252, 0.4)',
+ borderRadius: 8,
+ padding: '16px 12px',
+ textAlign: 'center',
+ cursor: 'pointer',
+ background: 'rgba(0,0,0,0.3)',
+ transition: 'all 150ms ease'
+ }}
+ >
+
handleFileSelect(e.target.files?.[0])}
+ />
+ {imageFile ? (
+
+ 🖼️ {imageFile.name} ({(imageFile.size / 1024).toFixed(0)} KB)
+
+ ) : (
+
+ Click to choose image or drag & drop here (.jpg, .png)
+
+ )}
+
+
+
+ {/* Text Prompts Input */}
+
+
+ Descriptive Text Prompts (Comma Separated):
+
+
+
+ {/* Threshold Sliders */}
+
+
+ Confidence Threshold:
+ {threshold}
+
+
setThreshold(parseFloat(e.target.value))}
+ style={{ width: '100%', cursor: 'pointer' }}
+ />
+
+
+
+
+ NMS IoU Threshold:
+ {iouThreshold}
+
+
setIouThreshold(parseFloat(e.target.value))}
+ style={{ width: '100%', cursor: 'pointer' }}
+ />
+
+
+ {/* Run Button */}
+
+ {loading ? 'Running SAM3 Engine…' : '🚀 Run SAM3 Test'}
+
+
+ {/* Detection Stats List */}
+ {result && (
+
+
+ Results: {result.detections?.length || 0} Object(s) Found
+
+
+ {result.detections.map((det, i) => {
+ const color = getColor(det.class_id)
+ return (
+
+ {det.prompt}
+ {(det.score * 100).toFixed(1)}%
+
+ )
+ })}
+
+
+ )}
+
+
+ {/* Right Canvas Viewer */}
+
+ {!previewUrl ? (
+
+
🖼️
+
Upload an image on the left to start testing SAM3 text prompts.
+
+ ) : (
+
+
+
+ {/* Overlay Canvas SVG */}
+ {result && result.width && result.height && (
+
+ {result.detections.map((det, i) => {
+ const color = getColor(det.class_id)
+ const [x0, y0, x1, y1] = [
+ det.box[0] * result.width,
+ det.box[1] * result.height,
+ det.box[2] * result.width,
+ det.box[3] * result.height
+ ]
+ const bw = x1 - x0
+ const bh = y1 - y0
+
+ return (
+
+ {/* Polygon Masks */}
+ {det.polygons?.map((poly, pIdx) => {
+ const ptsString = poly.map(([px, py]) => `${px * result.width},${py * result.height}`).join(' ')
+ return (
+
+ )
+ })}
+
+ {/* Bounding Box */}
+
+
+ {/* Label Badge */}
+
+
+ {det.prompt} ({Math.round(det.score * 100)}%)
+
+
+ )
+ })}
+
+ )}
+
+ )}
+
+
+
+ )
+}
diff --git a/scripts/restart_app.sh b/scripts/restart_app.sh
index e904d9f..d80a9da 100755
--- a/scripts/restart_app.sh
+++ b/scripts/restart_app.sh
@@ -7,17 +7,18 @@ cd "$REPO_DIR"
echo "🔄 Restarting app services..."
# Kill running uvicorn and vite processes
+fuser -k -9 8000/tcp 8080/tcp 5173/tcp || true
pkill -f "uvicorn backend.main:app" || true
pkill -f "vite" || true
-sleep 1
+sleep 2.5
export PATH="$HOME/.local/bin:$PATH"
# Run backend
-nohup uv run uvicorn backend.main:app --host 0.0.0.0 --port 8000 > "$REPO_DIR/data/backend.log" 2>&1 &
+setsid uv run uvicorn backend.main:app --host 0.0.0.0 --port 8000 > "$REPO_DIR/data/backend.log" 2>&1 &
# Run frontend
cd "$REPO_DIR/frontend"
-nohup npm run dev > "$REPO_DIR/data/frontend.log" 2>&1 &
+setsid npm run dev -- --host 0.0.0.0 --port 8080 > "$REPO_DIR/data/frontend.log" 2>&1 &
echo "✅ Services restarted."