update from asus 106

This commit is contained in:
asus committed 2026-08-05 15:56:11 +07:00
1 parent 6637fb1302
commit 8285400254
28 files changed
+3215 -459

No files matched your search

+485
View File
@@ -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()
+26
View File
@@ -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
+1
View File
@@ -0,0 +1 @@
"""Package marker."""
+390
View File
@@ -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)
+237
View File
@@ -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()
+251
View File
@@ -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,
)
+82
View File
@@ -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
+98
View File
@@ -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: ...
+126
View File
@@ -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()
+83
View File
@@ -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)
+151
View File
@@ -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
+43
View File
@@ -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
View File
@@ -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}
+2
View File
@@ -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))
+64 -69
View File
@@ -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,7 +200,10 @@ 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})
review.replace_auto(frame["id"], items) if job.params.get("append"):
review.append_auto(frame["id"], items)
else:
review.replace_auto(frame["id"], items)
written += len(items) 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)")
except Exception as exc: except Exception as exc:
+46 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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
+5
View File
@@ -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"]]
+50
View File
@@ -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
View File
@@ -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
View File
@@ -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>
+24
View File
@@ -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}`),
+5 -1
View File
@@ -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" />
+493 -273
View File
@@ -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 className="btn" disabled={isProcessing || batch.frame_count === 0}
title="Append new class detections using custom YOLO model or SAM3 text prompts"
onClick={() => setAppendChoiceBatch(batch)}>
+ Append
</button>
<button className="btn" disabled={isProcessing || batch.annotation_count === 0}
title="Clear all auto-generated shapes and reset frame review statuses to pending"
onClick={() => resetAutoAnnotations(batch)}>
🔄 Reset Auto
</button> </button>
{batch.annotation_count > 0 && (
<button className="btn" disabled={isProcessing || batch.frame_count === 0}
title="Skip frames that already have automatic shapes"
onClick={() => autolabel(batch, true)}>
Resume
</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 &gt;</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 &gt;</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>
) )
} }
+71 -94
View File
@@ -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,103 +178,40 @@ 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 baseline benchmark.
Used as fine-tuning starting point and benchmark baseline. </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 style={{ fontSize: '0.78rem', color: '#38bdf8', marginBottom: 10 }}>
<strong>Model 1 Classes:</strong> {project.classes?.map((c) => c.name).join(', ')}
</div>
<input
type="file"
accept=".pt"
id="upload-primary-model"
style={{ display: 'none' }}
onChange={async (e) => {
const file = e.target.files?.[0]
if (!file) return
try {
await api.uploadBaseModel(projectId, file)
load()
} catch (err) {
setError(err.message)
}
}}
/>
<label htmlFor="upload-primary-model" className="btn btn-ghost" style={{ cursor: 'pointer', padding: '4px 10px', fontSize: '0.8rem' }}>
Upload Primary Model 1 (.pt)
</label>
</div> </div>
<div style={{ fontSize: '0.78rem', color: '#38bdf8', marginBottom: 10 }}>
<div style={{ padding: 14, background: 'rgba(0,0,0,0.3)', borderRadius: 8, border: '1px solid rgba(255,255,255,0.1)' }}> <strong>Base Model Classes:</strong> {project.classes?.map((c) => c.name).join(', ')}
<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>
<input
type="file"
accept=".pt"
id="upload-primary-model"
style={{ display: 'none' }}
onChange={async (e) => {
const file = e.target.files?.[0]
if (!file) return
try {
await api.uploadBaseModel(projectId, file)
load()
} catch (err) {
setError(err.message)
}
}}
/>
<label htmlFor="upload-primary-model" className="btn btn-ghost" style={{ cursor: 'pointer', padding: '4px 10px', fontSize: '0.8rem' }}>
Upload Base Model (.pt)
</label>
</div> </div>
</div> </div>
@@ -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 && (
+328
View File
@@ -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>
)
}
+4 -3
View File
@@ -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."