CooLLegend's picture
Added fashionclip model
02f8392
Raw
History Blame Contribute Delete
2.62 kB
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from transformers import CLIPProcessor, CLIPModel
from PIL import Image
import requests
from io import BytesIO
import torch
# 1. Initialize the App
app = FastAPI()
# 2. Load the Fashion-CLIP Model
# We use 'CLIPModel' directly to get access to the vector layers.
# This runs ONCE when the server starts.
print("Loading Fashion-CLIP Model...")
model_id = "patrickjohncyh/fashion-clip"
model = CLIPModel.from_pretrained(model_id)
processor = CLIPProcessor.from_pretrained(model_id)
print("Model Loaded!")
# 3. Define Input Data Structures
class TextRequest(BaseModel):
text: str
class ImageRequest(BaseModel):
image_url: str
# 4. The Home Route (Health Check)
@app.get("/")
def home():
return {"status": "Online", "model": "Fashion-CLIP"}
# 5. Endpoint: Convert Text to Vector
@app.post("/embed-text")
def embed_text(req: TextRequest):
try:
# Process text
inputs = processor(text=[req.text], return_tensors="pt", padding=True)
# Calculate features
with torch.no_grad(): # Disable gradient calculation for CPU speed
text_features = model.get_text_features(**inputs)
# Normalize the vector (Crucial for Cosine Similarity!)
text_features = text_features / text_features.norm(p=2, dim=-1, keepdim=True)
# Convert to standard Python list
vector = text_features[0].tolist()
return {"vector": vector}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
# 6. Endpoint: Convert Image URL to Vector
@app.post("/embed-image")
def embed_image(req: ImageRequest):
try:
# Download image
response = requests.get(req.image_url)
if response.status_code != 200:
raise HTTPException(status_code=400, detail="Could not download image")
image = Image.open(BytesIO(response.content))
# Process image
inputs = processor(images=image, return_tensors="pt", padding=True)
# Calculate features
with torch.no_grad(): # Disable gradient calculation for CPU speed
image_features = model.get_image_features(**inputs)
# Normalize the vector (Crucial for Cosine Similarity!)
image_features = image_features / image_features.norm(p=2, dim=-1, keepdim=True)
# Convert to standard Python list
vector = image_features[0].tolist()
return {"vector": vector}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))