Merge branch 'develop' into feat/upgrade-python-3-9

This commit is contained in:
Onuralp SEZER 2025-07-15 15:27:05 +03:00
commit 3dd1ce4e1c
No known key found for this signature in database
GPG Key ID: CF0835DFDF14CA38
7 changed files with 511 additions and 36 deletions

View File

@ -27,3 +27,9 @@ comments: true
</div>
:::supervision.detection.overlap_filter.box_non_max_merge
<div class="md-typeset">
<h2><a href="#supervision.detection.overlap_filter.mask_non_max_merge">mask_non_max_merge</a></h2>
</div>
:::supervision.detection.overlap_filter.mask_non_max_merge

View File

@ -5,6 +5,12 @@ status: new
# Detection Utils
<div class="md-typeset">
<h2><a href="#supervision.detection.overlap_filter.OverlapMetric">OverlapMetric</a></h2>
</div>
:::supervision.detection.overlap_filter.OverlapMetric
<div class="md-typeset">
<h2><a href="#supervision.detection.utils.box_iou">box_iou</a></h2>
</div>

View File

@ -3,6 +3,7 @@ from __future__ import annotations
from collections.abc import Iterator
from dataclasses import dataclass, field
from enum import Enum
from functools import reduce
from typing import Any
import numpy as np
@ -14,6 +15,7 @@ from supervision.config import (
from supervision.detection.overlap_filter import (
box_non_max_merge,
box_non_max_suppression,
mask_non_max_merge,
mask_non_max_suppression,
)
from supervision.detection.tools.transformers import (
@ -22,12 +24,14 @@ from supervision.detection.tools.transformers import (
process_transformers_v5_segmentation_result,
)
from supervision.detection.utils import (
OverlapMetric,
box_iou_batch,
calculate_masks_centroids,
extract_ultralytics_masks,
get_data_item,
is_data_equal,
is_metadata_equal,
mask_iou_batch,
mask_to_xyxy,
merge_data,
merge_metadata,
@ -1788,7 +1792,10 @@ class Detections:
return (self.xyxy[:, 3] - self.xyxy[:, 1]) * (self.xyxy[:, 2] - self.xyxy[:, 0])
def with_nms(
self, threshold: float = 0.5, class_agnostic: bool = False
self,
threshold: float = 0.5,
class_agnostic: bool = False,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
) -> Detections:
"""
Performs non-max suppression on detection set. If the detections result
@ -1801,6 +1808,8 @@ class Detections:
class_agnostic (bool): Whether to perform class-agnostic
non-maximum suppression. If True, the class_id of each detection
will be ignored. Defaults to False.
overlap_metric (OverlapMetric): Metric used for measuring overlap between
detections in slices.
Returns:
Detections: A new Detections object containing the subset of detections
@ -1834,17 +1843,25 @@ class Detections:
if self.mask is not None:
indices = mask_non_max_suppression(
predictions=predictions, masks=self.mask, iou_threshold=threshold
predictions=predictions,
masks=self.mask,
iou_threshold=threshold,
overlap_metric=overlap_metric,
)
else:
indices = box_non_max_suppression(
predictions=predictions, iou_threshold=threshold
predictions=predictions,
iou_threshold=threshold,
overlap_metric=overlap_metric,
)
return self[indices]
def with_nmm(
self, threshold: float = 0.5, class_agnostic: bool = False
self,
threshold: float = 0.5,
class_agnostic: bool = False,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
) -> Detections:
"""
Perform non-maximum merging on the current set of object detections.
@ -1855,6 +1872,8 @@ class Detections:
class_agnostic (bool): Whether to perform class-agnostic
non-maximum merging. If True, the class_id of each detection
will be ignored. Defaults to False.
overlap_metric (OverlapMetric): Metric used for measuring overlap between
detections in slices.
Returns:
Detections: A new Detections object containing the subset of detections
@ -1888,15 +1907,25 @@ class Detections:
)
)
merge_groups = box_non_max_merge(
predictions=predictions, iou_threshold=threshold
)
if self.mask is not None:
merge_groups = mask_non_max_merge(
predictions=predictions,
masks=self.mask,
iou_threshold=threshold,
overlap_metric=overlap_metric,
)
else:
merge_groups = box_non_max_merge(
predictions=predictions,
iou_threshold=threshold,
overlap_metric=overlap_metric,
)
result = []
for merge_group in merge_groups:
unmerged_detections = [self[i] for i in merge_group]
merged_detections = merge_inner_detections_objects(
unmerged_detections, threshold
merged_detections = merge_inner_detections_objects_without_iou(
unmerged_detections
)
result.append(merged_detections)
@ -1996,7 +2025,9 @@ def merge_inner_detection_object_pair(
def merge_inner_detections_objects(
detections: list[Detections], threshold=0.5
detections: list[Detections],
threshold=0.5,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
) -> Detections:
"""
Given N detections each of length 1 (exactly one object inside), combine them into a
@ -2008,13 +2039,32 @@ def merge_inner_detections_objects(
"""
detections_1 = detections[0]
for detections_2 in detections[1:]:
box_iou = box_iou_batch(detections_1.xyxy, detections_2.xyxy)[0]
if box_iou < threshold:
if detections_1.mask is not None and detections_2.mask is not None:
iou = mask_iou_batch(detections_1.mask, detections_2.mask, overlap_metric)[
0
]
else:
iou = box_iou_batch(detections_1.xyxy, detections_2.xyxy, overlap_metric)[0]
if iou < threshold:
break
detections_1 = merge_inner_detection_object_pair(detections_1, detections_2)
return detections_1
def merge_inner_detections_objects_without_iou(
detections: List[Detections],
) -> Detections:
"""
Given N detections each of length 1 (exactly one object inside), combine them into a
single detection object of length 1. The contained inner object will be the merged
result of all the input detections.
For example, this lets you merge N boxes into one big box, N masks into one mask,
etc.
"""
return reduce(merge_inner_detection_object_pair, detections)
def validate_fields_both_defined_or_none(
detections_1: Detections, detections_2: Detections
) -> None:

View File

@ -5,7 +5,7 @@ from enum import Enum
import numpy as np
import numpy.typing as npt
from supervision.detection.utils import box_iou_batch, mask_iou_batch
from supervision.detection.utils import OverlapMetric, box_iou_batch, mask_iou_batch
def resize_masks(masks: np.ndarray, max_dimension: int = 640) -> np.ndarray:
@ -41,6 +41,7 @@ def mask_non_max_suppression(
predictions: np.ndarray,
masks: np.ndarray,
iou_threshold: float = 0.5,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
mask_dimension: int = 640,
) -> np.ndarray:
"""
@ -56,6 +57,7 @@ def mask_non_max_suppression(
dimensions of each mask.
iou_threshold (float): The intersection-over-union threshold
to use for non-maximum suppression.
overlap_metric (OverlapMetric): Metric used for matching detections in slices.
mask_dimension (int): The dimension to which the masks should be
resized before computing IOU values. Defaults to 640.
@ -80,7 +82,7 @@ def mask_non_max_suppression(
predictions = predictions[sort_index]
masks = masks[sort_index]
masks_resized = resize_masks(masks, mask_dimension)
ious = mask_iou_batch(masks_resized, masks_resized)
ious = mask_iou_batch(masks_resized, masks_resized, overlap_metric)
categories = predictions[:, 5]
keep = np.ones(rows, dtype=bool)
@ -93,7 +95,9 @@ def mask_non_max_suppression(
def box_non_max_suppression(
predictions: np.ndarray, iou_threshold: float = 0.5
predictions: np.ndarray,
iou_threshold: float = 0.5,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
) -> np.ndarray:
"""
Perform Non-Maximum Suppression (NMS) on object detection predictions.
@ -104,6 +108,7 @@ def box_non_max_suppression(
or `(x_min, y_min, x_max, y_max, score, class)`.
iou_threshold (float): The intersection-over-union threshold
to use for non-maximum suppression.
overlap_metric (OverlapMetric): Metric used for matching detections in slices.
Returns:
np.ndarray: A boolean array indicating which predictions to keep after n
@ -129,7 +134,7 @@ def box_non_max_suppression(
boxes = predictions[:, :4]
categories = predictions[:, 5]
ious = box_iou_batch(boxes, boxes)
ious = box_iou_batch(boxes, boxes, overlap_metric)
ious = ious - np.eye(rows)
keep = np.ones(rows, dtype=bool)
@ -147,8 +152,11 @@ def box_non_max_suppression(
def group_overlapping_boxes(
predictions: npt.NDArray[np.float64], iou_threshold: float = 0.5
predictions: npt.NDArray[np.float64],
iou_threshold: float = 0.5,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
) -> list[list[int]]:
"""
Apply greedy version of non-maximum merging to avoid detecting too many
overlapping bounding boxes for a given object.
@ -159,6 +167,7 @@ def group_overlapping_boxes(
and the confidence scores.
iou_threshold (float): The intersection-over-union threshold
to use for non-maximum suppression. Defaults to 0.5.
overlap_metric (OverlapMetric): Metric used for matching detections in slices.
Returns:
List[List[int]]: Groups of prediction indices be merged.
@ -178,7 +187,9 @@ def group_overlapping_boxes(
break
merge_candidate = np.expand_dims(predictions[idx], axis=0)
ious = box_iou_batch(predictions[order][:, :4], merge_candidate[:, :4])
ious = box_iou_batch(
predictions[order][:, :4], merge_candidate[:, :4], overlap_metric
)
ious = ious.flatten()
above_threshold = ious >= iou_threshold
@ -188,9 +199,70 @@ def group_overlapping_boxes(
return merge_groups
def mask_non_max_merge(
predictions: np.ndarray,
masks: np.ndarray,
iou_threshold: float = 0.5,
mask_dimension: int = 640,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
) -> List[List[int]]:
"""
Perform Non-Maximum Merging (NMM) on segmentation predictions.
Args:
predictions (np.ndarray): A 2D array of object detection predictions in
the format of `(x_min, y_min, x_max, y_max, score)`
or `(x_min, y_min, x_max, y_max, score, class)`. Shape: `(N, 5)` or
`(N, 6)`, where N is the number of predictions.
masks (np.ndarray): A 3D array of binary masks corresponding to the predictions.
Shape: `(N, H, W)`, where N is the number of predictions, and H, W are the
dimensions of each mask.
iou_threshold (float): The intersection-over-union threshold
to use for non-maximum suppression.
mask_dimension (int): The dimension to which the masks should be
resized before computing IOU values. Defaults to 640.
overlap_metric (OverlapMetric): Metric used for matching detections in slices.
Returns:
np.ndarray: A boolean array indicating which predictions to keep after
non-maximum suppression.
Raises:
AssertionError: If `iou_threshold` is not within the closed
range from `0` to `1`.
"""
masks_resized = resize_masks(masks, mask_dimension)
if predictions.shape[1] == 5:
return group_overlapping_masks(
predictions, masks_resized, iou_threshold, overlap_metric
)
category_ids = predictions[:, 5]
merge_groups = []
for category_id in np.unique(category_ids):
curr_indices = np.where(category_ids == category_id)[0]
merge_class_groups = group_overlapping_masks(
predictions[curr_indices],
masks_resized[curr_indices],
iou_threshold,
overlap_metric,
)
for merge_class_group in merge_class_groups:
merge_groups.append(curr_indices[merge_class_group].tolist())
for merge_group in merge_groups:
if len(merge_group) == 0:
raise ValueError(
f"Empty group detected when non-max-merging detections: {merge_groups}"
)
return merge_groups
def box_non_max_merge(
predictions: npt.NDArray[np.float64],
iou_threshold: float = 0.5,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
) -> list[list[int]]:
"""
Apply greedy version of non-maximum merging per category to avoid detecting
@ -203,20 +275,21 @@ def box_non_max_merge(
detections of different classes to be merged.
iou_threshold (float): The intersection-over-union threshold
to use for non-maximum suppression. Defaults to 0.5.
overlap_metric (OverlapMetric): Metric used for matching detections in slices.
Returns:
List[List[int]]: Groups of prediction indices be merged.
Each group may have 1 or more elements.
"""
if predictions.shape[1] == 5:
return group_overlapping_boxes(predictions, iou_threshold)
return group_overlapping_boxes(predictions, iou_threshold, overlap_metric)
category_ids = predictions[:, 5]
merge_groups = []
for category_id in np.unique(category_ids):
curr_indices = np.where(category_ids == category_id)[0]
merge_class_groups = group_overlapping_boxes(
predictions[curr_indices], iou_threshold
predictions[curr_indices], iou_threshold, overlap_metric
)
for merge_class_group in merge_class_groups:
@ -230,6 +303,62 @@ def box_non_max_merge(
return merge_groups
def group_overlapping_masks(
predictions: npt.NDArray[np.float64],
masks: npt.NDArray[np.float64],
iou_threshold: float = 0.5,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
) -> List[List[int]]:
"""
Apply greedy version of non-maximum merging to avoid detecting too many
Args:
predictions (npt.NDArray[np.float64]): An array of shape `(n, 5)` containing
the bounding boxes coordinates in format `[x1, y1, x2, y2]`
and the confidence scores.
masks (npt.NDArray[np.float64]): A 3D array of binary masks corresponding to
the predictions.
iou_threshold (float): The intersection-over-union threshold
to use for non-maximum suppression. Defaults to 0.5.
overlap_metric (OverlapMetric): Metric used for matching detections in slices.
Returns:
List[List[int]]: Groups of prediction indices be merged.
Each group may have 1 or more elements.
"""
merge_groups: List[List[int]] = []
scores = predictions[:, 4]
order = scores.argsort()
while len(order) > 0:
idx = int(order[-1])
order = order[:-1]
if len(order) == 0:
merge_groups.append([idx])
break
merge_candidate = masks[idx][None, ...]
candidate_groups = [idx]
while len(order) > 0:
ious = mask_iou_batch(masks[order], merge_candidate, overlap_metric)
above_threshold: np.ndarray = ious.flatten() >= iou_threshold
if not above_threshold.any():
break
above_idx = order[above_threshold]
merge_candidate = np.logical_or.reduce(
np.concatenate([masks[above_idx], merge_candidate]),
axis=0,
keepdims=True,
)
candidate_groups.extend(np.flip(above_idx).tolist())
order = order[~above_threshold]
merge_groups.append(candidate_groups)
return merge_groups
class OverlapFilter(Enum):
"""
Enum specifying the strategy for filtering overlapping detections.

View File

@ -72,6 +72,8 @@ class InferenceSlicer:
filtering or merging overlapping detections in slices.
iou_threshold (float): Intersection over Union (IoU) threshold
used when filtering by overlap.
match_metric (str): Metric used for matching detections in slices.
"IOU" or "IOS". Defaults "IOU".
callback (Callable): A function that performs inference on a given image
slice and returns detections.
thread_workers (int): Number of threads for parallel execution.
@ -91,6 +93,7 @@ class InferenceSlicer:
overlap_wh: Optional[tuple[int, int]] = None,
overlap_filter: Union[OverlapFilter, str] = OverlapFilter.NON_MAX_SUPPRESSION,
iou_threshold: float = 0.5,
match_metric: str = "IOU",
thread_workers: int = 1,
):
if overlap_ratio_wh is not None:
@ -106,6 +109,7 @@ class InferenceSlicer:
self.slice_wh = slice_wh
self.iou_threshold = iou_threshold
self.match_metric = match_metric
self.overlap_filter = OverlapFilter.from_value(overlap_filter)
self.callback = callback
self.thread_workers = thread_workers
@ -165,9 +169,13 @@ class InferenceSlicer:
if self.overlap_filter == OverlapFilter.NONE:
return merged
elif self.overlap_filter == OverlapFilter.NON_MAX_SUPPRESSION:
return merged.with_nms(threshold=self.iou_threshold)
return merged.with_nms(
threshold=self.iou_threshold, match_metric=self.match_metric
)
elif self.overlap_filter == OverlapFilter.NON_MAX_MERGE:
return merged.with_nmm(threshold=self.iou_threshold)
return merged.with_nmm(
threshold=self.iou_threshold, match_metric=self.match_metric
)
else:
warnings.warn(
f"Invalid overlap filter strategy: {self.overlap_filter}",

View File

@ -1,3 +1,6 @@
from __future__ import annotations
from enum import Enum
from itertools import chain
from typing import Any, Optional, Union
@ -11,6 +14,42 @@ from supervision.geometry.core import Vector
MIN_POLYGON_POINT_COUNT = 3
class OverlapMetric(Enum):
"""
Enum specifying the metric for measuring overlap between detections.
Attributes:
IOU: Intersection over Union. A region-overlap metric that compares
two shapes (usually bounding boxes or masks) by normalising the
shared area with the area of their union.
IOS: Intersection over Smaller, a region-overlap metric that compares
two shapes (usually bounding boxes or masks) by normalising the
shared area with the smaller of the two shapes.
"""
IOU = "IOU"
IOS = "IOS"
@classmethod
def list(cls):
return list(map(lambda c: c.value, cls))
@classmethod
def from_value(cls, value: Union[OverlapMetric, str]) -> OverlapMetric:
if isinstance(value, cls):
return value
if isinstance(value, str):
value = value.lower()
try:
return cls(value)
except ValueError:
raise ValueError(f"Invalid value: {value}. Must be one of {cls.list()}")
raise ValueError(
f"Invalid value type: {type(value)}. Must be an instance of "
f"{cls.__name__} or str."
)
def xyxy_to_polygons(box: np.ndarray) -> np.ndarray:
"""
Convert an array of boxes to an array of polygons.
@ -108,7 +147,11 @@ def box_iou(
return inter_area / union_area + 1e-6
def box_iou_batch(boxes_true: np.ndarray, boxes_detection: np.ndarray) -> np.ndarray:
def box_iou_batch(
boxes_true: np.ndarray,
boxes_detection: np.ndarray,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
) -> np.ndarray:
"""
Compute Intersection over Union (IoU) of two sets of bounding boxes -
`boxes_true` and `boxes_detection`. Both sets
@ -124,6 +167,7 @@ def box_iou_batch(boxes_true: np.ndarray, boxes_detection: np.ndarray) -> np.nda
`shape = (N, 4)` where `N` is number of true objects.
boxes_detection (np.ndarray): 2D `np.ndarray` representing detection boxes.
`shape = (M, 4)` where `M` is number of detected objects.
overlap_metric (OverlapMetric): Metric used for matching detections in slices.
Returns:
np.ndarray: Pairwise IoU of boxes from `boxes_true` and `boxes_detection`.
@ -162,13 +206,37 @@ def box_iou_batch(boxes_true: np.ndarray, boxes_detection: np.ndarray) -> np.nda
bottom_right = np.minimum(boxes_true[:, None, 2:], boxes_detection[:, 2:])
area_inter = np.prod(np.clip(bottom_right - top_left, a_min=0, a_max=None), 2)
ious = area_inter / (area_true[:, None] + area_detection - area_inter)
if overlap_metric == OverlapMetric.IOU:
union_area = area_true[:, None] + area_detection - area_inter
ious = np.divide(
area_inter,
union_area,
out=np.zeros_like(area_inter, dtype=float),
where=union_area != 0,
)
elif overlap_metric == OverlapMetric.IOS:
small_area = np.minimum(area_true[:, None], area_detection)
ious = np.divide(
area_inter,
small_area,
out=np.zeros_like(area_inter, dtype=float),
where=small_area != 0,
)
else:
raise ValueError(
f"overlap_metric {overlap_metric} is not supported, "
"only 'IOU' and 'IOS' are supported"
)
ious = np.nan_to_num(ious)
return ious
def _mask_iou_batch_split(
masks_true: np.ndarray, masks_detection: np.ndarray
masks_true: np.ndarray,
masks_detection: np.ndarray,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
) -> np.ndarray:
"""
Internal function.
@ -178,6 +246,7 @@ def _mask_iou_batch_split(
Args:
masks_true (np.ndarray): 3D `np.ndarray` representing ground-truth masks.
masks_detection (np.ndarray): 3D `np.ndarray` representing detection masks.
overlap_metric (OverlapMetric): Metric used for matching detections in slices.
Returns:
np.ndarray: Pairwise IoU of masks from `masks_true` and `masks_detection`.
@ -186,21 +255,40 @@ def _mask_iou_batch_split(
axis=(2, 3)
)
masks_true_area = masks_true.sum(axis=(1, 2))
masks_detection_area = masks_detection.sum(axis=(1, 2))
union_area = masks_true_area[:, None] + masks_detection_area - intersection_area
masks_true_area = masks_true.sum(axis=(1, 2)) # (area1, area2, ...)
masks_detection_area = masks_detection.sum(axis=(1, 2)) # (area1)
return np.divide(
intersection_area,
union_area,
out=np.zeros_like(intersection_area, dtype=float),
where=union_area != 0,
)
if overlap_metric == OverlapMetric.IOU:
union_area = masks_true_area[:, None] + masks_detection_area - intersection_area
ious = np.divide(
intersection_area,
union_area,
out=np.zeros_like(intersection_area, dtype=float),
where=union_area != 0,
)
elif overlap_metric == OverlapMetric.IOS:
# ios = intersection_area / min(area1, area2)
small_area = np.minimum(masks_true_area[:, None], masks_detection_area)
ious = np.divide(
intersection_area,
small_area,
out=np.zeros_like(intersection_area, dtype=float),
where=small_area != 0,
)
else:
raise ValueError(
f"overlap_metric {overlap_metric} is not supported, "
"only 'IOU' and 'IOS' are supported"
)
ious = np.nan_to_num(ious)
return ious
def mask_iou_batch(
masks_true: np.ndarray,
masks_detection: np.ndarray,
overlap_metric: OverlapMetric = OverlapMetric.IOU,
memory_limit: int = 1024 * 5,
) -> np.ndarray:
"""
@ -210,6 +298,7 @@ def mask_iou_batch(
Args:
masks_true (np.ndarray): 3D `np.ndarray` representing ground-truth masks.
masks_detection (np.ndarray): 3D `np.ndarray` representing detection masks.
overlap_metric (OverlapMetric): Metric used for matching detections in slices.
memory_limit (int): memory limit in MB, default is 1024 * 5 MB (5GB).
Returns:
@ -224,7 +313,7 @@ def mask_iou_batch(
/ 1024
)
if memory <= memory_limit:
return _mask_iou_batch_split(masks_true, masks_detection)
return _mask_iou_batch_split(masks_true, masks_detection, overlap_metric)
ious = []
step = max(
@ -239,7 +328,11 @@ def mask_iou_batch(
1,
)
for i in range(0, masks_true.shape[0], step):
ious.append(_mask_iou_batch_split(masks_true[i : i + step], masks_detection))
ious.append(
_mask_iou_batch_split(
masks_true[i : i + step], masks_detection, overlap_metric
)
)
return np.vstack(ious)

View File

@ -7,6 +7,7 @@ import pytest
from supervision.detection.overlap_filter import (
box_non_max_suppression,
group_overlapping_boxes,
mask_non_max_merge,
mask_non_max_suppression,
)
@ -447,3 +448,185 @@ def test_mask_non_max_suppression(
predictions=predictions, masks=masks, iou_threshold=iou_threshold
)
assert np.array_equal(result, expected_result)
@pytest.mark.parametrize(
"predictions, masks, iou_threshold, expected_result, exception",
[
(
np.empty((0, 6)),
np.empty((0, 5, 5)),
0.5,
[],
DoesNotRaise(),
), # empty predictions and masks
(
np.array([[0, 0, 0, 0, 0.8]]),
np.array(
[
[
[False, False, False, False, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, False, False, False, False],
]
]
),
0.5,
[[0]],
DoesNotRaise(),
), # single mask with no category
(
np.array([[0, 0, 0, 0, 0.8, 0]]),
np.array(
[
[
[False, False, False, False, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, False, False, False, False],
]
]
),
0.5,
[[0]],
DoesNotRaise(),
), # single mask with category
(
np.array([[0, 0, 0, 0, 0.8], [0, 0, 0, 0, 0.9]]),
np.array(
[
[
[False, False, False, False, False],
[False, True, True, False, False],
[False, True, True, False, False],
[False, False, False, False, False],
[False, False, False, False, False],
],
[
[False, False, False, False, False],
[False, False, False, False, False],
[False, False, False, True, True],
[False, False, False, True, True],
[False, False, False, False, False],
],
]
),
0.5,
[[0], [1]],
DoesNotRaise(),
), # two masks non-overlapping with no category
(
np.array([[0, 0, 0, 0, 0.8], [0, 0, 0, 0, 0.9]]),
np.array(
[
[
[False, False, False, False, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, False, False, False, False],
],
[
[False, False, False, False, False],
[False, False, True, True, True],
[False, False, True, True, True],
[False, False, True, True, True],
[False, False, False, False, False],
],
]
),
0.4,
[[0, 1]],
DoesNotRaise(),
), # two masks partially overlapping with no category, merge
(
np.array([[0, 0, 0, 0, 0.8], [0, 0, 0, 0, 0.9]]),
np.array(
[
[
[False, False, False, False, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, False, False, False, False],
],
[
[False, False, False, False, False],
[False, False, True, True, True],
[False, False, True, True, True],
[False, False, True, True, True],
[False, False, False, False, False],
],
]
),
0.6,
[[0, 1]],
DoesNotRaise(),
), # two masks partially overlapping with no category, no merge
(
np.array([[0, 0, 0, 0, 0.8, 0], [0, 0, 0, 0, 0.9, 1]]),
np.array(
[
[
[False, False, False, False, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, False, False, False, False],
],
[
[False, False, False, False, False],
[False, False, True, True, True],
[False, False, True, True, True],
[False, False, True, True, True],
[False, False, False, False, False],
],
]
),
0.4,
[[0], [1]],
DoesNotRaise(),
), # two masks partially overlapping with different categories
(
np.array([[0, 0, 0, 0, 0.8, 0], [0, 0, 0, 0, 0.9, 0]]),
np.array(
[
[
[False, False, False, False, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, True, True, True, False],
[False, False, False, False, False],
],
[
[False, False, False, False, False],
[False, False, True, True, True],
[False, False, True, True, True],
[False, False, True, True, True],
[False, False, False, False, False],
],
]
),
0.4,
[[0, 1]],
DoesNotRaise(),
), # two masks partially overlapping with same category
],
)
def test_mask_non_max_merge(
predictions: np.ndarray,
masks: np.ndarray,
iou_threshold: float,
expected_result: List[List[int]],
exception: Exception,
) -> None:
with exception:
result = mask_non_max_merge(
predictions=predictions, masks=masks, iou_threshold=iou_threshold
)
sorted_result = sorted([sorted(group) for group in result])
sorted_expected_result = sorted([sorted(group) for group in expected_result])
assert sorted_result == sorted_expected_result