File size: 3,769 Bytes
0dc14c7
 
 
 
 
 
 
 
5f1e317
 
 
0dc14c7
5f1e317
 
0dc14c7
 
0dd5b63
0dc14c7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
"""
YOLOv8 vehicle and license plate detector.
Handles frame-by-frame detection and crop extraction for OCR.
"""
import cv2
import numpy as np
from ultralytics import YOLO
from huggingface_hub import hf_hub_download
import torch
import torch.serialization
from ultralytics.nn.tasks import DetectionModel

# Allow YOLO models to load safely in PyTorch 2.6+
torch.serialization.add_safe_globals([DetectionModel])
VEHICLE_CLASSES = {2: "car", 3: "motorcycle", 5: "bus", 7: "truck"}
CONF_THRESHOLD = 0.45
HF_REPO_ID = "aaradhya233/license-plate-model"


class VehicleDetector:
    def __init__(self, model_path: str = "yolov8n.pt", plate_model_path: str | None = None):
        # Download YOLO vehicle model from Hugging Face Hub
        local_model_path = hf_hub_download(repo_id=HF_REPO_ID, filename=model_path)
        self.vehicle_model = YOLO(local_model_path)

        # Download plate model if specified
        if plate_model_path:
            local_plate_path = hf_hub_download(repo_id=HF_REPO_ID, filename=plate_model_path)
            self.plate_model = YOLO(local_plate_path)
        else:
            # Auto-load license_plate_detector.pt from HF Hub
            local_plate_path = hf_hub_download(repo_id=HF_REPO_ID, filename="license_plate_detector.pt")
            self.plate_model = YOLO(local_plate_path)

        print(f"[Detector] Loaded: {model_path}")

    def detect_frame(self, frame: np.ndarray) -> dict:
        """
        Run detection on a single BGR frame.
        Returns annotated frame + list of vehicle dicts with plate crops.
        """
        results = self.vehicle_model(frame, conf=CONF_THRESHOLD, verbose=False)[0]
        vehicles = []

        for box in results.boxes:
            cls_id = int(box.cls[0])
            if cls_id not in VEHICLE_CLASSES:
                continue
            x1, y1, x2, y2 = map(int, box.xyxy[0])
            vehicle_crop = frame[y1:y2, x1:x2]
            plate_crop = self._extract_plate(vehicle_crop, frame, x1, y1, x2, y2)
            vehicles.append({
                "bbox": [x1, y1, x2, y2],
                "class": VEHICLE_CLASSES[cls_id],
                "confidence": float(box.conf[0]),
                "plate_crop": plate_crop,
            })

        return {"annotated_frame": self._draw_boxes(frame.copy(), vehicles), "vehicles": vehicles}

    def detect_video(self, source: str | int = 0):
        """Generator yielding per-frame detection results. source=0 for webcam."""
        cap = cv2.VideoCapture(source)
        if not cap.isOpened():
            raise IOError(f"Cannot open: {source}")
        try:
            while True:
                ret, frame = cap.read()
                if not ret:
                    break
                yield self.detect_frame(frame)
        finally:
            cap.release()

    def _extract_plate(self, vehicle_crop, full_frame, vx1, vy1, vx2, vy2):
        if vehicle_crop.size == 0:
            return None
        if self.plate_model:
            res = self.plate_model(vehicle_crop, conf=0.4, verbose=False)[0]
            if len(res.boxes):
                px1, py1, px2, py2 = map(int, res.boxes[0].xyxy[0])
                return vehicle_crop[py1:py2, px1:px2]
        # Heuristic: bottom 35%, center 70% of vehicle bbox
        h, w = vy2 - vy1, vx2 - vx1
        return full_frame[vy1 + int(h * 0.65):vy2, vx1 + int(w * 0.15):vx2 - int(w * 0.15)]

    @staticmethod
    def _draw_boxes(frame, vehicles):
        for v in vehicles:
            x1, y1, x2, y2 = v["bbox"]
            cv2.rectangle(frame, (x1, y1), (x2, y2), (0, 200, 100), 2)
            cv2.putText(frame, f"{v['class']} {v['confidence']:.2f}",
                        (x1, max(y1 - 8, 12)), cv2.FONT_HERSHEY_SIMPLEX, 0.55, (0, 200, 100), 2)
        return frame