File size: 6,491 Bytes
a9d655b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
"""CLIP zero-shot classifier over taxonomy nodes.

This is what makes the system *open-vocabulary*: every taxonomy node (not just
the leaves) has a text prompt, and we can score how well any image crop matches
each node. That lets the abstraction logic ask "is this confidently a Truck? no?
then is it confidently a Transport Vehicle? ..." all the way up the tree, even
for objects the underlying detector was never trained on.
"""
from __future__ import annotations

from functools import lru_cache
from typing import Optional

import numpy as np

from .taxonomy import Node, Taxonomy


# Out-of-ODD and background prompts for an open-set reject (spike B).
# The idea is an ODD gate: before taxonomic classification, ask "is this even a
# road object?". If the crop matches one of these better than the best taxonomy
# leaf, it is out-of-distribution for our operational design domain and should be
# flagged, not forced into a leaf/floor. These are broad categories, not the exact
# failure classes, so the check is a gate rather than an enumerated blocklist.
NEGATIVE_PROMPTS = [
    "the empty sky",
    "an aircraft",
    "a boat on the water",
    "an indoor scene",
    "a sports ball or flying toy",
    "a plain textureless background",
    "a blurry out-of-focus region",
    "an unremarkable piece of the background",
]


class ClipClassifier:
    """Wraps open_clip and pre-computes text features for every taxonomy node."""

    def __init__(
        self,
        taxonomy: Taxonomy,
        # NOTE: OpenAI CLIP weights were trained with the QuickGELU activation.
        # The plain "ViT-B-32" config uses nn.GELU and silently degrades
        # accuracy, so we must pair the "-quickgelu" model with the openai tag.
        model_name: str = "ViT-B-32-quickgelu",
        pretrained: str = "openai",
        device: Optional[str] = None,
    ) -> None:
        self.taxonomy = taxonomy
        self.model_name = model_name
        self.pretrained = pretrained
        self._device = device
        self._model = None
        self._preprocess = None
        self._tokenizer = None
        self._text_feats: dict[str, np.ndarray] = {}
        self._neg_feats: Optional[np.ndarray] = None   # (K, D) negative prompts
        self._neg_labels: list[str] = []

    # ---- lazy model ---------------------------------------------------- #
    def _ensure_model(self) -> None:
        if self._model is not None:
            return
        import open_clip
        import torch

        if self._device is None:
            if torch.cuda.is_available():
                self._device = "cuda"
            elif torch.backends.mps.is_available():
                self._device = "mps"
            else:
                self._device = "cpu"

        model, _, preprocess = open_clip.create_model_and_transforms(
            self.model_name, pretrained=self.pretrained
        )
        model = model.to(self._device).eval()
        self._model = model
        self._preprocess = preprocess
        self._tokenizer = open_clip.get_tokenizer(self.model_name)
        self._encode_taxonomy_text()
        self._encode_negatives()

    def _encode_negatives(self) -> None:
        import torch

        prompts = [f"a photo of {p}" for p in NEGATIVE_PROMPTS]
        tokens = self._tokenizer(prompts).to(self._device)
        with torch.no_grad():
            feats = self._model.encode_text(tokens)
            feats = feats / feats.norm(dim=-1, keepdim=True)
        self._neg_feats = feats.cpu().numpy()
        self._neg_labels = list(NEGATIVE_PROMPTS)

    def negatives_max_sim(self, image_feat: np.ndarray) -> tuple[float, str]:
        """Best cosine to any out-of-ODD / background prompt, and its label."""
        self._ensure_model()
        sims = self._neg_feats @ image_feat
        i = int(sims.argmax())
        return float(sims[i]), self._neg_labels[i]

    def _encode_taxonomy_text(self) -> None:
        import torch

        nodes = [n for n in self.taxonomy.iter_nodes()]
        prompts = [self._templated(n) for n in nodes]
        tokens = self._tokenizer(prompts).to(self._device)
        with torch.no_grad():
            feats = self._model.encode_text(tokens)
            feats = feats / feats.norm(dim=-1, keepdim=True)
        feats_np = feats.cpu().numpy()
        for node, f in zip(nodes, feats_np):
            self._text_feats[node.name] = f

    @staticmethod
    def _templated(node: Node) -> str:
        # A light prompt template improves CLIP zero-shot separability.
        return f"a photo of {node.prompt}"

    # ---- inference ----------------------------------------------------- #
    def image_features(self, crop_rgb: np.ndarray) -> np.ndarray:
        """Return an L2-normalized CLIP embedding for one RGB crop."""
        self._ensure_model()
        import torch
        from PIL import Image

        if crop_rgb.size == 0:
            # Degenerate crop -> return a zero vector (matches nothing well).
            dim = next(iter(self._text_feats.values())).shape[0]
            return np.zeros(dim, dtype=np.float32)

        pil = Image.fromarray(crop_rgb)
        tensor = self._preprocess(pil).unsqueeze(0).to(self._device)
        with torch.no_grad():
            feat = self._model.encode_image(tensor)
            feat = feat / feat.norm(dim=-1, keepdim=True)
        return feat.cpu().numpy()[0]

    def similarities(self, crop_rgb: np.ndarray, nodes: list[Node]) -> dict[str, float]:
        """Cosine similarity of the crop against each of ``nodes``."""
        self._ensure_model()
        img = self.image_features(crop_rgb)
        return {n.name: float(np.dot(img, self._text_feats[n.name])) for n in nodes}

    def child_probs(
        self, image_feat: np.ndarray, children: list[Node], temperature: float = 0.01
    ) -> dict[str, float]:
        """Softmax distribution over a node's children given a precomputed crop
        embedding. Used by the top-down hierarchical descent."""
        sims = np.array(
            [float(np.dot(image_feat, self._text_feats[c.name])) for c in children]
        )
        scaled = sims / max(temperature, 1e-6)
        scaled -= scaled.max()
        exp = np.exp(scaled)
        probs = exp / exp.sum()
        return {c.name: float(p) for c, p in zip(children, probs)}


@lru_cache(maxsize=1)
def get_classifier(taxonomy_id: int) -> "ClipClassifier":  # pragma: no cover
    # Not used directly (taxonomy isn't hashable); see pipeline for wiring.
    raise NotImplementedError