initial version of advanced key points filtering and slicing + some naming refactor

This commit is contained in:
SkalskiP 2025-07-28 08:24:43 +02:00
parent 8873da4a30
commit beaf89cb22
12 changed files with 603 additions and 92 deletions

View File

@ -345,7 +345,7 @@ Supervision is versatile and compatible with various models. Check this [link](/
We will define a `callback` function, which will process each frame of the video by obtaining model predictions and then annotating the frame based on these predictions.
Let's immediately visualize the results with our [`EdgeAnnotator`](/latest/keypoint/annotators/#supervision.keypoint.annotators.EdgeAnnotator) and [`VertexAnnotator`](https://supervision.roboflow.com/latest/keypoint/annotators/#supervision.keypoint.annotators.VertexAnnotator).
Let's immediately visualize the results with our [`EdgeAnnotator`](/latest/keypoint/annotators/#supervision.key_points.annotators.EdgeAnnotator) and [`VertexAnnotator`](https://supervision.roboflow.com/latest/keypoint/annotators/#supervision.key_points.annotators.VertexAnnotator).
=== "Ultralytics"
@ -408,7 +408,7 @@ Let's immediately visualize the results with our [`EdgeAnnotator`](/latest/keypo
### Convert to Detections
Keypoint tracking is currently supported via the conversion of `KeyPoints` to `Detections`. This is achieved with the [`KeyPoints.as_detections()`](/latest/keypoint/core/#supervision.keypoint.core.KeyPoints.as_detections) function.
Keypoint tracking is currently supported via the conversion of `KeyPoints` to `Detections`. This is achieved with the [`KeyPoints.as_detections()`](/latest/keypoint/core/#supervision.key_points.core.KeyPoints.as_detections) function.
Let's convert to detections and visualize the results with our [`BoxAnnotator`](/latest/detection/annotators/#supervision.annotators.core.BoxAnnotator).

View File

@ -78,19 +78,19 @@ comments: true
</div>
<div class="md-typeset">
<h2><a href="#supervision.keypoint.annotators.VertexAnnotator">VertexAnnotator</a></h2>
<h2><a href="#supervision.key_points.annotators.VertexAnnotator">VertexAnnotator</a></h2>
</div>
:::supervision.keypoint.annotators.VertexAnnotator
:::supervision.key_points.annotators.VertexAnnotator
<div class="md-typeset">
<h2><a href="#supervision.keypoint.annotators.EdgeAnnotator">EdgeAnnotator</a></h2>
<h2><a href="#supervision.key_points.annotators.EdgeAnnotator">EdgeAnnotator</a></h2>
</div>
:::supervision.keypoint.annotators.EdgeAnnotator
:::supervision.key_points.annotators.EdgeAnnotator
<div class="md-typeset">
<h2><a href="#supervision.keypoint.annotators.VertexLabelAnnotator">VertexLabelAnnotator</a></h2>
<h2><a href="#supervision.key_points.annotators.VertexLabelAnnotator">VertexLabelAnnotator</a></h2>
</div>
:::supervision.keypoint.annotators.VertexLabelAnnotator
:::supervision.key_points.annotators.VertexLabelAnnotator

View File

@ -5,4 +5,4 @@ status: new
# Keypoint Detection
:::supervision.keypoint.core.KeyPoints
:::supervision.key_points.core.KeyPoints

View File

@ -108,12 +108,12 @@ 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 (
from supervision.key_points.annotators import (
EdgeAnnotator,
VertexAnnotator,
VertexLabelAnnotator,
)
from supervision.keypoint.core import KeyPoints
from supervision.key_points.core import KeyPoints
from supervision.metrics.detection import ConfusionMatrix, MeanAveragePrecision
from supervision.tracker.byte_tracker.core import ByteTrack
from supervision.utils.conversion import cv2_to_pillow, pillow_to_cv2

View File

@ -11,8 +11,8 @@ from supervision.detection.utils.boxes import pad_boxes, spread_out_boxes
from supervision.draw.color import Color
from supervision.draw.utils import draw_rounded_rectangle
from supervision.geometry.core import Rect
from supervision.keypoint.core import KeyPoints
from supervision.keypoint.skeletons import SKELETONS_BY_VERTEX_COUNT
from supervision.key_points.core import KeyPoints
from supervision.key_points.skeletons import SKELETONS_BY_VERTEX_COUNT
from supervision.utils.conversion import ensure_cv2_image_for_annotation

View File

@ -10,7 +10,7 @@ import numpy.typing as npt
from supervision.config import CLASS_NAME_DATA_FIELD
from supervision.detection.core import Detections
from supervision.detection.utils.internal import get_data_item, is_data_equal
from supervision.validators import validate_keypoints_fields
from supervision.validators import validate_key_points_fields
@dataclass
@ -23,7 +23,7 @@ class KeyPoints:
=== "Ultralytics"
Use [`sv.KeyPoints.from_ultralytics`](/latest/keypoint/core/#supervision.keypoint.core.KeyPoints.from_ultralytics)
Use [`sv.KeyPoints.from_ultralytics`](/latest/keypoint/core/#supervision.key_points.core.KeyPoints.from_ultralytics)
method, which accepts [YOLOv8-pose](https://docs.ultralytics.com/models/yolov8/), [YOLO11-pose](https://docs.ultralytics.com/models/yolo11/)
[pose](https://docs.ultralytics.com/tasks/pose/) result.
@ -41,7 +41,7 @@ class KeyPoints:
=== "Inference"
Use [`sv.KeyPoints.from_inference`](/latest/keypoint/core/#supervision.keypoint.core.KeyPoints.from_inference)
Use [`sv.KeyPoints.from_inference`](/latest/keypoint/core/#supervision.key_points.core.KeyPoints.from_inference)
method, which accepts [Inference](https://inference.roboflow.com/) pose result.
```python
@ -58,7 +58,7 @@ class KeyPoints:
=== "MediaPipe"
Use [`sv.KeyPoints.from_mediapipe`](/latest/keypoint/core/#supervision.keypoint.core.KeyPoints.from_mediapipe)
Use [`sv.KeyPoints.from_mediapipe`](/latest/keypoint/core/#supervision.key_points.core.KeyPoints.from_mediapipe)
method, which accepts [MediaPipe](https://github.com/google-ai-edge/mediapipe)
pose result.
@ -88,11 +88,62 @@ class KeyPoints:
key_points = sv.KeyPoints.from_mediapipe(
pose_landmarker_result, (image_width, image_height))
```
=== "Transformers"
Use [`sv.KeyPoints.from_transformers`](/latest/keypoint/core/#supervision.key_points.core.KeyPoints.from_transformers)
method, which accepts [ViTPose](https://huggingface.co/docs/transformers/en/model_doc/vitpose) result.
```python
from PIL import Image
import requests
import supervision as sv
import torch
from transformers import (
AutoProcessor,
RTDetrForObjectDetection,
VitPoseForPoseEstimation,
)
device = "cuda" if torch.cuda.is_available() else "cpu"
image = Image.open(<SOURCE_IMAGE_PATH>)
DETECTION_MODEL_ID = "PekingU/rtdetr_r50vd_coco_o365"
detection_processor = AutoProcessor.from_pretrained(DETECTION_MODEL_ID, use_fast=True)
detection_model = RTDetrForObjectDetection.from_pretrained(DETECTION_MODEL_ID, device_map=DEVICE)
inputs = detection_processor(images=frame, return_tensors="pt").to(DEVICE)
with torch.no_grad():
outputs = detection_model(**inputs)
target_size = torch.tensor([(frame.height, frame.width)])
results = detection_processor.post_process_object_detection(
outputs, target_sizes=target_size, threshold=0.3)
detections = sv.Detections.from_transformers(results[0])
boxes = sv.xyxy_to_xywh(detections[detections.class_id == 0].xyxy)
POSE_ESTIMATION_MODEL_ID = "usyd-community/vitpose-base-simple"
pose_estimation_processor = AutoProcessor.from_pretrained(POSE_ESTIMATION_MODEL_ID)
pose_estimation_model = VitPoseForPoseEstimation.from_pretrained(
POSE_ESTIMATION_MODEL_ID, device_map=DEVICE)
inputs = pose_estimation_processor(frame, boxes=[boxes], return_tensors="pt").to(DEVICE)
with torch.no_grad():
outputs = pose_estimation_model(**inputs)
results = pose_estimation_processor.post_process_pose_estimation(outputs, boxes=[boxes])
key_point = sv.KeyPoints.from_transformers(results[0])
```
Attributes:
xy (np.ndarray): An array of shape `(n, m, 2)` containing
`n` detected objects, each composed of `m` equally-sized
sets of keypoints, where each point is `[x, y]`.
sets of key points, where each point is `[x, y]`.
class_id (Optional[np.ndarray]): An array of shape
`(n,)` containing the class ids of the detected objects.
confidence (Optional[np.ndarray]): An array of shape
@ -109,7 +160,7 @@ class KeyPoints:
data: dict[str, npt.NDArray[Any] | list] = field(default_factory=dict)
def __post_init__(self):
validate_keypoints_fields(
validate_key_points_fields(
xy=self.xy,
confidence=self.confidence,
class_id=self.class_id,
@ -514,13 +565,13 @@ class KeyPoints:
return cls.empty()
@classmethod
def from_transformers(cls, transfomers_results: Any) -> KeyPoints:
def from_transformers(cls, transformers_results: Any) -> KeyPoints:
"""
Create a `sv.KeyPoints` object from the
[Transformers](https://github.com/huggingface/transformers) inference result.
Args:
transfomers_results (Any): The output of a
transformers_results (Any): The output of a
Transformers model containing instances with prediction data.
Returns:
@ -576,8 +627,8 @@ class KeyPoints:
""" # noqa: E501 // docs
if "keypoints" in transfomers_results[0]:
if transfomers_results[0]["keypoints"].cpu().numpy().size == 0:
if "keypoints" in transformers_results[0]:
if transformers_results[0]["keypoints"].cpu().numpy().size == 0:
return cls.empty()
result_data = [
@ -585,7 +636,7 @@ class KeyPoints:
result["keypoints"].cpu().numpy(),
result["scores"].cpu().numpy(),
)
for result in transfomers_results
for result in transformers_results
]
xy, scores = zip(*result_data)
@ -599,55 +650,72 @@ class KeyPoints:
return cls.empty()
def __getitem__(
self, index: int | slice | list[int] | np.ndarray | str
) -> KeyPoints | list | np.ndarray | None:
"""
Get a subset of the `sv.KeyPoints` object or access an item from its data field.
When provided with an integer, slice, list of integers, or a numpy array, this
method returns a new `sv.KeyPoints` object that represents a subset of the
original `sv.KeyPoints`. When provided with a string, it accesses the
corresponding item in the data dictionary.
Args:
index (Union[int, slice, List[int], np.ndarray, str]): The index, indices,
or key to access a subset of the `sv.KeyPoints` or an item from the
data.
Returns:
A subset of the `sv.KeyPoints` object or an item from the data field.
Examples:
```python
import supervision as sv
key_points = sv.KeyPoints()
# access the first keypoint using an integer index
key_points[0]
# access the first 10 keypoints using index slice
key_points[0:10]
# access selected keypoints using a list of indices
key_points[[0, 2, 4]]
# access keypoints with selected class_id
key_points[key_points.class_id == 0]
# access keypoints with confidence greater than 0.5
key_points[key_points.confidence > 0.5]
```
"""
self, index: int | slice | list[int] | np.ndarray | tuple | str
) -> KeyPoints | np.ndarray | list | None:
if isinstance(index, str):
return self.data.get(index)
if isinstance(index, int):
index = [index]
if not isinstance(index, tuple):
index = (index, slice(None))
i, j = index
if isinstance(i, int):
i = [i]
if isinstance(i, list) and all(isinstance(x, bool) for x in i):
i = np.array(i)
if isinstance(j, list) and all(isinstance(x, bool) for x in j):
j = np.array(j)
if isinstance(i, np.ndarray) and i.dtype == bool:
i = np.flatnonzero(i)
if isinstance(j, np.ndarray) and j.dtype == bool:
j = np.flatnonzero(j)
if (
isinstance(i, (list, np.ndarray))
and isinstance(j, (list, np.ndarray))
and not np.isscalar(i)
and not np.isscalar(j)
):
i, j = np.ix_(i, j)
xy_selected = self.xy[i, j]
conf_selected = (
self.confidence[i, j] if self.confidence is not None else None
)
class_id_selected = (
self.class_id[i] if self.class_id is not None else None
)
data_selected = get_data_item(self.data, i)
if xy_selected.ndim == 1:
xy_selected = xy_selected.reshape(1, 1, 2)
if conf_selected is not None:
conf_selected = conf_selected.reshape(1, 1)
elif xy_selected.ndim == 2:
if np.isscalar(index[0]) or (
isinstance(index[0], np.ndarray) and index[0].ndim == 0
):
xy_selected = xy_selected[np.newaxis, ...]
if conf_selected is not None:
conf_selected = conf_selected[np.newaxis, ...]
elif np.isscalar(index[1]) or (
isinstance(index[1], np.ndarray) and index[1].ndim == 0
):
xy_selected = xy_selected[:, np.newaxis, :]
if conf_selected is not None:
conf_selected = conf_selected[:, np.newaxis]
return KeyPoints(
xy=self.xy[index],
confidence=self.confidence[index] if self.confidence is not None else None,
class_id=self.class_id[index] if self.class_id is not None else None,
data=get_data_item(self.data, index),
xy=xy_selected,
confidence=conf_selected,
class_id=class_id_selected,
data=data_selected,
)
def __setitem__(self, key: str, value: np.ndarray | list):
@ -668,12 +736,12 @@ class KeyPoints:
model = YOLO('yolov8s.pt')
result = model(image)[0]
keypoints = sv.KeyPoints.from_ultralytics(result)
key_points = sv.KeyPoints.from_ultralytics(result)
keypoints['class_name'] = [
key_points['class_name'] = [
model.model.names[class_id]
for class_id
in keypoints.class_id
in key_points.class_id
]
```
"""
@ -688,7 +756,7 @@ class KeyPoints:
@classmethod
def empty(cls) -> KeyPoints:
"""
Create an empty Keypoints object with no keypoints.
Create an empty KeyPoints object with no key points.
Returns:
An empty `sv.KeyPoints` object.
@ -706,9 +774,9 @@ class KeyPoints:
"""
Returns `True` if the `KeyPoints` object is considered empty.
"""
empty_keypoints = KeyPoints.empty()
empty_keypoints.data = self.data
return self == empty_keypoints
empty_key_points = KeyPoints.empty()
empty_key_points.data = self.data
return self == empty_key_points
def as_detections(
self, selected_keypoint_indices: Iterable[int] | None = None
@ -716,21 +784,21 @@ class KeyPoints:
"""
Convert a KeyPoints object to a Detections object. This
approximates the bounding box of the detected object by
taking the bounding box that fits all keypoints.
taking the bounding box that fits all key points.
Arguments:
selected_keypoint_indices (Optional[Iterable[int]]): The
indices of the keypoints to include in the bounding box
calculation. This helps focus on a subset of keypoints,
e.g. when some are occluded. Captures all keypoints by default.
indices of the key points to include in the bounding box
calculation. This helps focus on a subset of key points,
e.g. when some are occluded. Captures all key points by default.
Returns:
detections (Detections): The converted detections object.
Examples:
```python
keypoints = sv.KeyPoints.from_inference(...)
detections = keypoints.as_detections()
key_points = sv.KeyPoints.from_inference(...)
detections = key_points.as_detections()
```
"""
if self.is_empty():

View File

@ -53,7 +53,7 @@ def validate_confidence(confidence: Any, n: int) -> None:
)
def validate_keypoint_confidence(confidence: Any, n: int, m: int) -> None:
def validate_key_point_confidence(confidence: Any, n: int, m: int) -> None:
expected_shape = f"({n, m})"
actual_shape = str(getattr(confidence, "shape", None))
@ -126,7 +126,7 @@ def validate_detections_fields(
validate_data(data, n)
def validate_keypoints_fields(
def validate_key_points_fields(
xy: Any,
class_id: Any,
confidence: Any,
@ -136,7 +136,7 @@ def validate_keypoints_fields(
m = len(xy[0]) if len(xy) > 0 else 0
validate_xy(xy, n, m)
validate_class_id(class_id, n)
validate_keypoint_confidence(confidence, n, m)
validate_key_point_confidence(confidence, n, m)
validate_data(data, n)

View File

View File

@ -0,0 +1,444 @@
import numpy as np
import pytest
from contextlib import nullcontext as DoesNotRaise
from supervision.key_points.core import KeyPoints
from test.test_utils import mock_key_points
KEY_POINTS = mock_key_points(
xy=[
[[0, 1],
[2, 3],
[4, 5],
[6, 7],
[8, 9]],
[[10, 11],
[12, 13],
[14, 15],
[16, 17],
[18, 19]],
[[20, 21],
[22, 23],
[24, 25],
[26, 27],
[28, 29]]
],
confidence=[
[0.8, 0.2, 0.6, 0.1, 0.5],
[0.7, 0.9, 0.3, 0.4, 0.0],
[0.1, 0.6, 0.8, 0.2, 0.7]
],
class_id=[0, 1, 2],
)
@pytest.mark.parametrize(
"key_points, index, expected_result, exception",
[
(
KeyPoints.empty(),
slice(None),
KeyPoints.empty(),
DoesNotRaise(),
), # slice all key points when key points object empty
(
KEY_POINTS,
slice(None),
KEY_POINTS,
DoesNotRaise(),
), # slice all key points when key points object nonempty
(
KEY_POINTS,
slice(0, 1),
mock_key_points(
xy=[
[[0, 1],
[2, 3],
[4, 5],
[6, 7],
[8, 9]]
],
confidence=[
[0.8, 0.2, 0.6, 0.1, 0.5]
],
class_id=[0],
),
DoesNotRaise(),
), # select the first skeleton by slice
(
KEY_POINTS,
slice(0, 2),
mock_key_points(
xy=[
[[0, 1],
[2, 3],
[4, 5],
[6, 7],
[8, 9]],
[[10, 11],
[12, 13],
[14, 15],
[16, 17],
[18, 19]],
],
confidence=[
[0.8, 0.2, 0.6, 0.1, 0.5],
[0.7, 0.9, 0.3, 0.4, 0.0],
],
class_id=[0, 1],
),
DoesNotRaise(),
), # select the first skeleton by slice
(
KEY_POINTS,
0,
mock_key_points(
xy=[
[[0, 1],
[2, 3],
[4, 5],
[6, 7],
[8, 9]]
],
confidence=[
[0.8, 0.2, 0.6, 0.1, 0.5]
],
class_id=[0],
),
DoesNotRaise(),
), # select the first skeleton by index
(
KEY_POINTS,
-1,
mock_key_points(
xy=[
[[20, 21],
[22, 23],
[24, 25],
[26, 27],
[28, 29]]
],
confidence=[
[0.1, 0.6, 0.8, 0.2, 0.7]
],
class_id=[2],
),
DoesNotRaise(),
), # select the last skeleton by index
(
KEY_POINTS,
[0, 1],
mock_key_points(
xy=[
[[0, 1],
[2, 3],
[4, 5],
[6, 7],
[8, 9]],
[[10, 11],
[12, 13],
[14, 15],
[16, 17],
[18, 19]],
],
confidence=[
[0.8, 0.2, 0.6, 0.1, 0.5],
[0.7, 0.9, 0.3, 0.4, 0.0],
],
class_id=[0, 1],
),
DoesNotRaise(),
), # select the first two skeletons by index; list
(
KEY_POINTS,
np.array([0, 1]),
mock_key_points(
xy=[
[[0, 1],
[2, 3],
[4, 5],
[6, 7],
[8, 9]],
[[10, 11],
[12, 13],
[14, 15],
[16, 17],
[18, 19]],
],
confidence=[
[0.8, 0.2, 0.6, 0.1, 0.5],
[0.7, 0.9, 0.3, 0.4, 0.0],
],
class_id=[0, 1],
),
DoesNotRaise(),
), # select the first two skeletons by index; np.array
(
KEY_POINTS,
[True, True, False],
mock_key_points(
xy=[
[[0, 1],
[2, 3],
[4, 5],
[6, 7],
[8, 9]],
[[10, 11],
[12, 13],
[14, 15],
[16, 17],
[18, 19]],
],
confidence=[
[0.8, 0.2, 0.6, 0.1, 0.5],
[0.7, 0.9, 0.3, 0.4, 0.0],
],
class_id=[0, 1],
),
DoesNotRaise(),
), # select only skeletons associated with positive filter; list
(
KEY_POINTS,
np.array([True, True, False]),
mock_key_points(
xy=[
[[0, 1],
[2, 3],
[4, 5],
[6, 7],
[8, 9]],
[[10, 11],
[12, 13],
[14, 15],
[16, 17],
[18, 19]],
],
confidence=[
[0.8, 0.2, 0.6, 0.1, 0.5],
[0.7, 0.9, 0.3, 0.4, 0.0],
],
class_id=[0, 1],
),
DoesNotRaise(),
), # select only skeletons associated with positive filter; list
(
KEY_POINTS,
(slice(None), slice(None)),
KEY_POINTS,
DoesNotRaise(),
), # slice all anchors from all skeletons
(
KEY_POINTS,
(slice(None), slice(0, 1)),
mock_key_points(
xy=[
[[0, 1]],
[[10, 11]],
[[20, 21]]
],
confidence=[
[0.8],
[0.7],
[0.1]
],
class_id=[0, 1, 2],
),
DoesNotRaise(),
), # slice the first anchor from every skeleton
(
KEY_POINTS,
(slice(None), slice(0, 2)),
mock_key_points(
xy=[
[[0, 1],
[2, 3]],
[[10, 11],
[12, 13]],
[[20, 21],
[22, 23]]
],
confidence=[
[0.8, 0.2],
[0.7, 0.9],
[0.1, 0.6]
],
class_id=[0, 1, 2],
),
DoesNotRaise(),
), # slice the first anchor two anchors from every skeleton
(
KEY_POINTS,
(slice(None), 0),
mock_key_points(
xy=[
[[0, 1]],
[[10, 11]],
[[20, 21]]
],
confidence=[
[0.8],
[0.7],
[0.1]
],
class_id=[0, 1, 2],
),
DoesNotRaise(),
), # select the first anchor from every skeleton by index
(
KEY_POINTS,
(slice(None), -1),
mock_key_points(
xy=[
[[8, 9]],
[[18, 19]],
[[28, 29]]
],
confidence=[
[0.5],
[0.0],
[0.7]
],
class_id=[0, 1, 2],
),
DoesNotRaise(),
), # select the last anchor from every skeleton by index
(
KEY_POINTS,
(slice(None), [0, 1]),
mock_key_points(
xy=[
[[0, 1],
[2, 3]],
[[10, 11],
[12, 13]],
[[20, 21],
[22, 23]]
],
confidence=[
[0.8, 0.2],
[0.7, 0.9],
[0.1, 0.6]
],
class_id=[0, 1, 2],
),
DoesNotRaise(),
), # select the first two anchors from every skeleton by index; list
(
KEY_POINTS,
(slice(None), np.array([0, 1])),
mock_key_points(
xy=[
[[0, 1],
[2, 3]],
[[10, 11],
[12, 13]],
[[20, 21],
[22, 23]]
],
confidence=[
[0.8, 0.2],
[0.7, 0.9],
[0.1, 0.6]
],
class_id=[0, 1, 2],
),
DoesNotRaise(),
), # select the first two anchors from every skeleton by index; np.array
(
KEY_POINTS,
(slice(None), [True, True, False, False, False]),
mock_key_points(
xy=[
[[0, 1],
[2, 3]],
[[10, 11],
[12, 13]],
[[20, 21],
[22, 23]]
],
confidence=[
[0.8, 0.2],
[0.7, 0.9],
[0.1, 0.6]
],
class_id=[0, 1, 2],
),
DoesNotRaise(),
), # select only anchors associated with positive filter; list
(
KEY_POINTS,
(slice(None), np.array([True, True, False, False, False])),
mock_key_points(
xy=[
[[0, 1],
[2, 3]],
[[10, 11],
[12, 13]],
[[20, 21],
[22, 23]]
],
confidence=[
[0.8, 0.2],
[0.7, 0.9],
[0.1, 0.6]
],
class_id=[0, 1, 2],
),
DoesNotRaise(),
), # select only anchors associated with positive filter; np.array
(
KEY_POINTS,
(0, 0),
mock_key_points(
xy=[
[[0, 1]],
],
confidence=[
[0.8],
],
class_id=[0],
),
DoesNotRaise(),
), # select the first anchor from the first skeleton by index
(
KEY_POINTS,
(0, -1),
mock_key_points(
xy=[
[[8, 9]],
],
confidence=[
[0.5],
],
class_id=[0],
),
DoesNotRaise(),
), # select the last anchor from the first skeleton by index
],
)
def test_key_points_getitem(key_points, index, expected_result, exception):
with exception:
result = key_points[index]
assert result == expected_result

View File

@ -3,14 +3,13 @@ from __future__ import annotations
from typing import Any
import numpy as np
import numpy.typing as npt
from supervision.detection.core import Detections
from supervision.keypoint.core import KeyPoints
from supervision.key_points.core import KeyPoints
def mock_detections(
xyxy: npt.NDArray[np.float32],
xyxy: list[list[float]],
mask: list[np.ndarray] | None = None,
confidence: list[float] | None = None,
class_id: list[int] | None = None,
@ -34,9 +33,9 @@ def mock_detections(
)
def mock_keypoints(
xy: npt.NDArray[np.float32],
confidence: list[float] | None = None,
def mock_key_points(
xy: list[list[list[float]]],
confidence: list[list[float]] | None = None,
class_id: list[int] | None = None,
data: dict[str, list[Any]] | None = None,
) -> KeyPoints: