update from asus 106
This commit is contained in:
1 parent
6637fb1302
commit
8285400254
28 files changed
+3189
-433
No files matched your search
@@ -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()
|
||||
@@ -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
|
||||
@@ -0,0 +1 @@
|
||||
"""Package marker."""
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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: ...
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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/"
|
||||
}
|
||||
+121
-6
@@ -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}
|
||||
|
||||
@@ -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))
|
||||
|
||||
+63
-68
@@ -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,6 +200,9 @@ 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})
|
||||
|
||||
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)")
|
||||
|
||||
+46
-2
@@ -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:
|
||||
|
||||
+3
-3
@@ -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."
|
||||
|
||||
+11
-4
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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"]]
|
||||
|
||||
@@ -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(
|
||||
|
||||
+9
-4
@@ -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,
|
||||
|
||||
@@ -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' && <Sam3PlaygroundPage />}
|
||||
</ErrorBoundary>
|
||||
</main>
|
||||
</div>
|
||||
|
||||
@@ -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}`),
|
||||
|
||||
|
||||
@@ -74,11 +74,15 @@ export default function Sidebar({ route, currentProject, theme, onToggleTheme })
|
||||
</div>
|
||||
|
||||
<div className="sidebar-section">
|
||||
{!collapsed && <div className="sidebar-section-title">MODELS</div>}
|
||||
{!collapsed && <div className="sidebar-section-title">MODELS & TOOLS</div>}
|
||||
<a href={`#/projects/${pId}/models`} onClick={(e) => handleNav(e, `/projects/${pId}/models`)} className={`sidebar-item ${route.name === 'models' ? 'active' : ''}`} title="Train & Select Engine">
|
||||
<span className="sidebar-icon"><RocketIcon size={16} /></span>
|
||||
{!collapsed && <span>Train & Select Engine</span>}
|
||||
</a>
|
||||
<a href="#/sam3-playground" onClick={(e) => handleNav(e, '/sam3-playground')} className={`sidebar-item ${route.name === 'sam3-playground' ? 'active' : ''}`} title="SAM3 Playground">
|
||||
<span className="sidebar-icon">🤖</span>
|
||||
{!collapsed && <span>SAM3 Playground</span>}
|
||||
</a>
|
||||
</div>
|
||||
|
||||
<div className="sidebar-spacer" />
|
||||
|
||||
+491
-271
@@ -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 }) {
|
||||
<thead>
|
||||
<tr>
|
||||
<th>Batch</th><th>Range</th><th>Frames</th><th>Reviewed</th>
|
||||
<th>Shapes</th><th>Model & Class Configuration</th><th>Status</th><th />
|
||||
<th>Shapes</th><th>Status</th><th />
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
{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 (
|
||||
<React.Fragment key={batch.id}>
|
||||
@@ -180,96 +170,6 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
|
||||
<td className="num">{batch.frame_count}</td>
|
||||
<td className="num">{batch.reviewed}/{batch.frame_count}</td>
|
||||
<td className="num">{batch.annotation_count}</td>
|
||||
<td style={{ minWidth: 260 }}>
|
||||
<div style={{ display: 'flex', flexDirection: 'column', gap: 6 }}>
|
||||
<div style={{ display: 'flex', gap: 6, flexWrap: 'wrap', alignItems: 'center' }}>
|
||||
<button
|
||||
type="button"
|
||||
className={`btn ${activeEngines.includes('sam3') ? 'btn-primary' : 'btn-ghost'}`}
|
||||
style={{ padding: '2px 6px', fontSize: '0.72rem', cursor: 'pointer', opacity: activeEngines.includes('sam3') ? 1 : 0.6 }}
|
||||
onClick={() => toggleEngine(batch.id, 'sam3')}
|
||||
disabled={isProcessing}
|
||||
title="SAM3 Zero-shot Text Prompt Engine"
|
||||
>
|
||||
🤖 SAM3
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
className={`btn ${activeEngines.includes('base_model') ? 'btn-primary' : 'btn-ghost'}`}
|
||||
style={{ padding: '2px 6px', fontSize: '0.72rem', cursor: 'pointer', opacity: activeEngines.includes('base_model') ? 1 : 0.6 }}
|
||||
onClick={() => toggleEngine(batch.id, 'base_model')}
|
||||
disabled={isProcessing}
|
||||
title="Primary Model 1 (Base / Trained)"
|
||||
>
|
||||
⚡ Model 1
|
||||
</button>
|
||||
{hasSecondaryModel && (
|
||||
<button
|
||||
type="button"
|
||||
className={`btn ${activeEngines.includes('secondary_model') ? 'btn-primary' : 'btn-ghost'}`}
|
||||
style={{ padding: '2px 6px', fontSize: '0.72rem', cursor: 'pointer', opacity: activeEngines.includes('secondary_model') ? 1 : 0.6 }}
|
||||
onClick={() => toggleEngine(batch.id, 'secondary_model')}
|
||||
disabled={isProcessing}
|
||||
title={project?.secondary_model_name || "Secondary Uploaded Model 2"}
|
||||
>
|
||||
🎯 Model 2
|
||||
</button>
|
||||
)}
|
||||
<button
|
||||
type="button"
|
||||
className="btn btn-ghost"
|
||||
style={{
|
||||
padding: '2px 6px',
|
||||
fontSize: '0.72rem',
|
||||
display: 'flex',
|
||||
alignItems: 'center',
|
||||
gap: 3,
|
||||
borderColor: isFilterExpanded ? '#c084fc' : 'rgba(255,255,255,0.2)',
|
||||
color: isFilterExpanded ? '#c084fc' : '#e4e4e7',
|
||||
}}
|
||||
onClick={() => setExpandedFilterBatchId(isFilterExpanded ? null : batch.id)}
|
||||
disabled={isProcessing}
|
||||
>
|
||||
⚙️ {isFilterExpanded ? 'Hide' : 'Per-Engine'}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 4, alignItems: 'center', fontSize: '0.72rem' }}>
|
||||
<span className="muted" style={{ fontSize: '0.7rem' }}>Classes:</span>
|
||||
{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 (
|
||||
<button
|
||||
key={clsName}
|
||||
type="button"
|
||||
className="tag"
|
||||
style={{
|
||||
padding: '1px 5px',
|
||||
fontSize: '0.7rem',
|
||||
cursor: 'pointer',
|
||||
background: isActiveAny ? 'rgba(56, 189, 248, 0.25)' : 'rgba(255,255,255,0.05)',
|
||||
color: isActiveAny ? '#38bdf8' : '#71717a',
|
||||
border: isActiveAny ? '1px solid rgba(56, 189, 248, 0.5)' : '1px solid rgba(255,255,255,0.1)',
|
||||
}}
|
||||
title={`Toggle ${clsName} filter`}
|
||||
disabled={isProcessing}
|
||||
onClick={() => {
|
||||
if (activeEngines.includes('base_model')) toggleEngineClass(batch.id, 'base_model', clsName, model1Classes)
|
||||
if (activeEngines.includes('sam3')) toggleEngineClass(batch.id, 'sam3', clsName, model1Classes)
|
||||
if (activeEngines.includes('secondary_model')) toggleEngineClass(batch.id, 'secondary_model', clsName, model2Classes)
|
||||
}}
|
||||
>
|
||||
{isActiveAny ? '✓ ' : ''}{clsName}
|
||||
</button>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
</td>
|
||||
<td>
|
||||
<span className={`tag ${batchJob ? 'info' : ''}`}>
|
||||
{batchJob ? `${batchJob.type} (${batchJob.status})` : batch.status}
|
||||
@@ -278,16 +178,20 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
|
||||
<td>
|
||||
<div className="row" style={{ gap: 6 }}>
|
||||
<button className="btn btn-primary" disabled={isProcessing || batch.frame_count === 0}
|
||||
onClick={() => autolabel(batch)}>
|
||||
{batchJob?.type === 'autolabel' ? 'Processing…' : `Auto-annotate (${activeEngines.map(e => e === 'sam3' ? 'SAM3' : e === 'base_model' ? 'Model 1' : 'Model 2').join('+')})`}
|
||||
title="Annotate using project base model with target class selection"
|
||||
onClick={() => openBaseModelAutolabelModal(batch)}>
|
||||
{batchJob?.type === 'autolabel' ? 'Processing…' : 'Auto-annotate'}
|
||||
</button>
|
||||
{batch.annotation_count > 0 && (
|
||||
<button className="btn" disabled={isProcessing || batch.frame_count === 0}
|
||||
title="Skip frames that already have automatic shapes"
|
||||
onClick={() => autolabel(batch, true)}>
|
||||
Resume
|
||||
title="Append new class detections using custom YOLO model or SAM3 text prompts"
|
||||
onClick={() => setAppendChoiceBatch(batch)}>
|
||||
+ Append
|
||||
</button>
|
||||
<button className="btn" disabled={isProcessing || batch.annotation_count === 0}
|
||||
title="Clear all auto-generated shapes and reset frame review statuses to pending"
|
||||
onClick={() => resetAutoAnnotations(batch)}>
|
||||
🔄 Reset Auto
|
||||
</button>
|
||||
)}
|
||||
<button className="btn" disabled={batch.frame_count === 0}
|
||||
onClick={() => navigate(`/projects/${project.id}/review?batch=${batch.id}`)}>
|
||||
Review ({batch.reviewed}/{batch.frame_count})
|
||||
@@ -299,122 +203,438 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
|
||||
</div>
|
||||
</td>
|
||||
</tr>
|
||||
|
||||
{/* Inline Per-Engine Class Filter Panel */}
|
||||
{isFilterExpanded && (
|
||||
<tr>
|
||||
<td colSpan={8} style={{ background: '#09090b', padding: '12px 16px', borderBottom: '1px solid rgba(168,85,247,0.3)' }}>
|
||||
<div style={{ display: 'flex', flexDirection: 'column', gap: 10 }}>
|
||||
<div style={{ fontSize: '0.82rem', fontWeight: 600, color: '#c084fc' }}>
|
||||
🛠️ Inline Per-Engine Class Filters (Select target classes per detector):
|
||||
</div>
|
||||
|
||||
<div style={{ display: 'grid', gridTemplateColumns: 'repeat(auto-fit, minmax(240px, 1fr))', gap: 12 }}>
|
||||
{/* SAM3 Section */}
|
||||
{activeEngines.includes('sam3') && (
|
||||
<div style={{ background: 'rgba(24, 24, 27, 0.8)', padding: 10, borderRadius: 6, border: '1px solid rgba(255,255,255,0.1)' }}>
|
||||
<div style={{ fontSize: '0.78rem', fontWeight: 600, color: '#a855f7', marginBottom: 6 }}>
|
||||
🤖 SAM3 Text Prompts
|
||||
</div>
|
||||
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 4 }}>
|
||||
{model1Classes.map((clsName) => {
|
||||
const isActive = sam3Active.includes(clsName)
|
||||
return (
|
||||
<button
|
||||
key={clsName}
|
||||
type="button"
|
||||
className="tag"
|
||||
style={{
|
||||
padding: '2px 6px',
|
||||
fontSize: '0.72rem',
|
||||
cursor: 'pointer',
|
||||
background: isActive ? 'rgba(168, 85, 247, 0.3)' : 'rgba(255,255,255,0.05)',
|
||||
color: isActive ? '#f3e8ff' : '#666',
|
||||
border: isActive ? '1px solid rgba(168, 85, 247, 0.6)' : '1px solid transparent',
|
||||
}}
|
||||
onClick={() => toggleEngineClass(batch.id, 'sam3', clsName, model1Classes)}
|
||||
>
|
||||
{isActive ? '✓ ' : ''}{clsName}
|
||||
</button>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Model 1 Section */}
|
||||
{activeEngines.includes('base_model') && (
|
||||
<div style={{ background: 'rgba(24, 24, 27, 0.8)', padding: 10, borderRadius: 6, border: '1px solid rgba(255,255,255,0.1)' }}>
|
||||
<div style={{ fontSize: '0.78rem', fontWeight: 600, color: '#38bdf8', marginBottom: 6 }}>
|
||||
⚡ Model 1 (Primary Base) Classes
|
||||
</div>
|
||||
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 4 }}>
|
||||
{model1Classes.map((clsName) => {
|
||||
const isActive = model1Active.includes(clsName)
|
||||
return (
|
||||
<button
|
||||
key={clsName}
|
||||
type="button"
|
||||
className="tag"
|
||||
style={{
|
||||
padding: '2px 6px',
|
||||
fontSize: '0.72rem',
|
||||
cursor: 'pointer',
|
||||
background: isActive ? 'rgba(56, 189, 248, 0.3)' : 'rgba(255,255,255,0.05)',
|
||||
color: isActive ? '#e0f2fe' : '#666',
|
||||
border: isActive ? '1px solid rgba(56, 189, 248, 0.6)' : '1px solid transparent',
|
||||
}}
|
||||
onClick={() => toggleEngineClass(batch.id, 'base_model', clsName, model1Classes)}
|
||||
>
|
||||
{isActive ? '✓ ' : ''}{clsName}
|
||||
</button>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Model 2 Section */}
|
||||
{activeEngines.includes('secondary_model') && hasSecondaryModel && (
|
||||
<div style={{ background: 'rgba(24, 24, 27, 0.8)', padding: 10, borderRadius: 6, border: '1px solid rgba(255,255,255,0.1)' }}>
|
||||
<div style={{ fontSize: '0.78rem', fontWeight: 600, color: '#f43f5e', marginBottom: 6 }}>
|
||||
🎯 Model 2 (Secondary Engine) Native Classes
|
||||
</div>
|
||||
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 4 }}>
|
||||
{model2Classes.map((clsName) => {
|
||||
const isActive = model2Active.includes(clsName)
|
||||
return (
|
||||
<button
|
||||
key={clsName}
|
||||
type="button"
|
||||
className="tag"
|
||||
style={{
|
||||
padding: '2px 6px',
|
||||
fontSize: '0.72rem',
|
||||
cursor: 'pointer',
|
||||
background: isActive ? 'rgba(244, 63, 94, 0.3)' : 'rgba(255,255,255,0.05)',
|
||||
color: isActive ? '#ffe4e6' : '#666',
|
||||
border: isActive ? '1px solid rgba(244, 63, 94, 0.6)' : '1px solid transparent',
|
||||
}}
|
||||
onClick={() => toggleEngineClass(batch.id, 'secondary_model', clsName, model2Classes)}
|
||||
>
|
||||
{isActive ? '✓ ' : ''}{clsName}
|
||||
</button>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
</td>
|
||||
</tr>
|
||||
)}
|
||||
</React.Fragment>
|
||||
)
|
||||
})}
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
{/* Choice Modal: SAM3 vs Custom YOLO */}
|
||||
{appendChoiceBatch && (
|
||||
<div style={{
|
||||
position: 'fixed', top: 0, left: 0, right: 0, bottom: 0,
|
||||
background: 'rgba(0,0,0,0.8)', display: 'flex', alignItems: 'center',
|
||||
justifyContent: 'center', zIndex: 9999
|
||||
}}>
|
||||
<div className="panel" style={{ width: 480, maxWidth: '92vw', padding: 22, border: '1px solid rgba(168, 85, 247, 0.4)', background: '#18181b' }}>
|
||||
<h3 style={{ margin: '0 0 6px 0', color: '#c084fc', fontSize: '1.05rem' }}>Select Engine to Append Annotations</h3>
|
||||
<p className="hint" style={{ fontSize: '0.82rem', marginBottom: 18 }}>
|
||||
Choose how you want to detect and append new classes to batch <strong>{appendChoiceBatch.batch_label}</strong>:
|
||||
</p>
|
||||
|
||||
<div style={{ display: 'grid', gridTemplateColumns: '1fr', gap: 12, marginBottom: 20 }}>
|
||||
{/* SAM3 Card */}
|
||||
<div
|
||||
style={{
|
||||
padding: 14, background: 'rgba(24,24,27,0.9)', borderRadius: 8,
|
||||
border: '1px solid rgba(168, 85, 247, 0.4)', cursor: 'pointer',
|
||||
transition: 'all 150ms ease'
|
||||
}}
|
||||
onClick={() => openSam3AppendModal(appendChoiceBatch)}
|
||||
>
|
||||
<div className="row" style={{ justifyContent: 'space-between', marginBottom: 4 }}>
|
||||
<h4 style={{ margin: 0, color: '#c084fc', fontSize: '0.92rem' }}>🤖 SAM3 Zero-Shot (Text Prompts)</h4>
|
||||
<span className="btn btn-ghost" style={{ padding: '2px 8px', fontSize: '0.75rem' }}>Select ></span>
|
||||
</div>
|
||||
<p className="hint" style={{ margin: 0, fontSize: '0.78rem' }}>
|
||||
Detect objects by typing any text prompt (e.g. "box", "sack", "person") without needing a pre-trained model.
|
||||
</p>
|
||||
</div>
|
||||
|
||||
{/* YOLO Card */}
|
||||
<div
|
||||
style={{
|
||||
padding: 14, background: 'rgba(24,24,27,0.9)', borderRadius: 8,
|
||||
border: '1px solid rgba(56, 189, 248, 0.4)', cursor: 'pointer',
|
||||
transition: 'all 150ms ease'
|
||||
}}
|
||||
onClick={() => openFilePickerForYolo(appendChoiceBatch)}
|
||||
>
|
||||
<div className="row" style={{ justifyContent: 'space-between', marginBottom: 4 }}>
|
||||
<h4 style={{ margin: 0, color: '#38bdf8', fontSize: '0.92rem' }}>⚡ Custom YOLO Model (.pt)</h4>
|
||||
<span className="btn btn-ghost" style={{ padding: '2px 8px', fontSize: '0.75rem' }}>Upload ></span>
|
||||
</div>
|
||||
<p className="hint" style={{ margin: 0, fontSize: '0.78rem' }}>
|
||||
Upload a custom trained `.pt` model file from your computer and select which classes to append.
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="row" style={{ justifyContent: 'flex-end' }}>
|
||||
<button className="btn btn-ghost" onClick={() => setAppendChoiceBatch(null)}>Cancel</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* SAM3 Append Modal */}
|
||||
{sam3AppendState && (
|
||||
<div style={{
|
||||
position: 'fixed', top: 0, left: 0, right: 0, bottom: 0,
|
||||
background: 'rgba(0,0,0,0.8)', display: 'flex', alignItems: 'center',
|
||||
justifyContent: 'center', zIndex: 9999
|
||||
}}>
|
||||
<div className="panel" style={{ width: 500, maxWidth: '92vw', padding: 22, border: '1px solid rgba(168, 85, 247, 0.4)', background: '#18181b' }}>
|
||||
<h3 style={{ margin: '0 0 8px 0', color: '#c084fc', fontSize: '1.05rem' }}>Append Annotations with SAM3</h3>
|
||||
<p className="hint" style={{ fontSize: '0.82rem', marginBottom: 16 }}>
|
||||
Select target text prompts to detect with SAM3:
|
||||
</p>
|
||||
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<div className="row" style={{ justifyContent: 'space-between', marginBottom: 6 }}>
|
||||
<span className="hint" style={{ fontSize: '0.82rem' }}>Confidence Threshold:</span>
|
||||
<strong style={{ color: '#c084fc', fontSize: '0.85rem' }}>{sam3AppendState.threshold}</strong>
|
||||
</div>
|
||||
<input
|
||||
type="range"
|
||||
min="0.05"
|
||||
max="0.95"
|
||||
step="0.05"
|
||||
value={sam3AppendState.threshold}
|
||||
onChange={(e) => setSam3AppendState({ ...sam3AppendState, threshold: parseFloat(e.target.value) })}
|
||||
style={{ width: '100%', cursor: 'pointer' }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<div className="row" style={{ justifyContent: 'space-between', marginBottom: 6 }}>
|
||||
<span className="hint" style={{ fontSize: '0.82rem' }}>NMS IoU Threshold:</span>
|
||||
<strong style={{ color: '#38bdf8', fontSize: '0.85rem' }}>{sam3AppendState.iouThreshold ?? 0.8}</strong>
|
||||
</div>
|
||||
<input
|
||||
type="range"
|
||||
min="0.1"
|
||||
max="0.9"
|
||||
step="0.05"
|
||||
value={sam3AppendState.iouThreshold ?? 0.8}
|
||||
onChange={(e) => setSam3AppendState({ ...sam3AppendState, iouThreshold: parseFloat(e.target.value) })}
|
||||
style={{ width: '100%', cursor: 'pointer' }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<span className="hint" style={{ fontSize: '0.82rem', display: 'block', marginBottom: 6 }}>Target Prompts to Detect:</span>
|
||||
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 6, maxHeight: 120, overflowY: 'auto', padding: 10, background: 'rgba(0,0,0,0.4)', borderRadius: 6, border: '1px solid rgba(255,255,255,0.1)', marginBottom: 10 }}>
|
||||
{sam3AppendState.selectedClasses.map((clsName) => (
|
||||
<button
|
||||
key={clsName}
|
||||
type="button"
|
||||
className="tag"
|
||||
style={{
|
||||
cursor: 'pointer', padding: '4px 9px', fontSize: '0.8rem',
|
||||
background: 'rgba(168, 85, 247, 0.25)', color: '#f3e8ff',
|
||||
border: '1px solid rgba(168, 85, 247, 0.6)'
|
||||
}}
|
||||
onClick={() => {
|
||||
const next = sam3AppendState.selectedClasses.filter((c) => c !== clsName)
|
||||
setSam3AppendState({ ...sam3AppendState, selectedClasses: next })
|
||||
}}
|
||||
>
|
||||
✕ {clsName}
|
||||
</button>
|
||||
))}
|
||||
</div>
|
||||
|
||||
{/* Add Custom SAM3 Prompt */}
|
||||
<div className="row" style={{ gap: 8 }}>
|
||||
<input
|
||||
type="text"
|
||||
placeholder="Type new SAM3 text prompt (e.g. box)..."
|
||||
value={customPromptInput}
|
||||
onChange={(e) => 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' }}
|
||||
/>
|
||||
<button
|
||||
type="button"
|
||||
className="btn"
|
||||
style={{ fontSize: '0.8rem' }}
|
||||
onClick={() => {
|
||||
if (customPromptInput.trim()) {
|
||||
const val = customPromptInput.trim().toLowerCase()
|
||||
if (!sam3AppendState.selectedClasses.includes(val)) {
|
||||
setSam3AppendState({
|
||||
...sam3AppendState,
|
||||
selectedClasses: [...sam3AppendState.selectedClasses, val]
|
||||
})
|
||||
}
|
||||
setCustomPromptInput('')
|
||||
}
|
||||
}}
|
||||
>
|
||||
+ Add Prompt
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="row" style={{ justifyContent: 'flex-end', gap: 10 }}>
|
||||
<button className="btn btn-ghost" onClick={() => setSam3AppendState(null)}>Cancel</button>
|
||||
<button
|
||||
className="btn btn-primary"
|
||||
disabled={sam3AppendState.selectedClasses.length === 0}
|
||||
onClick={async () => {
|
||||
const { batch, threshold, iouThreshold, selectedClasses } = sam3AppendState
|
||||
setSam3AppendState(null)
|
||||
setBusyId(batch.id)
|
||||
try {
|
||||
await api.startAutolabel(batch.id, {
|
||||
resume: false,
|
||||
append: true,
|
||||
engine: 'sam3',
|
||||
threshold,
|
||||
iou_threshold: iouThreshold ?? 0.8,
|
||||
engine_classes: { sam3: selectedClasses },
|
||||
class_ids: null
|
||||
})
|
||||
onChanged()
|
||||
} catch (err) {
|
||||
onError(err.message)
|
||||
} finally {
|
||||
setBusyId(null)
|
||||
}
|
||||
}}
|
||||
>
|
||||
Start SAM3 Append
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* YOLO Custom Model Append Modal */}
|
||||
{appendModalState && (
|
||||
<div style={{
|
||||
position: 'fixed', top: 0, left: 0, right: 0, bottom: 0,
|
||||
background: 'rgba(0,0,0,0.8)', display: 'flex', alignItems: 'center',
|
||||
justifyContent: 'center', zIndex: 9999
|
||||
}}>
|
||||
<div className="panel" style={{ width: 500, maxWidth: '92vw', padding: 22, border: '1px solid rgba(56, 189, 248, 0.4)', background: '#18181b' }}>
|
||||
<h3 style={{ margin: '0 0 8px 0', color: '#38bdf8', fontSize: '1.05rem' }}>Append Annotations with Custom Model</h3>
|
||||
<p className="hint" style={{ fontSize: '0.82rem', marginBottom: 16 }}>
|
||||
Model file: <strong style={{ color: '#fff' }}>{appendModalState.filename}</strong>
|
||||
</p>
|
||||
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<div className="row" style={{ justifyContent: 'space-between', marginBottom: 6 }}>
|
||||
<span className="hint" style={{ fontSize: '0.82rem' }}>Confidence Threshold:</span>
|
||||
<strong style={{ color: '#38bdf8', fontSize: '0.85rem' }}>{appendModalState.threshold}</strong>
|
||||
</div>
|
||||
<input
|
||||
type="range"
|
||||
min="0.05"
|
||||
max="0.95"
|
||||
step="0.05"
|
||||
value={appendModalState.threshold}
|
||||
onChange={(e) => setAppendModalState({ ...appendModalState, threshold: parseFloat(e.target.value) })}
|
||||
style={{ width: '100%', cursor: 'pointer' }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<div className="row" style={{ justifyContent: 'space-between', marginBottom: 6 }}>
|
||||
<span className="hint" style={{ fontSize: '0.82rem' }}>NMS IoU Threshold:</span>
|
||||
<strong style={{ color: '#c084fc', fontSize: '0.85rem' }}>{appendModalState.iouThreshold ?? 0.8}</strong>
|
||||
</div>
|
||||
<input
|
||||
type="range"
|
||||
min="0.1"
|
||||
max="0.9"
|
||||
step="0.05"
|
||||
value={appendModalState.iouThreshold ?? 0.8}
|
||||
onChange={(e) => setAppendModalState({ ...appendModalState, iouThreshold: parseFloat(e.target.value) })}
|
||||
style={{ width: '100%', cursor: 'pointer' }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div style={{ marginBottom: 20 }}>
|
||||
<div className="row" style={{ justifyContent: 'space-between', marginBottom: 6 }}>
|
||||
<span className="hint" style={{ fontSize: '0.82rem' }}>Select Classes to Append:</span>
|
||||
<span style={{ fontSize: '0.78rem', color: '#a1a1aa' }}>
|
||||
{appendModalState.selectedClasses.length} of {appendModalState.classes.length} selected
|
||||
</span>
|
||||
</div>
|
||||
|
||||
{appendModalState.classes.length === 0 ? (
|
||||
<p className="hint" style={{ fontStyle: 'italic', color: '#e4e4e7' }}>No embedded class names found in model file. All predictions will be appended.</p>
|
||||
) : (
|
||||
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 6, maxHeight: 150, overflowY: 'auto', padding: 10, background: 'rgba(0,0,0,0.4)', borderRadius: 6, border: '1px solid rgba(255,255,255,0.1)' }}>
|
||||
{appendModalState.classes.map((clsName) => {
|
||||
const isChecked = appendModalState.selectedClasses.includes(clsName)
|
||||
return (
|
||||
<button
|
||||
key={clsName}
|
||||
type="button"
|
||||
className="tag"
|
||||
style={{
|
||||
cursor: 'pointer',
|
||||
padding: '4px 9px',
|
||||
fontSize: '0.8rem',
|
||||
background: isChecked ? 'rgba(56, 189, 248, 0.25)' : 'rgba(255,255,255,0.05)',
|
||||
color: isChecked ? '#e0f2fe' : '#71717a',
|
||||
border: isChecked ? '1px solid rgba(56, 189, 248, 0.6)' : '1px solid rgba(255,255,255,0.1)',
|
||||
transition: 'all 150ms ease'
|
||||
}}
|
||||
onClick={() => {
|
||||
const next = isChecked
|
||||
? appendModalState.selectedClasses.filter((c) => c !== clsName)
|
||||
: [...appendModalState.selectedClasses, clsName]
|
||||
setAppendModalState({ ...appendModalState, selectedClasses: next })
|
||||
}}
|
||||
>
|
||||
{isChecked ? '✓ ' : ''}{clsName}
|
||||
</button>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="row" style={{ justifyContent: 'flex-end', gap: 10 }}>
|
||||
<button className="btn btn-ghost" onClick={() => setAppendModalState(null)}>Cancel</button>
|
||||
<button
|
||||
className="btn btn-primary"
|
||||
disabled={appendModalState.classes.length > 0 && appendModalState.selectedClasses.length === 0}
|
||||
onClick={async () => {
|
||||
const { batch, file, threshold, iouThreshold, selectedClasses } = appendModalState
|
||||
setAppendModalState(null)
|
||||
setBusyId(batch.id)
|
||||
try {
|
||||
await api.autolabelWithModel(batch.id, file, threshold, selectedClasses, iouThreshold ?? 0.8)
|
||||
onChanged()
|
||||
} catch (err) {
|
||||
onError(err.message)
|
||||
} finally {
|
||||
setBusyId(null)
|
||||
}
|
||||
}}
|
||||
>
|
||||
Start Append Auto-Annotation
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Base Model Auto-annotate Modal */}
|
||||
{baseModelModalState && (
|
||||
<div style={{
|
||||
position: 'fixed', top: 0, left: 0, right: 0, bottom: 0,
|
||||
background: 'rgba(0,0,0,0.8)', display: 'flex', alignItems: 'center',
|
||||
justifyContent: 'center', zIndex: 9999
|
||||
}}>
|
||||
<div className="panel" style={{ width: 480, maxWidth: '92vw', padding: 22, border: '1px solid rgba(56, 189, 248, 0.4)', background: '#18181b' }}>
|
||||
<h3 style={{ margin: '0 0 6px 0', color: '#38bdf8', fontSize: '1.05rem' }}>Auto-annotate Batch (Base Model)</h3>
|
||||
<p className="hint" style={{ fontSize: '0.82rem', marginBottom: 16 }}>
|
||||
Select target classes to detect using the project's base model:
|
||||
</p>
|
||||
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<div className="row" style={{ justifyContent: 'space-between', marginBottom: 6 }}>
|
||||
<span className="hint" style={{ fontSize: '0.82rem' }}>Confidence Threshold:</span>
|
||||
<strong style={{ color: '#38bdf8', fontSize: '0.85rem' }}>{baseModelModalState.threshold}</strong>
|
||||
</div>
|
||||
<input
|
||||
type="range"
|
||||
min="0.05"
|
||||
max="0.95"
|
||||
step="0.05"
|
||||
value={baseModelModalState.threshold}
|
||||
onChange={(e) => setBaseModelModalState({ ...baseModelModalState, threshold: parseFloat(e.target.value) })}
|
||||
style={{ width: '100%', cursor: 'pointer' }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<div className="row" style={{ justifyContent: 'space-between', marginBottom: 6 }}>
|
||||
<span className="hint" style={{ fontSize: '0.82rem' }}>NMS IoU Threshold:</span>
|
||||
<strong style={{ color: '#38bdf8', fontSize: '0.85rem' }}>{baseModelModalState.iouThreshold ?? 0.8}</strong>
|
||||
</div>
|
||||
<input
|
||||
type="range"
|
||||
min="0.1"
|
||||
max="0.9"
|
||||
step="0.05"
|
||||
value={baseModelModalState.iouThreshold ?? 0.8}
|
||||
onChange={(e) => setBaseModelModalState({ ...baseModelModalState, iouThreshold: parseFloat(e.target.value) })}
|
||||
style={{ width: '100%', cursor: 'pointer' }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div style={{ marginBottom: 18 }}>
|
||||
<span className="hint" style={{ fontSize: '0.82rem', display: 'block', marginBottom: 6 }}>Target Classes to Detect:</span>
|
||||
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 6, padding: 10, background: 'rgba(0,0,0,0.4)', borderRadius: 6, border: '1px solid rgba(255,255,255,0.1)' }}>
|
||||
{project?.classes?.map((cls) => {
|
||||
const isChecked = baseModelModalState.selectedClasses.includes(cls.name)
|
||||
return (
|
||||
<button
|
||||
key={cls.class_id}
|
||||
type="button"
|
||||
className="tag"
|
||||
style={{
|
||||
cursor: 'pointer',
|
||||
padding: '4px 9px',
|
||||
fontSize: '0.8rem',
|
||||
background: isChecked ? 'rgba(56, 189, 248, 0.25)' : 'rgba(255,255,255,0.05)',
|
||||
color: isChecked ? '#e0f2fe' : '#71717a',
|
||||
border: isChecked ? '1px solid rgba(56, 189, 248, 0.6)' : '1px solid rgba(255,255,255,0.1)',
|
||||
}}
|
||||
onClick={() => {
|
||||
const newSel = isChecked
|
||||
? baseModelModalState.selectedClasses.filter((n) => n !== cls.name)
|
||||
: [...baseModelModalState.selectedClasses, cls.name]
|
||||
setBaseModelModalState({ ...baseModelModalState, selectedClasses: newSel })
|
||||
}}
|
||||
>
|
||||
{isChecked ? '✓ ' : ''}{cls.name}
|
||||
</button>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="row" style={{ justifyContent: 'flex-end', gap: 10 }}>
|
||||
<button className="btn btn-ghost" onClick={() => setBaseModelModalState(null)}>Cancel</button>
|
||||
<button
|
||||
className="btn btn-primary"
|
||||
disabled={baseModelModalState.selectedClasses.length === 0}
|
||||
onClick={async () => {
|
||||
const { batch, threshold, iouThreshold, selectedClasses } = baseModelModalState
|
||||
setBaseModelModalState(null)
|
||||
setBusyId(batch.id)
|
||||
try {
|
||||
await api.startAutolabel(batch.id, {
|
||||
resume: false,
|
||||
append: false,
|
||||
engine: 'base_model',
|
||||
threshold,
|
||||
iou_threshold: iouThreshold ?? 0.8,
|
||||
target_class_names: selectedClasses,
|
||||
})
|
||||
onChanged()
|
||||
} catch (err) {
|
||||
onError(err.message)
|
||||
} finally {
|
||||
setBusyId(null)
|
||||
}
|
||||
}}
|
||||
>
|
||||
Start Auto-Annotation
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
@@ -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 <p className="error-banner"><AlertIcon size={14} /> {error}</p>
|
||||
if (!project || !summary) return <p className="empty">Loading…</p>
|
||||
|
||||
@@ -169,50 +178,20 @@ export default function ModelsPage({ projectId, onProject }) {
|
||||
<AlertIcon size={14} /> {error}
|
||||
</p>}
|
||||
|
||||
<div className="select-engine-section">
|
||||
<h2>Select Engine</h2>
|
||||
<p className="hint">
|
||||
Select how you want to train your model. Configure custom fine-tuning parameters or use SAM3 grounding backbone.
|
||||
</p>
|
||||
|
||||
<div className="engine-grid">
|
||||
<div className="engine-card selected">
|
||||
<div className="engine-card-header">
|
||||
<span className="project-badge">Selected</span>
|
||||
<span className="engine-card-title">Custom Training (YOLO11)</span>
|
||||
</div>
|
||||
<p className="engine-card-desc">
|
||||
Fine-tune on the merged master dataset using pre-configured hardware batch size and image resolution.
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div className="engine-card" title="Neural Architecture Search / Rapid Auto-train">
|
||||
<div className="engine-card-header">
|
||||
<span className="chip-badge info">Rapid NAS</span>
|
||||
<span className="engine-card-title">Neural Architecture Search</span>
|
||||
</div>
|
||||
<p className="engine-card-desc">
|
||||
Automated model selection optimized for latency and accuracy trade-offs on your specific project dataset.
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="panel side-panel" style={{ marginBottom: 20 }}>
|
||||
<h2>Uploaded Models for Training & Auto-Annotation</h2>
|
||||
<p className="hint">Upload up to 2 models to use for fine-tuning baseline or auto-annotating new video batches:</p>
|
||||
<h2>Base Model Configuration</h2>
|
||||
<p className="hint">Upload a base YOLO model checkpoint (.pt) to use for fine-tuning baseline and auto-annotation:</p>
|
||||
|
||||
<div style={{ display: 'grid', gridTemplateColumns: 'repeat(auto-fit, minmax(280px, 1fr))', gap: 16, marginTop: 12 }}>
|
||||
<div style={{ padding: 14, background: 'rgba(0,0,0,0.3)', borderRadius: 8, border: '1px solid rgba(255,255,255,0.1)' }}>
|
||||
<h3 style={{ fontSize: '0.9rem', color: '#c084fc', margin: '0 0 6px 0' }}>Primary Model 1 (Base Model)</h3>
|
||||
<div style={{ padding: 14, background: 'rgba(0,0,0,0.3)', borderRadius: 8, border: '1px solid rgba(255,255,255,0.1)', marginTop: 12 }}>
|
||||
<h3 style={{ fontSize: '0.9rem', color: '#c084fc', margin: '0 0 6px 0' }}>Base Model</h3>
|
||||
<p className="hint" style={{ fontSize: '0.78rem', margin: '0 0 10px 0' }}>
|
||||
Used as fine-tuning starting point and benchmark baseline.
|
||||
Used as fine-tuning starting point and baseline benchmark.
|
||||
</p>
|
||||
<div style={{ fontSize: '0.8rem', color: '#a1a1aa', marginBottom: 10 }}>
|
||||
Status: <strong style={{ color: '#fff' }}>{project.base_model_path ? 'Custom model.pt loaded' : 'Default yolo11n.pt'}</strong>
|
||||
</div>
|
||||
<div style={{ fontSize: '0.78rem', color: '#38bdf8', marginBottom: 10 }}>
|
||||
<strong>Model 1 Classes:</strong> {project.classes?.map((c) => c.name).join(', ')}
|
||||
<strong>Base Model Classes:</strong> {project.classes?.map((c) => c.name).join(', ')}
|
||||
</div>
|
||||
<input
|
||||
type="file"
|
||||
@@ -231,42 +210,9 @@ export default function ModelsPage({ projectId, onProject }) {
|
||||
}}
|
||||
/>
|
||||
<label htmlFor="upload-primary-model" className="btn btn-ghost" style={{ cursor: 'pointer', padding: '4px 10px', fontSize: '0.8rem' }}>
|
||||
Upload Primary Model 1 (.pt)
|
||||
Upload Base Model (.pt)
|
||||
</label>
|
||||
</div>
|
||||
|
||||
<div style={{ padding: 14, background: 'rgba(0,0,0,0.3)', borderRadius: 8, border: '1px solid rgba(255,255,255,0.1)' }}>
|
||||
<h3 style={{ fontSize: '0.9rem', color: '#38bdf8', margin: '0 0 6px 0' }}>Secondary Model 2 (Auto-Annotate Engine)</h3>
|
||||
<p className="hint" style={{ fontSize: '0.78rem', margin: '0 0 10px 0' }}>
|
||||
Used as an auxiliary engine choice for fast auto-annotation.
|
||||
</p>
|
||||
<div style={{ fontSize: '0.8rem', color: '#a1a1aa', marginBottom: 10 }}>
|
||||
Status: <strong style={{ color: '#fff' }}>{project.secondary_model_path ? (project.secondary_model_name || 'secondary_model.pt') : 'None uploaded'}</strong>
|
||||
</div>
|
||||
<div style={{ fontSize: '0.78rem', color: '#c084fc', marginBottom: 10 }}>
|
||||
<strong>Model 2 Native Classes:</strong> {project.secondary_model_classes?.length > 0 ? project.secondary_model_classes.join(', ') : 'Extracted on upload'}
|
||||
</div>
|
||||
<input
|
||||
type="file"
|
||||
accept=".pt"
|
||||
id="upload-secondary-model"
|
||||
style={{ display: 'none' }}
|
||||
onChange={async (e) => {
|
||||
const file = e.target.files?.[0]
|
||||
if (!file) return
|
||||
try {
|
||||
await api.uploadSecondaryModel(projectId, file)
|
||||
load()
|
||||
} catch (err) {
|
||||
setError(err.message)
|
||||
}
|
||||
}}
|
||||
/>
|
||||
<label htmlFor="upload-secondary-model" className="btn btn-ghost" style={{ cursor: 'pointer', padding: '4px 10px', fontSize: '0.8rem' }}>
|
||||
Upload Secondary Model 2 (.pt)
|
||||
</label>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="review">
|
||||
@@ -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).
|
||||
</p>
|
||||
|
||||
<div style={{ marginBottom: 12 }}>
|
||||
<label style={{ fontSize: '0.82rem', display: 'block', marginBottom: 6 }}>Target Classes to Train:</label>
|
||||
<div style={{ display: 'flex', flexWrap: 'wrap', gap: 6 }}>
|
||||
{project.classes?.map((cls) => {
|
||||
const isChecked = selectedClassIds.includes(cls.class_id)
|
||||
return (
|
||||
<button
|
||||
key={cls.class_id}
|
||||
type="button"
|
||||
className="tag"
|
||||
style={{
|
||||
cursor: 'pointer',
|
||||
padding: '4px 9px',
|
||||
fontSize: '0.8rem',
|
||||
background: isChecked ? 'rgba(192, 132, 252, 0.25)' : 'rgba(255,255,255,0.05)',
|
||||
color: isChecked ? '#f3e8ff' : '#71717a',
|
||||
border: isChecked ? '1px solid rgba(192, 132, 252, 0.6)' : '1px solid rgba(255,255,255,0.1)',
|
||||
}}
|
||||
onClick={() => toggleClassSelect(cls.class_id)}
|
||||
>
|
||||
{isChecked ? '✓ ' : ''}{cls.name}
|
||||
</button>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div>
|
||||
<label htmlFor="epochs">Epochs</label>
|
||||
<input id="epochs" type="number" min="1" max="500" value={epochs}
|
||||
@@ -303,7 +277,7 @@ export default function ModelsPage({ projectId, onProject }) {
|
||||
</p>
|
||||
)}
|
||||
<button className="btn btn-primary" onClick={train}
|
||||
disabled={running || summary.splits.train === 0 || selectedBatchIds.length === 0}>
|
||||
disabled={running || summary.splits.train === 0 || selectedBatchIds.length === 0 || selectedClassIds.length === 0}>
|
||||
{running ? 'Training…' : 'Start training'}
|
||||
</button>
|
||||
{summary.splits.train === 0 && (
|
||||
@@ -312,6 +286,9 @@ export default function ModelsPage({ projectId, onProject }) {
|
||||
{selectedBatchIds.length === 0 && summary.splits.train > 0 && (
|
||||
<p className="hint" style={{ color: '#ef4444' }}>Select at least one dataset batch to train.</p>
|
||||
)}
|
||||
{selectedClassIds.length === 0 && (
|
||||
<p className="hint" style={{ color: '#ef4444' }}>Select at least one class to train.</p>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{job && (
|
||||
|
||||
@@ -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 (
|
||||
<div className="roboflow-page" style={{ padding: 24, maxWidth: 1400, margin: '0 auto' }}>
|
||||
<div className="page-header" style={{ marginBottom: 20 }}>
|
||||
<div>
|
||||
<h1 style={{ margin: '0 0 6px 0', fontSize: '1.4rem', color: '#c084fc', display: 'flex', alignItems: 'center', gap: 8 }}>
|
||||
🤖 SAM3 Global Playground
|
||||
</h1>
|
||||
<p className="hint" style={{ margin: 0, fontSize: '0.85rem' }}>
|
||||
Test SAM3 zero-shot text grounding prompts on any image without affecting projects or database records.
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{error && (
|
||||
<div className="alert error" style={{ marginBottom: 16, padding: '10px 14px', borderRadius: 6, background: 'rgba(239, 68, 68, 0.15)', border: '1px solid rgba(239, 68, 68, 0.4)', color: '#fca5a5', fontSize: '0.85rem' }}>
|
||||
⚠️ {error}
|
||||
</div>
|
||||
)}
|
||||
|
||||
<div style={{ display: 'grid', gridTemplateColumns: '360px 1fr', gap: 20, alignItems: 'start' }}>
|
||||
{/* Left Control Panel */}
|
||||
<div className="panel" style={{ padding: 18, background: '#18181b', borderRadius: 8, border: '1px solid rgba(255,255,255,0.1)' }}>
|
||||
<h3 style={{ margin: '0 0 14px 0', fontSize: '0.95rem', color: '#fff' }}>Configuration</h3>
|
||||
|
||||
{/* Image Selection / Dropzone */}
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<label className="hint" style={{ fontSize: '0.82rem', display: 'block', marginBottom: 6 }}>
|
||||
Select Test Image:
|
||||
</label>
|
||||
<div
|
||||
onDragOver={(e) => 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'
|
||||
}}
|
||||
>
|
||||
<input
|
||||
ref={fileInputRef}
|
||||
type="file"
|
||||
accept="image/*"
|
||||
style={{ display: 'none' }}
|
||||
onChange={(e) => handleFileSelect(e.target.files?.[0])}
|
||||
/>
|
||||
{imageFile ? (
|
||||
<div style={{ fontSize: '0.82rem', color: '#38bdf8', wordBreak: 'break-all' }}>
|
||||
🖼️ {imageFile.name} ({(imageFile.size / 1024).toFixed(0)} KB)
|
||||
</div>
|
||||
) : (
|
||||
<div style={{ fontSize: '0.82rem', color: '#a1a1aa' }}>
|
||||
Click to choose image or drag & drop here (.jpg, .png)
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Text Prompts Input */}
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<label className="hint" style={{ fontSize: '0.82rem', display: 'block', marginBottom: 6 }}>
|
||||
Descriptive Text Prompts (Comma Separated):
|
||||
</label>
|
||||
<textarea
|
||||
rows={3}
|
||||
value={promptsText}
|
||||
onChange={(e) => setPromptsText(e.target.value)}
|
||||
placeholder="e.g. white plastic bag, forklift, pallet"
|
||||
style={{
|
||||
width: '100%',
|
||||
padding: '8px 10px',
|
||||
fontSize: '0.83rem',
|
||||
background: '#09090b',
|
||||
border: '1px solid rgba(255,255,255,0.15)',
|
||||
borderRadius: 6,
|
||||
color: '#fff',
|
||||
resize: 'vertical'
|
||||
}}
|
||||
/>
|
||||
<span className="hint" style={{ fontSize: '0.75rem', color: '#71717a' }}>
|
||||
Tip: Use specific descriptive phrases for custom objects (e.g. "woven white bag" instead of "sack").
|
||||
</span>
|
||||
</div>
|
||||
|
||||
{/* Threshold Sliders */}
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<div className="row" style={{ justifyContent: 'space-between', marginBottom: 4 }}>
|
||||
<span className="hint" style={{ fontSize: '0.82rem' }}>Confidence Threshold:</span>
|
||||
<strong style={{ color: '#c084fc', fontSize: '0.85rem' }}>{threshold}</strong>
|
||||
</div>
|
||||
<input
|
||||
type="range"
|
||||
min="0.05"
|
||||
max="0.95"
|
||||
step="0.05"
|
||||
value={threshold}
|
||||
onChange={(e) => setThreshold(parseFloat(e.target.value))}
|
||||
style={{ width: '100%', cursor: 'pointer' }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div style={{ marginBottom: 20 }}>
|
||||
<div className="row" style={{ justifyContent: 'space-between', marginBottom: 4 }}>
|
||||
<span className="hint" style={{ fontSize: '0.82rem' }}>NMS IoU Threshold:</span>
|
||||
<strong style={{ color: '#38bdf8', fontSize: '0.85rem' }}>{iouThreshold}</strong>
|
||||
</div>
|
||||
<input
|
||||
type="range"
|
||||
min="0.1"
|
||||
max="0.9"
|
||||
step="0.05"
|
||||
value={iouThreshold}
|
||||
onChange={(e) => setIouThreshold(parseFloat(e.target.value))}
|
||||
style={{ width: '100%', cursor: 'pointer' }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* Run Button */}
|
||||
<button
|
||||
className="btn btn-primary"
|
||||
disabled={loading || !imageFile}
|
||||
onClick={runTest}
|
||||
style={{
|
||||
width: '100%',
|
||||
padding: '10px 14px',
|
||||
fontSize: '0.9rem',
|
||||
fontWeight: 600,
|
||||
background: '#c084fc',
|
||||
borderColor: '#c084fc',
|
||||
color: '#000',
|
||||
cursor: loading || !imageFile ? 'not-allowed' : 'pointer'
|
||||
}}
|
||||
>
|
||||
{loading ? 'Running SAM3 Engine…' : '🚀 Run SAM3 Test'}
|
||||
</button>
|
||||
|
||||
{/* Detection Stats List */}
|
||||
{result && (
|
||||
<div style={{ marginTop: 20, paddingTop: 16, borderTop: '1px solid rgba(255,255,255,0.1)' }}>
|
||||
<h4 style={{ margin: '0 0 10px 0', fontSize: '0.88rem', color: '#34d399' }}>
|
||||
Results: {result.detections?.length || 0} Object(s) Found
|
||||
</h4>
|
||||
<div style={{ display: 'flex', flexDirection: 'column', gap: 6, maxHeight: 220, overflowY: 'auto' }}>
|
||||
{result.detections.map((det, i) => {
|
||||
const color = getColor(det.class_id)
|
||||
return (
|
||||
<div
|
||||
key={i}
|
||||
style={{
|
||||
padding: '6px 10px',
|
||||
background: 'rgba(0,0,0,0.4)',
|
||||
borderRadius: 6,
|
||||
borderLeft: `4px solid ${color}`,
|
||||
fontSize: '0.78rem',
|
||||
display: 'flex',
|
||||
justifyContent: 'space-between'
|
||||
}}
|
||||
>
|
||||
<strong style={{ color }}>{det.prompt}</strong>
|
||||
<span style={{ color: '#a1a1aa' }}>{(det.score * 100).toFixed(1)}%</span>
|
||||
</div>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* Right Canvas Viewer */}
|
||||
<div className="panel" style={{ padding: 18, background: '#18181b', borderRadius: 8, border: '1px solid rgba(255,255,255,0.1)', minHeight: 500 }}>
|
||||
{!previewUrl ? (
|
||||
<div style={{ display: 'flex', flexDirection: 'column', alignItems: 'center', justifyContent: 'center', minHeight: 450, color: '#71717a' }}>
|
||||
<div style={{ fontSize: '2.5rem', marginBottom: 10 }}>🖼️</div>
|
||||
<p style={{ margin: 0, fontSize: '0.9rem' }}>Upload an image on the left to start testing SAM3 text prompts.</p>
|
||||
</div>
|
||||
) : (
|
||||
<div style={{ position: 'relative', width: '100%', overflow: 'hidden', display: 'flex', justifyContent: 'center', background: '#09090b', borderRadius: 6 }}>
|
||||
<img
|
||||
src={previewUrl}
|
||||
alt="SAM3 Test Target"
|
||||
style={{ maxWidth: '100%', maxHeight: '75vh', objectFit: 'contain', display: 'block' }}
|
||||
/>
|
||||
|
||||
{/* Overlay Canvas SVG */}
|
||||
{result && result.width && result.height && (
|
||||
<svg
|
||||
viewBox={`0 0 ${result.width} ${result.height}`}
|
||||
style={{
|
||||
position: 'absolute',
|
||||
top: 0, left: 0, width: '100%', height: '100%',
|
||||
pointerEvents: 'none'
|
||||
}}
|
||||
>
|
||||
{result.detections.map((det, i) => {
|
||||
const color = getColor(det.class_id)
|
||||
const [x0, y0, x1, y1] = [
|
||||
det.box[0] * result.width,
|
||||
det.box[1] * result.height,
|
||||
det.box[2] * result.width,
|
||||
det.box[3] * result.height
|
||||
]
|
||||
const bw = x1 - x0
|
||||
const bh = y1 - y0
|
||||
|
||||
return (
|
||||
<g key={i}>
|
||||
{/* Polygon Masks */}
|
||||
{det.polygons?.map((poly, pIdx) => {
|
||||
const ptsString = poly.map(([px, py]) => `${px * result.width},${py * result.height}`).join(' ')
|
||||
return (
|
||||
<polygon
|
||||
key={pIdx}
|
||||
points={ptsString}
|
||||
fill={color}
|
||||
fillOpacity={0.35}
|
||||
stroke={color}
|
||||
strokeWidth={2}
|
||||
/>
|
||||
)
|
||||
})}
|
||||
|
||||
{/* Bounding Box */}
|
||||
<rect
|
||||
x={x0}
|
||||
y={y0}
|
||||
width={bw}
|
||||
height={bh}
|
||||
fill="none"
|
||||
stroke={color}
|
||||
strokeWidth={2}
|
||||
strokeDasharray="4 2"
|
||||
/>
|
||||
|
||||
{/* Label Badge */}
|
||||
<rect
|
||||
x={x0}
|
||||
y={Math.max(0, y0 - 22)}
|
||||
width={Math.max(80, det.prompt.length * 8 + 45)}
|
||||
height={20}
|
||||
fill={color}
|
||||
rx={3}
|
||||
/>
|
||||
<text
|
||||
x={x0 + 6}
|
||||
y={Math.max(14, y0 - 8)}
|
||||
fill="#000"
|
||||
fontSize={12}
|
||||
fontWeight="bold"
|
||||
fontFamily="sans-serif"
|
||||
>
|
||||
{det.prompt} ({Math.round(det.score * 100)}%)
|
||||
</text>
|
||||
</g>
|
||||
)
|
||||
})}
|
||||
</svg>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
@@ -7,17 +7,18 @@ cd "$REPO_DIR"
|
||||
echo "🔄 Restarting app services..."
|
||||
|
||||
# Kill running uvicorn and vite processes
|
||||
fuser -k -9 8000/tcp 8080/tcp 5173/tcp || true
|
||||
pkill -f "uvicorn backend.main:app" || true
|
||||
pkill -f "vite" || true
|
||||
sleep 1
|
||||
sleep 2.5
|
||||
|
||||
export PATH="$HOME/.local/bin:$PATH"
|
||||
|
||||
# Run backend
|
||||
nohup uv run uvicorn backend.main:app --host 0.0.0.0 --port 8000 > "$REPO_DIR/data/backend.log" 2>&1 &
|
||||
setsid uv run uvicorn backend.main:app --host 0.0.0.0 --port 8000 > "$REPO_DIR/data/backend.log" 2>&1 &
|
||||
|
||||
# Run frontend
|
||||
cd "$REPO_DIR/frontend"
|
||||
nohup npm run dev > "$REPO_DIR/data/frontend.log" 2>&1 &
|
||||
setsid npm run dev -- --host 0.0.0.0 --port 8080 > "$REPO_DIR/data/frontend.log" 2>&1 &
|
||||
|
||||
echo "✅ Services restarted."
|
||||
Reference in new issue
Block a user