File size: 9,522 Bytes
72534cf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
"""
detector.py
-----------
Core AI module for the Traffic Scene Interpretation System.

Responsibilities:
- Load a YOLO model (Ultralytics) once and reuse it.
- Run detection on a single image (numpy array / PIL image).
- Run detection on a video file, frame by frame, and write an annotated
  output video.
- Aggregate per-frame detections into simple traffic-scene statistics
  (vehicle counts, congestion level) — this is the "scene interpretation"
  layer on top of raw object detection.

Kept deliberately simple and well-commented so it's easy to explain
during a project defense.
"""

from __future__ import annotations

import time
from collections import Counter
from dataclasses import dataclass, field

import cv2
import numpy as np
from ultralytics import YOLO

# ---------------------------------------------------------------------------
# Config
# ---------------------------------------------------------------------------

# Classes from the COCO dataset (what pretrained YOLO already knows) that are
# relevant to a traffic scene. No custom training needed for the MVP.
VEHICLE_CLASSES = {
    "car": "Cars",
    "bus": "Buses",
    "truck": "Trucks",
    "motorcycle": "Motorcycles",
    "bicycle": "Bicycles",
    "person": "Pedestrians",
}

# Thresholds used to translate a raw vehicle count into a human-readable
# "scene interpretation" label. Tune these once you see real results.
CONGESTION_THRESHOLDS = {
    "free": 5,       # 0-5 vehicles  -> Free flowing
    "moderate": 15,  # 6-15 vehicles -> Moderate traffic
    # >15 vehicles -> Congested
}

IMAGE_MODEL = "yolov8m.pt"   # Used for single-image detection. Accuracy matters more than
                              # speed here since it only runs once per upload.
VIDEO_MODEL = "yolov8m.pt"   # Bumped up from yolov8s for better accuracy, per testing feedback.
                              # This is noticeably slower (roughly 2-3x the compute of yolov8s).
                              # Frame skipping + reduced imgsz below help offset that cost.
                              # If processing time becomes uncomfortable, drop back to "yolov8s.pt".
VIDEO_IMGSZ = 480             # Shrinking the frame before detection speeds video up further.
                              # Lower = faster but less accurate on small/distant objects.
                              # Try 384 if still too slow; try 640 (native) if you have room to spare.
VIDEO_FRAME_SKIP = 2          # Run detection on 1 out of every N frames; reuse the previous
                              # frame's boxes for the skipped ones. 2 = run detection on half
                              # the frames (~2x faster). Set to 1 to disable (detect every frame).
CONFIDENCE_THRESHOLD = 0.25    # Lowered from 0.35 to catch smaller/more distant vehicles.
                                # If you start seeing false detections (boxes on things that aren't
                                # vehicles), raise this back up toward 0.35-0.4.


@dataclass
class SceneStats:
    """Aggregated statistics for one image or one video."""
    counts: Counter = field(default_factory=Counter)
    total_frames: int = 1
    fps: float = 0.0

    def per_frame_average(self) -> Counter:
        if self.total_frames <= 0:
            return self.counts
        return Counter({k: round(v / self.total_frames, 1) for k, v in self.counts.items()})

    def congestion_label(self) -> str:
        vehicle_count = sum(
            v for k, v in self.per_frame_average().items() if k != "Pedestrians"
        )
        if vehicle_count <= CONGESTION_THRESHOLDS["free"]:
            return "Free flowing"
        elif vehicle_count <= CONGESTION_THRESHOLDS["moderate"]:
            return "Moderate traffic"
        else:
            return "Congested"


