From 5d8b378f35ea4112a0c8d9269bb285cd7d1c3531 Mon Sep 17 00:00:00 2001 From: James Gallagher Date: Tue, 13 Jun 2023 14:35:32 +0100 Subject: [PATCH] add test cases for get_top_k, add docs for Classifications and ClassificationsDataset --- docs/classification/core.md | 3 +++ docs/dataset/core.md | 6 ++++- supervision/.DS_Store | Bin 0 -> 6148 bytes supervision/__init__.py | 6 ++++- supervision/classification/.DS_Store | Bin 0 -> 6148 bytes supervision/dataset/core.py | 3 +-- test/classification/__init__.py | 0 test/classification/test_core.py | 39 +++++++++++++++++++++++++++ 8 files changed, 53 insertions(+), 4 deletions(-) create mode 100644 docs/classification/core.md create mode 100644 supervision/.DS_Store create mode 100644 supervision/classification/.DS_Store create mode 100644 test/classification/__init__.py create mode 100644 test/classification/test_core.py diff --git a/docs/classification/core.md b/docs/classification/core.md new file mode 100644 index 00000000..55ae57ea --- /dev/null +++ b/docs/classification/core.md @@ -0,0 +1,3 @@ +## Classifications + +:::supervision.classification.core.Classifications diff --git a/docs/dataset/core.md b/docs/dataset/core.md index 0713c9dc..618f671b 100644 --- a/docs/dataset/core.md +++ b/docs/dataset/core.md @@ -5,4 +5,8 @@ ## DetectionDataset -:::supervision.dataset.core.DetectionDataset \ No newline at end of file +:::supervision.dataset.core.DetectionDataset + +## ClassificationDataset + +:::supervision.dataset.core.ClassificationDataset \ No newline at end of file diff --git a/supervision/.DS_Store b/supervision/.DS_Store new file mode 100644 index 0000000000000000000000000000000000000000..766f9f819e1a43f28cb61a2819aed389c3d65a73 GIT binary patch literal 6148 zcmeHKOKQVF43(NJ0)^mZmUD&PU@++ke1VpfltOTzZoBp>=gQIY^r2v!KsMc!Cy?HZ ztT%(-!m>m}+wZq0kw!#Ta6>s+n43K}pV>oZ6bQ!|gM7#yzLVEk_4R~t*Qh^$F--Vh zIOiz+Pxsj$j{SSwas3cxsQ?wA0#twsP=UJ@u-*$>Jq9vT0V+TReig9qLxCIC#4*r6 z9SA-G0GCL+VePX7uvh|E6URVgU>a0lP&G#k4Lb5A>uTZ{7m>+I#M*B1B-+-lBnGpwD0 m;O!Xb?HC(t#~Uw-x?*cQuZd%z(~);NkUs;a3yli=wE`EWxEC=1 literal 0 HcmV?d00001 diff --git a/supervision/__init__.py b/supervision/__init__.py index b62c73c3..28bd6d8d 100644 --- a/supervision/__init__.py +++ b/supervision/__init__.py @@ -1,7 +1,11 @@ __version__ = "0.9.0" from supervision.classification.core import Classifications -from supervision.dataset.core import BaseDataset, DetectionDataset, ClassificationDataset +from supervision.dataset.core import ( + BaseDataset, + ClassificationDataset, + DetectionDataset, +) from supervision.detection.annotate import BoxAnnotator, MaskAnnotator from supervision.detection.core import Detections from supervision.detection.line_counter import LineZone, LineZoneAnnotator diff --git a/supervision/classification/.DS_Store b/supervision/classification/.DS_Store new file mode 100644 index 0000000000000000000000000000000000000000..08b4929c92be9ed0a9153a9701cf314ed6d168c7 GIT binary patch literal 6148 zcmeHK!A`?441Iw~Oxk5fjyZ6p(*B@K<-mc{J^RCvY>e;FQoCOSlutbL?w zuJC~G>wmqfn?+eSB|U_m3$|>hLJc^<0u44;>2AB;?;Y2Ce8Pq}L&+bx@j__d-?sP8=j_(#r>QysMNC_4S}gsq{>B6bZY`a?ho L@y;3e0|q_-Av{2o literal 0 HcmV?d00001 diff --git a/supervision/dataset/core.py b/supervision/dataset/core.py index d0564688..926672e8 100644 --- a/supervision/dataset/core.py +++ b/supervision/dataset/core.py @@ -1,5 +1,6 @@ from __future__ import annotations +import os from abc import ABC, abstractmethod from dataclasses import dataclass from pathlib import Path @@ -8,9 +9,7 @@ from typing import Dict, Iterator, List, Optional, Tuple import cv2 import numpy as np -import os from supervision.classification.core import Classifications - from supervision.dataset.formats.pascal_voc import ( detections_to_pascal_voc, load_pascal_voc_annotations, diff --git a/test/classification/__init__.py b/test/classification/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/test/classification/test_core.py b/test/classification/test_core.py new file mode 100644 index 00000000..e742720d --- /dev/null +++ b/test/classification/test_core.py @@ -0,0 +1,39 @@ +from typing import List, TypeVar, Optional, Tuple +from contextlib import ExitStack as DoesNotRaise + +import pytest +import numpy as np + +from supervision.classification.core import Classifications + +T = TypeVar("T") + + +@pytest.mark.parametrize( + 'class_id, confidence, expected_result, exception', + [ + ( + [0, 1, 2, 3, 4], + [0.1, 0.2, 0.9, 0.4, 0.5], + (np.array([2, 4, 3, 1, 0]), np.array([0.9, 0.5, 0.4, 0.2, 0.1])), + DoesNotRaise() + ), # class_id with 5 numbers and 5 confidences + ( + [0, 1, 2, 3, 4], + [0.1, 0.2, 0.3, 0.4], + None, + pytest.raises(ValueError) + ), # class_id with 5 numbers and 4 confidences + ] +) +def test_top_k( + class_id: List[T], + confidence: Optional[List[T]], + expected_result: Optional[Tuple[List[T], List[T]]], + exception: Exception +) -> None: + with exception: + result = Classifications(class_id=np.array(class_id), confidence=np.array(confidence)).get_top_k(len(class_id)) + + assert result[0].tolist() == expected_result[0].tolist() + assert result[1].tolist() == expected_result[1].tolist()