Download app.py from Sparsh141/Hand_Written_Digit_Classifier: direct link, hf CLI and curl.
- Browser
- Download file 1.16 kB
-
https://huggingface.co/Sparsh141/Hand_Written_Digit_Classifier/resolve/main/app.py
- Command line
-
hf download hf://Sparsh141/Hand_Written_Digit_Classifier/app.py
-
curl -L -o app.py https://huggingface.co/Sparsh141/Hand_Written_Digit_Classifier/resolve/main/app.py
1.16 kB
| 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() |