moda-pro-lite-plus / serving_ann.py
ArkidMitra's picture
MODA Pro Lite+: moda-pro-lite with its calibrated serving recipe
3eaff2e verified
Raw History Blame Contribute Delete
7.03 kB
"""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()