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 os
|
||||||
|
import shutil
|
||||||
|
import tempfile
|
||||||
from typing import Optional
|
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 fastapi.responses import FileResponse
|
||||||
from pydantic import BaseModel
|
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 batches as batch_store
|
||||||
from backend import review as review_store
|
from backend import review as review_store
|
||||||
from backend.api.common import project_or_404, thumbnail
|
from backend.api.common import project_or_404, thumbnail
|
||||||
@@ -27,10 +28,12 @@ class AutolabelRequest(BaseModel):
|
|||||||
engines: Optional[list[str]] = None
|
engines: Optional[list[str]] = None
|
||||||
class_ids: Optional[list[int]] = None
|
class_ids: Optional[list[int]] = None
|
||||||
engine_classes: Optional[dict[str, list[str]]] = None
|
engine_classes: Optional[dict[str, list[str]]] = None
|
||||||
|
target_class_names: Optional[list[str]] = None
|
||||||
threshold: float = autolabel.DEFAULT_THRESHOLD
|
threshold: float = autolabel.DEFAULT_THRESHOLD
|
||||||
iou_threshold: float = autolabel.DEFAULT_IOU
|
iou_threshold: float = autolabel.DEFAULT_IOU
|
||||||
min_box_frac: float = 0.0
|
min_box_frac: float = 0.0
|
||||||
resume: bool = False
|
resume: bool = False
|
||||||
|
append: bool = False
|
||||||
|
|
||||||
|
|
||||||
@router.post("/api/projects/{project_id}/batches")
|
@router.post("/api/projects/{project_id}/batches")
|
||||||
@@ -90,11 +93,115 @@ def start_autolabel(batch_id: int, request: AutolabelRequest) -> dict:
|
|||||||
try:
|
try:
|
||||||
engine_list = request.engines if (request.engines and len(request.engines) > 0) else [request.engine]
|
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,
|
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,
|
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:
|
except batch_store.BatchError as exc:
|
||||||
raise HTTPException(400, str(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)
|
deleted = review_store.clear_batch_class_annotations(batch_id, class_id)
|
||||||
return {"deleted": deleted}
|
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
|
imgsz: Optional[int] = None
|
||||||
device: Optional[Union[int, str]] = None
|
device: Optional[Union[int, str]] = None
|
||||||
batch_ids: Optional[list] = None
|
batch_ids: Optional[list] = None
|
||||||
|
class_ids: Optional[list] = None
|
||||||
|
|
||||||
|
|
||||||
@router.get("/api/hardware")
|
@router.get("/api/hardware")
|
||||||
@@ -34,6 +35,7 @@ def start_training(project_id: int, request: TrainRequest) -> dict:
|
|||||||
project_id, request.epochs,
|
project_id, request.epochs,
|
||||||
{"batch": request.batch, "imgsz": request.imgsz, "device": request.device},
|
{"batch": request.batch, "imgsz": request.imgsz, "device": request.device},
|
||||||
batch_ids=request.batch_ids,
|
batch_ids=request.batch_ids,
|
||||||
|
class_ids=request.class_ids,
|
||||||
)
|
)
|
||||||
except training.TrainingError as exc:
|
except training.TrainingError as exc:
|
||||||
raise HTTPException(400, str(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,
|
def start(batch_id: int, threshold: float = DEFAULT_THRESHOLD,
|
||||||
iou_threshold: float = DEFAULT_IOU, min_box_frac: float = 0.0,
|
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,
|
engines: Optional[List[str]] = None,
|
||||||
class_ids: Optional[List[int]] = 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)
|
batch = batches.get(batch_id)
|
||||||
if batch is None:
|
if batch is None:
|
||||||
raise batches.BatchError("No such batch")
|
raise batches.BatchError("No such batch")
|
||||||
@@ -34,8 +36,9 @@ def start(batch_id: int, threshold: float = DEFAULT_THRESHOLD,
|
|||||||
"autolabel",
|
"autolabel",
|
||||||
params={"batch_id": batch_id, "threshold": threshold,
|
params={"batch_id": batch_id, "threshold": threshold,
|
||||||
"iou_threshold": iou_threshold, "min_box_frac": min_box_frac,
|
"iou_threshold": iou_threshold, "min_box_frac": min_box_frac,
|
||||||
"resume": resume, "engine": active_engines[0], "engines": active_engines,
|
"resume": resume, "append": append, "engine": active_engines[0], "engines": active_engines,
|
||||||
"class_ids": class_ids, "engine_classes": engine_classes},
|
"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"],
|
project_id=batch["project_id"],
|
||||||
batch_id=batch_id,
|
batch_id=batch_id,
|
||||||
message=f"{batch['date_label']}/{batch['batch_label']} ({'+'.join(e.upper() for e in active_engines)})",
|
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")
|
raise batches.BatchError("The batch disappeared before labeling started")
|
||||||
project = projects.get(batch["project_id"])
|
project = projects.get(batch["project_id"])
|
||||||
|
|
||||||
raw_active = job.params.get("engines") or [job.params.get("engine", "sam3")]
|
selected_engine = job.params.get("engine", "base_model")
|
||||||
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))
|
|
||||||
|
|
||||||
frames = batches.frames(batch["id"])
|
frames = batches.frames(batch["id"])
|
||||||
batches.set_status(batch["id"], "labeling")
|
batches.set_status(batch["id"], "labeling")
|
||||||
@@ -86,42 +80,34 @@ def _run_autolabel(job) -> None:
|
|||||||
attempted = 0
|
attempted = 0
|
||||||
failures = []
|
failures = []
|
||||||
|
|
||||||
from ultralytics import YOLO
|
conf = job.params.get("threshold", DEFAULT_THRESHOLD)
|
||||||
m1_path = projects.training_start_point(project)
|
iou_thresh = job.params.get("iou_threshold", DEFAULT_IOU)
|
||||||
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 {}
|
|
||||||
|
|
||||||
|
yolo_model = None
|
||||||
sam3_target_classes = []
|
sam3_target_classes = []
|
||||||
if "sam3" in expanded_engines:
|
|
||||||
sam3_classes = engine_classes.get("sam3")
|
custom_path = job.params.get("custom_model_path")
|
||||||
if sam3_classes is not None:
|
target_class_names = job.params.get("target_class_names")
|
||||||
allowed_set = {c.strip().lower() for c in sam3_classes}
|
|
||||||
sam3_target_classes = [
|
if selected_engine == "sam3" and not custom_path:
|
||||||
c for c in project["classes"]
|
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
|
||||||
if c["name"].strip().lower() in allowed_set or c["prompt"].strip().lower() in allowed_set
|
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:
|
else:
|
||||||
sam3_target_classes = [
|
sam3_target_classes = [c for c in project["classes"]]
|
||||||
c for c in project["classes"]
|
|
||||||
if (allowed_class_ids is None or c["class_id"] in allowed_class_ids)
|
|
||||||
]
|
|
||||||
prompts = [c["prompt"] for c in sam3_target_classes]
|
prompts = [c["prompt"] for c in sam3_target_classes]
|
||||||
if prompts:
|
if prompts:
|
||||||
from backend.sam3_engine import engine_is_loaded, get_engine
|
from backend.sam3_engine import engine_is_loaded, get_engine
|
||||||
@@ -130,13 +116,27 @@ def _run_autolabel(job) -> None:
|
|||||||
engine = get_engine()
|
engine = get_engine()
|
||||||
job.log(f"SAM3 ready on {engine.device}; prompts: {', '.join(prompts)}")
|
job.log(f"SAM3 ready on {engine.device}; prompts: {', '.join(prompts)}")
|
||||||
else:
|
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"]}
|
name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
|
||||||
conf = job.params.get("threshold", DEFAULT_THRESHOLD)
|
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
|
||||||
iou_thresh = job.params.get("iou_threshold", DEFAULT_IOU)
|
|
||||||
|
|
||||||
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):
|
for index, frame in enumerate(frames):
|
||||||
if job.cancelled:
|
if job.cancelled:
|
||||||
@@ -151,29 +151,21 @@ def _run_autolabel(job) -> None:
|
|||||||
frame_file = os.path.join(directory, frame["filename"])
|
frame_file = os.path.join(directory, frame["filename"])
|
||||||
all_raw_detections = []
|
all_raw_detections = []
|
||||||
|
|
||||||
for eng_key, y_model in yolo_models.items():
|
if yolo_model is not None:
|
||||||
allowed_for_eng = engine_classes.get(eng_key)
|
results = yolo_model.predict(frame_file, conf=conf, verbose=False)
|
||||||
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 results and len(results) > 0:
|
if results and len(results) > 0:
|
||||||
model_names = results[0].names
|
model_names = results[0].names
|
||||||
for box in results[0].boxes:
|
for box in results[0].boxes:
|
||||||
cls_idx = int(box.cls[0].item())
|
cls_idx = int(box.cls[0].item())
|
||||||
cls_name = str(model_names.get(cls_idx, cls_idx)).strip().lower()
|
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
|
continue
|
||||||
if cls_name not in name_to_class_id:
|
|
||||||
try:
|
target_class_id = name_to_class_id.get(cls_name)
|
||||||
updated_proj = projects.add_class(project["id"], {"name": cls_name, "prompt": cls_name})
|
if target_class_id is None:
|
||||||
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:
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
score = float(box.conf[0].item())
|
score = float(box.conf[0].item())
|
||||||
xyxyn = box.xyxyn[0].tolist()
|
xyxyn = box.xyxyn[0].tolist()
|
||||||
all_raw_detections.append(labeling.Detection(
|
all_raw_detections.append(labeling.Detection(
|
||||||
@@ -184,7 +176,7 @@ def _run_autolabel(job) -> None:
|
|||||||
mask=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]
|
prompts = [c["prompt"] for c in sam3_target_classes]
|
||||||
res = labeling.label_image(
|
res = labeling.label_image(
|
||||||
frame_file, frame["filename"], prompts, conf,
|
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"]):
|
for geometry in _geometries(det, frame["width"], frame["height"], project["label_type"]):
|
||||||
items.append({"class_id": det.class_id, "geometry": geometry, "score": det.score})
|
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)
|
review.replace_auto(frame["id"], items)
|
||||||
written += len(items)
|
written += len(items)
|
||||||
job.progress(index + 1, len(frames), f"{frame['filename']}: {len(items)} shape(s)")
|
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 os
|
||||||
import shutil
|
import shutil
|
||||||
import time
|
import time
|
||||||
|
from typing import List, Optional
|
||||||
|
|
||||||
from backend import batches, config, db, jobs, projects, review
|
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"
|
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)."""
|
"""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"])
|
root = dataset_dir(project["slug"])
|
||||||
os.makedirs(root, exist_ok=True)
|
os.makedirs(root, exist_ok=True)
|
||||||
counts = summary(project["id"])["splits"]
|
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:
|
if batch_ids:
|
||||||
with db.cursor() as cur:
|
with db.cursor() as cur:
|
||||||
|
|||||||
+3
-3
@@ -44,9 +44,9 @@ def defaults(epochs: int = 50) -> dict:
|
|||||||
if info["device"] == "cpu":
|
if info["device"] == "cpu":
|
||||||
settings = {"batch": 4, "imgsz": 512, "device": "cpu", "workers": 2}
|
settings = {"batch": 4, "imgsz": 512, "device": "cpu", "workers": 2}
|
||||||
note = "No GPU visible — training on CPU will be very slow."
|
note = "No GPU visible — training on CPU will be very slow."
|
||||||
elif vram < 8:
|
elif vram < 6:
|
||||||
settings = {"batch": 8, "imgsz": 640, "device": 0, "workers": 2}
|
settings = {"batch": 16, "imgsz": 640, "device": 0, "workers": 4}
|
||||||
note = f"{vram} GB of VRAM: small batches, 640 px."
|
note = f"{vram} GB of VRAM: batch 16, 640 px."
|
||||||
elif vram <= 16:
|
elif vram <= 16:
|
||||||
settings = {"batch": 32, "imgsz": 640, "device": 0, "workers": 8}
|
settings = {"batch": 32, "imgsz": 640, "device": 0, "workers": 8}
|
||||||
note = f"{vram} GB of VRAM: optimized batch 32, 640 px."
|
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]:
|
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] = []
|
kept: List[Detection] = []
|
||||||
for det in sorted(detections, key=lambda d: d.score, reverse=True):
|
for cls_dets in by_class.values():
|
||||||
if all(_iou(det.box, k.box) < iou_threshold for k in kept):
|
cls_kept: List[Detection] = []
|
||||||
kept.append(det)
|
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
|
return kept
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -384,6 +384,11 @@ def _kept_prompts(project: dict, names: List[str]) -> List[str]:
|
|||||||
def training_start_point(project: dict) -> str:
|
def training_start_point(project: dict) -> str:
|
||||||
"""The weights a training run should start from (REQ-060, REQ-004)."""
|
"""The weights a training run should start from (REQ-060, REQ-004)."""
|
||||||
path = project["base_model_path"]
|
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):
|
if path and os.path.isfile(path):
|
||||||
return path
|
return path
|
||||||
return PRETRAINED[project["label_type"]]
|
return PRETRAINED[project["label_type"]]
|
||||||
|
|||||||
@@ -234,6 +234,38 @@ def replace_auto(frame_id: int, items: List[dict]) -> int:
|
|||||||
return len(items)
|
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:
|
def frames_with_auto(batch_id: int) -> set:
|
||||||
"""Frame ids that already carry automatic shapes — the resume skip-list
|
"""Frame ids that already carry automatic shapes — the resume skip-list
|
||||||
for REQ-035."""
|
for REQ-035."""
|
||||||
@@ -339,6 +371,24 @@ def clear_batch_class_annotations(batch_id: int, class_id: int) -> int:
|
|||||||
return cur.rowcount
|
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:
|
def _check_class(project_id: int, class_id: int) -> None:
|
||||||
with db.cursor() as cur:
|
with db.cursor() as cur:
|
||||||
cur.execute(
|
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")
|
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)
|
project = projects.get(project_id)
|
||||||
if project is None:
|
if project is None:
|
||||||
raise TrainingError("No such project")
|
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)
|
settings = hardware.resolve(overrides, epochs)
|
||||||
job = jobs.create(
|
job = jobs.create(
|
||||||
"train",
|
"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,
|
project_id=project_id,
|
||||||
message=f"{counts['train']} train / {counts['val']} val",
|
message=f"{counts['train']} train / {counts['val']} val",
|
||||||
)
|
)
|
||||||
@@ -102,7 +102,8 @@ def _run_train(job) -> None:
|
|||||||
project = projects.get(job.params["project_id"])
|
project = projects.get(job.params["project_id"])
|
||||||
settings = job.params["settings"]
|
settings = job.params["settings"]
|
||||||
batch_ids = job.params.get("batch_ids")
|
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).
|
# SAM3 and a training run must not hold VRAM at the same time (REQ-065).
|
||||||
from backend.sam3_engine import release_engine
|
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)
|
model.add_callback("on_fit_epoch_end", on_epoch)
|
||||||
job.progress(0, settings["epochs"])
|
job.progress(0, settings["epochs"])
|
||||||
|
|
||||||
|
import torch
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.backends.cudnn.benchmark = True
|
||||||
|
|
||||||
keep_run_dir = False
|
keep_run_dir = False
|
||||||
try:
|
try:
|
||||||
model.train(
|
model.train(
|
||||||
@@ -142,7 +147,7 @@ def _run_train(job) -> None:
|
|||||||
batch=settings["batch"],
|
batch=settings["batch"],
|
||||||
device=settings["device"],
|
device=settings["device"],
|
||||||
workers=settings.get("workers", 8),
|
workers=settings.get("workers", 8),
|
||||||
cache=False,
|
cache="ram",
|
||||||
project=os.path.join(out_dir, "runs"),
|
project=os.path.join(out_dir, "runs"),
|
||||||
name="train",
|
name="train",
|
||||||
exist_ok=True,
|
exist_ok=True,
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import LibraryPage from './pages/LibraryPage'
|
|||||||
import TrimPage from './pages/TrimPage'
|
import TrimPage from './pages/TrimPage'
|
||||||
import ReviewPage from './pages/ReviewPage'
|
import ReviewPage from './pages/ReviewPage'
|
||||||
import ModelsPage from './pages/ModelsPage'
|
import ModelsPage from './pages/ModelsPage'
|
||||||
|
import Sam3PlaygroundPage from './pages/Sam3PlaygroundPage'
|
||||||
import './roboflow.css'
|
import './roboflow.css'
|
||||||
|
|
||||||
function parseRoute(hash) {
|
function parseRoute(hash) {
|
||||||
@@ -14,6 +15,10 @@ function parseRoute(hash) {
|
|||||||
const query = new URLSearchParams(queryString || '')
|
const query = new URLSearchParams(queryString || '')
|
||||||
const parts = path.split('/').filter(Boolean)
|
const parts = path.split('/').filter(Boolean)
|
||||||
|
|
||||||
|
if (parts[0] === 'sam3-playground') {
|
||||||
|
return { name: 'sam3-playground' }
|
||||||
|
}
|
||||||
|
|
||||||
if (parts[0] === 'batches' && parts[1]) {
|
if (parts[0] === 'batches' && parts[1]) {
|
||||||
return { name: 'review', batchId: Number(parts[1]) }
|
return { name: 'review', batchId: Number(parts[1]) }
|
||||||
}
|
}
|
||||||
@@ -135,6 +140,7 @@ export default function App() {
|
|||||||
onProject={(p) => setCurrentProject(p)}
|
onProject={(p) => setCurrentProject(p)}
|
||||||
/>
|
/>
|
||||||
)}
|
)}
|
||||||
|
{route.name === 'sam3-playground' && <Sam3PlaygroundPage />}
|
||||||
</ErrorBoundary>
|
</ErrorBoundary>
|
||||||
</main>
|
</main>
|
||||||
</div>
|
</div>
|
||||||
|
|||||||
@@ -60,8 +60,32 @@ export const api = {
|
|||||||
|
|
||||||
startAutolabel: (batchId, body) =>
|
startAutolabel: (batchId, body) =>
|
||||||
request(`/batches/${batchId}/autolabel`, { method: 'POST', body: 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) =>
|
clearBatchClassAnnotations: (batchId, classId) =>
|
||||||
request(`/batches/${batchId}/classes/${classId}/annotations`, { method: 'DELETE' }),
|
request(`/batches/${batchId}/classes/${classId}/annotations`, { method: 'DELETE' }),
|
||||||
|
resetBatchAutoAnnotations: (batchId) =>
|
||||||
|
request(`/batches/${batchId}/reset-auto-annotations`, { method: 'POST' }),
|
||||||
nextPending: (batchId, afterIdx = -1) =>
|
nextPending: (batchId, afterIdx = -1) =>
|
||||||
request(`/batches/${batchId}/next-pending?after_idx=${afterIdx}`),
|
request(`/batches/${batchId}/next-pending?after_idx=${afterIdx}`),
|
||||||
|
|
||||||
|
|||||||
@@ -74,11 +74,15 @@ export default function Sidebar({ route, currentProject, theme, onToggleTheme })
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
<div className="sidebar-section">
|
<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">
|
<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>
|
<span className="sidebar-icon"><RocketIcon size={16} /></span>
|
||||||
{!collapsed && <span>Train & Select Engine</span>}
|
{!collapsed && <span>Train & Select Engine</span>}
|
||||||
</a>
|
</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>
|
||||||
|
|
||||||
<div className="sidebar-spacer" />
|
<div className="sidebar-spacer" />
|
||||||
|
|||||||
+491
-271
@@ -43,40 +43,27 @@ function ActiveJobsBanner({ jobs, onCancel }) {
|
|||||||
|
|
||||||
function BatchList({ project, batches, activeJobs, onChanged, onError }) {
|
function BatchList({ project, batches, activeJobs, onChanged, onError }) {
|
||||||
const [busyId, setBusyId] = useState(null)
|
const [busyId, setBusyId] = useState(null)
|
||||||
const [selectedEngine, setSelectedEngine] = useState({})
|
const [appendChoiceBatch, setAppendChoiceBatch] = useState(null)
|
||||||
const [engineClassMap, setEngineClassMap] = useState({})
|
const [appendModalState, setAppendModalState] = useState(null)
|
||||||
const [expandedFilterBatchId, setExpandedFilterBatchId] = useState(null)
|
const [sam3AppendState, setSam3AppendState] = useState(null)
|
||||||
|
const [customPromptInput, setCustomPromptInput] = useState('')
|
||||||
|
const [baseModelModalState, setBaseModelModalState] = useState(null)
|
||||||
|
|
||||||
const hasSecondaryModel = Boolean(project?.secondary_model_path)
|
function openBaseModelAutolabelModal(batch) {
|
||||||
const model1Classes = project?.classes?.map((c) => c.name) || []
|
const projectClasses = project?.classes?.map((c) => c.name) || []
|
||||||
const model2Classes = project?.secondary_model_classes?.length > 0
|
setBaseModelModalState({
|
||||||
? project.secondary_model_classes
|
batch,
|
||||||
: model1Classes
|
selectedClasses: [...projectClasses],
|
||||||
|
threshold: 0.35,
|
||||||
|
iouThreshold: 0.8,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
const defaultEngines = project?.base_model_path ? ['base_model'] : ['sam3']
|
async function resetAutoAnnotations(batch) {
|
||||||
const [selectedThreshold, setSelectedThreshold] = useState({})
|
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
|
||||||
|
|
||||||
async function autolabel(batch, resume = false) {
|
|
||||||
setBusyId(batch.id)
|
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 {
|
try {
|
||||||
await api.startAutolabel(batch.id, { resume, engines, engine_classes, threshold })
|
await api.resetBatchAutoAnnotations(batch.id)
|
||||||
onChanged()
|
onChanged()
|
||||||
} catch (exc) {
|
} catch (exc) {
|
||||||
onError(exc.message)
|
onError(exc.message)
|
||||||
@@ -85,35 +72,45 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const toggleEngine = (batchId, engineKey) => {
|
function openFilePickerForYolo(batch) {
|
||||||
const current = selectedEngine[batchId] || defaultEngines
|
setAppendChoiceBatch(null)
|
||||||
let next
|
const input = document.createElement('input')
|
||||||
if (current.includes(engineKey)) {
|
input.type = 'file'
|
||||||
if (current.length === 1) return // keep at least 1 engine selected
|
input.accept = '.pt'
|
||||||
next = current.filter((e) => e !== engineKey)
|
input.onchange = async (e) => {
|
||||||
} else {
|
const file = e.target.files?.[0]
|
||||||
next = [...current, engineKey]
|
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) => {
|
function openSam3AppendModal(batch) {
|
||||||
const batchFilters = engineClassMap[batchId] || {}
|
setAppendChoiceBatch(null)
|
||||||
const currentEngClasses = batchFilters[engineKey] || defaultClasses
|
const projectClasses = project?.classes?.map((c) => c.name) || []
|
||||||
let next
|
setSam3AppendState({
|
||||||
if (currentEngClasses.includes(className)) {
|
batch,
|
||||||
if (currentEngClasses.length === 1) return // keep at least 1 class
|
selectedClasses: [...projectClasses],
|
||||||
next = currentEngClasses.filter((c) => c !== className)
|
threshold: 0.35,
|
||||||
} else {
|
iouThreshold: 0.8,
|
||||||
next = [...currentEngClasses, className]
|
|
||||||
}
|
|
||||||
setEngineClassMap({
|
|
||||||
...engineClassMap,
|
|
||||||
[batchId]: {
|
|
||||||
...batchFilters,
|
|
||||||
[engineKey]: next,
|
|
||||||
},
|
|
||||||
})
|
})
|
||||||
|
setCustomPromptInput('')
|
||||||
}
|
}
|
||||||
|
|
||||||
async function deleteBatch(batch) {
|
async function deleteBatch(batch) {
|
||||||
@@ -151,20 +148,13 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
|
|||||||
<thead>
|
<thead>
|
||||||
<tr>
|
<tr>
|
||||||
<th>Batch</th><th>Range</th><th>Frames</th><th>Reviewed</th>
|
<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>
|
</tr>
|
||||||
</thead>
|
</thead>
|
||||||
<tbody>
|
<tbody>
|
||||||
{batches.map((batch) => {
|
{batches.map((batch) => {
|
||||||
const batchJob = activeJobs?.find((j) => j.batch_id === batch.id)
|
const batchJob = activeJobs?.find((j) => j.batch_id === batch.id)
|
||||||
const isProcessing = Boolean(batchJob || busyId === 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 (
|
return (
|
||||||
<React.Fragment key={batch.id}>
|
<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.frame_count}</td>
|
||||||
<td className="num">{batch.reviewed}/{batch.frame_count}</td>
|
<td className="num">{batch.reviewed}/{batch.frame_count}</td>
|
||||||
<td className="num">{batch.annotation_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>
|
<td>
|
||||||
<span className={`tag ${batchJob ? 'info' : ''}`}>
|
<span className={`tag ${batchJob ? 'info' : ''}`}>
|
||||||
{batchJob ? `${batchJob.type} (${batchJob.status})` : batch.status}
|
{batchJob ? `${batchJob.type} (${batchJob.status})` : batch.status}
|
||||||
@@ -278,16 +178,20 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
|
|||||||
<td>
|
<td>
|
||||||
<div className="row" style={{ gap: 6 }}>
|
<div className="row" style={{ gap: 6 }}>
|
||||||
<button className="btn btn-primary" disabled={isProcessing || batch.frame_count === 0}
|
<button className="btn btn-primary" disabled={isProcessing || batch.frame_count === 0}
|
||||||
onClick={() => autolabel(batch)}>
|
title="Annotate using project base model with target class selection"
|
||||||
{batchJob?.type === 'autolabel' ? 'Processing…' : `Auto-annotate (${activeEngines.map(e => e === 'sam3' ? 'SAM3' : e === 'base_model' ? 'Model 1' : 'Model 2').join('+')})`}
|
onClick={() => openBaseModelAutolabelModal(batch)}>
|
||||||
|
{batchJob?.type === 'autolabel' ? 'Processing…' : 'Auto-annotate'}
|
||||||
</button>
|
</button>
|
||||||
{batch.annotation_count > 0 && (
|
|
||||||
<button className="btn" disabled={isProcessing || batch.frame_count === 0}
|
<button className="btn" disabled={isProcessing || batch.frame_count === 0}
|
||||||
title="Skip frames that already have automatic shapes"
|
title="Append new class detections using custom YOLO model or SAM3 text prompts"
|
||||||
onClick={() => autolabel(batch, true)}>
|
onClick={() => setAppendChoiceBatch(batch)}>
|
||||||
Resume
|
+ 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>
|
||||||
)}
|
|
||||||
<button className="btn" disabled={batch.frame_count === 0}
|
<button className="btn" disabled={batch.frame_count === 0}
|
||||||
onClick={() => navigate(`/projects/${project.id}/review?batch=${batch.id}`)}>
|
onClick={() => navigate(`/projects/${project.id}/review?batch=${batch.id}`)}>
|
||||||
Review ({batch.reviewed}/{batch.frame_count})
|
Review ({batch.reviewed}/{batch.frame_count})
|
||||||
@@ -299,122 +203,438 @@ function BatchList({ project, batches, activeJobs, onChanged, onError }) {
|
|||||||
</div>
|
</div>
|
||||||
</td>
|
</td>
|
||||||
</tr>
|
</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>
|
</React.Fragment>
|
||||||
)
|
)
|
||||||
})}
|
})}
|
||||||
</tbody>
|
</tbody>
|
||||||
</table>
|
</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>
|
</div>
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -91,6 +91,7 @@ export default function ModelsPage({ projectId, onProject }) {
|
|||||||
const [error, setError] = useState('')
|
const [error, setError] = useState('')
|
||||||
|
|
||||||
const [selectedBatchIds, setSelectedBatchIds] = useState([])
|
const [selectedBatchIds, setSelectedBatchIds] = useState([])
|
||||||
|
const [selectedClassIds, setSelectedClassIds] = useState([])
|
||||||
|
|
||||||
const load = useCallback(async () => {
|
const load = useCallback(async () => {
|
||||||
const [loadedProject, loadedSummary, modelPayload, hw, jobsPayload] = await Promise.all([
|
const [loadedProject, loadedSummary, modelPayload, hw, jobsPayload] = await Promise.all([
|
||||||
@@ -103,6 +104,7 @@ export default function ModelsPage({ projectId, onProject }) {
|
|||||||
setModels(modelPayload.models)
|
setModels(modelPayload.models)
|
||||||
setHardware(hw)
|
setHardware(hw)
|
||||||
setSelectedBatchIds(loadedSummary.batches.map((b) => b.id))
|
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))
|
const activeJob = jobsPayload.jobs?.find((j) => ['running', 'queued'].includes(j.status))
|
||||||
if (activeJob) setJob(activeJob)
|
if (activeJob) setJob(activeJob)
|
||||||
}, [projectId])
|
}, [projectId])
|
||||||
@@ -126,6 +128,7 @@ export default function ModelsPage({ projectId, onProject }) {
|
|||||||
setJob(await api.startTraining(projectId, {
|
setJob(await api.startTraining(projectId, {
|
||||||
epochs: Number(epochs),
|
epochs: Number(epochs),
|
||||||
batch_ids: selectedBatchIds.length > 0 ? selectedBatchIds : null,
|
batch_ids: selectedBatchIds.length > 0 ? selectedBatchIds : null,
|
||||||
|
class_ids: selectedClassIds.length > 0 ? selectedClassIds : null,
|
||||||
}))
|
}))
|
||||||
} catch (exc) {
|
} catch (exc) {
|
||||||
setError(exc.message)
|
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 (error && !project) return <p className="error-banner"><AlertIcon size={14} /> {error}</p>
|
||||||
if (!project || !summary) return <p className="empty">Loading…</p>
|
if (!project || !summary) return <p className="empty">Loading…</p>
|
||||||
|
|
||||||
@@ -169,50 +178,20 @@ export default function ModelsPage({ projectId, onProject }) {
|
|||||||
<AlertIcon size={14} /> {error}
|
<AlertIcon size={14} /> {error}
|
||||||
</p>}
|
</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 }}>
|
<div className="panel side-panel" style={{ marginBottom: 20 }}>
|
||||||
<h2>Uploaded Models for Training & Auto-Annotation</h2>
|
<h2>Base Model Configuration</h2>
|
||||||
<p className="hint">Upload up to 2 models to use for fine-tuning baseline or auto-annotating new video batches:</p>
|
<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)', 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' }}>Base Model</h3>
|
||||||
<h3 style={{ fontSize: '0.9rem', color: '#c084fc', margin: '0 0 6px 0' }}>Primary Model 1 (Base Model)</h3>
|
|
||||||
<p className="hint" style={{ fontSize: '0.78rem', margin: '0 0 10px 0' }}>
|
<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>
|
</p>
|
||||||
<div style={{ fontSize: '0.8rem', color: '#a1a1aa', marginBottom: 10 }}>
|
<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>
|
Status: <strong style={{ color: '#fff' }}>{project.base_model_path ? 'Custom model.pt loaded' : 'Default yolo11n.pt'}</strong>
|
||||||
</div>
|
</div>
|
||||||
<div style={{ fontSize: '0.78rem', color: '#38bdf8', marginBottom: 10 }}>
|
<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>
|
</div>
|
||||||
<input
|
<input
|
||||||
type="file"
|
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' }}>
|
<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>
|
</label>
|
||||||
</div>
|
</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>
|
||||||
|
|
||||||
<div className="review">
|
<div className="review">
|
||||||
@@ -291,6 +237,34 @@ export default function ModelsPage({ projectId, onProject }) {
|
|||||||
: `${project.base_model_fallback} (no base model uploaded)`}{' '}
|
: `${project.base_model_fallback} (no base model uploaded)`}{' '}
|
||||||
on the selected dataset batches ({selectedBatchIds.length}/{summary.batches.length} selected).
|
on the selected dataset batches ({selectedBatchIds.length}/{summary.batches.length} selected).
|
||||||
</p>
|
</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>
|
<div>
|
||||||
<label htmlFor="epochs">Epochs</label>
|
<label htmlFor="epochs">Epochs</label>
|
||||||
<input id="epochs" type="number" min="1" max="500" value={epochs}
|
<input id="epochs" type="number" min="1" max="500" value={epochs}
|
||||||
@@ -303,7 +277,7 @@ export default function ModelsPage({ projectId, onProject }) {
|
|||||||
</p>
|
</p>
|
||||||
)}
|
)}
|
||||||
<button className="btn btn-primary" onClick={train}
|
<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'}
|
{running ? 'Training…' : 'Start training'}
|
||||||
</button>
|
</button>
|
||||||
{summary.splits.train === 0 && (
|
{summary.splits.train === 0 && (
|
||||||
@@ -312,6 +286,9 @@ export default function ModelsPage({ projectId, onProject }) {
|
|||||||
{selectedBatchIds.length === 0 && summary.splits.train > 0 && (
|
{selectedBatchIds.length === 0 && summary.splits.train > 0 && (
|
||||||
<p className="hint" style={{ color: '#ef4444' }}>Select at least one dataset batch to train.</p>
|
<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>
|
</div>
|
||||||
|
|
||||||
{job && (
|
{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..."
|
echo "🔄 Restarting app services..."
|
||||||
|
|
||||||
# Kill running uvicorn and vite processes
|
# 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 "uvicorn backend.main:app" || true
|
||||||
pkill -f "vite" || true
|
pkill -f "vite" || true
|
||||||
sleep 1
|
sleep 2.5
|
||||||
|
|
||||||
export PATH="$HOME/.local/bin:$PATH"
|
export PATH="$HOME/.local/bin:$PATH"
|
||||||
|
|
||||||
# Run backend
|
# 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
|
# Run frontend
|
||||||
cd "$REPO_DIR/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."
|
echo "✅ Services restarted."
|
||||||
Reference in new issue
Block a user