fashion-api / app.py
BigBawdi's picture
Update app.py
10bcd13 verified
Raw
History Blame Contribute Delete
10.8 kB
"""
FastAPI service for:
1. FashionCLIP zero-shot clothing detection.
2. YOLO person detection / normalized person bounding boxes.
The original uploaded file was a Markdown table containing Python code.
This version restores the Python syntax and fixes the missing/overwritten
application setup so both endpoints live on the same FastAPI app.
"""
from io import BytesIO
from typing import Optional
import requests
import torch
from fastapi import FastAPI, HTTPException, Request
from pydantic import BaseModel, Field
from PIL import Image
from transformers import CLIPModel, CLIPProcessor
from ultralytics import YOLO
# ---------------------------------------------------------------------------
# Configuration / models
# ---------------------------------------------------------------------------
FASHION_CLIP_MODEL = "patrickjohncyh/fashion-clip"
# CLIP cosine similarities are close together, so the same temperature is
# used to sharpen the within-group softmax competition.
SIGMOID_TEMPERATURE = 20.0
app = FastAPI(title="Fashion Detection API")
processor = CLIPProcessor.from_pretrained(FASHION_CLIP_MODEL)
model = CLIPModel.from_pretrained(FASHION_CLIP_MODEL)
model.eval()
# Loaded once at Space startup, not once per request.
yolo_person_model = YOLO("yolov8n.pt")
# ---------------------------------------------------------------------------
# Clothing labels
# ---------------------------------------------------------------------------
TOP_LABELS = [
"t-shirt",
"dress shirt",
"button-down shirt",
"polo shirt",
"blouse",
"crop top",
"tank top",
"sweater",
"hoodie",
"cardigan",
"turtleneck",
"jacket",
"blazer",
"suit jacket",
"coat",
"leather jacket",
"denim jacket",
"windbreaker",
]
BOTTOM_LABELS = [
"jeans",
"dress pants",
"chinos",
"shorts",
"skirt",
"mini skirt",
"maxi skirt",
"leggings",
"sweatpants",
"cargo pants",
]
FULL_OUTFIT_LABELS = [
"suit",
"tuxedo",
"dress",
"mini dress",
"maxi dress",
"cocktail dress",
"jumpsuit",
"romper",
"overalls",
]
FOOTWEAR_LABELS = [
"sneakers",
"dress shoes",
"loafers",
"boots",
"heels",
"sandals",
"flip flops",
"oxfords",
"chelsea boots",
]
NECKWEAR_LABELS = ["tie", "bow tie"]
CATEGORY_GROUPS = {
"top": TOP_LABELS,
"bottom": BOTTOM_LABELS,
"full_outfit": FULL_OUTFIT_LABELS,
"footwear": FOOTWEAR_LABELS,
"neckwear": NECKWEAR_LABELS,
}
# Unlike category groups, accessories are independently optional.
OPTIONAL_ACCESSORY_LABELS = [
"belt",
"hat",
"cap",
"sunglasses",
"scarf",
"handbag",
"backpack",
]
# A neutral competitor prevents softmax from always selecting a garment
# when that category is not actually visible.
NOT_VISIBLE_LABEL = "no clothing of this type visible in the photo"
# ---------------------------------------------------------------------------
# Input schema
# ---------------------------------------------------------------------------
class DetectRequest(BaseModel):
image_url: str = Field(..., min_length=1)
top_k: int = Field(default=8, ge=0, le=8)
threshold: float = Field(default=0.55, ge=0.0, le=1.0)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def download_image(image_url: str) -> Image.Image:
"""Download and decode an image URL into an RGB PIL image."""
try:
response = requests.get(
image_url,
timeout=20,
headers={"User-Agent": "FashionDetectionAPI/1.0"},
)
response.raise_for_status()
return Image.open(BytesIO(response.content)).convert("RGB")
except requests.RequestException as exc:
raise HTTPException(
status_code=400,
detail=f"Could not download image: {exc}",
) from exc
except Exception as exc:
raise HTTPException(
status_code=400,
detail=f"Could not decode image: {exc}",
) from exc
def clip_embeddings(image: Image.Image, labels: list[str]):
"""Return normalized image/text embeddings for one label group."""
inputs = processor(
text=labels,
images=image,
return_tensors="pt",
padding=True,
)
# The model and processor are CPU-safe by default. If the model is moved
# to another device later, the input tensors are moved with it here.
device = next(model.parameters()).device
inputs = {
key: value.to(device) if hasattr(value, "to") else value
for key, value in inputs.items()
}
with torch.no_grad():
outputs = model(**inputs)
image_embeds = outputs.image_embeds
text_embeds = outputs.text_embeds
image_embeds = image_embeds / image_embeds.norm(dim=-1, keepdim=True)
text_embeds = text_embeds / text_embeds.norm(dim=-1, keepdim=True)
cosine_sim = (image_embeds @ text_embeds.t())[0]
return cosine_sim
# ---------------------------------------------------------------------------
# Health check
# ---------------------------------------------------------------------------
@app.get("/")
async def health_check():
"""Health-check / keep-warm endpoint."""
return {
"status": "ok",
"service": "fashion-detection-api",
}
# ---------------------------------------------------------------------------
# FashionCLIP detection
# ---------------------------------------------------------------------------
@app.post("/detect")
def detect_clothing(req: DetectRequest):
"""
Take an image URL and return:
- one winner per mutually-exclusive clothing category
- independently detected optional accessories
- a readable summary
"""
image = download_image(req.image_url)
categories = {}
# One winner per mutually exclusive category.
for category, labels in CATEGORY_GROUPS.items():
labels_with_none = labels + [NOT_VISIBLE_LABEL]
cosine_sim = clip_embeddings(image, labels_with_none)
# Softmax is intentionally applied within each category rather than
# across all clothing labels.
group_probs = torch.softmax(
cosine_sim * SIGMOID_TEMPERATURE,
dim=0,
)
top_idx = int(torch.argmax(group_probs))
won_by_none = labels_with_none[top_idx] == NOT_VISIBLE_LABEL
categories[category] = (
None
if won_by_none
else {
"label": labels_with_none[top_idx],
"confidence": round(float(group_probs[top_idx]), 3),
}
)
# Optional accessories use an independent present-vs-absent competition.
accessories = []
for label in OPTIONAL_ACCESSORY_LABELS:
pair = [
label,
f"no {label} visible, not wearing one",
]
cosine_sim = clip_embeddings(image, pair)
pair_probs = torch.softmax(
cosine_sim * SIGMOID_TEMPERATURE,
dim=0,
)
present_confidence = round(float(pair_probs[0]), 3)
if present_confidence >= req.threshold:
accessories.append(
{
"label": label,
"confidence": present_confidence,
}
)
accessories = sorted(
accessories,
key=lambda item: item["confidence"],
reverse=True,
)[: req.top_k]
has_any_detection = any(categories.values()) or bool(accessories)
if not has_any_detection:
return {
"categories": categories,
"accessories": [],
"summary": "NO_CLOTHING_DETECTED",
}
# Build a readable summary for the interpreter/client.
summary_parts = []
label_names = {
"full_outfit": "Full outfit",
"top": "Top",
"bottom": "Bottom",
"footwear": "Footwear",
"neckwear": "Neckwear",
}
for category, name in label_names.items():
if categories.get(category):
summary_parts.append(
f"{name}: {categories[category]['label']}"
)
if accessories:
acc_text = ", ".join(
f"{item['label']} "
f"({round(item['confidence'] * 100)}% confidence)"
for item in accessories
)
summary_parts.append(f"Accessories: {acc_text}")
summary = ". ".join(summary_parts) + "."
return {
"categories": categories,
"accessories": accessories,
"summary": summary,
}
# ---------------------------------------------------------------------------
# YOLO person detection
# ---------------------------------------------------------------------------
@app.post("/crop-person")
async def crop_person(request: Request):
"""
Accept raw image bytes in the request body.
Returns:
{
"person_detected": true,
"bounding_box": {
"left": 0-1,
"top": 0-1,
"width": 0-1,
"height": 0-1
},
"confidence": 0-1
}
or:
{"person_detected": false}
This endpoint does NOT crop the image itself. The client can use the
normalized bounding box to perform the crop.
"""
image_bytes = await request.body()
try:
image = Image.open(BytesIO(image_bytes))
if image.mode != "RGB":
image = image.convert("RGB")
except Exception as exc:
return {
"person_detected": False,
"error": f"could not decode image: {exc}",
}
width, height = image.size
results = yolo_person_model(image, verbose=False)
boxes = results[0].boxes
# Filter to YOLO's person class (0) and require at least 0.5 confidence.
person_boxes = [
(box.xyxy[0].tolist(), float(box.conf[0]))
for box in boxes
if int(box.cls[0]) == 0 and float(box.conf[0]) >= 0.5
]
if not person_boxes:
return {"person_detected": False}
# If multiple people are present, use the highest-confidence person.
(x1, y1, x2, y2), confidence = max(
person_boxes,
key=lambda item: item[1],
)
# Normalized 0-1 fractions.
bounding_box = {
"left": x1 / width,
"top": y1 / height,
"width": (x2 - x1) / width,
"height": (y2 - y1) / height,
}
return {
"person_detected": True,
"bounding_box": bounding_box,
"confidence": confidence,
}
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)