shaked08's picture
Update app.py
e369dae verified
Raw
History Blame Contribute Delete
4.46 kB
import pandas as pd
import torch
import torch.nn.functional as F
from transformers import CLIPProcessor, CLIPModel
from PIL import Image
import numpy as np
from sklearn.metrics.pairwise import cosine_similarity
import gradio as gr
import io
# ─── 1. INITIALIZATION ───────────────────────────────────────────
device = "cuda" if torch.cuda.is_available() else "cpu"
MODEL_ID = "openai/clip-vit-base-patch32"
processor = CLIPProcessor.from_pretrained(MODEL_ID)
model = CLIPModel.from_pretrained(MODEL_ID).to(device)
model.eval()
df = pd.read_parquet("food_embeddings.parquet")
embeddings = np.stack(df["embedding"].values)
def idx_to_image(idx):
return Image.open(io.BytesIO(df.iloc[idx]["image_bytes"]))
print(f"βœ… Loaded {len(df)} embeddings from parquet")
# ─── 2. EMBEDDING FUNCTION ───────────────────────────────────────
def get_embedding(user_input, input_type):
with torch.no_grad():
if input_type == "image":
if isinstance(user_input, str):
image = Image.open(user_input).convert("RGB")
else:
image = user_input.convert("RGB")
inputs = processor(images=image, return_tensors="pt").to(device)
outputs = model.get_image_features(**inputs)
else:
inputs = processor(text=[user_input], return_tensors="pt", padding=True).to(device)
outputs = model.get_text_features(**inputs)
if isinstance(outputs, torch.Tensor):
features = outputs
else:
features = outputs.image_embeds if hasattr(outputs, "image_embeds") else outputs.pooler_output
features = F.normalize(features, p=2, dim=-1)
return features.cpu().numpy().flatten()
# ─── 3. RECOMMENDATION FUNCTION ─────────────────────────────────
def find_similar_foods(image_input, text_input):
try:
if image_input is not None:
user_vec = get_embedding(image_input, "image")
mode = "Image"
elif text_input and text_input.strip():
user_vec = get_embedding(text_input.strip(), "text")
mode = "Text"
else:
return None, None, None, "⚠️ Please upload an image or enter a text query."
scores = cosine_similarity(user_vec.reshape(1, -1), embeddings).flatten()
top3_idx = np.argsort(scores)[::-1][:3]
results = [idx_to_image(int(idx)) for idx in top3_idx]
status = f"βœ… [{mode} search] Top-3 matches found!"
return results[0], results[1], results[2], status
except Exception as e:
return None, None, None, f"❌ Error: {str(e)}"
# ─── 4. GRADIO UI ────────────────────────────────────────────────
with gr.Blocks(title="Plating Suggestions") as demo:
gr.Markdown("# 🍽️ Plating Suggestions")
gr.Markdown("Upload a food photo **or** describe what you're craving β€” get 3 meal serving suggestions!")
with gr.Row():
with gr.Column(scale=1):
img_input = gr.Image(label="Upload Food Image", type="pil")
text_input = gr.Textbox(
label="Or describe a food",
placeholder="e.g. crispy fried chicken with sauce"
)
with gr.Row():
search_btn = gr.Button("πŸ” Find Serving Suggestions", variant="primary")
clear_btn = gr.Button("βœ• Clear")
status_box = gr.Markdown("")
with gr.Column(scale=2):
gr.Markdown("### πŸ† Top-3 Suggestions")
with gr.Row():
out1 = gr.Image(label="#1 Best Match")
out2 = gr.Image(label="#2 Match")
out3 = gr.Image(label="#3 Match")
search_btn.click(
fn=find_similar_foods,
inputs=[img_input, text_input],
outputs=[out1, out2, out3, status_box]
)
clear_btn.click(
fn=lambda: (None, "", None, None, None, ""),
outputs=[img_input, text_input, out1, out2, out3, status_box]
)
demo.launch(server_name="0.0.0.0", server_port=7860)