use `CLASS_NAME_DATA_FIELD` data constant over 'class_name' string
This commit is contained in:
parent
0d10704370
commit
f4d2f2a900
|
|
@ -0,0 +1 @@
|
|||
CLASS_NAME_DATA_FIELD = 'class_name'
|
||||
|
|
@ -4,6 +4,8 @@ from typing import Any, Dict, List, Optional, Tuple, Union
|
|||
import cv2
|
||||
import numpy as np
|
||||
|
||||
from supervision.config import CLASS_NAME_DATA_FIELD
|
||||
|
||||
MIN_POLYGON_POINT_COUNT = 3
|
||||
|
||||
|
||||
|
|
@ -336,7 +338,7 @@ def process_roboflow_result(
|
|||
np.ndarray,
|
||||
np.ndarray,
|
||||
Optional[np.ndarray],
|
||||
np.ndarray,
|
||||
Optional[np.ndarray],
|
||||
Dict[str, List[np.ndarray]],
|
||||
]:
|
||||
if not roboflow_result["predictions"]:
|
||||
|
|
@ -346,7 +348,7 @@ def process_roboflow_result(
|
|||
np.empty(0),
|
||||
None,
|
||||
None,
|
||||
{"class_name": np.empty(0)},
|
||||
{CLASS_NAME_DATA_FIELD: np.empty(0)},
|
||||
)
|
||||
|
||||
xyxy = []
|
||||
|
|
@ -395,7 +397,7 @@ def process_roboflow_result(
|
|||
class_name = np.array(class_name) if len(class_name) > 0 else np.empty(0)
|
||||
masks = np.array(masks, dtype=bool) if len(masks) > 0 else None
|
||||
tracker_id = np.array(tracker_ids).astype(int) if len(tracker_ids) > 0 else None
|
||||
data = {"class_name": class_name}
|
||||
data = {CLASS_NAME_DATA_FIELD: class_name}
|
||||
|
||||
return xyxy, confidence, class_id, masks, tracker_id, data
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from typing import Any, Dict, List, Optional, Tuple
|
|||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from supervision.config import CLASS_NAME_DATA_FIELD
|
||||
from supervision.detection.utils import (
|
||||
calculate_masks_centroids,
|
||||
clip_boxes,
|
||||
|
|
@ -274,7 +275,7 @@ def test_filter_polygons_by_area(
|
|||
np.empty(0),
|
||||
None,
|
||||
None,
|
||||
{"class_name": np.empty(0)},
|
||||
{CLASS_NAME_DATA_FIELD: np.empty(0)},
|
||||
),
|
||||
DoesNotRaise(),
|
||||
), # empty result
|
||||
|
|
@ -299,7 +300,7 @@ def test_filter_polygons_by_area(
|
|||
np.array([0]),
|
||||
None,
|
||||
None,
|
||||
{"class_name": np.array(["person"])},
|
||||
{CLASS_NAME_DATA_FIELD: np.array(["person"])},
|
||||
),
|
||||
DoesNotRaise(),
|
||||
), # single correct object detection result
|
||||
|
|
@ -335,7 +336,7 @@ def test_filter_polygons_by_area(
|
|||
np.array([0, 7]),
|
||||
None,
|
||||
np.array([1, 2]),
|
||||
{"class_name": np.array(["person", "truck"])},
|
||||
{CLASS_NAME_DATA_FIELD: np.array(["person", "truck"])},
|
||||
),
|
||||
DoesNotRaise(),
|
||||
), # two correct object detection result
|
||||
|
|
@ -362,7 +363,7 @@ def test_filter_polygons_by_area(
|
|||
np.empty(0),
|
||||
None,
|
||||
None,
|
||||
{"class_name": np.empty(0)},
|
||||
{CLASS_NAME_DATA_FIELD: np.empty(0)},
|
||||
),
|
||||
DoesNotRaise(),
|
||||
), # single incorrect instance segmentation result with no points
|
||||
|
|
@ -388,7 +389,7 @@ def test_filter_polygons_by_area(
|
|||
np.empty(0),
|
||||
None,
|
||||
None,
|
||||
{"class_name": np.empty(0)},
|
||||
{CLASS_NAME_DATA_FIELD: np.empty(0)},
|
||||
),
|
||||
DoesNotRaise(),
|
||||
), # single incorrect instance segmentation result with no enough points
|
||||
|
|
@ -419,7 +420,7 @@ def test_filter_polygons_by_area(
|
|||
np.array([0]),
|
||||
TEST_MASK,
|
||||
None,
|
||||
{"class_name": np.array(["person"])},
|
||||
{CLASS_NAME_DATA_FIELD: np.array(["person"])},
|
||||
),
|
||||
DoesNotRaise(),
|
||||
), # single incorrect instance segmentation result with no enough points
|
||||
|
|
@ -460,7 +461,7 @@ def test_filter_polygons_by_area(
|
|||
np.array([0]),
|
||||
TEST_MASK,
|
||||
None,
|
||||
{"class_name": np.array(["person"])},
|
||||
{CLASS_NAME_DATA_FIELD: np.array(["person"])},
|
||||
),
|
||||
DoesNotRaise(),
|
||||
), # two instance segmentation results - one correct, one incorrect
|
||||
|
|
|
|||
Loading…
Reference in New Issue