Vehicle-Counting / core /detector.py
riezqidr's picture
Add initial project structure with Streamlit UI and utility functions
ab2f940
Raw
History Blame Contribute Delete
3.64 kB
"""
RT-DETR Vehicle Detector
Wrapper class untuk inference menggunakan model RT-DETR yang sudah di-fine-tune
pada dataset kendaraan. Nama kelas dibaca langsung dari model, jadi kompatibel
dengan dataset apapun yang dipakai saat training.
"""
import torch
from ultralytics import RTDETR
from pathlib import Path
class VehicleDetector:
"""
Wrapper untuk RT-DETR model inference.
Menangani loading model, inference per frame, dan filtering
berdasarkan confidence threshold.
"""
def __init__(self, model_path="models/best.pt", confidence=0.5):
"""
Args:
model_path: path ke file model .pt
confidence: minimum confidence score untuk deteksi
"""
self.model_path = Path(model_path)
self.confidence = confidence
self.device = "cuda" if torch.cuda.is_available() else "cpu"
self.model = None
self.class_names = {}
self._load_model()
def _load_model(self):
"""Load RT-DETR model dari file dan baca class names."""
if not self.model_path.exists():
raise FileNotFoundError(
f"Model tidak ditemukan di {self.model_path}. "
f"Pastikan file best.pt sudah ada di folder models/"
)
self.model = RTDETR(str(self.model_path))
self.model.to(self.device)
# baca class names langsung dari model
# model.names biasanya berupa dict {0: 'class0', 1: 'class1', ...}
if hasattr(self.model, "names") and self.model.names:
self.class_names = self.model.names
else:
# fallback kalau names tidak ada
self.class_names = {0: "Vehicle"}
print(f"Model loaded pada device: {self.device}")
print(f"Kelas terdeteksi: {self.class_names}")
def detect(self, frame):
"""
Jalankan deteksi pada satu frame.
Args:
frame: numpy array (BGR image dari OpenCV)
Returns:
list of dict, masing-masing berisi:
- bbox: [x1, y1, x2, y2]
- confidence: float
- class_id: int
- class_name: str
"""
results = self.model.predict(
frame,
conf=self.confidence,
device=self.device,
verbose=False
)
detections = []
for result in results:
boxes = result.boxes
if boxes is None or len(boxes) == 0:
continue
for i in range(len(boxes)):
bbox = boxes.xyxy[i].cpu().numpy().tolist()
conf = float(boxes.conf[i].cpu().numpy())
cls_id = int(boxes.cls[i].cpu().numpy())
# ambil nama kelas, fallback ke "Unknown" kalau index diluar mapping
cls_name = self.class_names.get(cls_id, "Unknown")
detections.append({
"bbox": bbox,
"confidence": conf,
"class_id": cls_id,
"class_name": cls_name
})
return detections
def set_confidence(self, confidence):
"""Update confidence threshold."""
self.confidence = confidence
def get_model_info(self):
"""Return informasi model untuk ditampilkan di UI."""
info = {
"model_path": str(self.model_path),
"device": self.device,
"confidence_threshold": self.confidence,
"num_classes": len(self.class_names),
"classes": list(self.class_names.values())
}
return info