"""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()