diff --git a/supervision/detection/vlm.py b/supervision/detection/vlm.py index acee6826..1294b1d4 100644 --- a/supervision/detection/vlm.py +++ b/supervision/detection/vlm.py @@ -421,9 +421,13 @@ def from_google_gemini_2_0( def from_google_gemini_2_5( result: str, resolution_wh: Tuple[int, int], - classes: Optional[List[str]] = None + classes: Optional[List[str]] = None, ) -> Tuple[ - np.ndarray, np.ndarray, np.ndarray, Optional[np.ndarray], Optional[np.ndarray] + np.ndarray, + Optional[np.ndarray], + np.ndarray, + Optional[np.ndarray], + Optional[np.ndarray], ]: """ Parse and scale bounding boxes and masks from Google Gemini 2.5 style @@ -444,6 +448,9 @@ def from_google_gemini_2_5( Args: result: String containing the JSON snippet enclosed by triple backticks. + resolution_wh: (output_width, output_height) to which we rescale the boxes. + classes: Optional list of valid class names. If provided, returned boxes/labels + are filtered to only those classes found here. Returns: xyxy (np.ndarray): An array of shape `(n, 4)` containing @@ -472,22 +479,21 @@ def from_google_gemini_2_5( except json.JSONDecodeError: return ( np.empty((0, 4)), - np.empty((0,), dtype=str), - np.empty((0,), dtype=int), + np.array([], dtype=int), + np.array([], dtype=str), + np.array([], dtype=float), None, - ) - xyxy: list = [] - class_id: list = [] - class_name: list = [] - confidence: list = [] - masks: list = [] + xyxy_list: list = [] + labels_list: list = [] + confidence_list: Optional[list] = [] + masks_list: Optional[list] = [] for item in data: if "box_2d" not in item or "label" not in item: continue - class_name.append(item["label"]) + labels_list.append(item["label"]) box = item["box_2d"] # Gemini bbox order is [y_min, x_min, y_max, x_max] absolute_bbox = denormalize_boxes( @@ -495,65 +501,83 @@ def from_google_gemini_2_5( resolution_wh=(w, h), normalization_factor=1000, ) - xyxy.append(absolute_bbox) + xyxy_list.append(absolute_bbox) if "mask" in item: - png_str = item["mask"] - if not png_str.startswith("data:image/png;base64,"): - masks.append(np.zeros((h, w), dtype=bool)) - continue + if masks_list is not None: + png_str = item["mask"] + if not png_str.startswith("data:image/png;base64,"): + masks_list.append(np.zeros((h, w), dtype=bool)) + continue - png_str = png_str.removeprefix("data:image/png;base64,") - png_str = base64.b64decode(png_str) - mask_img = Image.open(io.BytesIO(png_str)) + png_str = png_str.removeprefix("data:image/png;base64,") + png_str = base64.b64decode(png_str) + mask_img = Image.open(io.BytesIO(png_str)) - y_min, y_max = int(absolute_bbox[1]), int(absolute_bbox[3]) - x_min, x_max = int(absolute_bbox[0]), int(absolute_bbox[2]) + y_min, y_max = int(absolute_bbox[1]), int(absolute_bbox[3]) + x_min, x_max = int(absolute_bbox[0]), int(absolute_bbox[2]) - bbox_height = y_max - y_min - bbox_width = x_max - x_min + bbox_height = y_max - y_min + bbox_width = x_max - x_min - if bbox_height > 0 and bbox_width > 0: - mask_img = mask_img.resize( - (bbox_width, bbox_height), resample=Image.Resampling.BILINEAR - ) - np_mask = np.zeros((h, w), dtype=bool) - np_mask[y_min:y_max, x_min:x_max] = np.array(mask_img) > 0 - masks.append(np_mask) - else: - masks.append(np.zeros((h, w), dtype=bool)) + if bbox_height > 0 and bbox_width > 0: + mask_img = mask_img.resize( + (bbox_width, bbox_height), resample=Image.Resampling.BILINEAR + ) + np_mask = np.zeros((h, w), dtype=bool) + np_mask[y_min:y_max, x_min:x_max] = np.array(mask_img) > 0 + masks_list.append(np_mask) + else: + masks_list.append(np.zeros((h, w), dtype=bool)) else: - masks.append(np.zeros((h, w), dtype=bool)) + masks_list = None if "confidence" in item: - confidence.append(item["confidence"]) + if confidence_list is not None: + confidence_list.append(item["confidence"]) else: - confidence.append(0.0) + confidence_list = None - if not xyxy: + if not xyxy_list: return ( np.empty((0, 4)), - np.array([], dtype=int), + np.array([], dtype=int), np.array([], dtype=str), - np.array([], dtype=np.float32), + np.array([], dtype=float), None, ) + xyxy = np.array(xyxy_list, dtype=float) + class_name = np.array(labels_list) + class_id: np.ndarray if classes is not None: mask = np.array([name in classes for name in class_name], dtype=bool) xyxy = xyxy[mask] class_name = class_name[mask] - class_id = np.array([classes.index(name) for name in class_name], dtype=int) - masks = [masks[i] for i in range(len(masks)) if mask[i]] + class_id = np.array([classes.index(name) for name in class_name]) + if masks_list is not None: + masks_list = [masks_list[i] for i, m in enumerate(mask) if m] + if confidence_list is not None: + confidence_list = [c for c, m in zip(confidence_list, mask) if m] + else: + # When classes is None, generate class_id based on unique labels + unique_labels = sorted(list(set(class_name))) + label_to_id = {label: i for i, label in enumerate(unique_labels)} + class_id = np.array([label_to_id[name] for name in class_name]) + + confidence = ( + np.array(confidence_list, dtype=float) if confidence_list is not None else None + ) + masks = np.array(masks_list) if masks_list is not None else None return ( - np.array(xyxy, dtype=float), - np.array(class_id, dtype=int), - np.array(class_name, dtype=str), - np.array(confidence, dtype=float), - np.array(masks) if masks is not None else None, + xyxy, + class_id, + class_name, + confidence, + masks, ) @@ -583,7 +607,7 @@ def from_moondream( Args: result: Dictionary containing the JSON output from the model. resolution_wh: (output_width, output_height) to which we rescale the boxes. - + Returns: xyxy (np.ndarray): An array of shape `(n, 4)` containing the bounding boxes coordinates in format `[x1, y1, x2, y2]` @@ -596,7 +620,7 @@ def from_moondream( ) if "objects" not in result or not isinstance(result["objects"], list): - return np.empty((0, 4)) + return np.empty((0, 4), dtype=float) denormalize_xyxy = [] diff --git a/test/detection/test_vlm.py b/test/detection/test_vlm.py index c2762f5d..2135600b 100644 --- a/test/detection/test_vlm.py +++ b/test/detection/test_vlm.py @@ -8,6 +8,7 @@ import pytest from supervision.detection.vlm import ( from_florence_2, from_google_gemini_2_0, + from_google_gemini_2_5, from_moondream, from_paligemma, from_qwen_2_5_vl, @@ -883,3 +884,204 @@ def test_florence_2( assert result[3] is None else: np.testing.assert_array_equal(result[3], expected_results[3]) + + +@pytest.mark.parametrize( + "exception, result, resolution_wh, classes, expected_results", + [ + ( + does_not_raise(), + "random text", + (1000, 1000), + None, + ( + np.empty((0, 4)), + np.empty(0, dtype=int), + np.empty(0, dtype=str), + np.empty(0, dtype=float), + None, + ), + ), + ( + does_not_raise(), + "```json\ninvalid json\n```", + (1000, 1000), + None, + ( + np.empty((0, 4)), + np.empty(0, dtype=int), + np.empty(0, dtype=str), + np.empty(0, dtype=float), + None, + ), + ), + ( + does_not_raise(), + "```json\n[]\n```", + (1000, 1000), + None, + ( + np.empty((0, 4)), + np.empty(0, dtype=int), + np.empty(0, dtype=str), + np.empty(0, dtype=float), + None, + ), + ), + ( + does_not_raise(), + """```json + [ + {"box_2d": [100, 200, 300, 400], "label": "cat", "confidence": 0.8} + ] + ```""", + (1000, 500), + None, + ( + np.array([[200.0, 50.0, 400.0, 150.0]]), + np.array([0]), + np.array(["cat"], dtype=str), + np.array([0.8]), + None, + ), + ), + ( + does_not_raise(), + """```json + [ + {"box_2d": [10, 20, 110, 120], "label": "cat", "confidence": 0.8}, + {"box_2d": [50, 100, 150, 200], "label": "dog", "confidence": 0.9} + ] + ```""", + (640, 480), + None, + ( + np.array([[12.8, 4.8, 76.8, 52.8], [64.0, 24.0, 128.0, 72.0]]), + np.array([0, 1]), + np.array(["cat", "dog"], dtype=str), + np.array([0.8, 0.9]), + None, + ), + ), + ( + does_not_raise(), + """```json + [ + {"box_2d": [10, 20, 110, 120], "label": "cat", "confidence": 0.8} + ] + ```""", + (640, 480), + ["dog", "person"], + ( + np.empty((0, 4)), + np.empty(0, dtype=int), + np.empty(0, dtype=str), + np.empty(0, dtype=float), + None, + ), + ), + ( + does_not_raise(), + """```json + [ + {"box_2d": [10, 20, 110, 120], "label": "cat", "confidence": 0.8}, + {"box_2d": [50, 100, 150, 200], "label": "dog", "confidence": 0.9} + ] + ```""", + (640, 480), + ["person", "dog"], + ( + np.array([[64.0, 24.0, 128.0, 72.0]]), + np.array([1]), + np.array(["dog"], dtype=str), + np.array([0.9]), + None, + ), + ), + ( + does_not_raise(), + """```json + [ + {"box_2d": [10, 20, 110, 120], "label": "cat", "confidence": 0.8}, + {"box_2d": [50, 100, 150, 200], "label": "dog", "confidence": 0.9} + ] + ```""", + (640, 480), + ["cat", "dog"], + ( + np.array([[12.8, 4.8, 76.8, 52.8], [64.0, 24.0, 128.0, 72.0]]), + np.array([0, 1]), + np.array(["cat", "dog"]), + np.array([0.8, 0.9]), + None, + ), + ), + ( + pytest.raises(ValueError), + """```json + [ + {"box_2d": [10, 20, 110, 120], "label": "cat"} + ] + ```""", + (0, 480), + None, + None, + ), + ( + pytest.raises(ValueError), + """```json + [ + {"box_2d": [10, 20, 110, 120], "label": "cat"} + ] + ```""", + (640, -100), + None, + None, + ), + ], +) +def test_from_google_gemini_2_5( + exception, + result: str, + resolution_wh: Tuple[int, int], + classes: Optional[List[str]], + expected_results: Optional[ + Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray] + ], +): + with exception: + ( + xyxy, + class_id, + class_name, + confidence, + masks, + ) = from_google_gemini_2_5( + result=result, resolution_wh=resolution_wh, classes=classes + ) + + if expected_results is None: + return + + assert xyxy.shape == expected_results[0].shape + assert np.allclose(xyxy, expected_results[0]) + + assert class_id.shape == expected_results[1].shape + assert np.array_equal(class_id, expected_results[1]) + + assert class_name.shape == expected_results[2].shape + assert np.array_equal(class_name, expected_results[2]) + + if confidence is None: + assert expected_results[3] is None + else: + assert expected_results[3] is not None + assert confidence.shape == expected_results[3].shape + assert np.allclose(confidence, expected_results[3]) + + if masks is None: + assert expected_results[4] is None + else: + assert masks is not None + assert masks.shape == expected_results[4].shape + assert np.array_equal(masks, expected_results[4])