CoolFace
Apppublic

altera2015/open_clip_v2

sourceHugging Faceupdated 2y agoView on Hugging Face
1likes
main.py87 linesDownload Raw Back to root
1from fastapi import FastAPI, File, UploadFile2from fastapi.responses import HTMLResponse3import io4import json5import open_clip6from PIL import Image, ImageFile7import torch8from typing import Annotated, Union9import numpy10 11model, _, preprocess = open_clip.create_model_and_transforms('ViT-B-16-plus-240', pretrained="laion400m_e32")12tokenizer = open_clip.get_tokenizer('ViT-B-32')13device ="cuda" if torch.cuda.is_available() else "cpu"14model.to(device)15print(f"Using {device}")16 17def embedding_to_json(embedding_tensor:torch.Tensor):18    embedding = []19    for i in range(embedding_tensor.shape[1]):20        embedding.append(float(embedding_tensor[0][i]))21    return { "embedding" : embedding}22 23def text_encode(name:str):24    tokens = tokenizer([name]).to(device)25    text_features = model.encode_text(tokens).to(device)26    text_features = torch.nn.functional.normalize(text_features, p=2, dim=1)27    return embedding_to_json(text_features.to("cpu"))28 29def image_encoder(image_data: numpy.ndarray) -> torch.Tensor:30    img = Image.fromarray(image_data)#.convert('RGB')31    img = preprocess(img).unsqueeze(0).to(device) # type: ignore32    embedding = model.encode_image(img)33    return embedding34 35def image_encode_bytes(data:bytes):36    ImageFile.LOAD_TRUNCATED_IMAGES = True37    pil_img = Image.open(io.BytesIO(data)).convert("RGB")38    #pil_img = Image.open(data).convert("RGB")39    image_data = numpy.asarray(pil_img)40    pil_img.close()41 42    image_features = image_encoder(image_data)43    image_features = torch.nn.functional.normalize(image_features, p=2, dim=1)44    return embedding_to_json(image_features.to("cpu"))45 46app = FastAPI()47 48@app.post("/embedding/image")49async def encode_image(file: UploadFile):50    if not file:51        return {"message": "No file sent"}52    data = await file.read()53    return image_encode_bytes(data)54 55@app.get("/embedding/text")56def encode_text(q: Union[str, None] = None):57    if q is None:58        return {}59    return text_encode(q)60 61@app.get("/", response_class=HTMLResponse)62def home():63    return f"""64        <html>65            <head>66                <title>OpenCLIP ViT-B-16-plus-240 text embedding calculator (laion400m_e32)</title>67                <style>68                    body {{69                        font-family: sans-serif;70                    }}71                </style>72            </head>73            <body>74                <h1>OpenCLIP ViT-B-16-plus-240 (laion400m_e32)</h1>75                <h2>Text Embedding</h2>76                <form action="/embedding/text">77                Text: <input type="text" name="q" placeholder="Enter your text prompt to encode">78                </form>79                <h2>Image Embedding</h2>80                <form action="/embedding/image" enctype="multipart/form-data" method="post">81                Image: <input name="file" type="file">82                <input type="submit">83                </form>84            </body>85        </html>86    """87