diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml
index c0903f3a..9465c2af 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.4.2
+ rev: v0.4.4
hooks:
- id: ruff
args: [--fix, --exit-non-zero-on-fix]
diff --git a/docs/datasets.md b/docs/datasets/core.md
similarity index 97%
rename from docs/datasets.md
rename to docs/datasets/core.md
index 73931515..03d0c196 100644
--- a/docs/datasets.md
+++ b/docs/datasets/core.md
@@ -1,5 +1,6 @@
---
comments: true
+status: new
---
# Datasets
diff --git a/docs/datasets/utils.md b/docs/datasets/utils.md
new file mode 100644
index 00000000..6be56303
--- /dev/null
+++ b/docs/datasets/utils.md
@@ -0,0 +1,18 @@
+---
+comments: true
+status: new
+---
+
+# Datasets Utils
+
+
+
+:::supervision.dataset.utils.rle_to_mask
+
+
+
+:::supervision.dataset.utils.mask_to_rle
diff --git a/docs/detection/annotators.md b/docs/detection/annotators.md
index e1f4b115..958f2a74 100644
--- a/docs/detection/annotators.md
+++ b/docs/detection/annotators.md
@@ -285,6 +285,37 @@ status: new
+=== "RichLabel"
+
+ ```python
+ import supervision as sv
+
+ image = ...
+ detections = sv.Detections(...)
+
+ labels = [
+ f"{class_name} {confidence:.2f}"
+ for class_name, confidence
+ in zip(detections['class_name'], detections.confidence)
+ ]
+
+ rich_label_annotator = sv.RichLabelAnnotator(
+ font_path=".../font.ttf",
+ text_position=sv.Position.CENTER
+ )
+ annotated_frame = label_annotator.annotate(
+ scene=image.copy(),
+ detections=detections,
+ labels=labels
+ )
+ ```
+
+
+
+ { align=center width="800" }
+
+
+
=== "Crop"
```python
@@ -492,6 +523,12 @@ status: new
:::supervision.annotators.core.LabelAnnotator
+
+
+:::supervision.annotators.core.RichLabelAnnotator
+
diff --git a/docs/detection/utils.md b/docs/detection/utils.md
index abacdc21..f9c9473b 100644
--- a/docs/detection/utils.md
+++ b/docs/detection/utils.md
@@ -65,8 +65,38 @@ status: new
:::supervision.detection.utils.move_boxes
+
+
+:::supervision.detection.utils.move_masks
+
:::supervision.detection.utils.scale_boxes
+
+
+
+:::supervision.detection.utils.clip_boxes
+
+
+
+:::supervision.detection.utils.pad_boxes
+
+
+
+:::supervision.detection.utils.contains_holes
+
+
+
+:::supervision.detection.utils.contains_multiple_segments
diff --git a/docs/how_to/detect_small_objects.md b/docs/how_to/detect_small_objects.md
index e2d02328..175b4f36 100644
--- a/docs/how_to/detect_small_objects.md
+++ b/docs/how_to/detect_small_objects.md
@@ -6,7 +6,7 @@ status: new
# Detect Small Objects
This guide shows how to detect small objects
-with the [Inference](https://github.com/roboflow/inference),
+with the [Inference](https://github.com/roboflow/inference),
[Ultralytics](https://github.com/ultralytics/ultralytics) or
[Transformers](https://github.com/huggingface/transformers) packages using
[`InferenceSlicer`](/latest/detection/tools/inference_slicer/#supervision.detection.tools.inference_slicer.InferenceSlicer).
@@ -68,10 +68,10 @@ size relative to the image resolution.
import torch
import supervision as sv
from PIL import Image
- from transformers import DetrImageProcessor, DetrForObjectDetection
+ from transformers import DetrImageProcessor, DetrForSegmentation
processor = DetrImageProcessor.from_pretrained("facebook/detr-resnet-50")
- model = DetrForObjectDetection.from_pretrained("facebook/detr-resnet-50")
+ model = DetrForSegmentation.from_pretrained("facebook/detr-resnet-50")
image = Image.open()
inputs = processor(images=image, return_tensors="pt")
@@ -79,8 +79,8 @@ size relative to the image resolution.
with torch.no_grad():
outputs = model(**inputs)
- width, height = image.size
- target_size = torch.tensor([[height, width]])
+ width, height = image_slice.size
+ target_size = torch.tensor([[width, height]])
results = processor.post_process_object_detection(
outputs=outputs, target_sizes=target_size)[0]
detections = sv.Detections.from_transformers(results)
@@ -175,7 +175,7 @@ objects within each, and aggregating the results.
def callback(image_slice: np.ndarray) -> sv.Detections:
results = model.infer(image_slice)[0]
- detections = sv.Detections.from_inference(results)
+ return sv.Detections.from_inference(results)
slicer = sv.InferenceSlicer(callback = callback)
detections = slicer(image)
@@ -239,8 +239,8 @@ objects within each, and aggregating the results.
with torch.no_grad():
outputs = model(**inputs)
- width, height = image.size
- target_size = torch.tensor([[height, width]])
+ width, height = image_slice.size
+ target_size = torch.tensor([[width, height]])
results = processor.post_process_object_detection(
outputs=outputs, target_sizes=target_size)[0]
return sv.Detections.from_transformers(results)
@@ -264,3 +264,63 @@ objects within each, and aggregating the results.
```

