File size: 7,034 Bytes
3eaff2e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""MODA Pro Lite+ (moda-pro-lite with its calibrated serving harness) -- retrieval with ANN, end to end.

The harness below was selected on held-out development data only (OpenVTON + GLAMI); no target benchmark was touched during selection.
It adds ZERO parameters and stores ONE vector per item.

Serving cost
------------
  stored vectors per item : 1
  ANN queries per search  : 1
  image forwards at index : 3x (paid once, offline)
  text forwards per query : 2x (cheap next to the ANN probe)

The harness is a recipe for WHAT YOU ENCODE, not a model change: the views are
combined into a single unit vector before indexing, so nearest-neighbour search
costs exactly what it costs for the bare model. No extra routes, no re-ranking.

    pip install open_clip_torch pillow numpy hnswlib

Example
-------
    python serving_ann.py --demo
"""
from __future__ import annotations

import argparse
import math
from typing import Sequence

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

MODEL = "hf-hub:HopitAI/moda-pro-lite"
IMAGE_MIX = {'official': 1.0, 'pad': 0.25, 'foreground_pad': 0.25}          # view -> weight, combined then L2-normalised
PROMPT_MIX = {'raw': 1.0, 'photo': 0.25}        # prompt -> weight, combined then L2-normalised
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 catalog 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,
}


# -------------------------------------------------------------- encoding ----
def load(device: str = "cpu"):
    import open_clip
    model, _, preprocess = open_clip.create_model_and_transforms(MODEL)
    model.eval().to(device)
    for p in model.parameters():
        p.requires_grad = False
    return model, preprocess, open_clip.get_tokenizer(MODEL), device


@torch.inference_mode()
def encode_images(images: Sequence[Image.Image], enc, batch_size: int = 32) -> np.ndarray:
    """-> (n, 768) float32, unit norm. ONE vector per item."""
    model, preprocess, _, device = enc
    parts = {v: [] for v in IMAGE_MIX}
    for s in range(0, len(images), batch_size):
        chunk = images[s:s + batch_size]
        for view in 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 IMAGE_MIX.items())
    return F.normalize(fused, dim=-1).numpy().astype("float32")


@torch.inference_mode()
def encode_queries(texts: Sequence[str], enc, batch_size: int = 128) -> np.ndarray:
    """-> (n, 768) float32, unit norm. ONE vector per query."""
    model, _, tokenizer, device = enc
    parts = {p: [] for p in PROMPT_MIX}
    for s in range(0, len(texts), batch_size):
        chunk = texts[s:s + batch_size]
        for prompt in PROMPT_MIX:
            rendered = [PROMPT_TEMPLATES[prompt].format(query=t) for t in chunk]
            tok = tokenizer(rendered).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 PROMPT_MIX.items())
    return F.normalize(fused, dim=-1).numpy().astype("float32")


# ------------------------------------------------------------------ ANN ----
def build_index(vectors: np.ndarray, m: int = 32, ef_construction: int = 200):
    """Cosine similarity over unit vectors is inner product, so use space='ip'."""
    import hnswlib
    index = hnswlib.Index(space="ip", dim=vectors.shape[1])
    index.init_index(max_elements=len(vectors), ef_construction=ef_construction, M=m)
    index.add_items(vectors, np.arange(len(vectors)))
    return index


def search(index, queries: np.ndarray, k: int = 10, ef: int = 64):
    index.set_ef(max(ef, k))
    ids, distances = index.knn_query(queries, k=k)
    return ids, 1.0 - distances          # ip distance -> cosine similarity


def search_exact(vectors: np.ndarray, queries: np.ndarray, k: int = 10):
    """Ground truth, for checking ANN recall on your own corpus."""
    sims = queries @ vectors.T
    ids = np.argsort(-sims, axis=1)[:, :k]
    return ids, np.take_along_axis(sims, ids, axis=1)


def _demo() -> None:
    enc = load()
    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")]]
    doc = encode_images(corpus, enc)
    qry = encode_queries(["black leather ankle boots", "navy wool coat"], enc)
    print(f"documents {doc.shape}  queries {qry.shape}  (one vector each)")
    index = build_index(doc)
    ann_ids, ann_scores = search(index, qry, k=3)
    ex_ids, _ = search_exact(doc, qry, k=3)
    agree = float(np.mean([len(set(a) & set(b)) / len(b) for a, b in zip(ann_ids, ex_ids)]))
    print(f"ANN top-3 ids {ann_ids.tolist()}")
    print(f"exact top-3 ids {ex_ids.tolist()}")
    print(f"ANN/exact overlap@3: {agree:.3f} (1.000 expected on a corpus this small)")


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