File size: 9,213 Bytes
326e3fb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""MODA Duo -- two open constituents, one answer per query.

Every query is routed to whichever constituent suits its shape:

    catalogue titles      ->  MODA Pro Lite+   (HopitAI/moda-pro-lite + its recipe)
    long descriptions     ->  MODA             (Marqo/marqo-fashionSigLIP + its recipe)

Both constituents are open. Duo adds ZERO parameters. The default router is a
word-count rule frozen on development data; it is a plain callable, so any policy
that maps a query string to a constituent name may replace it.

Serving cost
------------
  indexes                 : 2      one per constituent, both built offline
  stored vectors per item : 2
  encoders run per query  : 1      only the routed constituent's text tower
  ANN queries per search  : 1
  re-ranking              : none

    pip install open_clip_torch pillow numpy hnswlib
    python serving_ann.py --demo
"""
from __future__ import annotations

import argparse
import math
from dataclasses import dataclass, field
from typing import Callable, Sequence

import numpy as np
import torch
import torch.nn.functional as F
from PIL import Image

THRESHOLD_WORDS = 36          # frozen on development data; see config.json

PROMPT_TEMPLATES = {
    "raw": "{query}",
    "photo": "a photo of {query}",
    "product": "a fashion product photo of {query}",
}


# ---------------------------------------------------------------- views ----
def square_pad(image: Image.Image, fill: int = 128) -> Image.Image:
    image = image.convert("RGB")
    side = max(image.size)
    canvas = Image.new("RGB", (side, side), (fill, fill, fill))
    canvas.paste(image, ((side - image.width) // 2, (side - image.height) // 2))
    return canvas


def center_square(image: Image.Image) -> Image.Image:
    side = min(image.size)
    left, top = (image.width - side) // 2, (image.height - side) // 2
    return image.convert("RGB").crop((left, top, left + side, top + side))


def foreground_square(image: Image.Image) -> Image.Image:
    """Crop a near-white catalogue border, then pad without distorting aspect."""
    image = image.convert("RGB")
    preview = image.copy()
    preview.thumbnail((256, 256), Image.Resampling.BILINEAR)
    mask = np.any(np.asarray(preview, dtype=np.uint8) < 242, axis=-1)
    if float(mask.mean()) < 0.01:
        return square_pad(image)
    ys, xs = np.nonzero(mask)
    sx, sy = image.width / preview.width, image.height / preview.height
    left = max(0, math.floor(float(xs.min()) * sx))
    right = min(image.width, math.ceil(float(xs.max() + 1) * sx))
    top = max(0, math.floor(float(ys.min()) * sy))
    bottom = min(image.height, math.ceil(float(ys.max() + 1) * sy))
    mx, my = max(1, round((right - left) * 0.05)), max(1, round((bottom - top) * 0.05))
    return square_pad(image.crop((max(0, left - mx), max(0, top - my),
                                  min(image.width, right + mx), min(image.height, bottom + my))))


VIEWS = {
    "official": lambda im: im.convert("RGB"),
    "pad": square_pad,
    "pad_white": lambda im: square_pad(im, fill=255),
    "center_crop": center_square,
    "foreground_pad": foreground_square,
}


# ---------------------------------------------------------- constituents ----
@dataclass
class Constituent:
    name: str
    repo: str
    image_mix: dict[str, float]
    prompt_mix: dict[str, float]
    enc: tuple = field(default=None, repr=False)

    def load(self, device: str = "cpu") -> "Constituent":
        import open_clip
        model, _, preprocess = open_clip.create_model_and_transforms(self.repo)
        model.eval().to(device)
        for p in model.parameters():
            p.requires_grad = False
        self.enc = (model, preprocess, open_clip.get_tokenizer(self.repo), device)
        return self

    @torch.inference_mode()
    def encode_images(self, images: Sequence[Image.Image], batch_size: int = 32) -> np.ndarray:
        model, preprocess, _, device = self.enc
        parts = {v: [] for v in self.image_mix}
        for s in range(0, len(images), batch_size):
            chunk = images[s:s + batch_size]
            for view in self.image_mix:
                px = torch.stack([preprocess(VIEWS[view](im)) for im in chunk]).to(device)
                parts[view].append(F.normalize(model.encode_image(px).float(), dim=-1).cpu())
        fused = sum(w * torch.cat(parts[v]) for v, w in self.image_mix.items())
        return F.normalize(fused, dim=-1).numpy().astype("float32")

    @torch.inference_mode()
    def encode_queries(self, texts: Sequence[str], batch_size: int = 128) -> np.ndarray:
        model, _, tokenizer, device = self.enc
        parts = {p: [] for p in self.prompt_mix}
        for s in range(0, len(texts), batch_size):
            chunk = texts[s:s + batch_size]
            for prompt in self.prompt_mix:
                tok = tokenizer([PROMPT_TEMPLATES[prompt].format(query=t) for t in chunk]).to(device)
                parts[prompt].append(F.normalize(model.encode_text(tok).float(), dim=-1).cpu())
        fused = sum(w * torch.cat(parts[p]) for p, w in self.prompt_mix.items())
        return F.normalize(fused, dim=-1).numpy().astype("float32")


CONSTITUENTS = {
    "moda": Constituent(
        "MODA", "hf-hub:Marqo/marqo-fashionSigLIP",
        {"official": 1.0, "pad_white": 0.25, "center_crop": 0.25},
        {"raw": 1.0, "product": 0.25}),
    "moda_pro_lite_plus": Constituent(
        "MODA Pro Lite+", "hf-hub:HopitAI/moda-pro-lite",
        {"official": 1.0, "pad": 0.25, "foreground_pad": 0.25},
        {"raw": 1.0, "photo": 0.25}),
}


# ---------------------------------------------------------------- router ----
def word_count_router(query: str, threshold: int = THRESHOLD_WORDS) -> str:
    """Default policy. Replace with any callable(query) -> constituent name."""
    return "moda_pro_lite_plus" if len(query.split()) <= threshold else "moda"


# ------------------------------------------------------------------ Duo ----
class Duo:
    def __init__(self, router: Callable[[str], str] = word_count_router, device: str = "cpu"):
        self.router = router
        self.c = {k: v.load(device) for k, v in CONSTITUENTS.items()}
        self.index: dict[str, object] = {}
        self.vectors: dict[str, np.ndarray] = {}

    def build(self, images: Sequence[Image.Image], m: int = 32, ef_construction: int = 200) -> None:
        """Index time: encode the catalogue with BOTH constituents, once."""
        import hnswlib
        for name, c in self.c.items():
            vec = c.encode_images(images)
            idx = hnswlib.Index(space="ip", dim=vec.shape[1])
            idx.init_index(max_elements=len(vec), ef_construction=ef_construction, M=m)
            idx.add_items(vec, np.arange(len(vec)))
            self.index[name], self.vectors[name] = idx, vec

    def search(self, queries: Sequence[str], k: int = 10, ef: int = 64):
        """Query time: ONE encoder, ONE ANN query."""
        routes = [self.router(q) for q in queries]
        ids = np.zeros((len(queries), k), dtype=np.int64)
        scores = np.zeros((len(queries), k), dtype=np.float32)
        for name in set(routes):
            rows = [i for i, r in enumerate(routes) if r == name]
            qv = self.c[name].encode_queries([queries[i] for i in rows])
            self.index[name].set_ef(max(ef, k))
            got_ids, dist = self.index[name].knn_query(qv, k=k)
            ids[rows], scores[rows] = got_ids, 1.0 - dist
        return ids, scores, routes

    def search_exact(self, queries: Sequence[str], k: int = 10):
        """Ground truth, for checking ANN recall on your own corpus."""
        routes = [self.router(q) for q in queries]
        ids = np.zeros((len(queries), k), dtype=np.int64)
        for name in set(routes):
            rows = [i for i, r in enumerate(routes) if r == name]
            qv = self.c[name].encode_queries([queries[i] for i in rows])
            sims = qv @ self.vectors[name].T
            ids[rows] = np.argsort(-sims, axis=1)[:, :k]
        return ids, routes


def _demo() -> None:
    duo = Duo()
    corpus = [Image.new("RGB", (w, h), c) for w, h, c in
              [(224, 300, "white"), (300, 224, "black"), (256, 256, "navy"),
               (400, 200, "beige"), (200, 400, "maroon")]]
    duo.build(corpus)
    queries = [
        "black leather ankle boots",
        "A woman is wearing a long navy wool coat with wide lapels, belted at the waist, "
        "over a cream turtleneck and dark trousers, styled for a cold city morning with a "
        "leather tote and ankle boots and a soft grey scarf wrapped twice",
    ]
    ids, scores, routes = duo.search(queries, k=3)
    exact, _ = duo.search_exact(queries, k=3)
    for q, r, a, b in zip(queries, routes, ids, exact):
        print(f"[{r:18s}] {q[:44]!r:48s} ANN {a.tolist()}  exact {b.tolist()}")
    agree = float(np.mean([len(set(a) & set(b)) / len(b) for a, b in zip(ids, exact)]))
    print(f"ANN/exact overlap@3: {agree:.3f}   indexes: {len(duo.index)}   encoders per query: 1")


if __name__ == "__main__":
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument("--demo", action="store_true")
    args = ap.parse_args()
    _demo() if args.demo else ap.print_help()