+
+## Small Object Segmentation
+
+[`InferenceSlicer`](/latest/detection/tools/inference_slicer/#supervision.detection.tools.inference_slicer.InferenceSlicer) can perform segmentation tasks too.
+
+=== "Inference"
+
+ ```{ .py hl_lines="6 16 19-20" }
+ import cv2
+ import numpy as np
+ import supervision as sv
+ from inference import get_model
+
+ model = get_model(model_id="yolov8x-seg-640")
+ image = cv2.imread()
+
+ def callback(image_slice: np.ndarray) -> sv.Detections:
+ results = model.infer(image_slice)[0]
+ return sv.Detections.from_inference(results)
+
+ slicer = sv.InferenceSlicer(callback = callback)
+ detections = slicer(image)
+
+ mask_annotator = sv.MaskAnnotator()
+ label_annotator = sv.LabelAnnotator()
+
+ annotated_image = mask_annotator.annotate(
+ scene=image, detections=detections)
+ annotated_image = label_annotator.annotate(
+ scene=annotated_image, detections=detections)
+ ```
+
+=== "Ultralytics"
+
+ ```{ .py hl_lines="6 16 19-20" }
+ import cv2
+ import numpy as np
+ import supervision as sv
+ from ultralytics import YOLO
+
+ model = YOLO("yolov8x-seg.pt")
+ image = cv2.imread()
+
+ def callback(image_slice: np.ndarray) -> sv.Detections:
+ result = model(image_slice)[0]
+ return sv.Detections.from_ultralytics(result)
+
+ slicer = sv.InferenceSlicer(callback = callback)
+ detections = slicer(image)
+
+ mask_annotator = sv.MaskAnnotator()
+ label_annotator = sv.LabelAnnotator()
+
+ annotated_image = mask_annotator.annotate(
+ scene=image, detections=detections)
+ annotated_image = label_annotator.annotate(
+ scene=annotated_image, detections=detections)
+ ```
+
+
diff --git a/docs/keypoint/annotators.md b/docs/keypoint/annotators.md
index b5f998bc..30a970ec 100644
--- a/docs/keypoint/annotators.md
+++ b/docs/keypoint/annotators.md
@@ -13,7 +13,10 @@ status: new
image = ...
key_points = sv.KeyPoints(...)
- vertex_annotator = sv.VertexAnnotator(color=sv.Color.GREEN, radius=10)
+ vertex_annotator = sv.VertexAnnotator(
+ color=sv.Color.GREEN,
+ radius=10
+ )
annotated_frame = vertex_annotator.annotate(
scene=image.copy(),
key_points=key_points
@@ -34,7 +37,10 @@ status: new
image = ...
key_points = sv.KeyPoints(...)
- edge_annotator = sv.EdgeAnnotator(color=sv.Color.GREEN, thickness=5)
+ edge_annotator = sv.EdgeAnnotator(
+ color=sv.Color.GREEN,
+ thickness=5
+ )
annotated_frame = edge_annotator.annotate(
scene=image.copy(),
key_points=key_points
@@ -47,6 +53,31 @@ status: new
+=== "VertexLabelAnnotator"
+
+ ```python
+ import supervision as sv
+
+ image = ...
+ key_points = sv.KeyPoints(...)
+
+ vertex_label_annotator = sv.VertexLabelAnnotator(
+ color=sv.Color.GREEN,
+ text_color=sv.Color.BLACK,
+ border_radius=5
+ )
+ annotated_frame = vertex_label_annotator.annotate(
+ scene=image.copy(),
+ key_points=key_points
+ )
+ ```
+
+
+
+ { align=center width="800" }
+
+
+
@@ -58,3 +89,9 @@ status: new
:::supervision.keypoint.annotators.EdgeAnnotator
+
+
+
+:::supervision.keypoint.annotators.VertexLabelAnnotator
diff --git a/docs/utils/draw.md b/docs/utils/draw.md
index 84758e06..f4b86a53 100644
--- a/docs/utils/draw.md
+++ b/docs/utils/draw.md
@@ -41,7 +41,7 @@ comments: true
:::supervision.draw.utils.draw_image
:::supervision.draw.utils.calculate_optimal_text_scale
diff --git a/examples/time_in_zone/README.md b/examples/time_in_zone/README.md
index 98587999..0a366a94 100644
--- a/examples/time_in_zone/README.md
+++ b/examples/time_in_zone/README.md
@@ -103,7 +103,7 @@ python scripts/draw_zones.py \
```bash
python scripts/draw_zones.py \
--source_path "data/traffic/video.mp4" \
---zone_configuration_path "data/traffic/custom_config.json"
+--zone_configuration_path "data/traffic/config.json"
```
https://github.com/roboflow/supervision/assets/26109316/9d514c9e-2a61-418b-ae49-6ac1ad6ae5ac
@@ -157,7 +157,7 @@ Script to run object detection on a video stream using the Roboflow Inference mo
- `--iou_threshold`: IOU threshold for non-max suppression. Default is `0.7`.
```bash
-python inference_file_example.py \
+python inference_stream_example.py \
--zone_configuration_path "data/checkout/config.json" \
--rtsp_url "rtsp://localhost:8554/live0.stream" \
--model_id "yolov8x-640" \
@@ -167,7 +167,7 @@ python inference_file_example.py \
```
```bash
-python inference_file_example.py \
+python inference_stream_example.py \
--zone_configuration_path "data/traffic/config.json" \
--rtsp_url "rtsp://localhost:8554/live0.stream" \
--model_id "yolov8x-640" \
@@ -192,7 +192,7 @@ Script to run object detection on a video file using the Ultralytics YOLOv8 mode
- `--iou_threshold`: IOU threshold for non-max suppression. Default is `0.7`.
```bash
-python inference_file_example.py \
+python ultralytics_file_example.py \
--zone_configuration_path "data/checkout/config.json" \
--source_video_path "data/checkout/video.mp4" \
--weights "yolov8x.pt" \
@@ -203,7 +203,7 @@ python inference_file_example.py \
```
```bash
-python inference_file_example.py \
+python ultralytics_file_example.py \
--zone_configuration_path "data/traffic/config.json" \
--source_video_path "data/traffic/video.mp4" \
--weights "yolov8x.pt" \
@@ -226,7 +226,7 @@ Script to run object detection on a video stream using the Ultralytics YOLOv8 mo
- `--iou_threshold`: IOU threshold for non-max suppression. Default is `0.7`.
```bash
-python inference_file_example.py \
+python ultralytics_stream_example.py \
--zone_configuration_path "data/checkout/config.json" \
--rtsp_url "rtsp://localhost:8554/live0.stream" \
--weights "yolov8x.pt" \
@@ -237,7 +237,7 @@ python inference_file_example.py \
```
```bash
-python inference_file_example.py \
+python ultralytics_stream_example.py \
--zone_configuration_path "data/traffic/config.json" \
--rtsp_url "rtsp://localhost:8554/live0.stream" \
--weights "yolov8x.pt" \
diff --git a/mkdocs.yml b/mkdocs.yml
index cf206a82..f257238d 100644
--- a/mkdocs.yml
+++ b/mkdocs.yml
@@ -41,7 +41,7 @@ nav:
- Save Detections: how_to/save_detections.md
- Filter Detections: how_to/filter_detections.md
- Detect Small Objects: how_to/detect_small_objects.md
- - Track Objects: how_to/track_objects.md
+ - Track Objects on Video: how_to/track_objects.md
- API:
- Detection and Segmentation:
@@ -61,7 +61,9 @@ nav:
- Detection Smoother: detection/tools/smoother.md
- Save Detections: detection/tools/save_detections.md
- Trackers: trackers.md
- - Datasets: datasets.md
+ - Datasets:
+ - Core: datasets/core.md
+ - Utils: datasets/utils.md
- Utils:
- Video: utils/video.md
- Image: utils/image.md
diff --git a/poetry.lock b/poetry.lock
index 2497d802..0a9276f8 100644
--- a/poetry.lock
+++ b/poetry.lock
@@ -1340,13 +1340,13 @@ trio = ["async_generator", "trio"]
[[package]]
name = "jinja2"
-version = "3.1.3"
+version = "3.1.4"
description = "A very fast and expressive template engine."
optional = false
python-versions = ">=3.7"
files = [
- {file = "Jinja2-3.1.3-py3-none-any.whl", hash = "sha256:7d6d50dd97d52cbc355597bd845fabfbac3f551e1f99619e39a35ce8c370b5fa"},
- {file = "Jinja2-3.1.3.tar.gz", hash = "sha256:ac8bd6544d4bb2c9792bf3a159e80bba8fda7f07e81bc3aed565432d5925ba90"},
+ {file = "jinja2-3.1.4-py3-none-any.whl", hash = "sha256:bc5dd2abb727a5319567b7a813e6a2e7318c39f4f487cfe6c89c6f9c7d25197d"},
+ {file = "jinja2-3.1.4.tar.gz", hash = "sha256:4a3aee7acbbe7303aede8e9648d13b8bf88a429282aa6122a993f0ac800cb369"},
]
[package.dependencies]
@@ -1566,13 +1566,13 @@ test = ["jupyter-server (>=2.0.0)", "pytest (>=7.0)", "pytest-jupyter[server] (>
[[package]]
name = "jupyterlab"
-version = "4.1.2"
+version = "4.2.0"
description = "JupyterLab computational environment"
optional = false
python-versions = ">=3.8"
files = [
- {file = "jupyterlab-4.1.2-py3-none-any.whl", hash = "sha256:aa88193f03cf4d3555f6712f04d74112b5eb85edd7d222c588c7603a26d33c5b"},
- {file = "jupyterlab-4.1.2.tar.gz", hash = "sha256:5d6348b3ed4085181499f621b7dfb6eb0b1f57f3586857aadfc8e3bf4c4885f9"},
+ {file = "jupyterlab-4.2.0-py3-none-any.whl", hash = "sha256:0dfe9278e25a145362289c555d9beb505697d269c10e99909766af7c440ad3cc"},
+ {file = "jupyterlab-4.2.0.tar.gz", hash = "sha256:356e9205a6a2ab689c47c8fe4919dba6c076e376d03f26baadc05748c2435dd5"},
]
[package.dependencies]
@@ -1580,23 +1580,24 @@ async-lru = ">=1.0.0"
httpx = ">=0.25.0"
importlib-metadata = {version = ">=4.8.3", markers = "python_version < \"3.10\""}
importlib-resources = {version = ">=1.4", markers = "python_version < \"3.9\""}
-ipykernel = "*"
+ipykernel = ">=6.5.0"
jinja2 = ">=3.0.3"
jupyter-core = "*"
jupyter-lsp = ">=2.0.0"
jupyter-server = ">=2.4.0,<3"
-jupyterlab-server = ">=2.19.0,<3"
+jupyterlab-server = ">=2.27.1,<3"
notebook-shim = ">=0.2"
packaging = "*"
-tomli = {version = "*", markers = "python_version < \"3.11\""}
+tomli = {version = ">=1.2.2", markers = "python_version < \"3.11\""}
tornado = ">=6.2.0"
traitlets = "*"
[package.extras]
-dev = ["build", "bump2version", "coverage", "hatch", "pre-commit", "pytest-cov", "ruff (==0.2.0)"]
+dev = ["build", "bump2version", "coverage", "hatch", "pre-commit", "pytest-cov", "ruff (==0.3.5)"]
docs = ["jsx-lexer", "myst-parser", "pydata-sphinx-theme (>=0.13.0)", "pytest", "pytest-check-links", "pytest-jupyter", "sphinx (>=1.8,<7.3.0)", "sphinx-copybutton"]
-docs-screenshots = ["altair (==5.2.0)", "ipython (==8.16.1)", "ipywidgets (==8.1.1)", "jupyterlab-geojson (==3.4.0)", "jupyterlab-language-pack-zh-cn (==4.0.post6)", "matplotlib (==3.8.2)", "nbconvert (>=7.0.0)", "pandas (==2.2.0)", "scipy (==1.12.0)", "vega-datasets (==0.9.0)"]
+docs-screenshots = ["altair (==5.3.0)", "ipython (==8.16.1)", "ipywidgets (==8.1.2)", "jupyterlab-geojson (==3.4.0)", "jupyterlab-language-pack-zh-cn (==4.1.post2)", "matplotlib (==3.8.3)", "nbconvert (>=7.0.0)", "pandas (==2.2.1)", "scipy (==1.12.0)", "vega-datasets (==0.9.0)"]
test = ["coverage", "pytest (>=7.0)", "pytest-check-links (>=0.7)", "pytest-console-scripts", "pytest-cov", "pytest-jupyter (>=0.5.3)", "pytest-timeout", "pytest-tornasync", "requests", "requests-cache", "virtualenv"]
+upgrade-extension = ["copier (>=8,<10)", "jinja2-time (<0.3)", "pydantic (<2.0)", "pyyaml-include (<2.0)", "tomli-w (<2.0)"]
[[package]]
name = "jupyterlab-pygments"
@@ -1611,13 +1612,13 @@ files = [
[[package]]
name = "jupyterlab-server"
-version = "2.25.3"
+version = "2.27.1"
description = "A set of server components for JupyterLab and JupyterLab like applications."
optional = false
python-versions = ">=3.8"
files = [
- {file = "jupyterlab_server-2.25.3-py3-none-any.whl", hash = "sha256:c48862519fded9b418c71645d85a49b2f0ec50d032ba8316738e9276046088c1"},
- {file = "jupyterlab_server-2.25.3.tar.gz", hash = "sha256:846f125a8a19656611df5b03e5912c8393cea6900859baa64fa515eb64a8dc40"},
+ {file = "jupyterlab_server-2.27.1-py3-none-any.whl", hash = "sha256:f5e26156e5258b24d532c84e7c74cc212e203bff93eb856f81c24c16daeecc75"},
+ {file = "jupyterlab_server-2.27.1.tar.gz", hash = "sha256:097b5ac709b676c7284ac9c5e373f11930a561f52cd5a86e4fc7e5a9c8a8631d"},
]
[package.dependencies]
@@ -1633,7 +1634,7 @@ requests = ">=2.31"
[package.extras]
docs = ["autodoc-traits", "jinja2 (<3.2.0)", "mistune (<4)", "myst-parser", "pydata-sphinx-theme", "sphinx", "sphinx-copybutton", "sphinxcontrib-openapi (>0.8)"]
openapi = ["openapi-core (>=0.18.0,<0.19.0)", "ruamel-yaml"]
-test = ["hatch", "ipykernel", "openapi-core (>=0.18.0,<0.19.0)", "openapi-spec-validator (>=0.6.0,<0.8.0)", "pytest (>=7.0)", "pytest-console-scripts", "pytest-cov", "pytest-jupyter[server] (>=0.6.2)", "pytest-timeout", "requests-mock", "ruamel-yaml", "sphinxcontrib-spelling", "strict-rfc3339", "werkzeug"]
+test = ["hatch", "ipykernel", "openapi-core (>=0.18.0,<0.19.0)", "openapi-spec-validator (>=0.6.0,<0.8.0)", "pytest (>=7.0,<8)", "pytest-console-scripts", "pytest-cov", "pytest-jupyter[server] (>=0.6.2)", "pytest-timeout", "requests-mock", "ruamel-yaml", "sphinxcontrib-spelling", "strict-rfc3339", "werkzeug"]
[[package]]
name = "jupyterlab-widgets"
@@ -1648,13 +1649,13 @@ files = [
[[package]]
name = "jupytext"
-version = "1.16.1"
+version = "1.16.2"
description = "Jupyter notebooks as Markdown documents, Julia, Python or R scripts"
optional = false
python-versions = ">=3.8"
files = [
- {file = "jupytext-1.16.1-py3-none-any.whl", hash = "sha256:796ec4f68ada663569e5d38d4ef03738a01284bfe21c943c485bc36433898bd0"},
- {file = "jupytext-1.16.1.tar.gz", hash = "sha256:68c7b68685e870e80e60fda8286fbd6269e9c74dc1df4316df6fe46eabc94c99"},
+ {file = "jupytext-1.16.2-py3-none-any.whl", hash = "sha256:197a43fef31dca612b68b311e01b8abd54441c7e637810b16b6cb8f2ab66065e"},
+ {file = "jupytext-1.16.2.tar.gz", hash = "sha256:8627dd9becbbebd79cc4a4ed4727d89d78e606b4b464eab72357b3b029023a14"},
]
[package.dependencies]
@@ -1663,16 +1664,16 @@ mdit-py-plugins = "*"
nbformat = "*"
packaging = "*"
pyyaml = "*"
-toml = "*"
+tomli = {version = "*", markers = "python_version < \"3.11\""}
[package.extras]
-dev = ["jupytext[test-cov,test-external]"]
+dev = ["autopep8", "black", "flake8", "gitpython", "ipykernel", "isort", "jupyter-fs (<0.4.0)", "jupyter-server (!=2.11)", "nbconvert", "pre-commit", "pytest", "pytest-cov (>=2.6.1)", "pytest-randomly", "pytest-xdist", "sphinx-gallery (<0.8)"]
docs = ["myst-parser", "sphinx", "sphinx-copybutton", "sphinx-rtd-theme"]
test = ["pytest", "pytest-randomly", "pytest-xdist"]
-test-cov = ["jupytext[test-integration]", "pytest-cov (>=2.6.1)"]
-test-external = ["autopep8", "black", "flake8", "gitpython", "isort", "jupyter-fs (<0.4.0)", "jupytext[test-integration]", "pre-commit", "sphinx-gallery (<0.8)"]
-test-functional = ["jupytext[test]"]
-test-integration = ["ipykernel", "jupyter-server (!=2.11)", "jupytext[test-functional]", "nbconvert"]
+test-cov = ["ipykernel", "jupyter-server (!=2.11)", "nbconvert", "pytest", "pytest-cov (>=2.6.1)", "pytest-randomly", "pytest-xdist"]
+test-external = ["autopep8", "black", "flake8", "gitpython", "ipykernel", "isort", "jupyter-fs (<0.4.0)", "jupyter-server (!=2.11)", "nbconvert", "pre-commit", "pytest", "pytest-randomly", "pytest-xdist", "sphinx-gallery (<0.8)"]
+test-functional = ["pytest", "pytest-randomly", "pytest-xdist"]
+test-integration = ["ipykernel", "jupyter-server (!=2.11)", "nbconvert", "pytest", "pytest-randomly", "pytest-xdist"]
test-ui = ["calysto-bash"]
[[package]]
@@ -2048,13 +2049,13 @@ files = [
[[package]]
name = "mike"
-version = "2.1.0"
+version = "2.1.1"
description = "Manage multiple versions of your MkDocs-powered documentation"
optional = false
python-versions = "*"
files = [
- {file = "mike-2.1.0-py3-none-any.whl", hash = "sha256:b3885f9b9e31fc4b0d61de473750d38ac170a6b291585076effb51a806245608"},
- {file = "mike-2.1.0.tar.gz", hash = "sha256:f0b8e51cbfae1273d648ffb602a4ab3061e57972ca1cd6836df1c51c01a36eb5"},
+ {file = "mike-2.1.1-py3-none-any.whl", hash = "sha256:0b1d01a397a423284593eeb1b5f3194e37169488f929b860c9bfe95c0d5efb79"},
+ {file = "mike-2.1.1.tar.gz", hash = "sha256:f39ed39f3737da83ad0adc33e9f885092ed27f8c9e7ff0523add0480352a2c22"},
]
[package.dependencies]
@@ -2064,6 +2065,7 @@ jinja2 = ">=2.7"
mkdocs = ">=1.0"
pyparsing = ">=3.0"
pyyaml = ">=5.1"
+pyyaml-env-tag = "*"
verspec = "*"
[package.extras]
@@ -2197,13 +2199,13 @@ pygments = ">2.12.0"
[[package]]
name = "mkdocs-material"
-version = "9.5.20"
+version = "9.5.24"
description = "Documentation that simply works"
optional = false
python-versions = ">=3.8"
files = [
- {file = "mkdocs_material-9.5.20-py3-none-any.whl", hash = "sha256:ad0094a7597bcb5d0cc3e8e543a10927c2581f7f647b9bb4861600f583180f9b"},
- {file = "mkdocs_material-9.5.20.tar.gz", hash = "sha256:986eef0250d22f70fb06ce0f4eac64cc92bd797a589ec3892ce31fad976fe3da"},
+ {file = "mkdocs_material-9.5.24-py3-none-any.whl", hash = "sha256:e12cd75954c535b61e716f359cf2a5056bf4514889d17161fdebd5df4b0153c6"},
+ {file = "mkdocs_material-9.5.24.tar.gz", hash = "sha256:02d5aaba0ee755e707c3ef6e748f9acb7b3011187c0ea766db31af8905078a34"},
]
[package.dependencies]
@@ -2239,13 +2241,13 @@ files = [
[[package]]
name = "mkdocstrings"
-version = "0.25.0"
+version = "0.25.1"
description = "Automatic documentation from sources, for MkDocs."
optional = false
python-versions = ">=3.8"
files = [
- {file = "mkdocstrings-0.25.0-py3-none-any.whl", hash = "sha256:df1b63f26675fcde8c1b77e7ea996cd2f93220b148e06455428f676f5dc838f1"},
- {file = "mkdocstrings-0.25.0.tar.gz", hash = "sha256:066986b3fb5b9ef2d37c4417255a808f7e63b40ff8f67f6cab8054d903fbc91d"},
+ {file = "mkdocstrings-0.25.1-py3-none-any.whl", hash = "sha256:da01fcc2670ad61888e8fe5b60afe9fee5781017d67431996832d63e887c2e51"},
+ {file = "mkdocstrings-0.25.1.tar.gz", hash = "sha256:c3a2515f31577f311a9ee58d089e4c51fc6046dbd9e9b4c3de4c3194667fe9bf"},
]
[package.dependencies]
@@ -2483,26 +2485,26 @@ setuptools = "*"
[[package]]
name = "notebook"
-version = "7.1.3"
+version = "7.2.0"
description = "Jupyter Notebook - A web-based notebook environment for interactive computing"
optional = false
python-versions = ">=3.8"
files = [
- {file = "notebook-7.1.3-py3-none-any.whl", hash = "sha256:919b911e59f41f6e3857ce93c9d93535ba66bb090059712770e5968c07e1004d"},
- {file = "notebook-7.1.3.tar.gz", hash = "sha256:41fcebff44cf7bb9377180808bcbae066629b55d8c7722f1ebbe75ca44f9cfc1"},
+ {file = "notebook-7.2.0-py3-none-any.whl", hash = "sha256:b4752d7407d6c8872fc505df0f00d3cae46e8efb033b822adacbaa3f1f3ce8f5"},
+ {file = "notebook-7.2.0.tar.gz", hash = "sha256:34a2ba4b08ad5d19ec930db7484fb79746a1784be9e1a5f8218f9af8656a141f"},
]
[package.dependencies]
jupyter-server = ">=2.4.0,<3"
-jupyterlab = ">=4.1.1,<4.2"
-jupyterlab-server = ">=2.22.1,<3"
+jupyterlab = ">=4.2.0,<4.3"
+jupyterlab-server = ">=2.27.1,<3"
notebook-shim = ">=0.2,<0.3"
tornado = ">=6.2.0"
[package.extras]
dev = ["hatch", "pre-commit"]
docs = ["myst-parser", "nbsphinx", "pydata-sphinx-theme", "sphinx (>=1.3.6)", "sphinxcontrib-github-alt", "sphinxcontrib-spelling"]
-test = ["importlib-resources (>=5.0)", "ipykernel", "jupyter-server[test] (>=2.4.0,<3)", "jupyterlab-server[test] (>=2.22.1,<3)", "nbval", "pytest (>=7.0)", "pytest-console-scripts", "pytest-timeout", "pytest-tornasync", "requests"]
+test = ["importlib-resources (>=5.0)", "ipykernel", "jupyter-server[test] (>=2.4.0,<3)", "jupyterlab-server[test] (>=2.27.1,<3)", "nbval", "pytest (>=7.0)", "pytest-console-scripts", "pytest-timeout", "pytest-tornasync", "requests"]
[[package]]
name = "notebook-shim"
@@ -3037,13 +3039,13 @@ tomli = {version = ">=1.1.0", markers = "python_version < \"3.11\""}
[[package]]
name = "pytest"
-version = "8.2.0"
+version = "8.2.1"
description = "pytest: simple powerful testing with Python"
optional = false
python-versions = ">=3.8"
files = [
- {file = "pytest-8.2.0-py3-none-any.whl", hash = "sha256:1733f0620f6cda4095bbf0d9ff8022486e91892245bb9e7d5542c018f612f233"},
- {file = "pytest-8.2.0.tar.gz", hash = "sha256:d507d4482197eac0ba2bae2e9babf0672eb333017bcedaa5fb1a3d42c1174b3f"},
+ {file = "pytest-8.2.1-py3-none-any.whl", hash = "sha256:faccc5d332b8c3719f40283d0d44aa5cf101cec36f88cde9ed8f2bc0538612b1"},
+ {file = "pytest-8.2.1.tar.gz", hash = "sha256:5046e5b46d8e4cac199c373041f26be56fdb81eb4e67dc11d4e10811fc3408fd"},
]
[package.dependencies]
@@ -3459,13 +3461,13 @@ files = [
[[package]]
name = "requests"
-version = "2.31.0"
+version = "2.32.2"
description = "Python HTTP for Humans."
optional = false
-python-versions = ">=3.7"
+python-versions = ">=3.8"
files = [
- {file = "requests-2.31.0-py3-none-any.whl", hash = "sha256:58cd2187c01e70e6e26505bca751777aa9f2ee0b7f4300988b709f44e013003f"},
- {file = "requests-2.31.0.tar.gz", hash = "sha256:942c5a758f98d790eaed1a29cb6eefc7ffb0d1cf7af05c3d2791656dbd6ad1e1"},
+ {file = "requests-2.32.2-py3-none-any.whl", hash = "sha256:fc06670dd0ed212426dfeb94fc1b983d917c4f9847c863f313c9dfaaffb7c23c"},
+ {file = "requests-2.32.2.tar.gz", hash = "sha256:dd951ff5ecf3e3b3aa26b40703ba77495dab41da839ae72ef3c8e5d8e2433289"},
]
[package.dependencies]
@@ -3660,28 +3662,28 @@ files = [
[[package]]
name = "ruff"
-version = "0.4.2"
+version = "0.4.5"
description = "An extremely fast Python linter and code formatter, written in Rust."
optional = false
python-versions = ">=3.7"
files = [
- {file = "ruff-0.4.2-py3-none-macosx_10_12_x86_64.macosx_11_0_arm64.macosx_10_12_universal2.whl", hash = "sha256:8d14dc8953f8af7e003a485ef560bbefa5f8cc1ad994eebb5b12136049bbccc5"},
- {file = "ruff-0.4.2-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:24016ed18db3dc9786af103ff49c03bdf408ea253f3cb9e3638f39ac9cf2d483"},
- {file = "ruff-0.4.2-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:0e2e06459042ac841ed510196c350ba35a9b24a643e23db60d79b2db92af0c2b"},
- {file = "ruff-0.4.2-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:3afabaf7ba8e9c485a14ad8f4122feff6b2b93cc53cd4dad2fd24ae35112d5c5"},
- {file = "ruff-0.4.2-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:799eb468ea6bc54b95527143a4ceaf970d5aa3613050c6cff54c85fda3fde480"},
- {file = "ruff-0.4.2-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:ec4ba9436a51527fb6931a8839af4c36a5481f8c19e8f5e42c2f7ad3a49f5069"},
- {file = "ruff-0.4.2-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:6a2243f8f434e487c2a010c7252150b1fdf019035130f41b77626f5655c9ca22"},
- {file = "ruff-0.4.2-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:8772130a063f3eebdf7095da00c0b9898bd1774c43b336272c3e98667d4fb8fa"},
- {file = "ruff-0.4.2-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:6ab165ef5d72392b4ebb85a8b0fbd321f69832a632e07a74794c0e598e7a8376"},
- {file = "ruff-0.4.2-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:1f32cadf44c2020e75e0c56c3408ed1d32c024766bd41aedef92aa3ca28eef68"},
- {file = "ruff-0.4.2-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:22e306bf15e09af45ca812bc42fa59b628646fa7c26072555f278994890bc7ac"},
- {file = "ruff-0.4.2-py3-none-musllinux_1_2_i686.whl", hash = "sha256:82986bb77ad83a1719c90b9528a9dd663c9206f7c0ab69282af8223566a0c34e"},
- {file = "ruff-0.4.2-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:652e4ba553e421a6dc2a6d4868bc3b3881311702633eb3672f9f244ded8908cd"},
- {file = "ruff-0.4.2-py3-none-win32.whl", hash = "sha256:7891ee376770ac094da3ad40c116258a381b86c7352552788377c6eb16d784fe"},
- {file = "ruff-0.4.2-py3-none-win_amd64.whl", hash = "sha256:5ec481661fb2fd88a5d6cf1f83403d388ec90f9daaa36e40e2c003de66751798"},
- {file = "ruff-0.4.2-py3-none-win_arm64.whl", hash = "sha256:cbd1e87c71bca14792948c4ccb51ee61c3296e164019d2d484f3eaa2d360dfaf"},
- {file = "ruff-0.4.2.tar.gz", hash = "sha256:33bcc160aee2520664bc0859cfeaebc84bb7323becff3f303b8f1f2d81cb4edc"},
+ {file = "ruff-0.4.5-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:8f58e615dec58b1a6b291769b559e12fdffb53cc4187160a2fc83250eaf54e96"},
+ {file = "ruff-0.4.5-py3-none-macosx_11_0_arm64.whl", hash = "sha256:84dd157474e16e3a82745d2afa1016c17d27cb5d52b12e3d45d418bcc6d49264"},
+ {file = "ruff-0.4.5-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:25f483ad9d50b00e7fd577f6d0305aa18494c6af139bce7319c68a17180087f4"},
+ {file = "ruff-0.4.5-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:63fde3bf6f3ad4e990357af1d30e8ba2730860a954ea9282c95fc0846f5f64af"},
+ {file = "ruff-0.4.5-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:78e3ba4620dee27f76bbcad97067766026c918ba0f2d035c2fc25cbdd04d9c97"},
+ {file = "ruff-0.4.5-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:441dab55c568e38d02bbda68a926a3d0b54f5510095c9de7f95e47a39e0168aa"},
+ {file = "ruff-0.4.5-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:1169e47e9c4136c997f08f9857ae889d614c5035d87d38fda9b44b4338909cdf"},
+ {file = "ruff-0.4.5-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:755ac9ac2598a941512fc36a9070a13c88d72ff874a9781493eb237ab02d75df"},
+ {file = "ruff-0.4.5-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:f4b02a65985be2b34b170025a8b92449088ce61e33e69956ce4d316c0fe7cce0"},
+ {file = "ruff-0.4.5-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:75a426506a183d9201e7e5664de3f6b414ad3850d7625764106f7b6d0486f0a1"},
+ {file = "ruff-0.4.5-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:6e1b139b45e2911419044237d90b60e472f57285950e1492c757dfc88259bb06"},
+ {file = "ruff-0.4.5-py3-none-musllinux_1_2_i686.whl", hash = "sha256:a6f29a8221d2e3d85ff0c7b4371c0e37b39c87732c969b4d90f3dad2e721c5b1"},
+ {file = "ruff-0.4.5-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:d6ef817124d72b54cc923f3444828ba24fa45c3164bc9e8f1813db2f3d3a8a11"},
+ {file = "ruff-0.4.5-py3-none-win32.whl", hash = "sha256:aed8166c18b1a169a5d3ec28a49b43340949e400665555b51ee06f22813ef062"},
+ {file = "ruff-0.4.5-py3-none-win_amd64.whl", hash = "sha256:b0b03c619d2b4350b4a27e34fd2ac64d0dabe1afbf43de57d0f9d8a05ecffa45"},
+ {file = "ruff-0.4.5-py3-none-win_arm64.whl", hash = "sha256:9d15de3425f53161b3f5a5658d4522e4eee5ea002bf2ac7aa380743dd9ad5fba"},
+ {file = "ruff-0.4.5.tar.gz", hash = "sha256:286eabd47e7d4d521d199cab84deca135557e6d1e0f0d01c29e757c3cb151b54"},
]
[[package]]
@@ -3913,17 +3915,6 @@ webencodings = ">=0.4"
doc = ["sphinx", "sphinx_rtd_theme"]
test = ["flake8", "isort", "pytest"]
-[[package]]
-name = "toml"
-version = "0.10.2"
-description = "Python Library for Tom's Obvious, Minimal Language"
-optional = false
-python-versions = ">=2.6, !=3.0.*, !=3.1.*, !=3.2.*"
-files = [
- {file = "toml-0.10.2-py2.py3-none-any.whl", hash = "sha256:806143ae5bfb6a3c6e736a764057db0e6a0e05e338b5630894a5f779cabb4f9b"},
- {file = "toml-0.10.2.tar.gz", hash = "sha256:b3bda1d108d5dd99f4a20d24d9c348e91c4db7ab1b749200bded2f839ccbe68f"},
-]
-
[[package]]
name = "tomli"
version = "2.0.1"
@@ -4019,13 +4010,13 @@ test = ["argcomplete (>=3.0.3)", "mypy (>=1.7.0)", "pre-commit", "pytest (>=7.0,
[[package]]
name = "twine"
-version = "5.0.0"
+version = "5.1.0"
description = "Collection of utilities for publishing packages on PyPI"
optional = false
python-versions = ">=3.8"
files = [
- {file = "twine-5.0.0-py3-none-any.whl", hash = "sha256:a262933de0b484c53408f9edae2e7821c1c45a3314ff2df9bdd343aa7ab8edc0"},
- {file = "twine-5.0.0.tar.gz", hash = "sha256:89b0cc7d370a4b66421cc6102f269aa910fe0f1861c124f573cf2ddedbc10cf4"},
+ {file = "twine-5.1.0-py3-none-any.whl", hash = "sha256:fe1d814395bfe50cfbe27783cb74efe93abeac3f66deaeb6c8390e4e92bacb43"},
+ {file = "twine-5.1.0.tar.gz", hash = "sha256:4d74770c88c4fcaf8134d2a6a9d863e40f08255ff7d8e2acb3cbbd57d25f6e9d"},
]
[package.dependencies]
@@ -4267,4 +4258,4 @@ desktop = ["opencv-python"]
[metadata]
lock-version = "2.0"
python-versions = "^3.8"
-content-hash = "29af5aa06f97e77a2dba94c5a6d77d7d1903448724df07416026a378d3c6a64d"
+content-hash = "ad8402ec1767f9427ab38bad7dab54b302a30f9e08b6489fad224c8481745b37"
diff --git a/pyproject.toml b/pyproject.toml
index 509c05b9..ff83f5fa 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -1,6 +1,6 @@
[tool.poetry]
name = "supervision"
-version = "0.21.0rc3"
+version = "0.21.0rc5"
description = "A set of easy-to-use utils that will come in handy in any Computer Vision project"
authors = ["Piotr Skalski "]
maintainers = ["Piotr Skalski "]
@@ -42,7 +42,7 @@ pyyaml = ">=5.3"
defusedxml = "^0.7.1"
opencv-python = { version = ">=4.5.5.64", optional = true }
opencv-python-headless = ">=4.5.5.64"
-requests = { version = ">=2.26.0,<=2.31.0", optional = true }
+requests = { version = ">=2.26.0,<=2.32.2", optional = true }
tqdm = { version = ">=4.62.3,<=4.66.4", optional = true }
pillow = ">=9.4"
diff --git a/supervision/__init__.py b/supervision/__init__.py
index bb526514..abe63390 100644
--- a/supervision/__init__.py
+++ b/supervision/__init__.py
@@ -23,6 +23,7 @@ from supervision.annotators.core import (
PercentageBarAnnotator,
PixelateAnnotator,
PolygonAnnotator,
+ RichLabelAnnotator,
RoundBoxAnnotator,
TraceAnnotator,
TriangleAnnotator,
@@ -34,6 +35,7 @@ from supervision.dataset.core import (
ClassificationDataset,
DetectionDataset,
)
+from supervision.dataset.utils import mask_to_rle, rle_to_mask
from supervision.detection.annotate import BoxAnnotator
from supervision.detection.core import Detections
from supervision.detection.line_zone import LineZone, LineZoneAnnotator
@@ -46,12 +48,17 @@ from supervision.detection.utils import (
box_iou_batch,
box_non_max_suppression,
calculate_masks_centroids,
+ clip_boxes,
+ contains_holes,
+ contains_multiple_segments,
filter_polygons_by_area,
mask_iou_batch,
mask_non_max_suppression,
mask_to_polygons,
mask_to_xyxy,
move_boxes,
+ move_masks,
+ pad_boxes,
polygon_to_mask,
polygon_to_xyxy,
scale_boxes,
@@ -69,7 +76,11 @@ from supervision.draw.utils import (
)
from supervision.geometry.core import Point, Position, Rect
from supervision.geometry.utils import get_polygon_center
-from supervision.keypoint.annotators import EdgeAnnotator, VertexAnnotator
+from supervision.keypoint.annotators import (
+ EdgeAnnotator,
+ VertexAnnotator,
+ VertexLabelAnnotator,
+)
from supervision.keypoint.core import KeyPoints
from supervision.metrics.detection import ConfusionMatrix, MeanAveragePrecision
from supervision.tracker.byte_tracker.core import ByteTrack
diff --git a/supervision/annotators/core.py b/supervision/annotators/core.py
index ac901862..1c0ad8af 100644
--- a/supervision/annotators/core.py
+++ b/supervision/annotators/core.py
@@ -3,9 +3,15 @@ from typing import List, Optional, Tuple, Union
import cv2
import numpy as np
+from PIL import Image, ImageDraw, ImageFont
from supervision.annotators.base import BaseAnnotator, ImageType
-from supervision.annotators.utils import ColorLookup, Trace, resolve_color
+from supervision.annotators.utils import (
+ ColorLookup,
+ Trace,
+ resolve_color,
+ resolve_text_background_xyxy,
+)
from supervision.config import CLASS_NAME_DATA_FIELD, ORIENTED_BOX_COORDINATES
from supervision.detection.core import Detections
from supervision.detection.utils import clip_boxes, mask_to_polygons
@@ -936,59 +942,6 @@ class LabelAnnotator:
self.text_anchor: Position = text_position
self.color_lookup: ColorLookup = color_lookup
- @staticmethod
- def resolve_text_background_xyxy(
- center_coordinates: Tuple[int, int],
- text_wh: Tuple[int, int],
- position: Position,
- ) -> Tuple[int, int, int, int]:
- center_x, center_y = center_coordinates
- text_w, text_h = text_wh
-
- if position == Position.TOP_LEFT:
- return center_x, center_y - text_h, center_x + text_w, center_y
- elif position == Position.TOP_RIGHT:
- return center_x - text_w, center_y - text_h, center_x, center_y
- elif position == Position.TOP_CENTER:
- return (
- center_x - text_w // 2,
- center_y - text_h,
- center_x + text_w // 2,
- center_y,
- )
- elif position == Position.CENTER or position == Position.CENTER_OF_MASS:
- return (
- center_x - text_w // 2,
- center_y - text_h // 2,
- center_x + text_w // 2,
- center_y + text_h // 2,
- )
- elif position == Position.BOTTOM_LEFT:
- return center_x, center_y, center_x + text_w, center_y + text_h
- elif position == Position.BOTTOM_RIGHT:
- return center_x - text_w, center_y, center_x, center_y + text_h
- elif position == Position.BOTTOM_CENTER:
- return (
- center_x - text_w // 2,
- center_y,
- center_x + text_w // 2,
- center_y + text_h,
- )
- elif position == Position.CENTER_LEFT:
- return (
- center_x - text_w,
- center_y - text_h // 2,
- center_x,
- center_y + text_h // 2,
- )
- elif position == Position.CENTER_RIGHT:
- return (
- center_x,
- center_y - text_h // 2,
- center_x + text_w,
- center_y + text_h // 2,
- )
-
@convert_for_annotation_method
def annotate(
self,
@@ -1056,9 +1009,11 @@ class LabelAnnotator:
color=self.color,
detections=detections,
detection_idx=detection_idx,
- color_lookup=self.color_lookup
- if custom_color_lookup is None
- else custom_color_lookup,
+ color_lookup=(
+ self.color_lookup
+ if custom_color_lookup is None
+ else custom_color_lookup
+ ),
)
if labels is not None:
@@ -1078,7 +1033,7 @@ class LabelAnnotator:
)[0]
text_w_padded = text_w + 2 * self.text_padding
text_h_padded = text_h + 2 * self.text_padding
- text_background_xyxy = self.resolve_text_background_xyxy(
+ text_background_xyxy = resolve_text_background_xyxy(
center_coordinates=tuple(center_coordinates),
text_wh=(text_w_padded, text_h_padded),
position=self.text_anchor,
@@ -1148,6 +1103,165 @@ class LabelAnnotator:
return scene
+class RichLabelAnnotator:
+ """
+ A class for annotating labels on an image using provided detections,
+ with support for Unicode characters by using a custom font.
+ """
+
+ def __init__(
+ self,
+ color: Union[Color, ColorPalette] = ColorPalette.DEFAULT,
+ text_color: Color = Color.WHITE,
+ font_path: str = None,
+ font_size: int = 10,
+ text_padding: int = 10,
+ text_position: Position = Position.TOP_LEFT,
+ color_lookup: ColorLookup = ColorLookup.CLASS,
+ border_radius: int = 0,
+ ):
+ """
+ Args:
+ color (Union[Color, ColorPalette]): The color or color palette to use for
+ annotating the text background.
+ text_color (Color): The color to use for the text.
+ font_path (str): Path to the font file (e.g., ".ttf" or ".otf") to use for
+ rendering text. If `None`, the default PIL font will be used.
+ font_size (int): Font size for the text.
+ text_padding (int): Padding around the text within its background box.
+ text_position (Position): Position of the text relative to the detection.
+ Possible values are defined in the `Position` enum.
+ color_lookup (ColorLookup): Strategy for mapping colors to annotations.
+ Options are `INDEX`, `CLASS`, `TRACK`.
+ border_radius (int): The radius to apply round edges. If the selected
+ value is higher than the lower dimension, width or height, is clipped.
+ """
+ self.color = color
+ self.text_color = text_color
+ self.text_padding = text_padding
+ self.text_anchor = text_position
+ self.color_lookup = color_lookup
+ self.border_radius = border_radius
+ if font_path is not None:
+ try:
+ self.font = ImageFont.truetype(font_path, font_size)
+ except OSError:
+ print(f"Font path '{font_path}' not found. Using PIL's default font.")
+ self.font = ImageFont.load_default(size=font_size)
+ else:
+ self.font = ImageFont.load_default(size=font_size)
+
+ def annotate(
+ self,
+ scene: ImageType,
+ detections: Detections,
+ labels: List[str] = None,
+ custom_color_lookup: Optional[np.ndarray] = None,
+ ) -> ImageType:
+ """
+ Annotates the given scene with labels based on the provided
+ detections, with support for Unicode characters.
+
+ Args:
+ scene (ImageType): The image where labels will be drawn.
+ `ImageType` is a flexible type, accepting either `numpy.ndarray`
+ or `PIL.Image.Image`.
+ detections (Detections): Object detections to annotate.
+ labels (List[str]): Optional. Custom labels for each detection.
+ custom_color_lookup (Optional[np.ndarray]): Custom color lookup array.
+ Allows to override the default color mapping strategy.
+
+ 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(...)
+
+ labels = [
+ f"{class_name} {confidence:.2f}"
+ for class_name, confidence
+ in zip(detections['class_name'], detections.confidence)
+ ]
+
+ rich_label_annotator = sv.RichLabelAnnotator(font_path="path/to/font.ttf")
+ annotated_frame = label_annotator.annotate(
+ scene=image.copy(),
+ detections=detections,
+ labels=labels
+ )
+ ```
+
+ """
+ if isinstance(scene, np.ndarray):
+ scene = Image.fromarray(cv2.cvtColor(scene, cv2.COLOR_BGR2RGB))
+ draw = ImageDraw.Draw(scene)
+ anchors_coordinates = detections.get_anchors_coordinates(
+ anchor=self.text_anchor
+ ).astype(int)
+ if labels is not None and len(labels) != len(detections):
+ raise ValueError(
+ f"The number of labels provided ({len(labels)}) does not match the "
+ f"number of detections ({len(detections)}). Each detection should have "
+ f"a corresponding label. This discrepancy can occur if the labels and "
+ f"detections are not aligned or if an incorrect number of labels has "
+ f"been provided. Please ensure that the labels array has the same "
+ f"length as the Detections object."
+ )
+ for detection_idx, center_coordinates in enumerate(anchors_coordinates):
+ color = resolve_color(
+ color=self.color,
+ detections=detections,
+ detection_idx=detection_idx,
+ color_lookup=(
+ self.color_lookup
+ if custom_color_lookup is None
+ else custom_color_lookup
+ ),
+ )
+ if labels is not None:
+ text = labels[detection_idx]
+ elif detections[CLASS_NAME_DATA_FIELD] is not None:
+ text = detections[CLASS_NAME_DATA_FIELD][detection_idx]
+ elif detections.class_id is not None:
+ text = str(detections.class_id[detection_idx])
+ else:
+ text = str(detection_idx)
+
+ left, top, right, bottom = draw.textbbox((0, 0), text, font=self.font)
+ text_width = right - left
+ text_height = bottom - top
+ text_w_padded = text_width + 2 * self.text_padding
+ text_h_padded = text_height + 2 * self.text_padding
+ text_background_xyxy = resolve_text_background_xyxy(
+ center_coordinates=tuple(center_coordinates),
+ text_wh=(text_w_padded, text_h_padded),
+ position=self.text_anchor,
+ )
+
+ text_x = text_background_xyxy[0] + self.text_padding - left
+ text_y = text_background_xyxy[1] + self.text_padding - top
+
+ draw.rounded_rectangle(
+ text_background_xyxy,
+ radius=self.border_radius,
+ fill=color.as_rgb(),
+ outline=None,
+ )
+ draw.text(
+ xy=(text_x, text_y),
+ text=text,
+ font=self.font,
+ fill=self.text_color.as_rgb(),
+ )
+
+ return scene
+
+
class BlurAnnotator(BaseAnnotator):
"""
A class for blurring regions in an image using provided detections.
diff --git a/supervision/annotators/utils.py b/supervision/annotators/utils.py
index e206c8cb..100b7874 100644
--- a/supervision/annotators/utils.py
+++ b/supervision/annotators/utils.py
@@ -1,5 +1,5 @@
from enum import Enum
-from typing import Optional, Union
+from typing import Optional, Tuple, Union
import numpy as np
@@ -34,14 +34,14 @@ def resolve_color_idx(
) -> int:
if detection_idx >= len(detections):
raise ValueError(
- f"Detection index {detection_idx}"
+ f"Detection index {detection_idx} "
f"is out of bounds for detections of length {len(detections)}"
)
if isinstance(color_lookup, np.ndarray):
if len(color_lookup) != len(detections):
raise ValueError(
- f"Length of color lookup {len(color_lookup)}"
+ f"Length of color lookup {len(color_lookup)} "
f"does not match length of detections {len(detections)}"
)
return color_lookup[detection_idx]
@@ -50,19 +50,72 @@ def resolve_color_idx(
elif color_lookup == ColorLookup.CLASS:
if detections.class_id is None:
raise ValueError(
- "Could not resolve color by class because"
+ "Could not resolve color by class because "
"Detections do not have class_id"
)
return detections.class_id[detection_idx]
elif color_lookup == ColorLookup.TRACK:
if detections.tracker_id is None:
raise ValueError(
- "Could not resolve color by track because"
+ "Could not resolve color by track because "
"Detections do not have tracker_id"
)
return detections.tracker_id[detection_idx]
+def resolve_text_background_xyxy(
+ center_coordinates: Tuple[int, int],
+ text_wh: Tuple[int, int],
+ position: Position,
+) -> Tuple[int, int, int, int]:
+ center_x, center_y = center_coordinates
+ text_w, text_h = text_wh
+
+ if position == Position.TOP_LEFT:
+ return center_x, center_y - text_h, center_x + text_w, center_y
+ elif position == Position.TOP_RIGHT:
+ return center_x - text_w, center_y - text_h, center_x, center_y
+ elif position == Position.TOP_CENTER:
+ return (
+ center_x - text_w // 2,
+ center_y - text_h,
+ center_x + text_w // 2,
+ center_y,
+ )
+ elif position == Position.CENTER or position == Position.CENTER_OF_MASS:
+ return (
+ center_x - text_w // 2,
+ center_y - text_h // 2,
+ center_x + text_w // 2,
+ center_y + text_h // 2,
+ )
+ elif position == Position.BOTTOM_LEFT:
+ return center_x, center_y, center_x + text_w, center_y + text_h
+ elif position == Position.BOTTOM_RIGHT:
+ return center_x - text_w, center_y, center_x, center_y + text_h
+ elif position == Position.BOTTOM_CENTER:
+ return (
+ center_x - text_w // 2,
+ center_y,
+ center_x + text_w // 2,
+ center_y + text_h,
+ )
+ elif position == Position.CENTER_LEFT:
+ return (
+ center_x - text_w,
+ center_y - text_h // 2,
+ center_x,
+ center_y + text_h // 2,
+ )
+ elif position == Position.CENTER_RIGHT:
+ return (
+ center_x,
+ center_y - text_h // 2,
+ center_x + text_w,
+ center_y + text_h // 2,
+ )
+
+
def get_color_by_index(color: Union[Color, ColorPalette], idx: int) -> Color:
if isinstance(color, ColorPalette):
return color.by_idx(idx)
diff --git a/supervision/dataset/core.py b/supervision/dataset/core.py
index 551e96da..c8863df3 100644
--- a/supervision/dataset/core.py
+++ b/supervision/dataset/core.py
@@ -116,13 +116,12 @@ class DetectionDataset(BaseDataset):
Tuple[DetectionDataset, DetectionDataset]: A tuple containing
the training and testing datasets.
- Example:
+ Examples:
```python
import supervision as sv
ds = sv.DetectionDataset(...)
- train_ds, test_ds = ds.split(split_ratio=0.7,
- random_state=42, shuffle=True)
+ train_ds, test_ds = ds.split(split_ratio=0.7, random_state=42, shuffle=True)
len(train_ds), len(test_ds)
# (700, 300)
```
@@ -229,7 +228,7 @@ class DetectionDataset(BaseDataset):
DetectionDataset: A DetectionDataset instance containing
the loaded images and annotations.
- Example:
+ Examples:
```python
import roboflow
from roboflow import Roboflow
@@ -286,7 +285,7 @@ class DetectionDataset(BaseDataset):
DetectionDataset: A DetectionDataset instance
containing the loaded images and annotations.
- Example:
+ Examples:
```python
import roboflow
from roboflow import Roboflow
@@ -391,7 +390,7 @@ class DetectionDataset(BaseDataset):
DetectionDataset: A DetectionDataset instance containing
the loaded images and annotations.
- Example:
+ Examples:
```python
import roboflow
from roboflow import Roboflow
@@ -431,6 +430,20 @@ class DetectionDataset(BaseDataset):
Exports the dataset to COCO format. This method saves the
images and their corresponding annotations in COCO format.
+ !!! tip
+
+ The format of the mask is determined automatically based on its structure:
+
+ - If a mask contains multiple disconnected components or holes, it will be
+ saved using the Run-Length Encoding (RLE) format for efficient storage and
+ processing.
+ - If a mask consists of a single, contiguous region without any holes, it
+ will be encoded as a polygon, preserving the outline of the object.
+
+ This automatic selection ensures that the masks are stored in the most
+ appropriate and space-efficient format, complying with COCO dataset
+ standards.
+
Args:
images_directory_path (Optional[str]): The path to the directory
where the images should be saved.
@@ -482,7 +495,7 @@ class DetectionDataset(BaseDataset):
(DetectionDataset): A single `DetectionDataset` object containing
the merged data from the input list.
- Example:
+ Examples:
```python
import supervision as sv
@@ -567,13 +580,12 @@ class ClassificationDataset(BaseDataset):
Tuple[ClassificationDataset, ClassificationDataset]: A tuple containing
the training and testing datasets.
- Example:
+ Examples:
```python
import supervision as sv
cd = sv.ClassificationDataset(...)
- train_cd,test_cd = cd.split(split_ratio=0.7,
- random_state=42,shuffle=True)
+ train_cd,test_cd = cd.split(split_ratio=0.7, random_state=42,shuffle=True)
len(train_cd), len(test_cd)
# (700, 300)
```
@@ -635,7 +647,7 @@ class ClassificationDataset(BaseDataset):
Returns:
ClassificationDataset: The dataset.
- Example:
+ Examples:
```python
import roboflow
from roboflow import Roboflow
diff --git a/supervision/dataset/formats/coco.py b/supervision/dataset/formats/coco.py
index 4f8679d5..353e33f5 100644
--- a/supervision/dataset/formats/coco.py
+++ b/supervision/dataset/formats/coco.py
@@ -5,13 +5,20 @@ from typing import Dict, List, Tuple
import cv2
import numpy as np
+import numpy.typing as npt
from supervision.dataset.utils import (
approximate_mask_with_polygons,
map_detections_class_id,
+ mask_to_rle,
+ rle_to_mask,
)
from supervision.detection.core import Detections
-from supervision.detection.utils import polygon_to_mask
+from supervision.detection.utils import (
+ contains_holes,
+ contains_multiple_segments,
+ polygon_to_mask,
+)
from supervision.utils.file import read_json_file, save_json_file
@@ -57,13 +64,24 @@ def group_coco_annotations_by_image_id(
return annotations
-def _polygons_to_masks(
- polygons: List[np.ndarray], resolution_wh: Tuple[int, int]
-) -> np.ndarray:
+def coco_annotations_to_masks(
+ image_annotations: List[dict], resolution_wh: Tuple[int, int]
+) -> npt.NDArray[np.bool_]:
return np.array(
[
- polygon_to_mask(polygon=polygon, resolution_wh=resolution_wh)
- for polygon in polygons
+ rle_to_mask(
+ rle=np.array(image_annotation["segmentation"]["counts"]),
+ resolution_wh=resolution_wh,
+ )
+ if image_annotation["iscrowd"]
+ else polygon_to_mask(
+ polygon=np.reshape(
+ np.asarray(image_annotation["segmentation"], dtype=np.int32),
+ (-1, 2),
+ ),
+ resolution_wh=resolution_wh,
+ )
+ for image_annotation in image_annotations
],
dtype=bool,
)
@@ -83,13 +101,9 @@ def coco_annotations_to_detections(
xyxy[:, 2:4] += xyxy[:, 0:2]
if with_masks:
- polygons = [
- np.reshape(
- np.asarray(image_annotation["segmentation"], dtype=np.int32), (-1, 2)
- )
- for image_annotation in image_annotations
- ]
- mask = _polygons_to_masks(polygons=polygons, resolution_wh=resolution_wh)
+ mask = coco_annotations_to_masks(
+ image_annotations=image_annotations, resolution_wh=resolution_wh
+ )
return Detections(
class_id=np.asarray(class_ids, dtype=int), xyxy=xyxy, mask=mask
)
@@ -108,24 +122,35 @@ def detections_to_coco_annotations(
coco_annotations = []
for xyxy, mask, _, class_id, _, _ in detections:
box_width, box_height = xyxy[2] - xyxy[0], xyxy[3] - xyxy[1]
- polygon = []
+ segmentation = []
+ iscrowd = 0
if mask is not None:
- polygon = list(
- approximate_mask_with_polygons(
- mask=mask,
- min_image_area_percentage=min_image_area_percentage,
- max_image_area_percentage=max_image_area_percentage,
- approximation_percentage=approximation_percentage,
- )[0].flatten()
- )
+ iscrowd = contains_holes(mask=mask) or contains_multiple_segments(mask=mask)
+
+ if iscrowd:
+ segmentation = {
+ "counts": mask_to_rle(mask=mask),
+ "size": list(mask.shape[:2]),
+ }
+ else:
+ segmentation = [
+ list(
+ approximate_mask_with_polygons(
+ mask=mask,
+ min_image_area_percentage=min_image_area_percentage,
+ max_image_area_percentage=max_image_area_percentage,
+ approximation_percentage=approximation_percentage,
+ )[0].flatten()
+ )
+ ]
coco_annotation = {
"id": annotation_id,
"image_id": image_id,
"category_id": int(class_id),
"bbox": [xyxy[0], xyxy[1], box_width, box_height],
"area": box_width * box_height,
- "segmentation": [polygon] if polygon else [],
- "iscrowd": 0,
+ "segmentation": segmentation,
+ "iscrowd": iscrowd,
}
coco_annotations.append(coco_annotation)
annotation_id += 1
diff --git a/supervision/dataset/utils.py b/supervision/dataset/utils.py
index 05ee3201..32ece6bf 100644
--- a/supervision/dataset/utils.py
+++ b/supervision/dataset/utils.py
@@ -2,10 +2,11 @@ import copy
import os
import random
from pathlib import Path
-from typing import Dict, List, Optional, Tuple, TypeVar
+from typing import Dict, List, Optional, Tuple, TypeVar, Union
import cv2
import numpy as np
+import numpy.typing as npt
from supervision.detection.core import Detections
from supervision.detection.utils import (
@@ -129,3 +130,123 @@ def train_test_split(
split_index = int(len(data) * train_ratio)
return data[:split_index], data[split_index:]
+
+
+def rle_to_mask(
+ rle: Union[npt.NDArray[np.int_], List[int]], resolution_wh: Tuple[int, int]
+) -> npt.NDArray[np.bool_]:
+ """
+ Converts run-length encoding (RLE) to a binary mask.
+
+ Args:
+ rle (Union[npt.NDArray[np.int_], List[int]]): The 1D RLE array, the format
+ used in the COCO dataset (column-wise encoding, values of an array with
+ even indices represent the number of pixels assigned as background,
+ values of an array with odd indices represent the number of pixels
+ assigned as foreground object).
+ resolution_wh (Tuple[int, int]): The width (w) and height (h)
+ of the desired binary mask.
+
+ Returns:
+ The generated 2D Boolean mask of shape `(h, w)`, where the foreground object is
+ marked with `True`'s and the rest is filled with `False`'s.
+
+ Raises:
+ AssertionError: If the sum of pixels encoded in RLE differs from the
+ number of pixels in the expected mask (computed based on resolution_wh).
+
+ Examples:
+ ```python
+ import supervision as sv
+
+ sv.rle_to_mask([5, 2, 2, 2, 5], (4, 4))
+ # array([
+ # [False, False, False, False],
+ # [False, True, True, False],
+ # [False, True, True, False],
+ # [False, False, False, False],
+ # ])
+ ```
+ """
+ if isinstance(rle, list):
+ rle = np.array(rle, dtype=int)
+
+ width, height = resolution_wh
+
+ assert width * height == np.sum(rle), (
+ "the sum of the number of pixels in the RLE must be the same "
+ "as the number of pixels in the expected mask"
+ )
+
+ zero_one_values = np.zeros(shape=(rle.size, 1), dtype=np.uint8)
+ zero_one_values[1::2] = 1
+
+ decoded_rle = np.repeat(zero_one_values, rle, axis=0)
+ decoded_rle = np.append(
+ decoded_rle, np.zeros(width * height - len(decoded_rle), dtype=np.uint8)
+ )
+ return decoded_rle.reshape((height, width), order="F")
+
+
+def mask_to_rle(mask: npt.NDArray[np.bool_]) -> List[int]:
+ """
+ Converts a binary mask into a run-length encoding (RLE).
+
+ Args:
+ mask (npt.NDArray[np.bool_]): 2D binary mask where `True` indicates foreground
+ object and `False` indicates background.
+
+ Returns:
+ The run-length encoded mask. Values of a list with even indices
+ represent the number of pixels assigned as background (`False`), values
+ of a list with odd indices represent the number of pixels assigned
+ as foreground object (`True`).
+
+ Raises:
+ AssertionError: If input mask is not 2D or is empty.
+
+ Examples:
+ ```python
+ import numpy as np
+ import supervision as sv
+
+ mask = np.array([
+ [True, True, True, True],
+ [True, True, True, True],
+ [True, True, True, True],
+ [True, True, True, True],
+ ])
+ sv.mask_to_rle(mask)
+ # [0, 16]
+
+ mask = np.array([
+ [False, False, False, False],
+ [False, True, True, False],
+ [False, True, True, False],
+ [False, False, False, False],
+ ])
+ sv.mask_to_rle(mask)
+ # [5, 2, 2, 2, 5]
+ ```
+
+ { align=center width="800" }
+ """ # noqa E501 // docs
+ assert mask.ndim == 2, "Input mask must be 2D"
+ assert mask.size != 0, "Input mask cannot be empty"
+
+ on_value_change_indices = np.where(
+ mask.ravel(order="F") != np.roll(mask.ravel(order="F"), 1)
+ )[0]
+
+ on_value_change_indices = np.append(on_value_change_indices, mask.size)
+ # need to add 0 at the beginning when the same value is in the first and
+ # last element of the flattened mask
+ if on_value_change_indices[0] != 0:
+ on_value_change_indices = np.insert(on_value_change_indices, 0, 0)
+
+ rle = np.diff(on_value_change_indices)
+
+ if mask[0][0] == 1:
+ rle = np.insert(rle, 0, 0)
+
+ return list(rle)
diff --git a/supervision/detection/core.py b/supervision/detection/core.py
index 1900954d..d939293f 100644
--- a/supervision/detection/core.py
+++ b/supervision/detection/core.py
@@ -7,6 +7,7 @@ from typing import Any, Dict, Iterator, List, Optional, Tuple, Union
import numpy as np
from supervision.config import CLASS_NAME_DATA_FIELD, ORIENTED_BOX_COORDINATES
+from supervision.detection.lmm import LMM, from_paligemma, validate_lmm_and_kwargs
from supervision.detection.utils import (
box_non_max_suppression,
calculate_masks_centroids,
@@ -240,7 +241,7 @@ class Detections:
Class names values can be accessed using `detections["class_name"]`.
""" # noqa: E501 // docs
- if "obb" in ultralytics_results and ultralytics_results.obb is not None:
+ if hasattr(ultralytics_results, "obb") and ultralytics_results.obb is not None:
class_id = ultralytics_results.obb.cls.cpu().numpy().astype(int)
class_names = np.array([ultralytics_results.names[i] for i in class_id])
oriented_box_coordinates = ultralytics_results.obb.xyxyxyxy.cpu().numpy()
@@ -418,6 +419,9 @@ class Detections:
xyxy=mmdet_results.pred_instances.bboxes.cpu().numpy(),
confidence=mmdet_results.pred_instances.scores.cpu().numpy(),
class_id=mmdet_results.pred_instances.labels.cpu().numpy().astype(int),
+ mask=mmdet_results.pred_instances.masks.cpu().numpy()
+ if "masks" in mmdet_results.pred_instances
+ else None,
)
@classmethod
@@ -802,6 +806,52 @@ class Detections:
class_id=paddledet_result["bbox"][:, 0].astype(int),
)
+ @classmethod
+ def from_lmm(cls, lmm: Union[LMM, str], result: str, **kwargs) -> Detections:
+ """
+ Creates a Detections object from the given result string based on the specified
+ Large Multimodal Model (LMM).
+
+ Args:
+ lmm (Union[LMM, str]): The type of LMM (Large Multimodal Model) to use.
+ result (str): The result string containing the detection data.
+ **kwargs: Additional keyword arguments required by the specified LMM.
+
+ Returns:
+ Detections: A new Detections object.
+
+ Raises:
+ ValueError: If the LMM is invalid, required arguments are missing, or
+ disallowed arguments are provided.
+ ValueError: If the specified LMM is not supported.
+
+ Examples:
+ ```python
+ import supervision as sv
+
+ paligemma_result = " cat"
+ detections = sv.Detections.from_lmm(
+ sv.LMM.PALIGEMMA,
+ paligemma_result,
+ resolution_wh=(1000, 1000),
+ classes=['cat', 'dog']
+ )
+ detections.xyxy
+ # array([[250., 250., 750., 750.]])
+
+ detections.class_id
+ # array([0])
+ ```
+ """
+ lmm = validate_lmm_and_kwargs(lmm, kwargs)
+
+ if lmm == LMM.PALIGEMMA:
+ xyxy, class_id, class_name = from_paligemma(result, **kwargs)
+ data = {CLASS_NAME_DATA_FIELD: class_name}
+ return cls(xyxy=xyxy, class_id=class_id, data=data)
+
+ raise ValueError(f"Unsupported LMM: {lmm}")
+
@classmethod
def empty(cls) -> Detections:
"""
@@ -831,9 +881,10 @@ class Detections:
This method takes a list of Detections objects and combines their
respective fields (`xyxy`, `mask`, `confidence`, `class_id`, and `tracker_id`)
- into a single Detections object. If all elements in a field are not
- `None`, the corresponding field will be stacked.
- Otherwise, the field will be set to `None`.
+ into a single Detections object.
+
+ For example, if merging Detections with 3 and 4 detected objects, this method
+ will return a Detections with 7 objects (7 entries in `xyxy`, `mask`, etc).
Args:
detections_list (List[Detections]): A list of Detections objects to merge.
@@ -891,13 +942,12 @@ class Detections:
def stack_or_none(name: str):
if all(d.__getattribute__(name) is None for d in detections_list):
return None
- if any(d.__getattribute__(name) is None for d in detections_list):
- raise ValueError(f"All or none of the '{name}' fields must be None")
- return (
- np.vstack([d.__getattribute__(name) for d in detections_list])
- if name == "mask"
- else np.hstack([d.__getattribute__(name) for d in detections_list])
- )
+ stack_list = [
+ d.__getattribute__(name)
+ for d in detections_list
+ if d.__getattribute__(name) is not None
+ ]
+ return np.vstack(stack_list) if name == "mask" else np.hstack(stack_list)
mask = stack_or_none("mask")
confidence = stack_or_none("confidence")
diff --git a/supervision/detection/line_zone.py b/supervision/detection/line_zone.py
index dc1751f4..761d27c0 100644
--- a/supervision/detection/line_zone.py
+++ b/supervision/detection/line_zone.py
@@ -1,3 +1,4 @@
+import warnings
from typing import Dict, Iterable, Optional, Tuple
import cv2
@@ -7,6 +8,7 @@ from supervision.detection.core import Detections
from supervision.draw.color import Color
from supervision.draw.utils import draw_text
from supervision.geometry.core import Point, Position, Vector
+from supervision.utils.internal import SupervisionWarnings
class LineZone:
@@ -142,6 +144,15 @@ class LineZone:
if len(detections) == 0:
return crossed_in, crossed_out
+ if detections.tracker_id is None:
+ warnings.warn(
+ "Line zone counting skipped. LineZone requires tracker_id. Refer to "
+ "https://supervision.roboflow.com/latest/trackers for more "
+ "information.",
+ category=SupervisionWarnings,
+ )
+ return crossed_in, crossed_out
+
all_anchors = np.array(
[
detections.get_anchors_coordinates(anchor)
@@ -150,9 +161,6 @@ class LineZone:
)
for i, tracker_id in enumerate(detections.tracker_id):
- if tracker_id is None:
- continue
-
box_anchors = [Point(x=x, y=y) for x, y in all_anchors[:, i, :]]
in_limits = all(
diff --git a/supervision/detection/lmm.py b/supervision/detection/lmm.py
new file mode 100644
index 00000000..0278fc00
--- /dev/null
+++ b/supervision/detection/lmm.py
@@ -0,0 +1,59 @@
+import re
+from enum import Enum
+from typing import Any, Dict, List, Optional, Tuple, Union
+
+import numpy as np
+
+
+class LMM(Enum):
+ PALIGEMMA = "paligemma"
+
+
+REQUIRED_ARGUMENTS: Dict[LMM, List[str]] = {LMM.PALIGEMMA: ["resolution_wh"]}
+
+ALLOWED_ARGUMENTS: Dict[LMM, List[str]] = {LMM.PALIGEMMA: ["resolution_wh", "classes"]}
+
+
+def validate_lmm_and_kwargs(lmm: Union[LMM, str], kwargs: Dict[str, Any]) -> LMM:
+ if isinstance(lmm, str):
+ try:
+ lmm = LMM(lmm.lower())
+ except ValueError:
+ raise ValueError(
+ f"Invalid lmm value: {lmm}. Must be one of {[e.value for e in LMM]}"
+ )
+
+ required_args = REQUIRED_ARGUMENTS.get(lmm, [])
+ for arg in required_args:
+ if arg not in kwargs:
+ raise ValueError(f"Missing required argument: {arg}")
+
+ allowed_args = ALLOWED_ARGUMENTS.get(lmm, [])
+ for arg in kwargs:
+ if arg not in allowed_args:
+ raise ValueError(f"Argument {arg} is not allowed for {lmm.name}")
+
+ return lmm
+
+
+def from_paligemma(
+ result: str, resolution_wh: Tuple[int, int], classes: Optional[List[str]] = None
+) -> Tuple[np.ndarray, Optional[np.ndarray], np.ndarray]:
+ w, h = resolution_wh
+ pattern = re.compile(
+ r"(?) ([\w\s]+)"
+ )
+ matches = pattern.findall(result)
+ matches = np.array(matches) if matches else np.empty((0, 5))
+
+ xyxy, class_name = matches[:, [1, 0, 3, 2]], matches[:, 4]
+ xyxy = xyxy.astype(int) / 1024 * np.array([w, h, w, h])
+ class_name = np.char.strip(class_name.astype(str))
+ class_id = None
+
+ if classes is not None:
+ mask = np.array([name in classes for name in class_name]).astype(bool)
+ xyxy, class_name = xyxy[mask], class_name[mask]
+ class_id = np.array([classes.index(name) for name in class_name])
+
+ return xyxy, class_id, class_name
diff --git a/supervision/detection/tools/inference_slicer.py b/supervision/detection/tools/inference_slicer.py
index 7157723f..82551434 100644
--- a/supervision/detection/tools/inference_slicer.py
+++ b/supervision/detection/tools/inference_slicer.py
@@ -4,20 +4,36 @@ from typing import Callable, Optional, Tuple
import numpy as np
from supervision.detection.core import Detections
-from supervision.detection.utils import move_boxes
+from supervision.detection.utils import move_boxes, move_masks
from supervision.utils.image import crop_image
-def move_detections(detections: Detections, offset: np.array) -> Detections:
+def move_detections(
+ detections: Detections,
+ offset: np.ndarray,
+ resolution_wh: Optional[Tuple[int, int]] = None,
+) -> Detections:
"""
Args:
detections (sv.Detections): Detections object to be moved.
- offset (np.array): An array of shape `(2,)` containing offset values in format
+ offset (np.ndarray): An array of shape `(2,)` containing offset values in format
is `[dx, dy]`.
+ resolution_wh (Tuple[int, int]): The width and height of the desired mask
+ resolution. Required for segmentation detections.
+
Returns:
(sv.Detections) repositioned Detections object.
"""
detections.xyxy = move_boxes(xyxy=detections.xyxy, offset=offset)
+ if detections.mask is not None:
+ if resolution_wh is None:
+ raise ValueError(
+ "Resolution width and height are required for moving segmentation "
+ "detections. This should be the same as (width, height) of image shape."
+ )
+ detections.mask = move_masks(
+ masks=detections.mask, offset=offset, resolution_wh=resolution_wh
+ )
return detections
@@ -126,7 +142,10 @@ class InferenceSlicer:
"""
image_slice = crop_image(image=image, xyxy=offset)
detections = self.callback(image_slice)
- detections = move_detections(detections=detections, offset=offset[:2])
+ resolution_wh = (image.shape[1], image.shape[0])
+ detections = move_detections(
+ detections=detections, offset=offset[:2], resolution_wh=resolution_wh
+ )
return detections
diff --git a/supervision/detection/tools/smoother.py b/supervision/detection/tools/smoother.py
index f58f3299..5768c3e8 100644
--- a/supervision/detection/tools/smoother.py
+++ b/supervision/detection/tools/smoother.py
@@ -1,3 +1,4 @@
+import warnings
from collections import defaultdict, deque
from copy import deepcopy
from typing import Optional
@@ -5,6 +6,7 @@ from typing import Optional
import numpy as np
from supervision.detection.core import Detections
+from supervision.utils.internal import SupervisionWarnings
class DetectionsSmoother:
@@ -70,16 +72,16 @@ class DetectionsSmoother:
"""
if detections.tracker_id is None:
- print(
+ warnings.warn(
"Smoothing skipped. DetectionsSmoother requires tracker_id. Refer to "
- "https://supervision.roboflow.com/latest/trackers for more information."
+ "https://supervision.roboflow.com/latest/trackers for more "
+ "information.",
+ category=SupervisionWarnings,
)
return detections
for detection_idx in range(len(detections)):
tracker_id = detections.tracker_id[detection_idx]
- if tracker_id is None:
- continue
self.tracks[tracker_id].append(detections[detection_idx])
diff --git a/supervision/detection/utils.py b/supervision/detection/utils.py
index 3eeba5b4..aac0d627 100644
--- a/supervision/detection/utils.py
+++ b/supervision/detection/utils.py
@@ -3,6 +3,7 @@ from typing import Dict, List, Optional, Tuple, Union
import cv2
import numpy as np
+import numpy.typing as npt
from supervision.config import CLASS_NAME_DATA_FIELD
@@ -56,7 +57,9 @@ 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)
- return area_inter / (area_true[:, None] + area_detection - area_inter)
+ ious = area_inter / (area_true[:, None] + area_detection - area_inter)
+ ious = np.nan_to_num(ious)
+ return ious
def _mask_iou_batch_split(
@@ -297,6 +300,35 @@ def clip_boxes(xyxy: np.ndarray, resolution_wh: Tuple[int, int]) -> np.ndarray:
return result
+def pad_boxes(xyxy: np.ndarray, px: int, py: Optional[int] = None) -> np.ndarray:
+ """
+ Pads bounding boxes coordinates with a constant padding.
+
+ Args:
+ xyxy (np.ndarray): A numpy array of shape `(N, 4)` where each
+ row corresponds to a bounding box in the format
+ `(x_min, y_min, x_max, y_max)`.
+ px (int): The padding value to be added to both the left and right sides of
+ each bounding box.
+ py (Optional[int]): The padding value to be added to both the top and bottom
+ sides of each bounding box. If not provided, `px` will be used for both
+ dimensions.
+
+ Returns:
+ np.ndarray: A numpy array of shape `(N, 4)` where each row corresponds to a
+ bounding box with coordinates padded according to the provided padding
+ values.
+ """
+ if py is None:
+ py = px
+
+ result = xyxy.copy()
+ result[:, [0, 1]] -= [px, py]
+ result[:, [2, 3]] += [px, py]
+
+ return result
+
+
def xywh_to_xyxy(boxes_xywh: np.ndarray) -> np.ndarray:
xyxy = boxes_xywh.copy()
xyxy[:, 2] = boxes_xywh[:, 0] + boxes_xywh[:, 2]
@@ -500,7 +532,7 @@ def process_roboflow_result(
np.ndarray,
Optional[np.ndarray],
Optional[np.ndarray],
- Dict[str, List[np.ndarray]],
+ Dict[str, Union[List[np.ndarray], np.ndarray]],
]:
if not roboflow_result["predictions"]:
return (
@@ -574,24 +606,61 @@ def move_boxes(xyxy: np.ndarray, offset: np.ndarray) -> np.ndarray:
Returns:
np.ndarray: Repositioned bounding boxes.
- Example:
+ Examples:
```python
import numpy as np
import supervision as sv
- boxes = np.array([[10, 10, 20, 20], [30, 30, 40, 40]])
+ xyxy = np.array([
+ [10, 10, 20, 20],
+ [30, 30, 40, 40]
+ ])
offset = np.array([5, 5])
- moved_box = sv.move_boxes(boxes, offset)
- print(moved_box)
- # np.array([
+
+ sv.move_boxes(xyxy=xyxy, offset=offset)
+ # array([
# [15, 15, 25, 25],
- # [35, 35, 45, 45]
+ # [35, 35, 45, 45]
# ])
```
"""
return xyxy + np.hstack([offset, offset])
+def move_masks(
+ masks: np.ndarray,
+ offset: np.ndarray,
+ resolution_wh: Tuple[int, int] = None,
+) -> np.ndarray:
+ """
+ Offset the masks in an array by the specified (x, y) amount.
+
+ Args:
+ 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.
+ offset (np.ndarray): An array of shape `(2,)` containing non-negative int values
+ `[dx, dy]`.
+ resolution_wh (Tuple[int, int]): The width and height of the desired mask
+ resolution.
+
+ Returns:
+ (np.ndarray) repositioned masks, optionally padded to the specified shape.
+ """
+
+ if offset[0] < 0 or offset[1] < 0:
+ raise ValueError(f"Offset values must be non-negative integers. Got: {offset}")
+
+ mask_array = np.full((masks.shape[0], resolution_wh[1], resolution_wh[0]), False)
+ mask_array[
+ :,
+ offset[1] : masks.shape[1] + offset[1],
+ offset[0] : masks.shape[2] + offset[0],
+ ] = masks
+
+ return mask_array
+
+
def scale_boxes(xyxy: np.ndarray, factor: float) -> np.ndarray:
"""
Scale the dimensions of bounding boxes.
@@ -606,16 +675,18 @@ def scale_boxes(xyxy: np.ndarray, factor: float) -> np.ndarray:
Returns:
np.ndarray: Scaled bounding boxes.
- Example:
+ Examples:
```python
import numpy as np
import supervision as sv
- boxes = np.array([[10, 10, 20, 20], [30, 30, 40, 40]])
- factor = 1.5
- scaled_bb = sv.scale_boxes(boxes, factor)
- print(scaled_bb)
- # np.array([
+ xyxy = np.array([
+ [10, 10, 20, 20],
+ [30, 30, 40, 40]
+ ])
+
+ scaled_bb = sv.scale_boxes(xyxy=xyxy, factor=1.5)
+ # array([
# [ 7.5, 7.5, 22.5, 22.5],
# [27.5, 27.5, 42.5, 42.5]
# ])
@@ -678,7 +749,9 @@ def merge_data(
Merges the data payloads of a list of Detections instances.
Args:
- data_list: The data payloads of the instances.
+ data_list: The data payloads of the Detections instances. Each data payload
+ is a dictionary with the same keys, and the values are either lists or
+ np.ndarray.
Returns:
A single data payload containing the merged data, preserving the original data
@@ -691,10 +764,6 @@ def merge_data(
if not data_list:
return {}
- all_keys_sets = [set(data.keys()) for data in data_list]
- if not all(keys_set == all_keys_sets[0] for keys_set in all_keys_sets):
- raise ValueError("All data dictionaries must have the same keys to merge.")
-
for data in data_list:
lengths = [len(value) for value in data.values()]
if len(set(lengths)) > 1:
@@ -702,10 +771,23 @@ def merge_data(
"All data values within a single object must have equal length."
)
- merged_data = {key: [] for key in all_keys_sets[0]}
+ keys_by_data = [set(data.keys()) for data in data_list]
+ keys_by_data = [keys for keys in keys_by_data if len(keys) > 0]
+ if not keys_by_data:
+ return {}
+ common_keys = set.intersection(*keys_by_data)
+ all_keys = set.union(*keys_by_data)
+ if common_keys != all_keys:
+ raise ValueError(
+ f"All sv.Detections.data dictionaries must have the same keys. Common "
+ f"keys: {common_keys}, but some dictionaries have additional keys: "
+ f"{all_keys.difference(common_keys)}."
+ )
+
+ merged_data = {key: [] for key in all_keys}
for data in data_list:
- for key in merged_data:
+ for key in data:
merged_data[key].append(data[key])
for key in merged_data:
@@ -766,3 +848,121 @@ def get_data_item(
raise TypeError(f"Unsupported data type for key '{key}': {type(value)}")
return subset_data
+
+
+def contains_holes(mask: npt.NDArray[np.bool_]) -> bool:
+ """
+ Checks if the binary mask contains holes (background pixels fully enclosed by
+ foreground pixels).
+
+ Args:
+ mask (npt.NDArray[np.bool_]): 2D binary mask where `True` indicates foreground
+ object and `False` indicates background.
+
+ Returns:
+ True if holes are detected, False otherwise.
+
+ Examples:
+ ```python
+ import numpy as np
+ import supervision as sv
+
+ mask = np.array([
+ [0, 0, 0, 0, 0],
+ [0, 1, 1, 1, 0],
+ [0, 1, 0, 1, 0],
+ [0, 1, 1, 1, 0],
+ [0, 0, 0, 0, 0]
+ ]).astype(bool)
+
+ sv.contains_holes(mask=mask)
+ # True
+
+ mask = np.array([
+ [0, 0, 0, 0, 0],
+ [0, 1, 1, 1, 0],
+ [0, 1, 1, 1, 0],
+ [0, 1, 1, 1, 0],
+ [0, 0, 0, 0, 0]
+ ]).astype(bool)
+
+ sv.contains_holes(mask=mask)
+ # False
+ ```
+
+ { align=center width="800" }
+ """ # noqa E501 // docs
+ mask_uint8 = mask.astype(np.uint8)
+ _, hierarchy = cv2.findContours(mask_uint8, cv2.RETR_CCOMP, cv2.CHAIN_APPROX_SIMPLE)
+
+ if hierarchy is not None:
+ parent_contour_index = 3
+ for h in hierarchy[0]:
+ if h[parent_contour_index] != -1:
+ return True
+ return False
+
+
+def contains_multiple_segments(
+ mask: npt.NDArray[np.bool_], connectivity: int = 4
+) -> bool:
+ """
+ Checks if the binary mask contains multiple unconnected foreground segments.
+
+ Args:
+ mask (npt.NDArray[np.bool_]): 2D binary mask where `True` indicates foreground
+ object and `False` indicates background.
+ connectivity (int) : Default: 4 is 4-way connectivity, which means that
+ foreground pixels are the part of the same segment/component
+ if their edges touch.
+ Alternatively: 8 for 8-way connectivity, when foreground pixels are
+ connected by their edges or corners touch.
+
+ Returns:
+ True when the mask contains multiple not connected components, False otherwise.
+
+ Raises:
+ ValueError: If connectivity(int) parameter value is not 4 or 8.
+
+ Examples:
+ ```python
+ import numpy as np
+ import supervision as sv
+
+ mask = np.array([
+ [0, 0, 0, 0, 0, 0],
+ [0, 1, 1, 0, 1, 1],
+ [0, 1, 1, 0, 1, 1],
+ [0, 0, 0, 0, 0, 0],
+ [0, 1, 1, 1, 0, 0],
+ [0, 1, 1, 1, 0, 0]
+ ]).astype(bool)
+
+ sv.contains_multiple_segments(mask=mask, connectivity=4)
+ # True
+
+ mask = np.array([
+ [0, 0, 0, 0, 0, 0],
+ [0, 1, 1, 1, 1, 1],
+ [0, 1, 1, 1, 1, 1],
+ [0, 1, 1, 1, 1, 1],
+ [0, 1, 1, 1, 1, 1],
+ [0, 0, 0, 0, 0, 0]
+ ]).astype(bool)
+
+ sv.contains_multiple_segments(mask=mask, connectivity=4)
+ # False
+ ```
+
+ { align=center width="800" }
+ """ # noqa E501 // docs
+ if connectivity != 4 and connectivity != 8:
+ raise ValueError(
+ "Incorrect connectivity value. Possible connectivity values: 4 or 8."
+ )
+ mask_uint8 = mask.astype(np.uint8)
+ labels = np.zeros_like(mask_uint8, dtype=np.int32)
+ number_of_labels, _ = cv2.connectedComponents(
+ mask_uint8, labels, connectivity=connectivity
+ )
+ return number_of_labels > 2
diff --git a/supervision/draw/utils.py b/supervision/draw/utils.py
index 638e6b75..6783ae25 100644
--- a/supervision/draw/utils.py
+++ b/supervision/draw/utils.py
@@ -81,6 +81,58 @@ def draw_filled_rectangle(scene: np.ndarray, rect: Rect, color: Color) -> np.nda
return scene
+def draw_rounded_rectangle(
+ scene: np.ndarray,
+ rect: Rect,
+ color: Color,
+ border_radius: int,
+) -> np.ndarray:
+ """
+ Draws a rounded rectangle on an image.
+
+ Parameters:
+ scene (np.ndarray): The image on which the rounded rectangle will be drawn.
+ rect (Rect): The rectangle to be drawn.
+ color (Color): The color of the rounded rectangle.
+ border_radius (int): The radius of the corner rounding.
+
+ Returns:
+ np.ndarray: The image with the rounded rectangle drawn on it.
+ """
+ x1, y1, x2, y2 = rect.as_xyxy_int_tuple()
+ width, height = x2 - x1, y2 - y1
+ border_radius = min(border_radius, min(width, height) // 2)
+
+ rectangle_coordinates = [
+ ((x1 + border_radius, y1), (x2 - border_radius, y2)),
+ ((x1, y1 + border_radius), (x2, y2 - border_radius)),
+ ]
+ circle_centers = [
+ (x1 + border_radius, y1 + border_radius),
+ (x2 - border_radius, y1 + border_radius),
+ (x1 + border_radius, y2 - border_radius),
+ (x2 - border_radius, y2 - border_radius),
+ ]
+
+ for coordinates in rectangle_coordinates:
+ cv2.rectangle(
+ img=scene,
+ pt1=coordinates[0],
+ pt2=coordinates[1],
+ color=color.as_bgr(),
+ thickness=-1,
+ )
+ for center in circle_centers:
+ cv2.circle(
+ img=scene,
+ center=center,
+ radius=border_radius,
+ color=color.as_bgr(),
+ thickness=-1,
+ )
+ return scene
+
+
def draw_polygon(
scene: np.ndarray, polygon: np.ndarray, color: Color, thickness: int = 2
) -> np.ndarray:
diff --git a/supervision/geometry/core.py b/supervision/geometry/core.py
index 39d42c60..a884a9da 100644
--- a/supervision/geometry/core.py
+++ b/supervision/geometry/core.py
@@ -98,6 +98,11 @@ class Rect:
width: float
height: float
+ @classmethod
+ def from_xyxy(cls, xyxy: Tuple[float, float, float, float]) -> Rect:
+ x1, y1, x2, y2 = xyxy
+ return cls(x=x1, y=y1, width=x2 - x1, height=y2 - y1)
+
@property
def top_left(self) -> Point:
return Point(x=self.x, y=self.y)
@@ -113,3 +118,11 @@ class Rect:
width=self.width + 2 * padding,
height=self.height + 2 * padding,
)
+
+ def as_xyxy_int_tuple(self) -> Tuple[int, int, int, int]:
+ return (
+ int(self.x),
+ int(self.y),
+ int(self.x + self.width),
+ int(self.y + self.height),
+ )
diff --git a/supervision/keypoint/annotators.py b/supervision/keypoint/annotators.py
index 4b43765c..e6ff1fcf 100644
--- a/supervision/keypoint/annotators.py
+++ b/supervision/keypoint/annotators.py
@@ -1,12 +1,14 @@
from abc import ABC, abstractmethod
from logging import warn
-from typing import List, Optional, Tuple
+from typing import List, Optional, Tuple, Union
import cv2
import numpy as np
+from supervision import Rect, pad_boxes
from supervision.annotators.base import ImageType
from supervision.draw.color import Color
+from supervision.draw.utils import draw_rounded_rectangle
from supervision.keypoint.core import KeyPoints
from supervision.keypoint.skeletons import SKELETONS_BY_VERTEX_COUNT
from supervision.utils.conversion import convert_for_annotation_method
@@ -46,8 +48,8 @@ class VertexAnnotator(BaseKeyPointAnnotator):
points. It draws circles at each key point location.
Args:
- scene (ImageType): The image where bounding boxes will be drawn. `ImageType`
- is a flexible type, accepting either `numpy.ndarray` or
+ scene (ImageType): The image where skeleton vertices will be drawn.
+ `ImageType` is a flexible type, accepting either `numpy.ndarray` or
`PIL.Image.Image`.
key_points (KeyPoints): A collection of key points where each key point
consists of x and y coordinates.
@@ -63,7 +65,10 @@ class VertexAnnotator(BaseKeyPointAnnotator):
image = ...
key_points = sv.KeyPoints(...)
- vertex_annotator = sv.VertexAnnotator(color=sv.Color.GREEN, radius=10)
+ vertex_annotator = sv.VertexAnnotator(
+ color=sv.Color.GREEN,
+ radius=10
+ )
annotated_frame = vertex_annotator.annotate(
scene=image.copy(),
key_points=key_points
@@ -119,7 +124,7 @@ class EdgeAnnotator(BaseKeyPointAnnotator):
edges.
Args:
- scene (ImageType): The image where bounding boxes will be drawn. `ImageType`
+ scene (ImageType): The image where skeleton edges will be drawn. `ImageType`
is a flexible type, accepting either `numpy.ndarray` or
`PIL.Image.Image`.
key_points (KeyPoints): A collection of key points where each key point
@@ -137,7 +142,10 @@ class EdgeAnnotator(BaseKeyPointAnnotator):
image = ...
key_points = sv.KeyPoints(...)
- edge_annotator = sv.EdgeAnnotator(color=sv.Color.GREEN, thickness=5)
+ edge_annotator = sv.EdgeAnnotator(
+ color=sv.Color.GREEN,
+ thickness=5
+ )
annotated_frame = edge_annotator.annotate(
scene=image.copy(),
key_points=key_points
@@ -175,3 +183,236 @@ class EdgeAnnotator(BaseKeyPointAnnotator):
)
return scene
+
+
+class VertexLabelAnnotator:
+ """
+ A class that draws labels of skeleton vertices on images. It uses specified key
+ points to determine the locations where the vertices should be drawn.
+ """
+
+ def __init__(
+ self,
+ color: Union[Color, List[Color]] = Color.ROBOFLOW,
+ text_color: Color = Color.WHITE,
+ text_scale: float = 0.5,
+ text_thickness: int = 1,
+ text_padding: int = 10,
+ border_radius: int = 0,
+ ):
+ """
+ Args:
+ 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_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.
+ border_radius (int, optional): The radius of the rounded corners of the
+ boxes. Set to a high value to produce circles.
+ """
+ self.border_radius: int = border_radius
+ self.color: Union[Color, List[Color]] = color
+ self.text_color: Color = text_color
+ self.text_scale: float = text_scale
+ self.text_thickness: int = text_thickness
+ self.text_padding: int = text_padding
+
+ def annotate(
+ self, scene: ImageType, key_points: KeyPoints, labels: List[str] = None
+ ) -> ImageType:
+ """
+ A class that draws labels of skeleton vertices on images. It uses specified key
+ points to determine the locations where the vertices should be drawn.
+
+ Args:
+ scene (ImageType): The image where vertex labels will be drawn. `ImageType`
+ is a flexible type, accepting either `numpy.ndarray` or
+ `PIL.Image.Image`.
+ key_points (KeyPoints): A collection of key points where each key point
+ consists of x and y coordinates.
+ labels (List[str], optional): A list of labels to be displayed on the
+ annotated image. If not provided, keypoint indices will be used.
+
+ Returns:
+ The annotated image, matching the type of `scene` (`numpy.ndarray`
+ or `PIL.Image.Image`)
+
+ Example:
+ ```python
+ import supervision as sv
+
+ image = ...
+ key_points = sv.KeyPoints(...)
+
+ vertex_label_annotator = sv.VertexLabelAnnotator(
+ color=sv.Color.GREEN,
+ text_color=sv.Color.BLACK,
+ border_radius=5
+ )
+ annotated_frame = vertex_label_annotator.annotate(
+ scene=image.copy(),
+ key_points=key_points
+ )
+ ```
+
+ 
+
+ !!! tip
+
+ `VertexLabelAnnotator` allows to customize the color of each keypoint label
+ values.
+
+ Example:
+ ```python
+ import supervision as sv
+
+ image = ...
+ key_points = sv.KeyPoints(...)
+
+ LABELS = [
+ "nose", "left eye", "right eye", "left ear",
+ "right ear", "left shoulder", "right shoulder", "left elbow",
+ "right elbow", "left wrist", "right wrist", "left hip",
+ "right hip", "left knee", "right knee", "left ankle",
+ "right ankle"
+ ]
+
+ COLORS = [
+ "#FF6347", "#FF6347", "#FF6347", "#FF6347",
+ "#FF6347", "#FF1493", "#00FF00", "#FF1493",
+ "#00FF00", "#FF1493", "#00FF00", "#FFD700",
+ "#00BFFF", "#FFD700", "#00BFFF", "#FFD700",
+ "#00BFFF"
+ ]
+ COLORS = [sv.Color.from_hex(color_hex=c) for c in COLORS]
+
+ vertex_label_annotator = sv.VertexLabelAnnotator(
+ color=COLORS,
+ text_color=sv.Color.BLACK,
+ border_radius=5
+ )
+ annotated_frame = vertex_label_annotator.annotate(
+ scene=image.copy(),
+ key_points=key_points,
+ labels=labels
+ )
+ ```
+ 
+ """
+ font = cv2.FONT_HERSHEY_SIMPLEX
+
+ skeletons_count, points_count, _ = key_points.xy.shape
+ if skeletons_count == 0:
+ return scene
+
+ anchors = key_points.xy.reshape(points_count * skeletons_count, 2).astype(int)
+ mask = np.all(anchors != 0, axis=1)
+
+ if not np.any(mask):
+ return scene
+
+ colors = self.preprocess_and_validate_colors(
+ colors=self.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]
+ labels = labels[mask]
+
+ xyxy = np.array(
+ [
+ self.get_text_bounding_box(
+ text=label,
+ font=font,
+ text_scale=self.text_scale,
+ text_thickness=self.text_thickness,
+ center_coordinates=tuple(anchor),
+ )
+ for anchor, label in zip(anchors, labels)
+ ]
+ )
+
+ xyxy_padded = pad_boxes(xyxy=xyxy, px=self.text_padding)
+
+ for text, color, box, box_padded in zip(labels, colors, xyxy, xyxy_padded):
+ draw_rounded_rectangle(
+ scene=scene,
+ rect=Rect.from_xyxy(box_padded),
+ color=color,
+ border_radius=self.border_radius,
+ )
+ cv2.putText(
+ img=scene,
+ text=text,
+ org=(box[0], box[1] + self.text_padding),
+ fontFace=font,
+ fontScale=self.text_scale,
+ color=self.text_color.as_rgb(),
+ thickness=self.text_thickness,
+ lineType=cv2.LINE_AA,
+ )
+
+ return scene
+
+ @staticmethod
+ def get_text_bounding_box(
+ text: str,
+ font: int,
+ text_scale: float,
+ text_thickness: int,
+ center_coordinates: Tuple[int, int],
+ ) -> Tuple[int, int, int, int]:
+ text_w, text_h = cv2.getTextSize(
+ text=text,
+ fontFace=font,
+ fontScale=text_scale,
+ thickness=text_thickness,
+ )[0]
+ center_x, center_y = center_coordinates
+ return (
+ center_x - text_w // 2,
+ center_y - text_h // 2,
+ center_x + text_w // 2,
+ center_y + text_h // 2,
+ )
+
+ @staticmethod
+ def preprocess_and_validate_labels(
+ labels: Optional[List[str]], points_count: int, skeletons_count: int
+ ) -> np.array:
+ if labels and len(labels) != points_count:
+ raise ValueError(
+ f"Number of labels ({len(labels)}) must match number of key points "
+ f"({points_count})."
+ )
+ if labels is None:
+ labels = [str(i) for i in range(points_count)]
+
+ return np.array(labels * skeletons_count)
+
+ @staticmethod
+ def preprocess_and_validate_colors(
+ colors: Optional[Union[Color, List[Color]]],
+ points_count: int,
+ skeletons_count: int,
+ ) -> np.array:
+ if isinstance(colors, list) and len(colors) != points_count:
+ raise ValueError(
+ f"Number of colors ({len(colors)}) must match number of key points "
+ f"({points_count})."
+ )
+ return (
+ np.array(colors * skeletons_count)
+ if isinstance(colors, list)
+ else np.array([colors] * points_count * skeletons_count)
+ )
diff --git a/supervision/tracker/byte_tracker/core.py b/supervision/tracker/byte_tracker/core.py
index 55db6293..ce3bbbbf 100644
--- a/supervision/tracker/byte_tracker/core.py
+++ b/supervision/tracker/byte_tracker/core.py
@@ -487,7 +487,7 @@ class ByteTrack:
self.lost_tracks = sub_tracks(self.lost_tracks, self.tracked_tracks)
self.lost_tracks.extend(lost_stracks)
self.lost_tracks = sub_tracks(self.lost_tracks, self.removed_tracks)
- self.removed_tracks.extend(removed_stracks)
+ self.removed_tracks = removed_stracks
self.tracked_tracks, self.lost_tracks = remove_duplicate_tracks(
self.tracked_tracks, self.lost_tracks
)
diff --git a/test/dataset/formats/test_coco.py b/test/dataset/formats/test_coco.py
index 62d1b75b..7e269dae 100644
--- a/test/dataset/formats/test_coco.py
+++ b/test/dataset/formats/test_coco.py
@@ -1,5 +1,5 @@
from contextlib import ExitStack as DoesNotRaise
-from typing import Dict, List, Tuple
+from typing import Dict, List, Tuple, Union
import numpy as np
import pytest
@@ -10,24 +10,30 @@ from supervision.dataset.formats.coco import (
classes_to_coco_categories,
coco_annotations_to_detections,
coco_categories_to_classes,
+ detections_to_coco_annotations,
group_coco_annotations_by_image_id,
)
-def mock_cock_coco_annotation(
+def mock_coco_annotation(
annotation_id: int = 0,
image_id: int = 0,
category_id: int = 0,
bbox: Tuple[float, float, float, float] = (0.0, 0.0, 0.0, 0.0),
area: float = 0.0,
+ segmentation: Union[List[list], Dict] = None,
+ iscrowd: bool = False,
) -> dict:
+ if not segmentation:
+ segmentation = []
return {
"id": annotation_id,
"image_id": image_id,
"category_id": category_id,
"bbox": list(bbox),
"area": area,
- "iscrowd": 0,
+ "segmentation": segmentation,
+ "iscrowd": int(iscrowd),
}
@@ -101,74 +107,46 @@ def test_classes_to_coco_categories_and_back_to_classes(
[
([], {}, DoesNotRaise()), # empty coco annotations
(
- [mock_cock_coco_annotation(annotation_id=0, image_id=0, category_id=0)],
- {
- 0: [
- mock_cock_coco_annotation(
- annotation_id=0, image_id=0, category_id=0
- )
- ]
- },
+ [mock_coco_annotation(annotation_id=0, image_id=0, category_id=0)],
+ {0: [mock_coco_annotation(annotation_id=0, image_id=0, category_id=0)]},
DoesNotRaise(),
), # single coco annotation
(
[
- mock_cock_coco_annotation(annotation_id=0, image_id=0, category_id=0),
- mock_cock_coco_annotation(annotation_id=1, image_id=1, category_id=0),
+ mock_coco_annotation(annotation_id=0, image_id=0, category_id=0),
+ mock_coco_annotation(annotation_id=1, image_id=1, category_id=0),
],
{
- 0: [
- mock_cock_coco_annotation(
- annotation_id=0, image_id=0, category_id=0
- )
- ],
- 1: [
- mock_cock_coco_annotation(
- annotation_id=1, image_id=1, category_id=0
- )
- ],
+ 0: [mock_coco_annotation(annotation_id=0, image_id=0, category_id=0)],
+ 1: [mock_coco_annotation(annotation_id=1, image_id=1, category_id=0)],
},
DoesNotRaise(),
), # two coco annotations
(
[
- mock_cock_coco_annotation(annotation_id=0, image_id=0, category_id=0),
- mock_cock_coco_annotation(annotation_id=1, image_id=1, category_id=1),
- mock_cock_coco_annotation(annotation_id=2, image_id=1, category_id=2),
- mock_cock_coco_annotation(annotation_id=3, image_id=2, category_id=3),
- mock_cock_coco_annotation(annotation_id=4, image_id=3, category_id=1),
- mock_cock_coco_annotation(annotation_id=5, image_id=3, category_id=2),
- mock_cock_coco_annotation(annotation_id=5, image_id=3, category_id=3),
+ mock_coco_annotation(annotation_id=0, image_id=0, category_id=0),
+ mock_coco_annotation(annotation_id=1, image_id=1, category_id=1),
+ mock_coco_annotation(annotation_id=2, image_id=1, category_id=2),
+ mock_coco_annotation(annotation_id=3, image_id=2, category_id=3),
+ mock_coco_annotation(annotation_id=4, image_id=3, category_id=1),
+ mock_coco_annotation(annotation_id=5, image_id=3, category_id=2),
+ mock_coco_annotation(annotation_id=5, image_id=3, category_id=3),
],
{
0: [
- mock_cock_coco_annotation(
- annotation_id=0, image_id=0, category_id=0
- ),
+ mock_coco_annotation(annotation_id=0, image_id=0, category_id=0),
],
1: [
- mock_cock_coco_annotation(
- annotation_id=1, image_id=1, category_id=1
- ),
- mock_cock_coco_annotation(
- annotation_id=2, image_id=1, category_id=2
- ),
+ mock_coco_annotation(annotation_id=1, image_id=1, category_id=1),
+ mock_coco_annotation(annotation_id=2, image_id=1, category_id=2),
],
2: [
- mock_cock_coco_annotation(
- annotation_id=3, image_id=2, category_id=3
- ),
+ mock_coco_annotation(annotation_id=3, image_id=2, category_id=3),
],
3: [
- mock_cock_coco_annotation(
- annotation_id=4, image_id=3, category_id=1
- ),
- mock_cock_coco_annotation(
- annotation_id=5, image_id=3, category_id=2
- ),
- mock_cock_coco_annotation(
- annotation_id=5, image_id=3, category_id=3
- ),
+ mock_coco_annotation(annotation_id=4, image_id=3, category_id=1),
+ mock_coco_annotation(annotation_id=5, image_id=3, category_id=2),
+ mock_coco_annotation(annotation_id=5, image_id=3, category_id=3),
],
},
DoesNotRaise(),
@@ -195,7 +173,7 @@ def test_group_coco_annotations_by_image_id(
), # empty image annotations
(
[
- mock_cock_coco_annotation(
+ mock_coco_annotation(
category_id=0, bbox=(0, 0, 100, 100), area=100 * 100
)
],
@@ -209,10 +187,10 @@ def test_group_coco_annotations_by_image_id(
), # single image annotations
(
[
- mock_cock_coco_annotation(
+ mock_coco_annotation(
category_id=0, bbox=(0, 0, 100, 100), area=100 * 100
),
- mock_cock_coco_annotation(
+ mock_coco_annotation(
category_id=0, bbox=(100, 100, 100, 100), area=100 * 100
),
],
@@ -226,6 +204,156 @@ def test_group_coco_annotations_by_image_id(
),
DoesNotRaise(),
), # two image annotations
+ (
+ [
+ mock_coco_annotation(
+ category_id=0,
+ bbox=(0, 0, 5, 5),
+ area=5 * 5,
+ segmentation=[[0, 0, 2, 0, 2, 2, 4, 2, 4, 4, 0, 4]],
+ )
+ ],
+ (5, 5),
+ True,
+ Detections(
+ xyxy=np.array([[0, 0, 5, 5]], dtype=np.float32),
+ class_id=np.array([0], dtype=int),
+ mask=np.array(
+ [
+ [
+ [1, 1, 1, 0, 0],
+ [1, 1, 1, 0, 0],
+ [1, 1, 1, 1, 1],
+ [1, 1, 1, 1, 1],
+ [1, 1, 1, 1, 1],
+ ]
+ ]
+ ),
+ ),
+ DoesNotRaise(),
+ ), # single image annotations with mask as polygon
+ (
+ [
+ mock_coco_annotation(
+ category_id=0,
+ bbox=(0, 0, 5, 5),
+ area=5 * 5,
+ segmentation={
+ "size": [5, 5],
+ "counts": [0, 15, 2, 3, 2, 3],
+ },
+ iscrowd=True,
+ )
+ ],
+ (5, 5),
+ True,
+ Detections(
+ xyxy=np.array([[0, 0, 5, 5]], dtype=np.float32),
+ class_id=np.array([0], dtype=int),
+ mask=np.array(
+ [
+ [
+ [1, 1, 1, 0, 0],
+ [1, 1, 1, 0, 0],
+ [1, 1, 1, 1, 1],
+ [1, 1, 1, 1, 1],
+ [1, 1, 1, 1, 1],
+ ]
+ ]
+ ),
+ ),
+ DoesNotRaise(),
+ ), # single image annotations with mask, RLE segmentation mask
+ (
+ [
+ mock_coco_annotation(
+ category_id=0,
+ bbox=(0, 0, 5, 5),
+ area=5 * 5,
+ segmentation=[[0, 0, 2, 0, 2, 2, 4, 2, 4, 4, 0, 4]],
+ ),
+ mock_coco_annotation(
+ category_id=0,
+ bbox=(3, 0, 2, 2),
+ area=2 * 2,
+ segmentation={
+ "size": [5, 5],
+ "counts": [15, 2, 3, 2, 3],
+ },
+ iscrowd=True,
+ ),
+ ],
+ (5, 5),
+ True,
+ Detections(
+ xyxy=np.array([[0, 0, 5, 5], [3, 0, 5, 2]], dtype=np.float32),
+ class_id=np.array([0, 0], dtype=int),
+ mask=np.array(
+ [
+ [
+ [1, 1, 1, 0, 0],
+ [1, 1, 1, 0, 0],
+ [1, 1, 1, 1, 1],
+ [1, 1, 1, 1, 1],
+ [1, 1, 1, 1, 1],
+ ],
+ [
+ [0, 0, 0, 1, 1],
+ [0, 0, 0, 1, 1],
+ [0, 0, 0, 0, 0],
+ [0, 0, 0, 0, 0],
+ [0, 0, 0, 0, 0],
+ ],
+ ]
+ ),
+ ),
+ DoesNotRaise(),
+ ), # two image annotations with mask, one mask as polygon ans second as RLE
+ (
+ [
+ mock_coco_annotation(
+ category_id=0,
+ bbox=(3, 0, 2, 2),
+ area=2 * 2,
+ segmentation={
+ "size": [5, 5],
+ "counts": [15, 2, 3, 2, 3],
+ },
+ iscrowd=True,
+ ),
+ mock_coco_annotation(
+ category_id=1,
+ bbox=(0, 0, 5, 5),
+ area=5 * 5,
+ segmentation=[[0, 0, 2, 0, 2, 2, 4, 2, 4, 4, 0, 4]],
+ ),
+ ],
+ (5, 5),
+ True,
+ Detections(
+ xyxy=np.array([[3, 0, 5, 2], [0, 0, 5, 5]], dtype=np.float32),
+ class_id=np.array([0, 1], dtype=int),
+ mask=np.array(
+ [
+ [
+ [0, 0, 0, 1, 1],
+ [0, 0, 0, 1, 1],
+ [0, 0, 0, 0, 0],
+ [0, 0, 0, 0, 0],
+ [0, 0, 0, 0, 0],
+ ],
+ [
+ [1, 1, 1, 0, 0],
+ [1, 1, 1, 0, 0],
+ [1, 1, 1, 1, 1],
+ [1, 1, 1, 1, 1],
+ [1, 1, 1, 1, 1],
+ ],
+ ]
+ ),
+ ),
+ DoesNotRaise(),
+ ), # two image annotations with mask, first mask as RLE and second as polygon
],
)
def test_coco_annotations_to_detections(
@@ -301,3 +429,131 @@ def test_build_coco_class_index_mapping(
coco_categories=coco_categories, target_classes=target_classes
)
assert result == expected_result
+
+
+@pytest.mark.parametrize(
+ "detections, image_id, annotation_id, expected_result, exception",
+ [
+ (
+ Detections(
+ xyxy=np.array([[0, 0, 100, 100]], dtype=np.float32),
+ class_id=np.array([0], dtype=int),
+ ),
+ 0,
+ 0,
+ [
+ mock_coco_annotation(
+ category_id=0, bbox=(0, 0, 100, 100), area=100 * 100
+ )
+ ],
+ DoesNotRaise(),
+ ), # no segmentation mask
+ (
+ Detections(
+ xyxy=np.array([[0, 0, 4, 5]], dtype=np.float32),
+ class_id=np.array([0], dtype=int),
+ mask=np.array(
+ [
+ [
+ [1, 1, 1, 1, 0],
+ [1, 1, 1, 1, 0],
+ [1, 1, 1, 1, 0],
+ [1, 1, 1, 1, 0],
+ [1, 1, 1, 1, 0],
+ ]
+ ]
+ ),
+ ),
+ 0,
+ 0,
+ [
+ mock_coco_annotation(
+ category_id=0,
+ bbox=(0, 0, 4, 5),
+ area=4 * 5,
+ segmentation=[[0, 0, 0, 4, 3, 4, 3, 0]],
+ )
+ ],
+ DoesNotRaise(),
+ ), # segmentation mask in single component,no holes in mask,
+ # expects polygon mask
+ (
+ Detections(
+ xyxy=np.array([[0, 0, 5, 5]], dtype=np.float32),
+ class_id=np.array([0], dtype=int),
+ mask=np.array(
+ [
+ [
+ [1, 1, 1, 0, 0],
+ [1, 1, 1, 0, 0],
+ [1, 1, 1, 0, 0],
+ [0, 0, 0, 1, 1],
+ [0, 0, 0, 1, 1],
+ ]
+ ]
+ ),
+ ),
+ 0,
+ 0,
+ [
+ mock_coco_annotation(
+ category_id=0,
+ bbox=(0, 0, 5, 5),
+ area=5 * 5,
+ segmentation={
+ "size": [5, 5],
+ "counts": [0, 3, 2, 3, 2, 3, 5, 2, 3, 2],
+ },
+ iscrowd=True,
+ )
+ ],
+ DoesNotRaise(),
+ ), # segmentation mask with 2 components, no holes in mask, expects RLE mask
+ (
+ Detections(
+ xyxy=np.array([[0, 0, 5, 5]], dtype=np.float32),
+ class_id=np.array([0], dtype=int),
+ mask=np.array(
+ [
+ [
+ [0, 1, 1, 1, 1],
+ [0, 1, 1, 1, 1],
+ [1, 1, 0, 0, 1],
+ [1, 1, 0, 0, 1],
+ [1, 1, 1, 1, 1],
+ ]
+ ]
+ ),
+ ),
+ 0,
+ 0,
+ [
+ mock_coco_annotation(
+ category_id=0,
+ bbox=(0, 0, 5, 5),
+ area=5 * 5,
+ segmentation={
+ "size": [5, 5],
+ "counts": [2, 10, 2, 3, 2, 6],
+ },
+ iscrowd=True,
+ )
+ ],
+ DoesNotRaise(),
+ ), # seg mask in single component, with holes in mask, expects RLE mask
+ ],
+)
+def test_detections_to_coco_annotations(
+ detections: Detections,
+ image_id: int,
+ annotation_id: int,
+ expected_result: List[Dict],
+ exception: Exception,
+) -> None:
+ with exception:
+ result, _ = detections_to_coco_annotations(
+ detections=detections,
+ image_id=image_id,
+ annotation_id=annotation_id,
+ )
+ assert result == expected_result
diff --git a/test/dataset/test_utils.py b/test/dataset/test_utils.py
index 5ca96ca5..41e1da5b 100644
--- a/test/dataset/test_utils.py
+++ b/test/dataset/test_utils.py
@@ -2,13 +2,17 @@ from contextlib import ExitStack as DoesNotRaise
from test.test_utils import mock_detections
from typing import Dict, List, Optional, Tuple, TypeVar
+import numpy as np
+import numpy.typing as npt
import pytest
from supervision import Detections
from supervision.dataset.utils import (
build_class_index_mapping,
map_detections_class_id,
+ mask_to_rle,
merge_class_lists,
+ rle_to_mask,
train_test_split,
)
@@ -229,3 +233,131 @@ def test_map_detections_class_id(
source_to_target_mapping=source_to_target_mapping, detections=detections
)
assert result == expected_result
+
+
+@pytest.mark.parametrize(
+ "mask, expected_rle, exception",
+ [
+ (
+ np.zeros((3, 3)).astype(bool),
+ [9],
+ DoesNotRaise(),
+ ), # mask with background only (mask with only False values)
+ (
+ np.ones((3, 3)).astype(bool),
+ [0, 9],
+ DoesNotRaise(),
+ ), # mask with foreground only (mask with only True values)
+ (
+ np.array(
+ [
+ [0, 0, 0, 0, 0],
+ [0, 1, 1, 1, 0],
+ [0, 1, 0, 1, 0],
+ [0, 1, 1, 1, 0],
+ [0, 0, 0, 0, 0],
+ ]
+ ).astype(bool),
+ [6, 3, 2, 1, 1, 1, 2, 3, 6],
+ DoesNotRaise(),
+ ), # mask where foreground object has hole
+ (
+ np.array(
+ [
+ [1, 0, 1, 0, 1],
+ [1, 0, 1, 0, 1],
+ [1, 0, 1, 0, 1],
+ [1, 0, 1, 0, 1],
+ [1, 0, 1, 0, 1],
+ ]
+ ).astype(bool),
+ [0, 5, 5, 5, 5, 5],
+ DoesNotRaise(),
+ ), # mask where foreground consists of 3 separate components
+ (
+ np.array([[[]]]).astype(bool),
+ None,
+ pytest.raises(AssertionError),
+ ), # raises AssertionError because mask dimentionality is not 2D
+ (
+ np.array([[]]).astype(bool),
+ None,
+ pytest.raises(AssertionError),
+ ), # raises AssertionError because mask is empty
+ ],
+)
+def test_mask_to_rle(
+ mask: npt.NDArray[np.bool_], expected_rle: List[int], exception: Exception
+) -> None:
+ with exception:
+ result = mask_to_rle(mask=mask)
+ assert result == expected_rle
+
+
+@pytest.mark.parametrize(
+ "rle, resolution_wh, expected_mask, exception",
+ [
+ (
+ np.array([9]),
+ [3, 3],
+ np.zeros((3, 3)).astype(bool),
+ DoesNotRaise(),
+ ), # mask with background only (mask with only False values); rle as array
+ (
+ [9],
+ [3, 3],
+ np.zeros((3, 3)).astype(bool),
+ DoesNotRaise(),
+ ), # mask with background only (mask with only False values); rle as list
+ (
+ np.array([0, 9]),
+ [3, 3],
+ np.ones((3, 3)).astype(bool),
+ DoesNotRaise(),
+ ), # mask with foreground only (mask with only True values)
+ (
+ np.array([6, 3, 2, 1, 1, 1, 2, 3, 6]),
+ [5, 5],
+ np.array(
+ [
+ [0, 0, 0, 0, 0],
+ [0, 1, 1, 1, 0],
+ [0, 1, 0, 1, 0],
+ [0, 1, 1, 1, 0],
+ [0, 0, 0, 0, 0],
+ ]
+ ).astype(bool),
+ DoesNotRaise(),
+ ), # mask where foreground object has hole
+ (
+ np.array([0, 5, 5, 5, 5, 5]),
+ [5, 5],
+ np.array(
+ [
+ [1, 0, 1, 0, 1],
+ [1, 0, 1, 0, 1],
+ [1, 0, 1, 0, 1],
+ [1, 0, 1, 0, 1],
+ [1, 0, 1, 0, 1],
+ ]
+ ).astype(bool),
+ DoesNotRaise(),
+ ), # mask where foreground consists of 3 separate components
+ (
+ np.array([0, 5, 5, 5, 5, 5]),
+ [2, 2],
+ None,
+ pytest.raises(AssertionError),
+ ), # raises AssertionError because number of pixels in RLE does not match
+ # number of pixels in expected mask (width x height).
+ ],
+)
+def test_rle_to_mask(
+ rle: npt.NDArray[np.int_],
+ resolution_wh: Tuple[int, int],
+ expected_mask: npt.NDArray[np.bool_],
+ exception: Exception,
+) -> None:
+ with exception:
+ result = rle_to_mask(rle=rle, resolution_wh=resolution_wh)
+ assert np.all(result == expected_mask)
diff --git a/test/detection/test_core.py b/test/detection/test_core.py
index f3b739e8..12f3de28 100644
--- a/test/detection/test_core.py
+++ b/test/detection/test_core.py
@@ -30,6 +30,84 @@ DETECTIONS = Detections(
)
+# Merge test
+TEST_MASK = np.zeros((1000, 1000), dtype=bool)
+TEST_MASK[300:351, 200:251] = True
+TEST_DET_1 = Detections(
+ xyxy=np.array([[10, 10, 20, 20], [30, 30, 40, 40], [50, 50, 60, 60]]),
+ mask=np.array([TEST_MASK, TEST_MASK, TEST_MASK]),
+ confidence=np.array([0.1, 0.2, 0.3]),
+ class_id=np.array([1, 2, 3]),
+ tracker_id=np.array([1, 2, 3]),
+ data={
+ "some_key": [1, 2, 3],
+ "other_key": [["1", "2"], ["3", "4"], ["5", "6"]],
+ },
+)
+TEST_DET_2 = Detections(
+ xyxy=np.array([[70, 70, 80, 80], [90, 90, 100, 100]]),
+ mask=np.array([TEST_MASK, TEST_MASK]),
+ confidence=np.array([0.4, 0.5]),
+ class_id=np.array([4, 5]),
+ tracker_id=np.array([4, 5]),
+ data={
+ "some_key": [4, 5],
+ "other_key": [["7", "8"], ["9", "10"]],
+ },
+)
+TEST_DET_1_2 = Detections(
+ xyxy=np.array(
+ [
+ [10, 10, 20, 20],
+ [30, 30, 40, 40],
+ [50, 50, 60, 60],
+ [70, 70, 80, 80],
+ [90, 90, 100, 100],
+ ]
+ ),
+ mask=np.array([TEST_MASK, TEST_MASK, TEST_MASK, TEST_MASK, TEST_MASK]),
+ confidence=np.array([0.1, 0.2, 0.3, 0.4, 0.5]),
+ class_id=np.array([1, 2, 3, 4, 5]),
+ tracker_id=np.array([1, 2, 3, 4, 5]),
+ data={
+ "some_key": [1, 2, 3, 4, 5],
+ "other_key": [["1", "2"], ["3", "4"], ["5", "6"], ["7", "8"], ["9", "10"]],
+ },
+)
+TEST_DET_ZERO_LENGTH = Detections(
+ xyxy=np.empty((0, 4), dtype=np.float32),
+ mask=np.empty((0, *TEST_MASK.shape), dtype=bool),
+ confidence=np.empty((0,)),
+ class_id=np.empty((0,)),
+ tracker_id=np.empty((0,)),
+ data={
+ "some_key": [],
+ "other_key": [],
+ },
+)
+TEST_DET_NONE = Detections(
+ xyxy=np.empty((0, 4), dtype=np.float32),
+)
+TEST_DET_DIFFERENT_FIELDS = Detections(
+ xyxy=np.array([[88, 88, 99, 99]]),
+ mask=np.array([np.logical_not(TEST_MASK)]),
+ confidence=None,
+ class_id=None,
+ tracker_id=np.array([9]),
+ data={"some_key": [9], "other_key": [["11", "12"]]},
+)
+TEST_DET_DIFFERENT_DATA = Detections(
+ xyxy=np.array([[88, 88, 99, 99]]),
+ mask=np.array([np.logical_not(TEST_MASK)]),
+ confidence=np.array([0.9]),
+ class_id=np.array([9]),
+ tracker_id=np.array([9]),
+ data={
+ "never_seen_key": [9],
+ },
+)
+
+
@pytest.mark.parametrize(
"detections, index, expected_result, exception",
[
@@ -148,52 +226,58 @@ def test_getitem(
DoesNotRaise(),
), # single empty detections
(
- [mock_detections(xyxy=[[10, 10, 20, 20]])],
- mock_detections(xyxy=[[10, 10, 20, 20]]),
+ [Detections.empty(), Detections.empty()],
+ Detections.empty(),
DoesNotRaise(),
- ), # single detection with xyxy field
+ ), # two empty detections
+ (
+ [TEST_DET_1],
+ TEST_DET_1,
+ DoesNotRaise(),
+ ), # single detection with fields
+ (
+ [TEST_DET_NONE],
+ TEST_DET_NONE,
+ DoesNotRaise(),
+ ), # Single weakly-defined detection
+ (
+ [TEST_DET_1, TEST_DET_2],
+ TEST_DET_1_2,
+ DoesNotRaise(),
+ ), # Fields with same keys
+ # Fields and empty
+ (
+ [TEST_DET_1, Detections.empty()],
+ TEST_DET_1,
+ DoesNotRaise(),
+ ), # single detection with fields
(
[
- mock_detections(xyxy=[[10, 10, 20, 20]]),
- mock_detections(xyxy=np.empty((0, 4), dtype=np.float32)),
+ TEST_DET_1,
+ TEST_DET_ZERO_LENGTH,
],
- mock_detections(xyxy=[[10, 10, 20, 20]]),
+ TEST_DET_1,
DoesNotRaise(),
- ), # single detection with xyxy field + empty detection
+ ), # Single detection and empty-array fields
(
[
- mock_detections(xyxy=[[10, 10, 20, 20]]),
- mock_detections(xyxy=[[20, 20, 30, 30]]),
+ TEST_DET_1,
+ TEST_DET_NONE,
],
- mock_detections(xyxy=[[10, 10, 20, 20], [20, 20, 30, 30]]),
+ TEST_DET_1,
DoesNotRaise(),
- ), # two detections with xyxy field
+ ), # Single detection and None fields (+ missing Dict keys)
+ # Errors: Non-zero-length differently defined keys & data
(
- [
- mock_detections(xyxy=[[10, 10, 20, 20]], class_id=[0]),
- mock_detections(xyxy=[[20, 20, 30, 30]]),
- ],
- mock_detections(xyxy=[[10, 10, 20, 20], [20, 20, 30, 30]]),
+ [TEST_DET_1, TEST_DET_DIFFERENT_FIELDS],
+ None,
pytest.raises(ValueError),
- ), # detection with xyxy, class_id fields + detection with xyxy field
+ ), # Non-empty detections with different fields
(
- [
- mock_detections(xyxy=[[10, 10, 20, 20]], class_id=[0]),
- mock_detections(xyxy=[[20, 20, 30, 30]], class_id=[1]),
- ],
- mock_detections(xyxy=[[10, 10, 20, 20], [20, 20, 30, 30]], class_id=[0, 1]),
- DoesNotRaise(),
- ), # two detections with xyxy, class_id fields
- (
- [
- mock_detections(xyxy=[[10, 10, 20, 20]], data={"test": [1]}),
- mock_detections(xyxy=[[20, 20, 30, 30]], data={"test": [2]}),
- ],
- mock_detections(
- xyxy=[[10, 10, 20, 20], [20, 20, 30, 30]], data={"test": [1, 2]}
- ),
- DoesNotRaise(),
- ), # two detections with xyxy, data fields
+ [TEST_DET_1, TEST_DET_DIFFERENT_DATA],
+ None,
+ pytest.raises(ValueError),
+ ), # Non-empty detections with different data keys
],
)
def test_merge(
diff --git a/test/detection/test_lmm.py b/test/detection/test_lmm.py
new file mode 100644
index 00000000..129aa44b
--- /dev/null
+++ b/test/detection/test_lmm.py
@@ -0,0 +1,131 @@
+from typing import List, Optional, Tuple
+
+import numpy as np
+import pytest
+
+from supervision.detection.lmm import from_paligemma
+
+
+@pytest.mark.parametrize(
+ "result, resolution_wh, classes, expected_results",
+ [
+ (
+ "",
+ (1000, 1000),
+ None,
+ (np.empty((0, 4)), None, np.empty(0).astype(str)),
+ ), # empty response
+ (
+ "",
+ (1000, 1000),
+ ["cat", "dog"],
+ (np.empty((0, 4)), None, np.empty(0).astype(str)),
+ ), # empty response with classes
+ (
+ "\n",
+ (1000, 1000),
+ None,
+ (np.empty((0, 4)), None, np.empty(0).astype(str)),
+ ), # new line response
+ (
+ "the quick brown fox jumps over the lazy dog.",
+ (1000, 1000),
+ None,
+ (np.empty((0, 4)), None, np.empty(0).astype(str)),
+ ), # response with no location
+ (
+ " cat",
+ (1000, 1000),
+ None,
+ (np.empty((0, 4)), None, np.empty(0).astype(str)),
+ ), # response with missing location
+ (
+ " cat",
+ (1000, 1000),
+ None,
+ (np.empty((0, 4)), None, np.empty(0).astype(str)),
+ ), # response with extra location
+ (
+ "",
+ (1000, 1000),
+ None,
+ (np.empty((0, 4)), None, np.empty(0).astype(str)),
+ ), # response with no class
+ (
+ " catt",
+ (1000, 1000),
+ ["cat", "dog"],
+ (np.empty((0, 4)), np.empty(0), np.empty(0).astype(str)),
+ ), # response with invalid class
+ (
+ " cat",
+ (1000, 1000),
+ None,
+ (
+ np.array([[250.0, 250.0, 750.0, 750.0]]),
+ None,
+ np.array(["cat"]).astype(str),
+ ),
+ ), # correct response; no classes
+ (
+ " black cat",
+ (1000, 1000),
+ None,
+ (
+ np.array([[250.0, 250.0, 750.0, 750.0]]),
+ None,
+ np.array(["black cat"]).astype(np.dtype("U")),
+ ),
+ ), # correct response; no classes
+ (
+ " cat ;",
+ (1000, 1000),
+ ["cat", "dog"],
+ (
+ np.array([[250.0, 250.0, 750.0, 750.0]]),
+ np.array([0]),
+ np.array(["cat"]).astype(str),
+ ),
+ ), # correct response; with classes
+ (
+ " cat ; dog", # noqa: E501
+ (1000, 1000),
+ ["cat", "dog"],
+ (
+ np.array([[250.0, 250.0, 750.0, 750.0], [250.0, 250.0, 750.0, 750.0]]),
+ np.array([0, 1]),
+ np.array(["cat", "dog"]).astype(np.dtype("U")),
+ ),
+ ), # correct response; with classes
+ (
+ " cat ; cat", # noqa: E501
+ (1000, 1000),
+ ["cat", "dog"],
+ (
+ np.array([[250.0, 250.0, 750.0, 750.0]]),
+ np.array([0]),
+ np.array(["cat"]).astype(str),
+ ),
+ ), # partially correct response; with classes
+ (
+ " cat ; cat", # noqa: E501
+ (1000, 1000),
+ ["cat", "dog"],
+ (
+ np.array([[250.0, 250.0, 750.0, 750.0]]),
+ np.array([0]),
+ np.array(["cat"]).astype(str),
+ ),
+ ), # partially correct response; with classes
+ ],
+)
+def test_from_paligemma(
+ result: str,
+ resolution_wh: Tuple[int, int],
+ classes: Optional[List[str]],
+ expected_results: Tuple[np.ndarray, Optional[np.ndarray], np.ndarray],
+) -> None:
+ result = from_paligemma(result=result, resolution_wh=resolution_wh, classes=classes)
+ np.testing.assert_array_equal(result[0], expected_results[0])
+ np.testing.assert_array_equal(result[1], expected_results[1])
+ np.testing.assert_array_equal(result[2], expected_results[2])
diff --git a/test/detection/test_utils.py b/test/detection/test_utils.py
index 1c4a1d34..6a2070da 100644
--- a/test/detection/test_utils.py
+++ b/test/detection/test_utils.py
@@ -2,6 +2,7 @@ from contextlib import ExitStack as DoesNotRaise
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
+import numpy.typing as npt
import pytest
from supervision.config import CLASS_NAME_DATA_FIELD
@@ -9,6 +10,8 @@ from supervision.detection.utils import (
box_non_max_suppression,
calculate_masks_centroids,
clip_boxes,
+ contains_holes,
+ contains_multiple_segments,
filter_polygons_by_area,
get_data_item,
mask_non_max_suppression,
@@ -911,6 +914,14 @@ def test_calculate_masks_centroids(
{"test_1": []},
DoesNotRaise(),
), # single data dict with a single field name and empty list values
+ (
+ [
+ {"test_1": []},
+ {"test_1": []},
+ ],
+ {"test_1": []},
+ DoesNotRaise(),
+ ), # two data dicts with the same field name and empty list values
(
[
{"test_1": np.array([])},
@@ -918,6 +929,14 @@ def test_calculate_masks_centroids(
{"test_1": np.array([])},
DoesNotRaise(),
), # single data dict with a single field name and empty np.array values
+ (
+ [
+ {"test_1": np.array([])},
+ {"test_1": np.array([])},
+ ],
+ {"test_1": np.array([])},
+ DoesNotRaise(),
+ ), # two data dicts with the same field name and empty np.array values
(
[
{"test_1": [1, 2, 3]},
@@ -932,7 +951,7 @@ def test_calculate_masks_centroids(
],
{"test_1": [3, 2, 1]},
DoesNotRaise(),
- ), # two data dicts with the same field name and empty and list values
+ ), # two data dicts with the same field name; one of with empty list as value
(
[
{"test_1": [1, 2, 3]},
@@ -1012,6 +1031,49 @@ def test_calculate_masks_centroids(
None,
pytest.raises(ValueError),
), # two data dicts with the same field name and different length arrays values
+ (
+ [{}, {"test_1": [1, 2, 3]}],
+ {"test_1": [1, 2, 3]},
+ DoesNotRaise(),
+ ), # two data dicts; one empty and one non-empty dict
+ (
+ [{"test_1": [], "test_2": []}, {"test_1": [1, 2, 3], "test_2": [1, 2, 3]}],
+ {"test_1": [1, 2, 3], "test_2": [1, 2, 3]},
+ DoesNotRaise(),
+ ), # two data dicts; one empty and one non-empty dict; same keys
+ (
+ [{"test_1": []}, {"test_1": [1, 2, 3], "test_2": [4, 5, 6]}],
+ None,
+ pytest.raises(ValueError),
+ ), # two data dicts; one empty and one non-empty dict; different keys
+ (
+ [
+ {
+ "test_1": [1, 2, 3],
+ "test_2": [4, 5, 6],
+ "test_3": [7, 8, 9],
+ },
+ {"test_1": [1, 2, 3], "test_2": [4, 5, 6]},
+ ],
+ None,
+ pytest.raises(ValueError),
+ ), # two data dicts; one with three keys, one with two keys
+ (
+ [
+ {"test_1": [1, 2, 3]},
+ {"test_1": [1, 2, 3], "test_2": [1, 2, 3]},
+ ],
+ None,
+ pytest.raises(ValueError),
+ ), # some keys missing in one dict
+ (
+ [
+ {"test_1": [1, 2, 3], "test_2": ["a", "b"]},
+ {"test_1": [4, 5], "test_2": ["c", "d", "e"]},
+ ],
+ None,
+ pytest.raises(ValueError),
+ ), # different value lengths for the same key
],
)
def test_merge_data(
@@ -1021,6 +1083,9 @@ def test_merge_data(
):
with exception:
result = merge_data(data_list=data_list)
+ if expected_result is None:
+ assert False, f"Expected an error, but got result {result}"
+
for key in result:
if isinstance(result[key], np.ndarray):
assert np.array_equal(
@@ -1203,3 +1268,138 @@ def test_get_data_item(
assert (
result[key] == expected_result[key]
), f"Mismatch in non-array data for key {key}"
+
+
+@pytest.mark.parametrize(
+ "mask, expected_result, exception",
+ [
+ (
+ np.array([[0, 0, 0, 0], [0, 1, 1, 0], [0, 1, 0, 0], [0, 1, 1, 0]]).astype(
+ bool
+ ),
+ False,
+ DoesNotRaise(),
+ ), # foreground object in one continuous piece
+ (
+ np.array([[1, 0, 0, 0], [1, 0, 0, 0], [0, 0, 0, 0], [0, 1, 1, 0]]).astype(
+ bool
+ ),
+ False,
+ DoesNotRaise(),
+ ), # foreground object in 2 seperate elements
+ (
+ np.array([[0, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0]]).astype(
+ bool
+ ),
+ False,
+ DoesNotRaise(),
+ ), # no foreground pixels in mask
+ (
+ np.array([[1, 1, 1, 1], [1, 1, 1, 1], [1, 1, 1, 1], [1, 1, 1, 1]]).astype(
+ bool
+ ),
+ False,
+ DoesNotRaise(),
+ ), # only foreground pixels in mask
+ (
+ np.array([[1, 1, 1, 0], [1, 0, 1, 0], [1, 1, 1, 0], [0, 0, 0, 0]]).astype(
+ bool
+ ),
+ True,
+ DoesNotRaise(),
+ ), # foreground object has 1 hole
+ (
+ np.array([[1, 1, 1, 0], [1, 0, 1, 1], [1, 1, 0, 1], [0, 1, 1, 1]]).astype(
+ bool
+ ),
+ True,
+ DoesNotRaise(),
+ ), # foreground object has 2 holes
+ ],
+)
+def test_contains_holes(
+ mask: npt.NDArray[np.bool_], expected_result: bool, exception: Exception
+) -> None:
+ with exception:
+ result = contains_holes(mask)
+ assert result == expected_result
+
+
+@pytest.mark.parametrize(
+ "mask, connectivity, expected_result, exception",
+ [
+ (
+ np.array([[0, 0, 0, 0], [0, 1, 1, 0], [0, 1, 0, 0], [0, 1, 1, 0]]).astype(
+ bool
+ ),
+ 4,
+ False,
+ DoesNotRaise(),
+ ), # foreground object in one continuous piece
+ (
+ np.array([[1, 0, 0, 0], [1, 0, 0, 0], [0, 0, 0, 0], [0, 1, 1, 0]]).astype(
+ bool
+ ),
+ 4,
+ True,
+ DoesNotRaise(),
+ ), # foreground object in 2 seperate elements
+ (
+ np.array([[0, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0]]).astype(
+ bool
+ ),
+ 4,
+ False,
+ DoesNotRaise(),
+ ), # no foreground pixels in mask
+ (
+ np.array([[1, 1, 1, 1], [1, 1, 1, 1], [1, 1, 1, 1], [1, 1, 1, 1]]).astype(
+ bool
+ ),
+ 4,
+ False,
+ DoesNotRaise(),
+ ), # only foreground pixels in mask
+ (
+ np.array([[1, 1, 1, 0], [1, 0, 1, 1], [1, 1, 0, 1], [0, 1, 1, 1]]).astype(
+ bool
+ ),
+ 4,
+ False,
+ DoesNotRaise(),
+ ), # foreground object has 2 holes, but is in single piece
+ (
+ np.array([[1, 1, 0, 0], [1, 1, 0, 1], [1, 0, 1, 1], [0, 0, 1, 1]]).astype(
+ bool
+ ),
+ 4,
+ True,
+ DoesNotRaise(),
+ ), # foreground object in 2 elements with respect to 4-way connectivity
+ (
+ np.array([[1, 1, 0, 0], [1, 1, 0, 1], [1, 0, 1, 1], [0, 0, 1, 1]]).astype(
+ bool
+ ),
+ 8,
+ False,
+ DoesNotRaise(),
+ ), # foreground object in single piece with respect to 8-way connectivity
+ (
+ np.array([[1, 1, 0, 0], [1, 1, 0, 1], [1, 0, 1, 1], [0, 0, 1, 1]]).astype(
+ bool
+ ),
+ 5,
+ None,
+ pytest.raises(ValueError),
+ ), # Incorrect connectivity parameter value, raises ValueError
+ ],
+)
+def test_contains_multiple_segments(
+ mask: npt.NDArray[np.bool_],
+ connectivity: int,
+ expected_result: bool,
+ exception: Exception,
+) -> None:
+ with exception:
+ result = contains_multiple_segments(mask=mask, connectivity=connectivity)
+ assert result == expected_result