class TrafficDetector:
    """
    Wraps two YOLO models:
      - self.image_model: larger/more accurate, used for single-image detection.
      - self.video_model: smaller/faster, used for video (and would be used for
        webcam too, if that's added later) since it runs once per frame.
    Both are loaded once at startup and reused.
    """

    def __init__(
        self,
        image_model_path: str = IMAGE_MODEL,
        video_model_path: str = VIDEO_MODEL,
        conf: float = CONFIDENCE_THRESHOLD,
    ):
        self.image_model = YOLO(image_model_path)
        # Avoid loading the same weights twice if someone sets both to the same file.
        self.video_model = (
            self.image_model if video_model_path == image_model_path else YOLO(video_model_path)
        )
        self.conf = conf

    # -- Image -------------------------------------------------------------

    def detect_image(self, image: np.ndarray) -> tuple[np.ndarray, SceneStats]:
        """
        Run detection on a single BGR image (as read by cv2).
        Returns (annotated_image, stats).
        """
        results = self.image_model.predict(image, conf=self.conf, verbose=False)
        result = results[0]

        counts = self._count_from_result(result)
        annotated = result.plot()  # draws boxes + labels + confidence

        stats = SceneStats(counts=counts, total_frames=1)
        return annotated, stats

    # -- Video ---------------------------------------------------------------

    def detect_video(self, input_path: str, output_path: str, progress_callback=None) -> SceneStats:
        """
        Process a video file frame-by-frame:
          read frame -> YOLO detect -> draw boxes -> write frame to output.
        Uses the faster video_model + a reduced inference size (VIDEO_IMGSZ) to
        keep processing time reasonable without a GPU.

        Frame skipping (VIDEO_FRAME_SKIP): to save time, detection only runs on
        every Nth frame. For the frames in between, we reuse the last detected
        boxes and re-draw them onto the new frame. Since consecutive frames are
        1/25th-1/30th of a second apart, objects barely move between them, so
        this looks smooth while cutting detection calls (the expensive part)
        roughly in half.

        progress_callback(current_frame, total_frames) is called after each
        frame if provided, so a Streamlit progress bar can be updated.
        """
        cap = cv2.VideoCapture(input_path)
        if not cap.isOpened():
            raise RuntimeError(f"Could not open video: {input_path}")

        fps = cap.get(cv2.CAP_PROP_FPS) or 25.0
        width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
        height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
        total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) or None

        fourcc = cv2.VideoWriter_fourcc(*"mp4v")
        writer = cv2.VideoWriter(output_path, fourcc, fps, (width, height))

        total_counts: Counter = Counter()
        frame_idx = 0
        detected_frame_count = 0  # frames actually run through YOLO (for stats averaging)
        start_time = time.time()

        last_result = None  # cached YOLO result, reused on skipped frames

        while True:
            ok, frame = cap.read()
            if not ok:
                break

            run_detection = (frame_idx % VIDEO_FRAME_SKIP == 0) or (last_result is None)

            if run_detection:
                results = self.video_model.predict(frame, conf=self.conf, imgsz=VIDEO_IMGSZ, verbose=False)
                last_result = results[0]
                total_counts.update(self._count_from_result(last_result))
                detected_frame_count += 1
                annotated = last_result.plot()
            else:
                # Re-draw the previous frame's boxes onto the current frame.
                # plot(img=...) lets us reuse a YOLO result's boxes on a new image.
                annotated = last_result.plot(img=frame)

            writer.write(annotated)

            frame_idx += 1
            if progress_callback:
                progress_callback(frame_idx, total_frames)

        cap.release()
        writer.release()

        elapsed = max(time.time() - start_time, 1e-6)
        processing_fps = frame_idx / elapsed

        return SceneStats(
            counts=total_counts,
            total_frames=max(detected_frame_count, 1),
            fps=round(processing_fps, 1),
        )

    # -- Webcam (single-frame step, called repeatedly by the UI layer) -----

    def detect_frame(self, frame: np.ndarray) -> tuple[np.ndarray, Counter]:
        """
        Used for live webcam mode: process exactly one frame and return it
        annotated, plus its own counts. Uses the fast video_model, same
        reasoning as detect_video above. The caller (app.py) is responsible
        for the capture loop, since Streamlit needs to own that loop to
        keep the UI responsive.
        """
        results = self.video_model.predict(frame, conf=self.conf, imgsz=VIDEO_IMGSZ, verbose=False)
        result = results[0]
        counts = self._count_from_result(result)
        return result.plot(), counts

    # -- Helpers -------------------------------------------------------------

    def _count_from_result(self, result) -> Counter:
        """Turn one YOLO result into a Counter of {readable_label: count}."""
        counts: Counter = Counter()
        names = result.names
        if result.boxes is None:
            return counts
        for cls_id in result.boxes.cls.tolist():
            raw_name = names[int(cls_id)]
            label = VEHICLE_CLASSES.get(raw_name)
            if label:
                counts[label] += 1
        return counts