File size: 10,879 Bytes
f656a1a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b467a01
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f656a1a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
"""Fusion Perception 1 v0.1 — landmark / place retrieval descriptor.

A frozen DINOv2-L ViT-L/14 backbone (Apache-2.0) plus a small trained projection head
(~3.1M parameters, this repo's only weights). One image goes in, one 512-d L2-normalized
descriptor comes out; rank a gallery by cosine similarity.

    from PIL import Image
    from inference import FusionPerceptionRetrieval

    fp = FusionPerceptionRetrieval.from_pretrained("EximiusLabs/fusion-perception-1-preview")

    gallery = fp.embed([Image.open(p) for p in paths])   # [N, 512]
    query   = fp.embed(Image.open("query.jpg"))          # [512]
    scores, idx = fp.search(query, gallery, topk=10)

The descriptor is multi-scale: the backbone runs at three resolutions (short side 224,
322, 448 px, aspect preserved, long side capped at 1022), the CLS token is L2-normalized
per scale, the three are averaged and re-normalized, and the head maps that 1024-d vector
to 512-d. The backbone is never fine-tuned.

Two heads ship. `standard` is trained on all of GLDv2-clean, the same data published
systems use. `decon` excludes the 87 classes our audit matched to ROxford/RParis query
landmarks. See README_hf.md for what each is for and for the measured difference.

Requires: torch, transformers>=4.46, pillow, numpy, huggingface_hub (only for
from_pretrained on a repo id). The DINOv2 backbone downloads from Meta's repository.

Reported numbers were produced in float16 on CUDA. CPU runs default to float32 and will
differ in the last decimal.
"""
from __future__ import annotations

import json
import os
from typing import Iterable, Optional, Sequence, Tuple, Union

import torch
import torch.nn as nn
import torch.nn.functional as F

DEFAULT_BACKBONE = "facebook/dinov2-large"
SCALES = (1.0, 1.414, 2.0)
BASE_SHORT = 224
MAX_LONG = 1022
PATCH = 14
IMAGENET_MEAN = (0.485, 0.456, 0.406)
IMAGENET_STD = (0.229, 0.224, 0.225)


def _snap(v: float) -> int:
    """Round a side length to a multiple of the patch size (minimum two patches)."""
    return max(PATCH * 2, int(round(v / PATCH)) * PATCH)


def target_size(w: int, h: int, scale: float) -> Tuple[int, int]:
    """Resize geometry for one scale: (width, height), aspect preserved.

    Short side is 224*scale snapped to a multiple of 14; the long side follows the
    aspect ratio, is snapped the same way, and is clamped to 1022.
    """
    short = _snap(BASE_SHORT * scale)
    if w <= h:
        return short, min(_snap(h * short / w), MAX_LONG)
    return min(_snap(w * short / h), MAX_LONG), short


def make_head(emb: int = 512) -> nn.Sequential:
    """The trained head: 1024-d multi-scale CLS -> 512-d descriptor (pre-normalization)."""
    return nn.Sequential(nn.Linear(1024, 2048), nn.GELU(),
                         nn.Linear(2048, emb, bias=False), nn.BatchNorm1d(emb))


