ehsanash/foodvision
0
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 