Improvements in thread #818: documentation, handle unit test, improvements pattern directory

This commit is contained in:
Adonai Vera 2024-02-01 17:53:15 -05:00
parent 91260bdfc9
commit 788bc317fa
5 changed files with 402 additions and 142 deletions

View File

@ -1,8 +0,0 @@
---
comments: true
status: new
---
## Save CSV Detection
:::supervision.detection.tools.csv_sink.CSVSink

View File

@ -0,0 +1,12 @@
---
comments: true
status: new
---
# Save Detections
<div class="md-typeset">
<h2>CSV Sink</h2>
</div>
:::supervision.detection.tools.csv_sink.CSVSink

View File

@ -41,7 +41,7 @@ nav:
- Polygon Zone: detection/tools/polygon_zone.md
- Inference Slicer: detection/tools/inference_slicer.md
- Detection Smoother: detection/tools/smoother.md
- Save CSV Detection: detection/tools/csv_sink.md
- Save Detections: detection/tools/save_detections.md
- Annotators: annotators.md
- Trackers: trackers.md
- Datasets: datasets.md

View File

@ -1,8 +1,11 @@
from __future__ import annotations
import csv
import os
from typing import Any, Dict, List, Optional
import numpy as np
from supervision.detection.core import Detections
BASE_HEADER = [
@ -15,7 +18,6 @@ BASE_HEADER = [
"tracker_id",
]
class CSVSink:
"""
A utility class for saving detection data to a CSV file. This class is designed to
@ -26,60 +28,50 @@ class CSVSink:
providing flexibility for logging various types of information.
Args:
filename (str): The name of the CSV file where the detections will be stored.
file_name (str): The name of the CSV file where the detections will be stored.
Defaults to 'output.csv'.
Example:
```python
import numpy as np
import cv2
import supervision as sv
from ultralytics import YOLO
import time
model = YOLO("yolov8n.pt")
tracker = sv.ByteTrack()
box_annotator = sv.BoundingBoxAnnotator()
label_annotator = sv.LabelAnnotator()
csv_sink = sv.CSVSink(...)
image = cv2.imread(<SOURCE_IMAGE_PATH>)
model = YOLO(<SOURCE_MODEL_PATH>)
def callback(frame: np.ndarray, _: int) -> np.ndarray:
start_time = time.time()
results = model(frame)[0]
detections = sv.Detections.from_ultralytics(results)
detections = tracker.update_with_detections(detections)
csv_sink = sv.CSVSink(<RESULT_CSV_FILE_PATH>)
labels = [
f"#{tracker_id} {results.names[class_id]}"
for class_id, tracker_id
in zip(detections.class_id, detections.tracker_id)
]
time_frame = (time.time() - start_time)
csv_sink.append(detections, custom_data={"processing_time": time_frame})
annotated_frame = box_annotator.annotate(
frame.copy(), detections=detections)
return label_annotator.annotate(
annotated_frame, detections=detections, labels=labels)
csv_sink.open()
sv.process_video(
source_path="people-walking.mp4",
target_path="result.mp4",
callback=callback
)
csv_sink.close()
result = model(image)[0]
detections = sv.Detections.from_ultralytics(result)
with csv_sink as sink:
sink.append(detections, custom_data={'<CUSTOM_LABEL>':'<CUSTOM_DATA>'})
```
""" # noqa: E501 // docs
def __init__(self, filename: str = "output.csv"):
self.filename = filename
def __init__(self, file_name: str = "output.csv"):
"""
Initialize the CSVSink instance.
Args:
file_name (str): The name of the CSV file.
Returns:
None
"""
self.file_name = file_name
self.file: Optional[open] = None
self.writer: Optional[csv.writer] = None
self.header_written = False
self.fieldnames = [] # To keep track of header names
self.field_names = []
def __enter__(self) -> CSVSink:
"""
Enter the context manager.
Returns:
CSVSink: The CSVSink instance.
"""
self.open()
return self
@ -89,13 +81,40 @@ class CSVSink:
exc_val: Optional[Exception],
exc_tb: Optional[Any],
) -> None:
"""
Exit the context manager.
Args:
exc_type (Optional[type]): The type of exception.
exc_val (Optional[Exception]): The exception instance.
exc_tb (Optional[Any]): The traceback.
Returns:
None
"""
self.close()
def open(self) -> None:
self.file = open(self.filename, "w", newline="")
"""
Open the CSV file for writing.
Returns:
None
"""
parent_directory = os.path.dirname(self.file_name)
if parent_directory and not os.path.exists(parent_directory):
os.makedirs(parent_directory)
self.file = open(self.file_name, "w", newline="")
self.writer = csv.writer(self.file)
def close(self) -> None:
"""
Close the CSV file.
Returns:
None
"""
if self.file:
self.file.close()
@ -103,6 +122,16 @@ class CSVSink:
def parse_detection_data(
detections: Detections, custom_data: Dict[str, Any] = None
) -> List[Dict[str, Any]]:
"""
Parse detection data into a list of dictionaries.
Args:
detections (Detections): The detection data.
custom_data (Dict[str, Any]): Custom data to include.
Returns:
List[Dict[str, Any]]: A list of dictionaries representing the data.
"""
parsed_rows = []
for i in range(len(detections.xyxy)):
row = {
@ -110,13 +139,18 @@ class CSVSink:
"y_min": detections.xyxy[i][1],
"x_max": detections.xyxy[i][2],
"y_max": detections.xyxy[i][3],
"class_id": detections.class_id[i],
"confidence": detections.confidence[i],
"tracker_id": detections.tracker_id[i],
"class_id": "" if detections.class_id is None else str(detections.class_id[i]),
"confidence": "" if detections.confidence is None else str(detections.confidence[i]),
"tracker_id": "" if detections.tracker_id is None else str(detections.tracker_id[i]),
}
if hasattr(detections, "data"):
for key, value in detections.data.items():
row[key] = value[i]
if value.ndim == 0:
row[key] = value
else:
row[key] = value[i]
if custom_data:
row.update(custom_data)
parsed_rows.append(row)
@ -125,9 +159,19 @@ class CSVSink:
def append(
self, detections: Detections, custom_data: Dict[str, Any] = None
) -> None:
"""
Append detection data to the CSV file.
Args:
detections (Detections): The detection data.
custom_data (Dict[str, Any]): Custom data to include.
Returns:
None
"""
if not self.writer:
raise Exception(
f"Cannot append to CSV: The file '{self.filename}' is not open."
f"Cannot append to CSV: The file '{self.file_name}' is not open."
)
if not self.header_written:
self.write_header(detections, custom_data)
@ -135,13 +179,23 @@ class CSVSink:
parsed_rows = CSVSink.parse_detection_data(detections, custom_data)
for row in parsed_rows:
self.writer.writerow(
[row.get(fieldname, "") for fieldname in self.fieldnames]
[row.get(field_name, "") for field_name in self.field_names]
)
def write_header(self, detections: Detections, custom_data: Dict[str, Any]) -> None:
"""
Write the CSV header based on the provided detection and custom data.
Args:
detections (Detections): The detection data.
custom_data (Dict[str, Any]): Custom data to include in the header.
Returns:
None
"""
dynamic_header = sorted(
set(custom_data.keys()) | set(getattr(detections, "data", {}).keys())
)
self.fieldnames = BASE_HEADER + dynamic_header
self.writer.writerow(self.fieldnames)
self.field_names = BASE_HEADER + dynamic_header
self.writer.writerow(self.field_names)
self.header_written = True

