CoolFace
Apppublic

taesiri/CLIPScore

sourceHugging Facemitupdated 10mo agoView on Hugging Face
27likes
app.py133 linesDownload Raw Back to root
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()