416 lines
13 KiB
Python
416 lines
13 KiB
Python
import csv
|
|
import os
|
|
from test.test_utils import mock_detections
|
|
from typing import Any, Dict, List
|
|
|
|
import pytest
|
|
|
|
import supervision as sv
|
|
|
|
|
|
@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(
|
|
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:
|
|
with sv.CSVSink(file_name) as sink:
|
|
sink.append(detections, custom_data)
|
|
sink.append(second_detections, second_custom_data)
|
|
|
|
assert_csv_equal(file_name, expected_result)
|
|
|
|
|
|
@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(
|
|
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()
|
|
|
|
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]}"
|
|
|
|
os.remove(file_name)
|