From 2706304a73e44651014d8a43771f388321417da3 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Mon, 4 Sep 2023 20:18:46 +0000 Subject: [PATCH] =?UTF-8?q?fix(pre=5Fcommit):=20=F0=9F=8E=A8=20auto=20form?= =?UTF-8?q?at=20pre-commit=20hooks?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- examples/traffic_analysis/.gitignore | 2 +- examples/traffic_analysis/README.md | 12 ++-- examples/traffic_analysis/requirements.txt | 2 +- examples/traffic_analysis/script.py | 79 +++++++++++++--------- supervision/detection/annotate.py | 26 +++---- 5 files changed, 67 insertions(+), 54 deletions(-) diff --git a/examples/traffic_analysis/.gitignore b/examples/traffic_analysis/.gitignore index adbb97d2..8fce6030 100644 --- a/examples/traffic_analysis/.gitignore +++ b/examples/traffic_analysis/.gitignore @@ -1 +1 @@ -data/ \ No newline at end of file +data/ diff --git a/examples/traffic_analysis/README.md b/examples/traffic_analysis/README.md index 1f844033..8aed1ebc 100644 --- a/examples/traffic_analysis/README.md +++ b/examples/traffic_analysis/README.md @@ -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 -``` \ No newline at end of file +``` diff --git a/examples/traffic_analysis/requirements.txt b/examples/traffic_analysis/requirements.txt index 985a52bc..b57c1159 100644 --- a/examples/traffic_analysis/requirements.txt +++ b/examples/traffic_analysis/requirements.txt @@ -1,4 +1,4 @@ supervision tqdm ultralytics -gdown \ No newline at end of file +gdown diff --git a/examples/traffic_analysis/script.py b/examples/traffic_analysis/script.py index 8171f74d..8cf3a0c4 100644 --- a/examples/traffic_analysis/script.py +++ b/examples/traffic_analysis/script.py @@ -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() diff --git a/supervision/detection/annotate.py b/supervision/detection/annotate.py index 2a06a1a6..9c628937 100644 --- a/supervision/detection/annotate.py +++ b/supervision/detection/annotate.py @@ -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)