moda-duo / serving_ann.py
ArkidMitra's picture
MODA Duo: routes each query to the constituent that suits its shape
326e3fb verified
Raw History Blame Contribute Delete
9.21 kB
"""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()