aayanb09's picture
Update app.py
5eab86a verified
Raw
History Blame Contribute Delete
3.49 kB
import gradio as gr
import torch
import torch.nn as nn
from torchvision import transforms, models
from PIL import Image
import numpy as np
import pickle
import os
# Model definition
class FoodIngredientClassifier(nn.Module):
def __init__(self, num_classes):
super().__init__()
self.backbone = models.vit_b_16(weights=None)
num_features = self.backbone.heads.head.in_features
self.backbone.heads = nn.Sequential(
nn.Dropout(0.5),
nn.Linear(num_features, 1024),
nn.BatchNorm1d(1024),
nn.ReLU(),
nn.Dropout(0.4),
nn.Linear(1024, 512),
nn.BatchNorm1d(512),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(512, num_classes)
)
def forward(self, x):
return self.backbone(x)
# Load model
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
THRESHOLD = 0.5
def load_model():
checkpoint = torch.load("model.pth", map_location=DEVICE, weights_only=False)
mlb = checkpoint["mlb"]
num_classes = len(mlb.classes_)
model = FoodIngredientClassifier(num_classes)
model.load_state_dict(checkpoint["model_state_dict"])
model.to(DEVICE)
model.eval()
return model, mlb
model, mlb = load_model()
# Transform
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
# Inference Function
def predict(image, threshold):
if image is None:
return "Please upload an image."
img = image.convert("RGB")
img_tensor = transform(img).unsqueeze(0).to(DEVICE)
with torch.no_grad():
output = model(img_tensor)
probs = torch.sigmoid(output).cpu().numpy()[0]
pred_indices = np.where(probs > threshold)[0]
if len(pred_indices) == 0:
return "No ingredients detected above the threshold. Try lowering it."
results = sorted(
zip(mlb.classes_[pred_indices], probs[pred_indices]),
key=lambda x: x[1],
reverse=True
)
output_lines = ["### 🍽️ Detected Ingredients\n"]
for ingredient, confidence in results:
bar = "█" * int(confidence * 20)
output_lines.append(f"**{ingredient}** — {confidence:.1%} `{bar}`")
return "\n\n".join(output_lines)
# UI
with gr.Blocks(title="Food Ingredient Detector") as demo:
gr.Markdown("""
# Food Ingredient Detector
Upload a photo of food and the model will identify its ingredients.
Built with a **ViT-B/16** backbone fine-tuned for multi-label ingredient classification.
""")
with gr.Row():
with gr.Column():
image_input = gr.Image(type="pil", label="Upload Food Image")
threshold_slider = gr.Slider(
minimum=0.1, maximum=0.9, value=0.5, step=0.05,
label="Detection Threshold",
info="Lower = more ingredients detected, higher = only confident predictions"
)
predict_btn = gr.Button("Detect Ingredients", variant="primary")
with gr.Column():
output = gr.Markdown(label="Results")
predict_btn.click(
fn=predict,
inputs=[image_input, threshold_slider],
outputs=output
)
gr.Examples(
examples=[], # Add example image paths here if you include sample images
inputs=image_input
)
if __name__ == "__main__":
demo.launch()