File size: 8,175 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
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
"""Open-vocabulary semantic segmenter -- the second, independent perception path.

Where ``detector.py`` answers "there is an object, here is its box", this module
answers "what stuff is at every pixel?". It is deliberately a *different* model
family (CLIPSeg, not YOLO) so its output is genuine corroborating evidence for
the box path rather than a correlated echo of it. The segmentation taxonomy
lives in ``segmentation.yaml`` and, like the box taxonomy, is expressed as
open-vocabulary text prompts -- no fixed-class training required.

Loaded lazily, exactly like the detector and the CLIP classifier: importing this
module is cheap; the ~150 MB CLIPSeg weights are only pulled the first time
``segment`` is called.

Backend note: CLIPSeg is the default because it keeps the whole system
open-vocabulary and needs no dataset-specific fine-tuning. A Cityscapes-trained
closed-set model (e.g. ``nvidia/segformer-b0-finetuned-cityscapes-1024-1024``)
would give crisper masks; it could be dropped in behind the same ``SegResult``
interface without touching the rest of the pipeline.
"""
from __future__ import annotations

from dataclasses import dataclass, field
from functools import lru_cache
from pathlib import Path
from typing import Callable, Optional

import numpy as np
import yaml

from .detector import Box

_SEG_TAXONOMY_PATH = Path(__file__).resolve().parent.parent / "segmentation.yaml"


@dataclass
class SegClass:
    """One entry of the segmentation taxonomy (see segmentation.yaml)."""

    name: str
    prompt: str
    role: str                       # cityscapes-style super-category
    thing: bool                     # discrete object (True) vs. background stuff
    maps_to: str                    # hpercept taxonomy node name, or "" for stuff
    color: tuple[int, int, int]

    @property
    def is_sky(self) -> bool:
        return self.role == "sky"


@dataclass
class SegResult:
    """A dense semantic segmentation of one image.

    ``label_map`` holds, per pixel, an index into ``classes``. Kept intentionally
    simple (a single argmax label per pixel) so the cross-validation logic is a
    handful of transparent array operations, matching the paper's "simple,
    inspectable rules" stance for the validation layer.
    """

    label_map: np.ndarray           # (H, W) int; index into ``classes``
    classes: list[SegClass]

    @property
    def shape(self) -> tuple[int, int]:
        return self.label_map.shape  # type: ignore[return-value]

    def class_of(self, idx: int) -> SegClass:
        return self.classes[idx]

    def _region(self, box: Box) -> np.ndarray:
        """The label sub-array under a box, clipped to the image bounds."""
        h, w = self.label_map.shape
        x1, y1, x2, y2 = box.xyxy
        x1 = max(0, min(x1, w))
        x2 = max(0, min(x2, w))
        y1 = max(0, min(y1, h))
        y2 = max(0, min(y2, h))
        return self.label_map[y1:y2, x1:x2]

    def histogram_in(self, box: Box) -> dict[str, float]:
        """Fraction of the box's pixels assigned to each seg class (name -> frac)."""
        region = self._region(box)
        if region.size == 0:
            return {}
        counts = np.bincount(region.ravel(), minlength=len(self.classes))
        total = float(region.size)
        return {c.name: counts[i] / total for i, c in enumerate(self.classes)}

    def dominant_in(self, box: Box) -> tuple[Optional[SegClass], float]:
        """The most common seg class under a box and its pixel fraction."""
        region = self._region(box)
        if region.size == 0:
            return None, 0.0
        counts = np.bincount(region.ravel(), minlength=len(self.classes))
        idx = int(counts.argmax())
        return self.classes[idx], float(counts[idx]) / float(region.size)

    def fraction_in(self, box: Box, predicate: Callable[[SegClass], bool]) -> float:
        """Fraction of the box's pixels whose class satisfies ``predicate``."""
        region = self._region(box)
        if region.size == 0:
            return 0.0
        keep = np.array([predicate(c) for c in self.classes], dtype=bool)
        return float(keep[region].sum()) / float(region.size)

    def color_map(self) -> np.ndarray:
        """Render the label map to an (H, W, 3) uint8 RGB image."""
        palette = np.array([c.color for c in self.classes], dtype=np.uint8)
        return palette[self.label_map]


def load_seg_taxonomy(path: str | Path = _SEG_TAXONOMY_PATH) -> list[SegClass]:
    data = yaml.safe_load(Path(path).read_text(encoding="utf-8"))
    classes: list[SegClass] = []
    for spec in data["classes"]:
        classes.append(
            SegClass(
                name=spec["name"],
                prompt=spec.get("prompt", spec["name"]),
                role=spec.get("role", "object"),
                thing=bool(spec.get("thing", False)),
                maps_to=str(spec.get("maps_to", "") or ""),
                color=tuple(spec.get("color", [128, 128, 128])),  # type: ignore[arg-type]
            )
        )
    return classes


class Segmenter:
    """Thin wrapper around CLIPSeg with lazy model loading.

    One forward pass scores every taxonomy prompt against the image and we take
    a per-pixel argmax. CLIPSeg has no explicit background class, so the prompt
    set in ``segmentation.yaml`` is kept broad enough (road, building, sky, ...)
    that "nothing here" is rare -- the argmax then just picks the closest stuff
    class, which is the intended behaviour for a dense labelling.
    """

    def __init__(
        self,
        model_name: str = "CIDAS/clipseg-rd64-refined",
        classes: Optional[list[SegClass]] = None,
        device: Optional[str] = None,
    ) -> None:
        self.model_name = model_name
        self.classes = classes or load_seg_taxonomy()
        self._device = device
        self._model = None
        self._processor = None

    # ---- lazy model ---------------------------------------------------- #
    def _ensure_model(self) -> None:
        if self._model is not None:
            return
        # Imported lazily so the app (and the box-only pipeline) can start
        # without paying the transformers import until segmentation is asked for.
        import torch
        from transformers import CLIPSegForImageSegmentation, CLIPSegProcessor

        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"

        self._processor = CLIPSegProcessor.from_pretrained(self.model_name)
        model = CLIPSegForImageSegmentation.from_pretrained(self.model_name)
        self._model = model.to(self._device).eval()

    def segment(self, image_rgb: np.ndarray) -> SegResult:
        """Densely label an RGB image into the segmentation taxonomy."""
        self._ensure_model()
        import torch
        import torch.nn.functional as F
        from PIL import Image

        h, w = image_rgb.shape[:2]
        pil = Image.fromarray(image_rgb)
        prompts = [c.prompt for c in self.classes]

        inputs = self._processor(
            text=prompts,
            images=[pil] * len(prompts),
            padding=True,
            return_tensors="pt",
        ).to(self._device)

        with torch.no_grad():
            logits = self._model(**inputs).logits  # (C, h', w') or (h', w') if C==1
        if logits.dim() == 2:
            logits = logits.unsqueeze(0)

        # Upsample every class heatmap back to the original resolution, then take
        # the per-pixel argmax to get a single dense label map.
        up = F.interpolate(
            logits.unsqueeze(0), size=(h, w), mode="bilinear", align_corners=False
        )[0]
        label_map = up.argmax(dim=0).to("cpu").numpy().astype(np.int32)
        return SegResult(label_map=label_map, classes=self.classes)


@lru_cache(maxsize=1)
def get_segmenter(model_name: str = "CIDAS/clipseg-rd64-refined") -> Segmenter:
    """Process-wide singleton so the segmentation model is loaded at most once."""
    return Segmenter(model_name=model_name)