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