diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 96b40ccb..58958200 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -45,7 +45,7 @@ repos: - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.5.6 + rev: v0.6.0 hooks: - id: ruff args: [--fix, --exit-non-zero-on-fix] diff --git a/docs/detection/annotators.md b/docs/detection/annotators.md index 22486c16..eec45b00 100644 --- a/docs/detection/annotators.md +++ b/docs/detection/annotators.md @@ -433,6 +433,27 @@ status: new +=== "Background Color" + + ```python + import supervision as sv + + image = ... + detections = sv.Detections(...) + + background_overlay_annotator = sv.BackgroundOverlayAnnotator() + annotated_frame = background_overlay_annotator.annotate( + scene=image.copy(), + detections=detections + ) + ``` + +
+ + ![background-overlay-annotator-example](https://media.roboflow.com/supervision-annotator-examples/background-color-annotator-example-purple.png) + +
+

BoxAnnotator

@@ -553,6 +574,12 @@ status: new :::supervision.annotators.core.CropAnnotator +
+

BackgroundOverlayAnnotator

+
+ +:::supervision.annotators.core.BackgroundOverlayAnnotator +

ColorLookup

diff --git a/poetry.lock b/poetry.lock index ab9e0591..f2526808 100644 --- a/poetry.lock +++ b/poetry.lock @@ -2063,13 +2063,13 @@ files = [ [[package]] name = "mike" -version = "2.1.2" +version = "2.1.3" description = "Manage multiple versions of your MkDocs-powered documentation" optional = false python-versions = "*" files = [ - {file = "mike-2.1.2-py3-none-any.whl", hash = "sha256:d61d9b423ab412d634ca2bd520136d5114e3cc73f4bbd1aa6a0c6625c04918c0"}, - {file = "mike-2.1.2.tar.gz", hash = "sha256:d59cc8054c50f9c8a046cfd47f9b700cf9ff1b2b19f420bd8812ca6f94fa8bd3"}, + {file = "mike-2.1.3-py3-none-any.whl", hash = "sha256:d90c64077e84f06272437b464735130d380703a76a5738b152932884c60c062a"}, + {file = "mike-2.1.3.tar.gz", hash = "sha256:abd79b8ea483fb0275b7972825d3082e5ae67a41820f8d8a0dc7a3f49944e810"}, ] [package.dependencies] @@ -3838,29 +3838,29 @@ files = [ [[package]] name = "ruff" -version = "0.5.7" +version = "0.6.1" description = "An extremely fast Python linter and code formatter, written in Rust." optional = false python-versions = ">=3.7" files = [ - {file = "ruff-0.5.7-py3-none-linux_armv6l.whl", hash = "sha256:548992d342fc404ee2e15a242cdbea4f8e39a52f2e7752d0e4cbe88d2d2f416a"}, - {file = "ruff-0.5.7-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:00cc8872331055ee017c4f1071a8a31ca0809ccc0657da1d154a1d2abac5c0be"}, - {file = "ruff-0.5.7-py3-none-macosx_11_0_arm64.whl", hash = "sha256:eaf3d86a1fdac1aec8a3417a63587d93f906c678bb9ed0b796da7b59c1114a1e"}, - {file = "ruff-0.5.7-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:a01c34400097b06cf8a6e61b35d6d456d5bd1ae6961542de18ec81eaf33b4cb8"}, - {file = "ruff-0.5.7-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:fcc8054f1a717e2213500edaddcf1dbb0abad40d98e1bd9d0ad364f75c763eea"}, - {file = "ruff-0.5.7-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:7f70284e73f36558ef51602254451e50dd6cc479f8b6f8413a95fcb5db4a55fc"}, - {file = "ruff-0.5.7-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:a78ad870ae3c460394fc95437d43deb5c04b5c29297815a2a1de028903f19692"}, - {file = "ruff-0.5.7-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:9ccd078c66a8e419475174bfe60a69adb36ce04f8d4e91b006f1329d5cd44bcf"}, - {file = "ruff-0.5.7-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:7e31c9bad4ebf8fdb77b59cae75814440731060a09a0e0077d559a556453acbb"}, - {file = "ruff-0.5.7-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:8d796327eed8e168164346b769dd9a27a70e0298d667b4ecee6877ce8095ec8e"}, - {file = "ruff-0.5.7-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:4a09ea2c3f7778cc635e7f6edf57d566a8ee8f485f3c4454db7771efb692c499"}, - {file = "ruff-0.5.7-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:a36d8dcf55b3a3bc353270d544fb170d75d2dff41eba5df57b4e0b67a95bb64e"}, - {file = "ruff-0.5.7-py3-none-musllinux_1_2_i686.whl", hash = "sha256:9369c218f789eefbd1b8d82a8cf25017b523ac47d96b2f531eba73770971c9e5"}, - {file = "ruff-0.5.7-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:b88ca3db7eb377eb24fb7c82840546fb7acef75af4a74bd36e9ceb37a890257e"}, - {file = "ruff-0.5.7-py3-none-win32.whl", hash = "sha256:33d61fc0e902198a3e55719f4be6b375b28f860b09c281e4bdbf783c0566576a"}, - {file = "ruff-0.5.7-py3-none-win_amd64.whl", hash = "sha256:083bbcbe6fadb93cd86709037acc510f86eed5a314203079df174c40bbbca6b3"}, - {file = "ruff-0.5.7-py3-none-win_arm64.whl", hash = "sha256:2dca26154ff9571995107221d0aeaad0e75a77b5a682d6236cf89a58c70b76f4"}, - {file = "ruff-0.5.7.tar.gz", hash = "sha256:8dfc0a458797f5d9fb622dd0efc52d796f23f0a1493a9527f4e49a550ae9a7e5"}, + {file = "ruff-0.6.1-py3-none-linux_armv6l.whl", hash = "sha256:b4bb7de6a24169dc023f992718a9417380301b0c2da0fe85919f47264fb8add9"}, + {file = "ruff-0.6.1-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:45efaae53b360c81043e311cdec8a7696420b3d3e8935202c2846e7a97d4edae"}, + {file = "ruff-0.6.1-py3-none-macosx_11_0_arm64.whl", hash = "sha256:bc60c7d71b732c8fa73cf995efc0c836a2fd8b9810e115be8babb24ae87e0850"}, + {file = "ruff-0.6.1-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:2c7477c3b9da822e2db0b4e0b59e61b8a23e87886e727b327e7dcaf06213c5cf"}, + {file = "ruff-0.6.1-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:3a0af7ab3f86e3dc9f157a928e08e26c4b40707d0612b01cd577cc84b8905cc9"}, + {file = "ruff-0.6.1-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:392688dbb50fecf1bf7126731c90c11a9df1c3a4cdc3f481b53e851da5634fa5"}, + {file = "ruff-0.6.1-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:5278d3e095ccc8c30430bcc9bc550f778790acc211865520f3041910a28d0024"}, + {file = "ruff-0.6.1-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:fe6d5f65d6f276ee7a0fc50a0cecaccb362d30ef98a110f99cac1c7872df2f18"}, + {file = "ruff-0.6.1-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:b2e0dd11e2ae553ee5c92a81731d88a9883af8db7408db47fc81887c1f8b672e"}, + {file = "ruff-0.6.1-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:d812615525a34ecfc07fd93f906ef5b93656be01dfae9a819e31caa6cfe758a1"}, + {file = "ruff-0.6.1-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:faaa4060f4064c3b7aaaa27328080c932fa142786f8142aff095b42b6a2eb631"}, + {file = "ruff-0.6.1-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:99d7ae0df47c62729d58765c593ea54c2546d5de213f2af2a19442d50a10cec9"}, + {file = "ruff-0.6.1-py3-none-musllinux_1_2_i686.whl", hash = "sha256:9eb18dfd7b613eec000e3738b3f0e4398bf0153cb80bfa3e351b3c1c2f6d7b15"}, + {file = "ruff-0.6.1-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:c62bc04c6723a81e25e71715aa59489f15034d69bf641df88cb38bdc32fd1dbb"}, + {file = "ruff-0.6.1-py3-none-win32.whl", hash = "sha256:9fb4c4e8b83f19c9477a8745e56d2eeef07a7ff50b68a6998f7d9e2e3887bdc4"}, + {file = "ruff-0.6.1-py3-none-win_amd64.whl", hash = "sha256:c2ebfc8f51ef4aca05dad4552bbcf6fe8d1f75b2f6af546cc47cc1c1ca916b5b"}, + {file = "ruff-0.6.1-py3-none-win_arm64.whl", hash = "sha256:3bc81074971b0ffad1bd0c52284b22411f02a11a012082a76ac6da153536e014"}, + {file = "ruff-0.6.1.tar.gz", hash = "sha256:af3ffd8c6563acb8848d33cd19a69b9bfe943667f0419ca083f8ebe4224a3436"}, ] [[package]] @@ -4125,13 +4125,13 @@ files = [ [[package]] name = "tox" -version = "4.17.1" +version = "4.18.0" description = "tox is a generic virtualenv management and test command line tool" optional = false python-versions = ">=3.8" files = [ - {file = "tox-4.17.1-py3-none-any.whl", hash = "sha256:2974597c0353577126ab014f52d1a399fb761049e165ff34427f84e8cfe6c990"}, - {file = "tox-4.17.1.tar.gz", hash = "sha256:2c41565a571e34480bd401d668a4899806169a4633e972ac296c54406d2ded8a"}, + {file = "tox-4.18.0-py3-none-any.whl", hash = "sha256:0a457400cf70615dc0627eb70d293e80cd95d8ce174bb40ac011011f0c03a249"}, + {file = "tox-4.18.0.tar.gz", hash = "sha256:5dfa1cab9f146becd6e351333a82f9e0ade374451630ba65ee54584624c27b58"}, ] [package.dependencies] diff --git a/supervision/__init__.py b/supervision/__init__.py index 4d0e8284..0641bf23 100644 --- a/supervision/__init__.py +++ b/supervision/__init__.py @@ -7,6 +7,7 @@ except importlib_metadata.PackageNotFoundError: __version__ = "development" from supervision.annotators.core import ( + BackgroundOverlayAnnotator, BlurAnnotator, BoundingBoxAnnotator, BoxAnnotator, diff --git a/supervision/annotators/core.py b/supervision/annotators/core.py index c6ae7d82..e0002e86 100644 --- a/supervision/annotators/core.py +++ b/supervision/annotators/core.py @@ -927,6 +927,7 @@ class DotAnnotator(BaseAnnotator): position: Position = Position.CENTER, color_lookup: ColorLookup = ColorLookup.CLASS, outline_thickness: int = 0, + outline_color: Union[Color, ColorPalette] = Color.BLACK, ): """ Args: @@ -937,12 +938,16 @@ class DotAnnotator(BaseAnnotator): color_lookup (ColorLookup): Strategy for mapping colors to annotations. Options are `INDEX`, `CLASS`, `TRACK`. outline_thickness (int): Thickness of the outline of the dot. + outline_color (Union[Color, ColorPalette]): The color or color palette to + use for outline. It is activated by setting outline_thickness to a value + greater than 0. """ self.color: Union[Color, ColorPalette] = color self.radius: int = radius self.position: Position = position self.color_lookup: ColorLookup = color_lookup self.outline_thickness = outline_thickness + self.outline_color: Union[Color, ColorPalette] = outline_color @ensure_cv2_image_for_annotation def annotate( @@ -997,8 +1002,20 @@ class DotAnnotator(BaseAnnotator): cv2.circle(scene, center, self.radius, color.as_bgr(), -1) if self.outline_thickness: + outline_color = resolve_color( + color=self.outline_color, + detections=detections, + detection_idx=detection_idx, + color_lookup=self.color_lookup + if custom_color_lookup is None + else custom_color_lookup, + ) cv2.circle( - scene, center, self.radius, (0, 0, 0), self.outline_thickness + scene, + center, + self.radius, + outline_color.as_bgr(), + self.outline_thickness, ) return scene @@ -1744,6 +1761,7 @@ class TriangleAnnotator(BaseAnnotator): position: Position = Position.TOP_CENTER, color_lookup: ColorLookup = ColorLookup.CLASS, outline_thickness: int = 0, + outline_color: Union[Color, ColorPalette] = Color.BLACK, ): """ Args: @@ -1755,6 +1773,9 @@ class TriangleAnnotator(BaseAnnotator): color_lookup (ColorLookup): Strategy for mapping colors to annotations. Options are `INDEX`, `CLASS`, `TRACK`. outline_thickness (int): Thickness of the outline of the triangle. + outline_color (Union[Color, ColorPalette]): The color or color palette to + use for outline. It is activated by setting outline_thickness to a value + greater than 0. """ self.color: Union[Color, ColorPalette] = color self.base: int = base @@ -1762,6 +1783,7 @@ class TriangleAnnotator(BaseAnnotator): self.position: Position = position self.color_lookup: ColorLookup = color_lookup self.outline_thickness: int = outline_thickness + self.outline_color: Union[Color, ColorPalette] = outline_color @ensure_cv2_image_for_annotation def annotate( @@ -1824,8 +1846,20 @@ class TriangleAnnotator(BaseAnnotator): cv2.fillPoly(scene, [vertices], color.as_bgr()) if self.outline_thickness: + outline_color = resolve_color( + color=self.outline_color, + detections=detections, + detection_idx=detection_idx, + color_lookup=self.color_lookup + if custom_color_lookup is None + else custom_color_lookup, + ) cv2.polylines( - scene, [vertices], True, (0, 0, 0), thickness=self.outline_thickness + scene, + [vertices], + True, + outline_color.as_bgr(), + thickness=self.outline_thickness, ) return scene @@ -2283,3 +2317,84 @@ class CropAnnotator(BaseAnnotator): ) elif position == Position.BOTTOM_RIGHT: return (anchor_x, anchor_y), (anchor_x + width, anchor_y + height) + + +class BackgroundOverlayAnnotator(BaseAnnotator): + """ + A class for drawing a colored overlay on the background of an image outside + the region of detections. + + If masks are provided, the background is colored outside the masks. + If masks are not provided, the background is colored outside the bounding boxes. + + You can use the `force_box` parameter to force the annotator to use bounding boxes. + + !!! warning + + This annotator uses `sv.Detections.mask`. + """ + + def __init__( + self, + color: Color = Color.BLACK, + opacity: float = 0.5, + force_box: bool = False, + ): + """ + Args: + color (Color): The color to use for annotating detections. + opacity (float): Opacity of the overlay mask. Must be between `0` and `1`. + force_box (bool): If `True`, forces the annotator to use bounding boxes when + masks are provided in the supplied sv.Detections. + """ + self.color: Color = color + self.opacity = opacity + self.force_box = force_box + + @ensure_cv2_image_for_annotation + def annotate(self, scene: ImageType, detections: Detections) -> ImageType: + """ + Applies a colored overlay to the scene outside of the detected regions. + + Args: + scene (ImageType): The image where masks will be drawn. + `ImageType` is a flexible type, accepting either `numpy.ndarray` + or `PIL.Image.Image`. + detections (Detections): Object detections to annotate. + + Returns: + The annotated image, matching the type of `scene` (`numpy.ndarray` + or `PIL.Image.Image`) + + Example: + ```python + import supervision as sv + + image = ... + detections = sv.Detections(...) + + background_overlay_annotator = sv.BackgroundOverlayAnnotator() + annotated_frame = background_overlay_annotator.annotate( + scene=image.copy(), + detections=detections + ) + ``` + + ![background-overlay-annotator-example](https://media.roboflow.com/ + supervision-annotator-examples/background-color-annotator-example-purple.png) + """ + colored_mask = np.full_like(scene, self.color.as_bgr(), dtype=np.uint8) + + cv2.addWeighted( + scene, 1 - self.opacity, colored_mask, self.opacity, 0, dst=colored_mask + ) + + if detections.mask is None or self.force_box: + for x1, y1, x2, y2 in detections.xyxy.astype(int): + colored_mask[y1:y2, x1:x2] = scene[y1:y2, x1:x2] + else: + for mask in detections.mask: + colored_mask[mask] = scene[mask] + + np.copyto(scene, colored_mask) + return scene diff --git a/supervision/detection/core.py b/supervision/detection/core.py index 73daadda..7216b612 100644 --- a/supervision/detection/core.py +++ b/supervision/detection/core.py @@ -29,6 +29,7 @@ from supervision.detection.utils import ( extract_ultralytics_masks, get_data_item, is_data_equal, + mask_to_xyxy, merge_data, process_roboflow_result, xywh_to_xyxy, @@ -267,6 +268,14 @@ class Detections: }, ) + if hasattr(ultralytics_results, "boxes") and ultralytics_results.boxes is None: + masks = extract_ultralytics_masks(ultralytics_results) + return cls( + xyxy=mask_to_xyxy(masks), + mask=masks, + class_id=np.arange(len(ultralytics_results)), + ) + class_id = ultralytics_results.boxes.cls.cpu().numpy().astype(int) class_names = np.array([ultralytics_results.names[i] for i in class_id]) return cls( diff --git a/supervision/detection/overlap_filter.py b/supervision/detection/overlap_filter.py index ab4408d1..f51f1dce 100644 --- a/supervision/detection/overlap_filter.py +++ b/supervision/detection/overlap_filter.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from enum import Enum from typing import List, Union @@ -248,16 +250,21 @@ class OverlapFilter(Enum): NON_MAX_SUPPRESSION = "non_max_suppression" NON_MAX_MERGE = "non_max_merge" + @classmethod + def list(cls): + return list(map(lambda c: c.value, cls)) -def validate_overlap_filter( - strategy: Union[OverlapFilter, str], -) -> OverlapFilter: - if isinstance(strategy, str): - try: - strategy = OverlapFilter(strategy.lower()) - except ValueError: - raise ValueError( - f"Invalid strategy value: {strategy}. Must be one of " - f"{[e.value for e in OverlapFilter]}" - ) - return strategy + @classmethod + def from_value(cls, value: Union[OverlapFilter, str]) -> OverlapFilter: + 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." + ) diff --git a/supervision/detection/tools/inference_slicer.py b/supervision/detection/tools/inference_slicer.py index 38d85faa..52f45140 100644 --- a/supervision/detection/tools/inference_slicer.py +++ b/supervision/detection/tools/inference_slicer.py @@ -6,10 +6,14 @@ import numpy as np from supervision.config import ORIENTED_BOX_COORDINATES from supervision.detection.core import Detections -from supervision.detection.overlap_filter import OverlapFilter, validate_overlap_filter +from supervision.detection.overlap_filter import OverlapFilter from supervision.detection.utils import move_boxes, move_masks, move_oriented_boxes from supervision.utils.image import crop_image -from supervision.utils.internal import SupervisionWarnings +from supervision.utils.internal import ( + SupervisionWarnings, + deprecated_parameter, + warn_deprecated, +) def move_detections( @@ -54,11 +58,16 @@ class InferenceSlicer: slice, and then merging the detections. Args: - slice_wh (Tuple[int, int]): Dimensions of each slice in the format - `(width, height)`. - overlap_ratio_wh (Tuple[float, float]): Overlap ratio between consecutive - slices in the format `(width_ratio, height_ratio)`. - overlap_filter_strategy (Union[OverlapFilter, str]): Strategy for + slice_wh (Tuple[int, int]): Dimensions of each slice measured in pixels. The + tuple should be in the format `(width, height)`. + overlap_ratio_wh (Optional[Tuple[float, float]]): A tuple representing the + desired overlap ratio for width and height between consecutive slices. + Each value should be in the range [0, 1), where 0 means no overlap and + a value close to 1 means high overlap. + overlap_wh (Optional[Tuple[int, int]]): A tuple representing the desired + overlap for width and height between consecutive slices measured in pixels. + Each value should be greater than or equal to 0. + overlap_filter (Union[OverlapFilter, str]): Strategy for filtering or merging overlapping detections in slices. iou_threshold (float): Intersection over Union (IoU) threshold used when filtering by overlap. @@ -73,23 +82,37 @@ class InferenceSlicer: not a multiple of the slice's width or height minus the overlap. """ + @deprecated_parameter( + old_parameter="overlap_filter_strategy", + new_parameter="overlap_filter", + map_function=lambda x: x, + warning_message="`{old_parameter}` in `{function_name}` is deprecated and will " + "be removed in `supervision-0.27.0`. Use '{new_parameter}' " + "instead.", + ) def __init__( self, callback: Callable[[np.ndarray], Detections], slice_wh: Tuple[int, int] = (320, 320), - overlap_ratio_wh: Tuple[float, float] = (0.2, 0.2), - overlap_filter_strategy: Union[ - OverlapFilter, str - ] = OverlapFilter.NON_MAX_SUPPRESSION, + overlap_ratio_wh: Optional[Tuple[float, float]] = (0.2, 0.2), + overlap_wh: Optional[Tuple[int, int]] = None, + overlap_filter: Union[OverlapFilter, str] = OverlapFilter.NON_MAX_SUPPRESSION, iou_threshold: float = 0.5, thread_workers: int = 1, ): - overlap_filter_strategy = validate_overlap_filter(overlap_filter_strategy) + if overlap_ratio_wh is not None: + warn_deprecated( + "`overlap_ratio_wh` in `InferenceSlicer.__init__` is deprecated and " + "will be removed in `supervision-0.27.0`. Use `overlap_wh` instead." + ) + + self._validate_overlap(overlap_ratio_wh, overlap_wh) + self.overlap_ratio_wh = overlap_ratio_wh + self.overlap_wh = overlap_wh self.slice_wh = slice_wh - self.overlap_ratio_wh = overlap_ratio_wh self.iou_threshold = iou_threshold - self.overlap_filter_strategy = overlap_filter_strategy + self.overlap_filter = OverlapFilter.from_value(overlap_filter) self.callback = callback self.thread_workers = thread_workers @@ -134,6 +157,7 @@ class InferenceSlicer: resolution_wh=resolution_wh, slice_wh=self.slice_wh, overlap_ratio_wh=self.overlap_ratio_wh, + overlap_wh=self.overlap_wh, ) with ThreadPoolExecutor(max_workers=self.thread_workers) as executor: @@ -144,15 +168,15 @@ class InferenceSlicer: detections_list.append(future.result()) merged = Detections.merge(detections_list=detections_list) - if self.overlap_filter_strategy == OverlapFilter.NONE: + if self.overlap_filter == OverlapFilter.NONE: return merged - elif self.overlap_filter_strategy == OverlapFilter.NON_MAX_SUPPRESSION: + elif self.overlap_filter == OverlapFilter.NON_MAX_SUPPRESSION: return merged.with_nms(threshold=self.iou_threshold) - elif self.overlap_filter_strategy == OverlapFilter.NON_MAX_MERGE: + elif self.overlap_filter == OverlapFilter.NON_MAX_MERGE: return merged.with_nmm(threshold=self.iou_threshold) else: warnings.warn( - f"Invalid overlap filter strategy: {self.overlap_filter_strategy}", + f"Invalid overlap filter strategy: {self.overlap_filter}", category=SupervisionWarnings, ) return merged @@ -182,7 +206,8 @@ class InferenceSlicer: def _generate_offset( resolution_wh: Tuple[int, int], slice_wh: Tuple[int, int], - overlap_ratio_wh: Tuple[float, float], + overlap_ratio_wh: Optional[Tuple[float, float]], + overlap_wh: Optional[Tuple[int, int]], ) -> np.ndarray: """ Generate offset coordinates for slicing an image based on the given resolution, @@ -191,12 +216,15 @@ class InferenceSlicer: Args: resolution_wh (Tuple[int, int]): A tuple representing the width and height of the image to be sliced. - slice_wh (Tuple[int, int]): A tuple representing the desired width and - height of each slice. - overlap_ratio_wh (Tuple[float, float]): A tuple representing the desired - overlap ratio for width and height between consecutive slices. Each - value should be in the range [0, 1), where 0 means no overlap and a - value close to 1 means high overlap. + slice_wh (Tuple[int, int]): Dimensions of each slice measured in pixels. The + tuple should be in the format `(width, height)`. + overlap_ratio_wh (Optional[Tuple[float, float]]): A tuple representing the + desired overlap ratio for width and height between consecutive slices. + Each value should be in the range [0, 1), where 0 means no overlap and + a value close to 1 means high overlap. + overlap_wh (Optional[Tuple[int, int]]): A tuple representing the desired + overlap for width and height between consecutive slices measured in + pixels. Each value should be greater than or equal to 0. Returns: np.ndarray: An array of shape `(n, 4)` containing coordinates for each @@ -211,10 +239,19 @@ class InferenceSlicer: """ slice_width, slice_height = slice_wh image_width, image_height = resolution_wh - overlap_ratio_width, overlap_ratio_height = overlap_ratio_wh + overlap_width = ( + overlap_wh[0] + if overlap_wh is not None + else int(overlap_ratio_wh[0] * slice_width) + ) + overlap_height = ( + overlap_wh[1] + if overlap_wh is not None + else int(overlap_ratio_wh[1] * slice_height) + ) - width_stride = slice_width - int(overlap_ratio_width * slice_width) - height_stride = slice_height - int(overlap_ratio_height * slice_height) + width_stride = slice_width - overlap_width + height_stride = slice_height - overlap_height ws = np.arange(0, image_width, width_stride) hs = np.arange(0, image_height, height_stride) @@ -226,3 +263,32 @@ class InferenceSlicer: offsets = np.stack([xmin, ymin, xmax, ymax], axis=-1).reshape(-1, 4) return offsets + + @staticmethod + def _validate_overlap( + overlap_ratio_wh: Optional[Tuple[float, float]], + overlap_wh: Optional[Tuple[int, int]], + ) -> None: + if overlap_ratio_wh is not None and overlap_wh is not None: + raise ValueError( + "Both `overlap_ratio_wh` and `overlap_wh` cannot be provided. " + "Please provide only one of them." + ) + if overlap_ratio_wh is None and overlap_wh is None: + raise ValueError( + "Either `overlap_ratio_wh` or `overlap_wh` must be provided. " + "Please provide one of them." + ) + + if overlap_ratio_wh is not None: + if not (0 <= overlap_ratio_wh[0] < 1 and 0 <= overlap_ratio_wh[1] < 1): + raise ValueError( + "Overlap ratios must be in the range [0, 1). " + f"Received: {overlap_ratio_wh}" + ) + if overlap_wh is not None: + if not (overlap_wh[0] >= 0 and overlap_wh[1] >= 0): + raise ValueError( + "Overlap values must be greater than or equal to 0. " + f"Received: {overlap_wh}" + ) diff --git a/supervision/keypoint/annotators.py b/supervision/keypoint/annotators.py index d2024665..38dea23c 100644 --- a/supervision/keypoint/annotators.py +++ b/supervision/keypoint/annotators.py @@ -194,7 +194,7 @@ class VertexLabelAnnotator: def __init__( self, color: Union[Color, List[Color]] = Color.ROBOFLOW, - text_color: Color = Color.WHITE, + text_color: Union[Color, List[Color]] = Color.WHITE, text_scale: float = 0.5, text_thickness: int = 1, text_padding: int = 10, @@ -205,7 +205,9 @@ class VertexLabelAnnotator: color (Union[Color, List[Color]], optional): The color to use for each keypoint label. If a list is provided, the colors will be used in order for each keypoint. - text_color (Color, optional): The color to use for the labels. + text_color (Union[Color, List[Color]], optional): The color to use + for the labels. If a list is provided, the colors will be used in order + for each keypoint. text_scale (float, optional): The scale of the text. text_thickness (int, optional): The thickness of the text. text_padding (int, optional): The padding around the text. @@ -214,7 +216,7 @@ class VertexLabelAnnotator: """ self.border_radius: int = border_radius self.color: Union[Color, List[Color]] = color - self.text_color: Color = text_color + self.text_color: Union[Color, List[Color]] = text_color self.text_scale: float = text_scale self.text_thickness: int = text_thickness self.text_padding: int = text_padding @@ -321,12 +323,19 @@ class VertexLabelAnnotator: skeletons_count=skeletons_count, ) + text_colors = self.preprocess_and_validate_colors( + colors=self.text_color, + points_count=points_count, + skeletons_count=skeletons_count, + ) + labels = self.preprocess_and_validate_labels( labels=labels, points_count=points_count, skeletons_count=skeletons_count ) anchors = anchors[mask] colors = colors[mask] + text_colors = text_colors[mask] labels = labels[mask] xyxy = np.array( @@ -344,7 +353,9 @@ class VertexLabelAnnotator: xyxy_padded = pad_boxes(xyxy=xyxy, px=self.text_padding) - for text, color, box, box_padded in zip(labels, colors, xyxy, xyxy_padded): + for text, color, text_color, box, box_padded in zip( + labels, colors, text_colors, xyxy, xyxy_padded + ): draw_rounded_rectangle( scene=scene, rect=Rect.from_xyxy(box_padded), @@ -357,7 +368,7 @@ class VertexLabelAnnotator: org=(box[0], box[3]), fontFace=font, fontScale=self.text_scale, - color=self.text_color.as_rgb(), + color=text_color.as_bgr(), thickness=self.text_thickness, lineType=cv2.LINE_AA, )