View File

@ -1,111 +1,313 @@
import csv
import os
from contextlib import ExitStack as DoesNotRaise
from test.test_utils import mock_detections
from typing import Any, Dict, List
import numpy as np
import pytest
import supervision as sv
from supervision.detection.core import Detections
@pytest.fixture(scope="module")
def detection_instances():
# Setup detection instances as per the provided example
detections = Detections(
xyxy=np.array([[10, 20, 30, 40], [50, 60, 70, 80]]),
confidence=np.array([0.7, 0.8]),
class_id=np.array([0, 0]),
tracker_id=np.array([0, 1]),
data={"class_name": np.array(["person", "person"])},
)
@pytest.mark.parametrize(
"detections, custom_data, second_detections, second_custom_data, file_name, expected_result",
[
(
mock_detections(
xyxy=[[10, 20, 30, 40], [50, 60, 70, 80]],
confidence=[0.7, 0.8],
class_id=[0, 0],
tracker_id=[0, 1],
data={"class_name": ["person", "person"]}
),
{"frame_number": 42},
mock_detections(
xyxy=[[15, 25, 35, 45], [55, 65, 75, 85]],
confidence=[0.6, 0.9],
class_id=[1, 1],
tracker_id=[2, 3],
data={"class_name": ["car", "car"]}
),
{"frame_number": 43},
"test_detections.csv",
[
[
"x_min",
"y_min",
"x_max",
"y_max",
"class_id",
"confidence",
"tracker_id",
"class_name",
"frame_number",
],
["10.0", "20.0", "30.0", "40.0", "0", "0.7", "0", "person", "42"],
["50.0", "60.0", "70.0", "80.0", "0", "0.8", "1", "person", "42"],
["15.0", "25.0", "35.0", "45.0", "1", "0.6", "2", "car", "43"],
["55.0", "65.0", "75.0", "85.0", "1", "0.9", "3", "car", "43"],
]
), # multiple detections
(
mock_detections(
xyxy=[[60, 70, 80, 90], [100, 110, 120, 130]],
tracker_id=[4, 5],
data={"class_name": ["bike", "dog"]}
),
{"frame_number": 44},
mock_detections(
xyxy=[[65, 75, 85, 95], [105, 115, 125, 135]],
confidence=[0.5, 0.4],
data={"class_name": ["tree", "cat"]}
),
{"frame_number": 45},
"test_detections_missing_fields.csv",
[
[
"x_min",
"y_min",
"x_max",
"y_max",
"class_id",
"confidence",
"tracker_id",
"class_name",
"frame_number",
],
["60.0", "70.0", "80.0", "90.0", "", "", "4", "bike", "44"],
["100.0", "110.0", "120.0", "130.0", "", "", "5", "dog", "44"],
["65.0", "75.0", "85.0", "95.0", "", "0.5", "", "tree", "45"],
["105.0", "115.0", "125.0", "135.0", "", "0.4", "", "cat", "45"],
]
), # missing fields
(
mock_detections(
xyxy=[[10, 11, 12, 13]],
confidence=[0.95],
data={"class_name": "unknown", "is_detected": True, "score": 1}
),
{"frame_number": 46},
mock_detections(
xyxy=[[14, 15, 16, 17]],
data={"class_name": "artifact", "is_detected": False, "score": 0.85}
),
{"frame_number": 47},
"test_detections_varied_data.csv",
[
[
"x_min",
"y_min",
"x_max",
"y_max",
"class_id",
"confidence",
"tracker_id",
"class_name",
"frame_number",
"is_detected",
"score",
],
["10.0", "11.0", "12.0", "13.0", "", "0.95", "", "unknown", "46", "True", "1"],
["14.0", "15.0", "16.0", "17.0", "", "", "", "artifact", "47", "False", "0.85"],
]
), # Inconsistent Data Types
(
mock_detections(
xyxy=[[20, 21, 22, 23]],
),
{"metadata": {"sensor_id": 101, "location": "north"}, "tags": ["urgent", "review"]},
mock_detections(
xyxy=[[14, 15, 16, 17]],
),
{"metadata": {"sensor_id": 104, "location": "west"}, "tags": ["not-urgent", "done"]},
"test_detections_complex_data.csv",
[
[
"x_min",
"y_min",
"x_max",
"y_max",
"class_id",
"confidence",
"tracker_id",
"metadata",
"tags",
],
["20.0", "21.0", "22.0", "23.0", "", "", "", "{'sensor_id': 101, 'location': 'north'}", "['urgent', 'review']"],
["14.0", "15.0", "16.0", "17.0", "", "", "", "{'sensor_id': 104, 'location': 'west'}", "['not-urgent', 'done']"],
]
), # Complex Data
],
)
second_detections = Detections(
xyxy=np.array([[15, 25, 35, 45], [55, 65, 75, 85]]),
confidence=np.array([0.6, 0.9]),
class_id=np.array([1, 1]),
tracker_id=np.array([2, 3]),
data={"class_name": np.array(["car", "car"])},
)
def test_csv_sink(
detections: mock_detections,
custom_data: Dict[str, Any],
second_detections: mock_detections,
second_custom_data: Dict[str, Any],
file_name: str,
expected_result: List[List[Any]]
) -> None:
custom_data = {"frame_number": 42}
second_custom_data = {"frame_number": 43}
return detections, custom_data, second_detections, second_custom_data
def test_csv_sink(detection_instances):
detections, custom_data, second_detections, second_custom_data = detection_instances
csv_filename = "test_detections.csv"
expected_rows = [
[
"x_min",
"y_min",
"x_max",
"y_max",
"class_id",
"confidence",
"tracker_id",
"class_name",
"frame_number",
],
[10, 20, 30, 40, 0, 0.7, 0, "person", 42],
[50, 60, 70, 80, 0, 0.8, 1, "person", 42],
[15, 25, 35, 45, 1, 0.6, 2, "car", 43],
[55, 65, 75, 85, 1, 0.9, 3, "car", 43],
]
# Using the CSVSink class to write the detection data to a CSV file
with sv.CSVSink(filename=csv_filename) as sink:
with sv.CSVSink(file_name) as sink:
sink.append(detections, custom_data)
sink.append(second_detections, second_custom_data)
# Read back the CSV file and verify its contents
with open(csv_filename, mode="r", newline="") as file:
reader = csv.reader(file)
for i, row in enumerate(reader):
assert (
[str(item) for item in expected_rows[i]] == row
), f"Row in CSV didn't match expected output: {row} != {expected_rows[i]}"
assert_csv_equal(file_name, expected_result)
# Clean up by removing the test CSV file
os.remove(csv_filename)
@pytest.mark.parametrize(
"detections, custom_data, second_detections, second_custom_data, file_name, expected_result",
[
(
mock_detections(
xyxy=[[10, 20, 30, 40], [50, 60, 70, 80]],
confidence=[0.7, 0.8],
class_id=[0, 0],
tracker_id=[0, 1],
data={"class_name": ["person", "person"]}
),
{"frame_number": 42},
mock_detections(
xyxy=[[15, 25, 35, 45], [55, 65, 75, 85]],
confidence=[0.6, 0.9],
class_id=[1, 1],
tracker_id=[2, 3],
data={"class_name": ["car", "car"]}
),
{"frame_number": 43},
"test_detections.csv",
[
[
"x_min",
"y_min",
"x_max",
"y_max",
"class_id",
"confidence",
"tracker_id",
"class_name",
"frame_number",
],
["10.0", "20.0", "30.0", "40.0", "0", "0.7", "0", "person", "42"],
["50.0", "60.0", "70.0", "80.0", "0", "0.8", "1", "person", "42"],
["15.0", "25.0", "35.0", "45.0", "1", "0.6", "2", "car", "43"],
["55.0", "65.0", "75.0", "85.0", "1", "0.9", "3", "car", "43"],
]
), # multiple detections
(
mock_detections(
xyxy=[[60, 70, 80, 90], [100, 110, 120, 130]],
tracker_id=[4, 5],
data={"class_name": ["bike", "dog"]}
),
{"frame_number": 44},
mock_detections(
xyxy=[[65, 75, 85, 95], [105, 115, 125, 135]],
confidence=[0.5, 0.4],
data={"class_name": ["tree", "cat"]}
),
{"frame_number": 45},
"test_detections_missing_fields.csv",
[
[
"x_min",
"y_min",
"x_max",
"y_max",
"class_id",
"confidence",
"tracker_id",
"class_name",
"frame_number",
],
["60.0", "70.0", "80.0", "90.0", "", "", "4", "bike", "44"],
["100.0", "110.0", "120.0", "130.0", "", "", "5", "dog", "44"],
["65.0", "75.0", "85.0", "95.0", "", "0.5", "", "tree", "45"],
["105.0", "115.0", "125.0", "135.0", "", "0.4", "", "cat", "45"],
]
), # missing fields
(
mock_detections(
xyxy=[[10, 11, 12, 13]],
confidence=[0.95],
data={"class_name": "unknown", "is_detected": True, "score": 1}
),
{"frame_number": 46},
mock_detections(
xyxy=[[14, 15, 16, 17]],
data={"class_name": "artifact", "is_detected": False, "score": 0.85}
),
{"frame_number": 47},
"test_detections_varied_data.csv",
[
[
"x_min",
"y_min",
"x_max",
"y_max",
"class_id",
"confidence",
"tracker_id",
"class_name",
"frame_number",
"is_detected",
"score",
],
["10.0", "11.0", "12.0", "13.0", "", "0.95", "", "unknown", "46", "True", "1"],
["14.0", "15.0", "16.0", "17.0", "", "", "", "artifact", "47", "False", "0.85"],
]
), # Inconsistent Data Types
(
mock_detections(
xyxy=[[20, 21, 22, 23]],
),
{"metadata": {"sensor_id": 101, "location": "north"}, "tags": ["urgent", "review"]},
mock_detections(
xyxy=[[14, 15, 16, 17]],
),
{"metadata": {"sensor_id": 104, "location": "west"}, "tags": ["not-urgent", "done"]},
"test_detections_complex_data.csv",
[
[
"x_min",
"y_min",
"x_max",
"y_max",
"class_id",
"confidence",
"tracker_id",
"metadata",
"tags",
],
["20.0", "21.0", "22.0", "23.0", "", "", "", "{'sensor_id': 101, 'location': 'north'}", "['urgent', 'review']"],
["14.0", "15.0", "16.0", "17.0", "", "", "", "{'sensor_id': 104, 'location': 'west'}", "['not-urgent', 'done']"],
]
), # Complex Data
],
)
def test_csv_sink_manual(detection_instances):
detections, custom_data, second_detections, second_custom_data = detection_instances
csv_filename = "test_detections.csv"
expected_rows = [
[
"x_min",
"y_min",
"x_max",
"y_max",
"class_id",
"confidence",
"tracker_id",
"class_name",
"frame_number",
],
[10, 20, 30, 40, 0, 0.7, 0, "person", 42],
[50, 60, 70, 80, 0, 0.8, 1, "person", 42],
[15, 25, 35, 45, 1, 0.6, 2, "car", 43],
[55, 65, 75, 85, 1, 0.9, 3, "car", 43],
]
# Using the CSVSink class to write the detection data to a CSV file
sink = sv.CSVSink(filename=csv_filename)
def test_csv_sink_manual(
detections: mock_detections,
custom_data: Dict[str, Any],
second_detections: mock_detections,
second_custom_data: Dict[str, Any],
file_name: str,
expected_result: List[List[Any]],
) -> None:
sink = sv.CSVSink(file_name)
sink.open()
sink.append(detections, custom_data)
sink.append(second_detections, second_custom_data)
sink.close()
# Read back the CSV file and verify its contents
with open(csv_filename, mode="r", newline="") as file:
assert_csv_equal(file_name, expected_result)
def assert_csv_equal(file_name, expected_rows):
with open(file_name, mode="r", newline="") as file:
reader = csv.reader(file)
for i, row in enumerate(reader):
assert (
[str(item) for item in expected_rows[i]] == row
), f"Row in CSV didn't match expected output: {row} != {expected_rows[i]}"
# Clean up by removing the test CSV file
os.remove(csv_filename)
#os.remove(file_name)