import torch import torch.nn as nn import gradio as gr import numpy as np class CNN(nn.Module): def __init__(self): super().__init__() self.conv = nn.Sequential( nn.Conv2d(1,32,3,padding=1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32,64,3,padding=1), nn.ReLU(), nn.MaxPool2d(2) ) self.fc = nn.Sequential( nn.Linear(64*7*7,128), nn.ReLU(), nn.Linear(128,10) ) def forward(self,x): x = self.conv(x) x = x.view(x.size(0),-1) x = self.fc(x) return x model = CNN() model.load_state_dict(torch.load("model.pth",map_location="cpu")) model.eval() def predict(image): image = np.array(image) image = image/255.0 image = image.reshape(1,1,28,28) image = torch.tensor(image,dtype=torch.float32) output = model(image) pred = torch.argmax(output,1).item() return pred interface = gr.Interface( fn=predict, inputs=gr.Image(shape=(28,28),image_mode="L"), outputs="label", title="Handwritten Digit Classifier" ) interface.launch()