From 82ed322fba5a59a6360acfa5257fd68f8e3e04c0 Mon Sep 17 00:00:00 2001 From: James Gallagher Date: Fri, 3 Nov 2023 18:25:09 +0000 Subject: [PATCH 1/4] add timm data loader --- supervision/classification/core.py | 41 ++++++++++++++++++++++++++++++ 1 file changed, 41 insertions(+) diff --git a/supervision/classification/core.py b/supervision/classification/core.py index 157275dd..b4d776b6 100644 --- a/supervision/classification/core.py +++ b/supervision/classification/core.py @@ -69,6 +69,47 @@ class Classifications: confidence = ultralytics_results.probs.data.cpu().numpy() return cls(class_id=np.arange(confidence.shape[0]), confidence=confidence) + @classmethod + def from_timm(cls, timm_results) -> Classifications: + """ + Creates a Classifications instance from a + timm (https://huggingface.co/docs/hub/timm) model. + + Args: + timm: The output Results instance from timm model + + Returns: + Classifications: A new Classifications object. + + Example: + ```python + >>> import timm + >>> from PIL import Image + >>> from timm.data import resolve_data_config + >>> from timm.data.transforms_factory import create_transform + + >>> model = timm.create_model('hf-hub:nateraw/resnet50-oxford-iiit-pet', pretrained=True) + >>> model.eval() + + >>> config = resolve_data_config({}, model=model) + >>> transform = create_transform(**config) + + >>> image = Image.open('../image.jpg').convert('RGB') + >>> x = transform(image).unsqueeze(0) + + >>> output = model(x) + + >>> predictions = sv.Classifications.from_timm(output) + ``` + """ + confidence = timm_results.data.cpu().numpy()[0] + class_ids = list(range(len(confidence))) + + if len(class_ids) == 0: + return cls(class_id=np.array([]), confidence=np.array([])) + + return cls(class_id=np.array(class_ids), confidence=confidence) + def get_top_k(self, k: int) -> Tuple[np.ndarray, np.ndarray]: """ Retrieve the top k class IDs and confidences, From ce92069ba6279ef278fdb77aebe5b9962877647c Mon Sep 17 00:00:00 2001 From: James Gallagher Date: Mon, 27 Nov 2023 10:57:00 +0000 Subject: [PATCH 2/4] fix precommit --- supervision/classification/core.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/supervision/classification/core.py b/supervision/classification/core.py index b4d776b6..85310a30 100644 --- a/supervision/classification/core.py +++ b/supervision/classification/core.py @@ -88,7 +88,10 @@ class Classifications: >>> from timm.data import resolve_data_config >>> from timm.data.transforms_factory import create_transform - >>> model = timm.create_model('hf-hub:nateraw/resnet50-oxford-iiit-pet', pretrained=True) + >>> model = timm.create_model( + ... 'hf-hub:nateraw/resnet50-oxford-iiit-pet', + ... pretrained=True + ... ) >>> model.eval() >>> config = resolve_data_config({}, model=model) From 05602cd2a0ce83903000a9b5ca45035087caff39 Mon Sep 17 00:00:00 2001 From: SkalskiP Date: Mon, 27 Nov 2023 12:51:05 +0100 Subject: [PATCH 3/4] Refactor code for clarity and better variable use The variable names and usage within the 'ultralytics' and 'timm' functions were updated for improved code readability. The docstrings were also refined to better explain the function parameters. Instances where YOLO was redundantly called were also eliminated. --- supervision/classification/core.py | 30 +++++++++++++++--------------- 1 file changed, 15 insertions(+), 15 deletions(-) diff --git a/supervision/classification/core.py b/supervision/classification/core.py index 85310a30..befc0736 100644 --- a/supervision/classification/core.py +++ b/supervision/classification/core.py @@ -47,7 +47,7 @@ class Classifications: Args: ultralytics_results (ultralytics.engine.results.Results): - The output Results instance from ultralytics model + The inference result from ultralytics model. Returns: Classifications: A new Classifications object. @@ -60,10 +60,9 @@ class Classifications: >>> image = cv2.imread(SOURCE_IMAGE_PATH) >>> model = YOLO('yolov8n-cls.pt') - >>> model = YOLO('yolov8s-cls.pt') - >>> result = model(image)[0] - >>> classifications = sv.Classifications.from_ultralytics(result) + >>> output = model(image)[0] + >>> classifications = sv.Classifications.from_ultralytics(output) ``` """ confidence = ultralytics_results.probs.data.cpu().numpy() @@ -73,10 +72,10 @@ class Classifications: def from_timm(cls, timm_results) -> Classifications: """ Creates a Classifications instance from a - timm (https://huggingface.co/docs/hub/timm) model. + timm (https://huggingface.co/docs/hub/timm) inference result. Args: - timm: The output Results instance from timm model + timm_results: The inference result from timm model. Returns: Classifications: A new Classifications object. @@ -87,31 +86,32 @@ class Classifications: >>> from PIL import Image >>> from timm.data import resolve_data_config >>> from timm.data.transforms_factory import create_transform + >>> import supervision as sv >>> model = timm.create_model( - ... 'hf-hub:nateraw/resnet50-oxford-iiit-pet', + ... model_name='hf-hub:nateraw/resnet50-oxford-iiit-pet', ... pretrained=True - ... ) - >>> model.eval() + ... ).eval() >>> config = resolve_data_config({}, model=model) >>> transform = create_transform(**config) - >>> image = Image.open('../image.jpg').convert('RGB') + >>> image = Image.open(SOURCE_IMAGE_PATH).convert('RGB') >>> x = transform(image).unsqueeze(0) >>> output = model(x) - >>> predictions = sv.Classifications.from_timm(output) + >>> classifications = sv.Classifications.from_timm(output) ``` """ - confidence = timm_results.data.cpu().numpy()[0] - class_ids = list(range(len(confidence))) + confidence = timm_results.cpu().detach().numpy()[0] - if len(class_ids) == 0: + if len(confidence) == 0: return cls(class_id=np.array([]), confidence=np.array([])) - return cls(class_id=np.array(class_ids), confidence=confidence) + class_id = np.arange(len(confidence)) + + return cls(class_id=class_id, confidence=confidence) def get_top_k(self, k: int) -> Tuple[np.ndarray, np.ndarray]: """ From 1d66840cadd80f13af5a21aed41336292534f3e4 Mon Sep 17 00:00:00 2001 From: SkalskiP Date: Mon, 27 Nov 2023 12:58:34 +0100 Subject: [PATCH 4/4] Optimize imports in core.py Adjustments have been made to the import statements in the 'core.py' file of the 'classification' module. --- supervision/classification/core.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/supervision/classification/core.py b/supervision/classification/core.py index befc0736..9f9393e9 100644 --- a/supervision/classification/core.py +++ b/supervision/classification/core.py @@ -82,10 +82,9 @@ class Classifications: Example: ```python - >>> import timm >>> from PIL import Image - >>> from timm.data import resolve_data_config - >>> from timm.data.transforms_factory import create_transform + >>> import timm + >>> from timm.data import resolve_data_config, create_transform >>> import supervision as sv >>> model = timm.create_model(