dcarpintero/mlp-digit-classifier
0
1import gradio as gr2import torch3from model import *4 5from PIL import Image6import torchvision.transforms as transforms7 8title = "Digit Classifier"9description = (10 "Multilayer-Perceptron built for the fast.ai 'Deep Learning' course "11 "to classify handwritten digits from the MNIST dataset. "12)13inputs = gr.components.Image()14outputs = gr.components.Label()15examples = "examples"16 17model = torch.load("model/digit_classifier.pt", map_location=torch.device("cpu"))18labels = [str(i) for i in range(10)]19 20transform = transforms.Compose(21 [22 transforms.Resize((28, 28)),23 transforms.Grayscale(),24 transforms.ToTensor(),25 transforms.Lambda(lambda x: x[0]),26 transforms.Lambda(lambda x: x.unsqueeze(0)),27 ]28)29 30 31def predict_digit(img):32 img = transform(Image.fromarray(img))33 output = model(img)34 probs = torch.nn.functional.softmax(output, dim=1)35 return dict(zip(labels, map(float, probs.flatten()[:10])))36 37 38with gr.Blocks() as demo:39 with gr.Tab("Digit Prediction"):40 gr.Interface(41 fn=predict_digit,42 inputs=inputs,43 outputs=outputs,44 examples=examples,45 title=title,46 description=description,47 ).queue(default_concurrency_limit=5)48 49demo.launch()50 