fix(pre_commit): 🎨 auto format pre-commit hooks

This commit is contained in:
pre-commit-ci[bot] 2023-09-04 20:18:46 +00:00
parent 8538c26a22
commit 2706304a73
5 changed files with 67 additions and 54 deletions

View File

@ -1 +1 @@
data/
data/

View File

@ -1,6 +1,6 @@
## 👋 hello
This script performs traffic flow analysis using YOLOv8, an object-detection method and ByteTrack, a simple yet effective online multi-object tracking method. It uses the supervision package for multiple tasks such as tracking, annotations, etc.
This script performs traffic flow analysis using YOLOv8, an object-detection method and ByteTrack, a simple yet effective online multi-object tracking method. It uses the supervision package for multiple tasks such as tracking, annotations, etc.
## 💻 install
@ -17,15 +17,15 @@ This script performs traffic flow analysis using YOLOv8, an object-detection met
python3 -m venv venv
source venv/bin/activate
```
- install required dependencies
```bash
pip install -r requirements.txt
```
- download `traffic_analysis.pt` and `traffic_analysis.mov` files
```bash
./setup.sh
```
@ -39,4 +39,4 @@ python script.py \
--confidence_threshold 0.3 \
--iou_threshold 0.5 \
--target_video_path data/traffic_analysis_result.mov
```
```

View File

@ -1,4 +1,4 @@
supervision
tqdm
ultralytics
gdown
gdown

View File

@ -1,5 +1,5 @@
import argparse
from typing import Tuple, List, Dict, Set
from typing import Dict, List, Set, Tuple
import cv2
import numpy as np
@ -21,12 +21,11 @@ ZONE_OUT_POLYGONS = [
np.array([[950, 282], [1250, 282], [1250, 82], [950, 82]]),
np.array([[592, 860], [900, 860], [900, 1060], [592, 1060]]),
np.array([[592, 282], [592, 550], [392, 550], [392, 282]]),
np.array([[1250, 860], [1250, 560], [1450, 560], [1450, 860]])
np.array([[1250, 860], [1250, 560], [1450, 560], [1450, 860]]),
]
class DetectionsManager:
def __init__(self) -> None:
self.tracker_id_to_zone_id: Dict[int, int] = {}
self.counts: Dict[int, Dict[int, Set[int]]] = {}
@ -35,7 +34,7 @@ class DetectionsManager:
self,
detections_all: sv.Detections,
detections_in_zones: List[sv.Detections],
detections_out_zones: List[sv.Detections]
detections_out_zones: List[sv.Detections],
) -> sv.Detections:
for zone_in_id, detections_in_zone in enumerate(detections_in_zones):
for tracker_id in detections_in_zone.tracker_id:
@ -77,7 +76,7 @@ class VideoProcessor:
source_video_path: str,
target_video_path: str = None,
confidence_threshold: float = 0.3,
iou_threshold: float = 0.7
iou_threshold: float = 0.7,
) -> None:
self.conf_threshold = confidence_threshold
self.iou_threshold = iou_threshold
@ -89,18 +88,22 @@ class VideoProcessor:
self.video_info = sv.VideoInfo.from_video_path(source_video_path)
self.zones_in = initiate_polygon_zones(
ZONE_IN_POLYGONS, self.video_info.resolution_wh, sv.Position.CENTER)
ZONE_IN_POLYGONS, self.video_info.resolution_wh, sv.Position.CENTER
)
self.zones_out = initiate_polygon_zones(
ZONE_OUT_POLYGONS, self.video_info.resolution_wh, sv.Position.CENTER)
ZONE_OUT_POLYGONS, self.video_info.resolution_wh, sv.Position.CENTER
)
self.box_annotator = sv.BoxAnnotator(color=COLORS)
self.trace_annotator = sv.TraceAnnotator(
color=COLORS, position=sv.Position.CENTER, trace_length=100, thickness=2)
color=COLORS, position=sv.Position.CENTER, trace_length=100, thickness=2
)
self.detections_manager = DetectionsManager()
def process_video(self):
frame_generator = sv.get_video_frames_generator(
source_path=self.source_video_path)
source_path=self.source_video_path
)
if self.target_video_path:
with sv.VideoSink(self.target_video_path, self.video_info) as sink:
@ -110,24 +113,28 @@ class VideoProcessor:
else:
for frame in tqdm(frame_generator, total=self.video_info.total_frames):
annotated_frame = self.process_frame(frame)
cv2.imshow('Processed Video', annotated_frame)
if cv2.waitKey(1) & 0xFF == ord('q'):
cv2.imshow("Processed Video", annotated_frame)
if cv2.waitKey(1) & 0xFF == ord("q"):
break
cv2.destroyAllWindows()
def annotate_frame(self, frame: np.ndarray, detections: sv.Detections) -> np.ndarray:
def annotate_frame(
self, frame: np.ndarray, detections: sv.Detections
) -> np.ndarray:
annotated_frame = frame.copy()
for i, (zone_in, zone_out) in enumerate(zip(self.zones_in, self.zones_out)):
annotated_frame = sv.draw_polygon(
annotated_frame, zone_in.polygon, COLORS.colors[i])
annotated_frame, zone_in.polygon, COLORS.colors[i]
)
annotated_frame = sv.draw_polygon(
annotated_frame, zone_out.polygon, COLORS.colors[i])
annotated_frame, zone_out.polygon, COLORS.colors[i]
)
labels = [f"#{tracker_id}" for tracker_id in detections.tracker_id]
annotated_frame = self.trace_annotator.annotate(
annotated_frame, detections)
annotated_frame = self.trace_annotator.annotate(annotated_frame, detections)
annotated_frame = self.box_annotator.annotate(
annotated_frame, detections, labels)
annotated_frame, detections, labels
)
for zone_out_id, zone_out in enumerate(self.zones_out):
zone_center = sv.get_polygon_center(polygon=zone_out.polygon)
@ -140,14 +147,15 @@ class VideoProcessor:
scene=annotated_frame,
text=str(count),
text_anchor=text_anchor,
background_color=COLORS.colors[zone_in_id]
background_color=COLORS.colors[zone_in_id],
)
return annotated_frame
def process_frame(self, frame: np.ndarray) -> np.ndarray:
results = self.model(
frame, verbose=False, conf=self.conf_threshold, iou=self.iou_threshold)[0]
frame, verbose=False, conf=self.conf_threshold, iou=self.iou_threshold
)[0]
detections = sv.Detections.from_ultralytics(results)
detections.class_id = np.zeros(len(detections))
detections = self.tracker.update_with_detections(detections)
@ -162,33 +170,42 @@ class VideoProcessor:
detections_out_zones.append(detections_out_zone)
detections = self.detections_manager.update(
detections, detections_in_zones, detections_out_zones)
detections, detections_in_zones, detections_out_zones
)
return self.annotate_frame(frame, detections)
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="Traffic Flow Analysis with YOLO and ByteTrack")
description="Traffic Flow Analysis with YOLO and ByteTrack"
)
parser.add_argument(
"--source_weights_path", required=True,
help="Path to the source weights file", type=str
"--source_weights_path",
required=True,
help="Path to the source weights file",
type=str,
)
parser.add_argument(
"--source_video_path", required=True,
help="Path to the source video file", type=str
"--source_video_path",
required=True,
help="Path to the source video file",
type=str,
)
parser.add_argument(
"--target_video_path", default=None,
help="Path to the target video file (output)", type=str
"--target_video_path",
default=None,
help="Path to the target video file (output)",
type=str,
)
parser.add_argument(
"--confidence_threshold", default=0.3,
help="Confidence threshold for the model", type=float
"--confidence_threshold",
default=0.3,
help="Confidence threshold for the model",
type=float,
)
parser.add_argument(
"--iou_threshold", default=0.7,
help="IOU threshold for the model", type=float
"--iou_threshold", default=0.7, help="IOU threshold for the model", type=float
)
args = parser.parse_args()

View File

@ -3,9 +3,9 @@ from typing import List, Optional, Union
import cv2
import numpy as np
from supervision.geometry.core import Position
from supervision.detection.core import Detections
from supervision.draw.color import Color, ColorPalette
from supervision.geometry.core import Position
class BoxAnnotator:
@ -206,12 +206,11 @@ class MaskAnnotator:
class Trace:
def __init__(
self,
max_size: Optional[int] = None,
start_frame_id: int = 0,
anchor: Position = Position.CENTER
anchor: Position = Position.CENTER,
) -> None:
self.current_frame_id = start_frame_id
self.max_size = max_size
@ -224,8 +223,9 @@ class Trace:
def put(self, detections: Detections) -> None:
frame_id = np.full(len(detections), self.current_frame_id, dtype=int)
self.frame_id = np.concatenate([self.frame_id, frame_id])
self.xy = np.concatenate([
self.xy, detections.get_anchor_coordinates(self.anchor)])
self.xy = np.concatenate(
[self.xy, detections.get_anchor_coordinates(self.anchor)]
)
self.tracker_id = np.concatenate([self.tracker_id, detections.tracker_id])
unique_frame_id = np.unique(self.frame_id)
@ -244,14 +244,12 @@ class Trace:
class TraceAnnotator:
def __init__(
self,
color: Union[Color, ColorPalette] = ColorPalette.default(),
position: Optional[Position] = Position.CENTER,
trace_length: int = 30,
thickness: int = 2,
self,
color: Union[Color, ColorPalette] = ColorPalette.default(),
position: Optional[Position] = Position.CENTER,
trace_length: int = 30,
thickness: int = 2,
):
self.color: Union[Color, ColorPalette] = color
self.position = position
@ -262,9 +260,7 @@ class TraceAnnotator:
self.trace.put(detections)
for i, (xyxy, mask, confidence, class_id, tracker_id) in enumerate(detections):
class_id = (
detections.class_id[i] if class_id is not None else None
)
class_id = detections.class_id[i] if class_id is not None else None
idx = class_id if class_id is not None else i
color = (
self.color.by_idx(idx)