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 })
- {!collapsed &&
MODELS
} + {!collapsed &&
MODELS & TOOLS
} handleNav(e, `/projects/${pId}/models`)} className={`sidebar-item ${route.name === 'models' ? 'active' : ''}`} title="Train & Select Engine"> {!collapsed && Train & Select Engine} + handleNav(e, '/sam3-playground')} className={`sidebar-item ${route.name === 'sam3-playground' ? 'active' : ''}`} title="SAM3 Playground"> + 🤖 + {!collapsed && SAM3 Playground} +
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 }) { BatchRangeFramesReviewed - ShapesModel & Class ConfigurationStatus + ShapesStatus {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} - -
-
- - - {hasSecondaryModel && ( - - )} - -
- -
- 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 ( - - ) - })} -
-
- {batchJob ? `${batchJob.type} (${batchJob.status})` : batch.status} @@ -278,16 +178,20 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
+ + - {batch.annotation_count > 0 && ( - - )}
- - {/* 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 ( - - ) - })} -
-
- )} - - {/* Model 1 Section */} - {activeEngines.includes('base_model') && ( -
-
- ⚡ Model 1 (Primary Base) Classes -
-
- {model1Classes.map((clsName) => { - const isActive = model1Active.includes(clsName) - return ( - - ) - })} -
-
- )} - - {/* Model 2 Section */} - {activeEngines.includes('secondary_model') && hasSecondaryModel && ( -
-
- 🎯 Model 2 (Secondary Engine) Native Classes -
-
- {model2Classes.map((clsName) => { - const isActive = model2Active.includes(clsName) - return ( - - ) - })} -
-
- )} -
-
- - - )}
) })} + + {/* 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. +

+
+
+ +
+ +
+
+
+ )} + + {/* 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) => ( + + ))} +
+ + {/* 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' }} + /> + +
+
+ +
+ + +
+
+
+ )} + + {/* 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 ( + + ) + })} +
+ )} +
+ +
+ + +
+
+
+ )} + + {/* 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 ( + + ) + })} +
+
+ +
+ + +
+
+
+ )}
) } 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) - } - }} - /> - +
+

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) - } - }} - /> - +
+ 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) + } + }} + /> +
@@ -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).

+ +
+ +
+ {project.classes?.map((cls) => { + const isChecked = selectedClassIds.includes(cls.class_id) + return ( + + ) + })} +
+
+
)} {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 */} +
+ +
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 */} +
+ +