taesiri/CLIPScore
27
1import os2import torch3import torch.nn.functional as F4import gradio as gr5import spaces # ← keep this!6from transformers import (7 CLIPProcessor,8 CLIPModel,9 SiglipProcessor, # transformers ≥ 4.4010 SiglipModel,11)12 13# ---------------------------------------------------------------------14# 1. CONFIG15# ---------------------------------------------------------------------16DEVICE = "cuda" if torch.cuda.is_available() else "cpu"17 18MODELS = {19 "CLIP ViT-B/32": ("openai/clip-vit-base-patch32", 224, "clip"),20 "CLIP ViT-B/16": ("openai/clip-vit-base-patch16", 224, "clip"),21 "CLIP ViT-L/14": ("openai/clip-vit-large-patch14", 224, "clip"),22 "CLIP ViT-L/14@336": ("openai/clip-vit-large-patch14-336", 336, "clip"),23 "SigLIP Large-256": ("google/siglip-large-patch16-256", 256, "siglip"),24 "SigLIP Base-384": ("google/siglip-base-patch16-384", 384, "siglip"),25 "SigLIP Large-384": ("google/siglip-large-patch16-384", 384, "siglip"),26}27 28# ---------------------------------------------------------------------29# 2. LAZY MODEL LOADING30# ---------------------------------------------------------------------31_models, _processors = {}, {}32 33def _load_model(name: str):34 path, _, kind = MODELS[name]35 36 kwargs = dict(37 low_cpu_mem_usage=False, # avoid meta-device bug38 torch_dtype=torch.float16, # faster & smaller39 )40 41 if kind == "clip":42 model = CLIPModel.from_pretrained(path, **kwargs).to(DEVICE)43 processor = CLIPProcessor.from_pretrained(path)44 else:45 model = SiglipModel.from_pretrained(path, **kwargs).to(DEVICE)46 processor = SiglipProcessor.from_pretrained(path)47 48 model.eval()49 return model, processor50 51def get_model(name: str):52 if name not in _models:53 _models[name], _processors[name] = _load_model(name)54 return _models[name], _processors[name]55 56# ---------------------------------------------------------------------57# 3. SCORING FUNCTION (runs on GPU in Spaces)58# ---------------------------------------------------------------------59@spaces.GPU60def calculate_score(image, text: str, model_name: str):61 labels = [t.strip() for t in text.split(";") if t.strip()]62 if not labels:63 return {}64 65 model, processor = get_model(model_name)66 kind = MODELS[model_name][2]67 68 inputs = processor(69 text=labels,70 images=image,71 padding=True,72 return_tensors="pt",73 ).to(DEVICE)74 75 with torch.no_grad():76 if kind == "clip":77 out = model(**inputs)78 img_emb = out.image_embeds79 txt_emb = out.text_embeds80 else:81 img_emb = model.get_image_features(pixel_values=inputs["pixel_values"])82 txt_emb = model.get_text_features(83 input_ids=inputs["input_ids"],84 attention_mask=inputs["attention_mask"],85 )86 87 img_emb = F.normalize(img_emb, p=2, dim=-1)88 txt_emb = F.normalize(txt_emb, p=2, dim=-1)89 90 scores = (txt_emb @ img_emb.T).squeeze(1) # cosine91 if kind == "siglip":92 scores = torch.sigmoid(scores) # paper’s choice93 94 return {lbl: float(score.clamp(0, 1)) for lbl, score in zip(labels, scores.cpu())}95 96# ---------------------------------------------------------------------97# 4. GRADIO UI98# ---------------------------------------------------------------------99with gr.Blocks(title="CLIP / SigLIP Image-Text Similarity") as demo:100 gr.Markdown("## Compare an image with multiple text prompts")101 102 with gr.Row():103 image_in = gr.Image(type="pil", label="Image")104 score_out = gr.Label(label="Similarity (0‒1)")105 106 with gr.Row():107 text_in = gr.Textbox(108 label="Text prompts (use ‘;’ to separate)",109 placeholder="a cat; a flying cat; a dog",110 )111 model_in = gr.Dropdown(112 choices=list(MODELS.keys()),113 value="CLIP ViT-B/16",114 label="Model",115 )116 117 def infer(img, txt, mdl):118 return calculate_score(img, txt, mdl) if img and txt.strip() else {}119 120 for comp in (image_in, text_in, model_in):121 comp.change(infer, [image_in, text_in, model_in], score_out)122 123 gr.Examples(124 examples=[125 ["cat.jpg",126 "a cat stuck in a door; a cat jumping; a dog",127 "CLIP ViT-B/16"],128 ],129 inputs=[image_in, text_in, model_in],130 outputs=score_out,131 )132 133demo.launch()