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