MagicLens-B (CLIP-B) β€” PyTorch

PyTorch weights for MagicLens (Zhang et al., ICML 2024), converted from Google's official JAX/Flax release.

The conversion is numerically faithful: on identical inputs this model reproduces the original's embeddings to cosine similarity 1.00000000 (max absolute difference ~4e-07, over three seeds).

What it does

Composed image retrieval: given a reference image and a text instruction ("same dress but in black"), retrieve the image satisfying both. 166,433,025 parameters, 512-d embeddings.

Usage

git clone https://github.com/mat0k/magiclens-pytorch && cd magiclens-pytorch
pip install -e .
import torch
from huggingface_hub import hf_hub_download
from PIL import Image
from magiclens import MagicLens, MagicLensImagePreprocess

weights = hf_hub_download("mat0k/magiclens-base-pytorch", "magic_lens_clip_base.pt")

model = MagicLens("base")
model.load_state_dict(torch.load(weights, map_location="cpu"))
model.eval()

preprocess = MagicLensImagePreprocess(224, is_train=False)
images = torch.stack([preprocess(Image.open(p)) for p in ["a.jpg", "b.jpg"]])

with torch.no_grad():
    query = model.encode_mm(images[:1], ["make it blue"])  # reference + instruction
    candidates = model.encode_image(images)                # gallery
    scores = query @ candidates.T                          # both L2-normalised

encode_image is not a plain CLIP image embedding β€” candidates pass through the same fusion stack as queries with an empty instruction, matching the official evaluation protocol. Use MagicLensImagePreprocess; the official pipeline scales pixels by each image's own maximum and squashes aspect ratio rather than cropping, so standard CLIP preprocessing gives different results.

Verification

Parity with the original JAX model β€” identical inputs through both implementations:

min cosine 1.00000000   max|diff| 4.263e-07

Preprocessing parity on real photos at native resolution: 7.8e-06.

Benchmarks, zero-shot, against the paper's MagicLens-B / CLIP-B row:

Benchmark Metric Paper This port
Fashion-IQ (val) overall R@10 26.3 25.90
Fashion-IQ (val) overall R@50 47.4 48.41
CIRR (test) R@1 27.0 29.52
CIRR (test) R@5 58.0 59.61
CIRR (test) R@10 70.9 72.63
CIRR (test) R@50 91.1 91.74
CIRR (test) R_subset@1 66.7 67.35

CIRR was scored by the official evaluation server, not by us. Mean absolute difference 0.77 (Fashion-IQ) and 1.11 (CIRR). Since the model is provably identical, these residuals come from evaluation-pipeline differences, not the weights β€” the released code contains no CIRR evaluation at all, so the paper's CIRR numbers came from an internal pipeline that cannot be matched exactly.

Fine-tuning

Unlike the official release, this port supports training β€” the model returns an InfoNCE loss, so an ordinary PyTorch loop works. See examples/finetune_demo.py in the GitHub repo. Note the batch size is the negative pool: MagicLens was pretrained at batch 2048, and fine-tuning at batch 32 measurably degraded the model in our experiments.

License and citation

Code Apache 2.0. These weights are derived from Google DeepMind's release and remain under its CC-BY 4.0 terms for model materials.

@inproceedings{zhang2024magiclens,
  title     = {{M}agic{L}ens: Self-Supervised Image Retrieval with Open-Ended Instructions},
  author    = {Zhang, Kai and Luan, Yi and Hu, Hexiang and Lee, Kenton and Qiao, Siyuan
               and Chen, Wenhu and Su, Yu and Chang, Ming-Wei},
  booktitle = {Proceedings of the 41st International Conference on Machine Learning},
  pages     = {59403--59420},
  year      = {2024},
  volume    = {235},
  series    = {Proceedings of Machine Learning Research},
  publisher = {PMLR},
  url       = {https://proceedings.mlr.press/v235/zhang24an.html}
}
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support