CoolFace
Apppublic

ehsanash/foodvision

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
app.py78 linesDownload Raw Back to root
1### 1. Imports and class names setup ### 2import gradio as gr3import os4import torch5 6from model import create_effnetb2_model7from timeit import default_timer as timer8from typing import Tuple, Dict9 10# Setup class names11class_names = ["pizza", "steak", "sushi"]12 13### 2. Model and transforms preparation ###14 15# Create EffNetB2 model16effnetb2, effnetb2_transforms = create_effnetb2_model(17    num_classes=3, # len(class_names) would also work18)19 20# Load saved weights21effnetb2.load_state_dict(22    torch.load(23        f="effnet_b2_20percent_96%.pth",24        map_location=torch.device("cpu"),  # load to CPU25    )26)27 28### 3. Predict function ###29 30# Create predict function31def predict(img) -> Tuple[Dict, float]:32    """Transforms and performs a prediction on img and returns prediction and time taken.33    """34    # Start the timer35    start_time = timer()36    37    # Transform the target image and add a batch dimension38    img = effnetb2_transforms(img).unsqueeze(0)39    40    # Put model into evaluation mode and turn on inference mode41    effnetb2.eval()42    with torch.inference_mode():43        # Pass the transformed image through the model and turn the prediction logits into prediction probabilities44        pred_probs = torch.softmax(effnetb2(img), dim=1)45    46    # Create a prediction label and prediction probability dictionary for each prediction class (this is the required format for Gradio's output parameter)47    pred_labels_and_probs = {class_names[i]: float(pred_probs[0][i]) for i in range(len(class_names))}48    49    # Calculate the prediction time50    pred_time = round(timer() - start_time, 5)51    52    # Return the prediction dictionary and prediction time 53    return pred_labels_and_probs, pred_time54 55### 4. Gradio app ###56 57# Create title, description and article strings58title = "FoodVision Mini ๐Ÿ•๐Ÿฅฉ๐Ÿฃ"59description = "An EfficientNetB2 feature extractor computer vision model to classify images of food as pizza, steak or sushi."60article = "Created by: EHSAN ASHRAFZADEH"61 62# Create examples list from "examples/" directory63example_list = [["examples/" + example] for example in os.listdir("examples")]64 65# Create the Gradio demo66demo = gr.Interface(fn=predict, # mapping function from input to output67                    inputs=gr.Image(type="pil"), # what are the inputs?68                    outputs=[gr.Label(num_top_classes=3, label="Predictions"), # what are the outputs?69                             gr.Number(label="Prediction time (s)")], # our fn has two outputs, therefore we have two outputs70                    # Create examples list from "examples/" directory71                    examples=example_list, 72                    title=title,73                    description=description,74                    article=article)75 76# Launch the demo!77demo.launch()78