Spaces:
Running
Running
File size: 4,649 Bytes
ad91e86 | 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 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 | """Frozen Grounding DINO and SAM2 target-category inspection."""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path
import numpy as np
import yaml
Detection = tuple[str, float, tuple[float, float, float, float]]
Detector = Callable[[np.ndarray, tuple[str, ...]], tuple[Detection, ...]]
Segmenter = Callable[[np.ndarray, tuple[float, float, float, float]], np.ndarray]
@dataclass(frozen=True)
class InstanceDetection:
category: str
confidence: float
mask: np.ndarray
def detect_instances(detector: Detector, segmenter: Segmenter,
rgb: np.ndarray, category: str) -> tuple[InstanceDetection, ...]:
image = np.asarray(rgb)[..., :3]
found = []
for label, score, box in detector(image, (category,)):
if label != category:
continue
mask = np.asarray(segmenter(image, box), dtype=bool)
if mask.shape == image.shape[:2] and mask.any():
found.append(InstanceDetection(label, float(score), mask))
return tuple(found)
def detect_category(
detector: Detector, segmenter: Segmenter, rgb: np.ndarray, category: str
) -> bool:
"""Use only the rendered RGB image and public target category for STOP."""
return bool(detect_instances(detector, segmenter, rgb, category))
class GroundedSAMInspector:
"""Frozen public models; the private semantic image never enters these heads."""
def __init__(self, config_path: Path, *, dino_model: str, sam_model: str) -> None:
config = yaml.safe_load(config_path.read_text(encoding="utf-8"))
self.detector, self.segmenter = _load_models(
dino_model, sam_model, config
)
def __call__(self, rgb: np.ndarray, _depth: np.ndarray, category: str) -> bool:
return detect_category(self.detector, self.segmenter, rgb, category)
def detect_instances(self, rgb: np.ndarray, category: str) -> tuple[InstanceDetection, ...]:
return detect_instances(self.detector, self.segmenter, rgb, category)
def close(self) -> None:
del self.detector, self.segmenter
def _load_models(dino_model: str, sam_model: str, config: dict):
import torch
from PIL import Image
from transformers import (
AutoModelForZeroShotObjectDetection,
AutoProcessor,
Sam2Model,
Sam2Processor,
)
device = config["device"]
dino_revision = config["grounding_dino"]["revision"]
sam_revision = config["sam"]["revision"]
dino_kwargs = {"revision": dino_revision} if dino_model == config["grounding_dino"]["model_id"] else {}
sam_kwargs = {"revision": sam_revision} if sam_model == config["sam"]["model_id"] else {}
dino_processor = AutoProcessor.from_pretrained(dino_model, **dino_kwargs)
dino = AutoModelForZeroShotObjectDetection.from_pretrained(
dino_model, **dino_kwargs
).to(device).eval()
sam_processor = Sam2Processor.from_pretrained(sam_model, **sam_kwargs)
sam = Sam2Model.from_pretrained(sam_model, **sam_kwargs).to(device).eval()
def detector(rgb: np.ndarray, categories: tuple[str, ...]) -> tuple[Detection, ...]:
image = Image.fromarray(rgb)
inputs = dino_processor(
images=image, text=[list(categories)], return_tensors="pt"
).to(device)
with torch.inference_mode():
outputs = dino(**inputs)
processed = dino_processor.post_process_grounded_object_detection(
outputs,
inputs.input_ids,
threshold=config["grounding_dino"]["box_threshold"],
text_threshold=config["grounding_dino"]["text_threshold"],
target_sizes=[image.size[::-1]],
)[0]
labels = processed.get("text_labels", processed.get("labels", ()))
return tuple(
(str(label), float(score), tuple(float(value) for value in box))
for label, score, box in zip(
labels, processed["scores"], processed["boxes"], strict=True
)
)
def segmenter(rgb: np.ndarray, box: tuple[float, float, float, float]) -> np.ndarray:
image = Image.fromarray(rgb)
inputs = sam_processor(
images=image, input_boxes=[[[*box]]], return_tensors="pt"
).to(device)
with torch.inference_mode():
outputs = sam(**inputs)
masks = sam_processor.post_process_masks(
outputs.pred_masks.cpu(), inputs["original_sizes"]
)[0]
scores = outputs.iou_scores[0, 0].cpu()
return masks[0, int(scores.argmax())].numpy().astype(bool)
return detector, segmenter
|