class FusionPerceptionRetrieval:
    """Frozen DINOv2-L multi-scale CLS + trained projection head -> 512-d descriptor."""

    def __init__(self, root: str, device: Optional[str] = None,
                 protocol: Optional[str] = None, backbone: Optional[str] = None,
                 dtype: Optional[torch.dtype] = None):
        from transformers import AutoModel

        self.root = root
        with open(os.path.join(root, "config.json")) as fh:
            self.cfg = json.load(fh)
        self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
        self.dtype = dtype or (torch.float16 if self.device.startswith("cuda")
                               else torch.float32)
        self.protocol = protocol or self.cfg.get("default_head", "standard")
        heads = self.cfg["heads"]
        if self.protocol not in heads:
            raise ValueError(f"unknown head '{self.protocol}'; available: {list(heads)}")

        self.backbone_id = backbone or self.cfg.get("backbone", DEFAULT_BACKBONE)
        self.model = (AutoModel.from_pretrained(self.backbone_id)
                      .to(device=self.device, dtype=self.dtype).eval())
        for p in self.model.parameters():
            p.requires_grad_(False)

        self.head = make_head(int(self.cfg["descriptor"]["dim"]))
        entry = heads[self.protocol]
        st_name = entry.get("safetensors")
        st_path = os.path.join(root, st_name) if st_name else None
        if st_path and os.path.exists(st_path):
            from safetensors import safe_open

            with safe_open(st_path, framework="pt", device="cpu") as f:
                state = {k: f.get_tensor(k) for k in f.keys()}
                meta = dict(f.metadata() or {})
            self.head.load_state_dict(state)
            # metadata values are strings in safetensors; the JSON-encoded ones decode
            self.head_meta = {}
            for k, v in meta.items():
                try:
                    self.head_meta[k] = json.loads(v)
                except (ValueError, TypeError):
                    self.head_meta[k] = v
        else:
            ck = torch.load(os.path.join(root, entry["file"]),
                            map_location="cpu", weights_only=False)
            self.head.load_state_dict(ck["state_dict"])
            self.head_meta = {k: v for k, v in ck.items() if k != "state_dict"}
        self.head = self.head.to(self.device).float().eval()

        pre = self.cfg.get("preprocess", {})
        self.scales = tuple(pre.get("scales", SCALES))
        mean = torch.tensor(pre.get("mean", IMAGENET_MEAN)).view(1, 3, 1, 1)
        std = torch.tensor(pre.get("std", IMAGENET_STD)).view(1, 3, 1, 1)
        self._mean = mean.to(self.device)
        self._std = std.to(self.device)

    @classmethod
    def from_pretrained(cls, repo_or_path: str, **kw) -> "FusionPerceptionRetrieval":
        if os.path.isdir(repo_or_path):
            root = repo_or_path
        else:
            from huggingface_hub import snapshot_download
            root = snapshot_download(repo_or_path)
        return cls(root, **kw)

    # ------------------------------------------------------------------ backbone
    @torch.no_grad()
    def _cls_at_scale(self, pils: Sequence, scale: float) -> torch.Tensor:
        """L2-normalized CLS token for a batch of PIL images at one scale: [B, 1024].

        Images are grouped by resize shape so a batch of same-shaped images runs in one
        forward pass; mixed aspect ratios simply produce more groups.
        """
        import numpy as np
        from PIL import Image
        groups: dict = {}
        for i, im in enumerate(pils):
            groups.setdefault(target_size(im.size[0], im.size[1], scale), []).append(i)
        out = [None] * len(pils)
        for (nw, nh), idxs in groups.items():
            budget = max(1, 46000 // ((nw // PATCH) * (nh // PATCH) + 5))
            for s0 in range(0, len(idxs), budget):
                grp = idxs[s0:s0 + budget]
                arr = [torch.from_numpy(
                    np.array(pils[i].resize((nw, nh), Image.BICUBIC))
                ).permute(2, 0, 1) for i in grp]
                x = torch.stack(arr).float().to(self.device) / 255.0
                x = ((x - self._mean) / self._std).to(self.dtype)
                hs = self.model(pixel_values=x).last_hidden_state
                c = F.normalize(hs[:, 0].float(), dim=-1).cpu()
                for j, i in enumerate(grp):
                    out[i] = c[j]
        return torch.stack(out)

    @torch.no_grad()
    def embed_backbone(self, images, bbox: Optional[Iterable] = None) -> torch.Tensor:
        """Multi-scale CLS descriptor before the head: [B, 1024], L2-normalized.

        bbox, when given, is one (x1, y1, x2, y2) crop box per image, applied before
        resizing. The revisitop benchmarks crop queries this way; gallery images are
        never cropped.
        """
        one = not isinstance(images, (list, tuple))
        pils = [images] if one else list(images)
        pils = [im.convert("RGB") for im in pils]
        if bbox is not None:
            boxes = [bbox] if one else list(bbox)
            if len(boxes) != len(pils):
                raise ValueError("bbox must have one box per image")
            pils = [im if b is None else im.crop(tuple(int(v) for v in b))
                    for im, b in zip(pils, boxes)]
        acc = torch.stack([self._cls_at_scale(pils, s) for s in self.scales]).mean(0)
        v = F.normalize(acc, dim=-1)
        return v[0] if one else v

    # ------------------------------------------------------------------ descriptor
    @torch.no_grad()
    def embed(self, images, bbox: Optional[Iterable] = None) -> torch.Tensor:
        """512-d L2-normalized retrieval descriptor: [512] for one image, else [B, 512].

        Index these and rank a gallery by cosine similarity (a plain dot product, since
        the vectors are unit norm).
        """
        one = not isinstance(images, (list, tuple))
        ms = self.embed_backbone(images, bbox=bbox)
        if one:
            ms = ms.unsqueeze(0)
        e = self.head(ms.float().to(self.device))
        e = F.normalize(e, dim=-1).cpu()
        return e[0] if one else e

    # ------------------------------------------------------------------ ranking
    @staticmethod
    def search(query: torch.Tensor, gallery: torch.Tensor,
               topk: int = 10) -> Tuple[torch.Tensor, torch.Tensor]:
        """Rank a gallery by cosine similarity. Returns (scores, indices).

        query: [512] or [Q, 512]; gallery: [N, 512]. Both must be L2-normalized, which
        is what embed() returns. topk is clamped to the gallery size.
        """
        q = query.unsqueeze(0) if query.dim() == 1 else query
        sims = q.float() @ gallery.float().T
        scores, idx = torch.topk(sims, k=min(topk, gallery.shape[0]), dim=1)
        if query.dim() == 1:
            return scores[0], idx[0]
        return scores, idx


# Backwards-compatible alias.
FusionPerception = FusionPerceptionRetrieval


if __name__ == "__main__":
    import numpy as np
    from PIL import Image

    root = os.path.dirname(os.path.abspath(__file__))
    fp = FusionPerceptionRetrieval.from_pretrained(root)
    print(f"backbone={fp.backbone_id} head={fp.protocol} "
          f"device={fp.device} dtype={fp.dtype}")

    rng = np.random.default_rng(0)
    imgs = [Image.fromarray((rng.random((h, w, 3)) * 255).astype("uint8"))
            for h, w in ((240, 320), (200, 200), (320, 240))]
    gallery = fp.embed(imgs)
    query = fp.embed(imgs[0])
    scores, idx = fp.search(query, gallery, topk=3)
    print("gallery:", tuple(gallery.shape),
          "| unit norm:", bool(torch.allclose(gallery.norm(dim=-1),
                                              torch.ones(len(imgs)), atol=1e-3)))
    print("query:", tuple(query.shape), "| top-3 idx:", idx.tolist(),
          "| top-3 scores:", [round(s, 4) for s in scores.tolist()])
    assert int(idx[0]) == 0, "an image must retrieve itself first"
    print("OK")