CoolFace
Apppublic

dcarpintero/mlp-digit-classifier

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
app.py50 linesDownload Raw Back to root
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