Download digits.py from ByteJoseph/realtime-digit-draw: direct link, hf CLI and curl.
- Browser
- Download file 4.03 kB
-
https://huggingface.co/ByteJoseph/realtime-digit-draw/resolve/main/digits.py
- Command line
-
hf download hf://ByteJoseph/realtime-digit-draw/digits.py
-
curl -L -o digits.py https://huggingface.co/ByteJoseph/realtime-digit-draw/resolve/main/digits.py
4.03 kB
| import torch | |
| import torch.nn as nn | |
| import torch.optim as optim | |
| import torchvision | |
| import torchvision.transforms as transforms | |
| from torch.utils.data import DataLoader | |
| import pygame | |
| import numpy as np | |
| import cv2 | |
| class ANN(nn.Module): | |
| def __init__(self): | |
| super(ANN,self).__init__() | |
| self.fc1 = nn.Linear(28*28,128) | |
| self.fc2 = nn.Linear(128,128) | |
| self.fc3 = nn.Linear(128,128) | |
| self.fc4 = nn.Linear(128,10) | |
| self.activation = nn.ReLU() | |
| def forward(self,x): | |
| x = x.view(-1,28*28) | |
| x = self.activation(self.fc1(x)) | |
| x = self.activation(self.fc2(x)) | |
| x = self.activation(self.fc3(x)) | |
| x = self.fc4(x) | |
| return x | |
| # load dataset | |
| transform = transforms.Compose([ | |
| transforms.ToTensor(), | |
| transforms.Normalize((0.5,),(0.5,))] | |
| ) | |
| trainset = torchvision.datasets.MNIST(root='./data',train=True,download=True,transform=transform) | |
| testset = torchvision.datasets.MNIST(root='./data',train=False,download=True,transform=transform) | |
| trainloader = DataLoader(trainset,batch_size=64,shuffle=True) | |
| testloader = DataLoader(testset,batch_size=64,shuffle=False) | |
| model = ANN() | |
| los_fn = nn.CrossEntropyLoss() | |
| optimizer = optim.Adam(model.parameters(),lr=0.001) | |
| # training loop | |
| episodes = 10 | |
| for epoch in range(episodes): | |
| running_loss = 0.0 | |
| for images, labels in trainloader: | |
| optimizer.zero_grad() | |
| outputs = model(images) | |
| loss = los_fn(outputs,labels) | |
| loss.backward() | |
| optimizer.step() | |
| running_loss += loss.item() | |
| print(f"Epoch {epoch+1}, Loss: {running_loss/len(trainloader)}") | |
| # evaluation | |
| correct = 0 | |
| total = 0 | |
| model.eval() | |
| with torch.no_grad(): | |
| for images, labels in testloader: | |
| outputs = model(images) | |
| _, predicted = torch.max(outputs.data,1) | |
| total += labels.size(0) | |
| correct += (predicted == labels).sum().item() | |
| print(f"Test Accuracy: {100 * correct / total}%") | |
| def draw_digi(): | |
| pygame.init() | |
| window_size = 280 | |
| display_height = window_size+50 | |
| screen = pygame.display.set_mode((window_size,display_height)) | |
| pygame.display.set_caption("Draw a digit") | |
| clock = pygame.time.Clock() | |
| screen.fill((0,0,0)) | |
| drawing = False | |
| prediction = None | |
| font = pygame.font.Font(None,36) | |
| while True: | |
| # screen.fill((0,0,0)) | |
| for event in pygame.event.get(): | |
| if event.type == pygame.QUIT: | |
| pygame.quit() | |
| return | |
| if event.type == pygame.MOUSEBUTTONDOWN: | |
| drawing = True | |
| if event.type == pygame.MOUSEBUTTONUP: | |
| drawing = False | |
| prediction = predict_digit(screen) | |
| if event.type == pygame.KEYDOWN: | |
| if event.key == pygame.K_c: | |
| screen.fill((0,0,0)) | |
| prediction = None | |
| if event.type == pygame.MOUSEMOTION and drawing: | |
| pygame.draw.circle(screen,(255,255,255),event.pos,8) | |
| if prediction is not None: | |
| text = font.render(f"Prediction: {prediction}",True,(0,255,0)) | |
| screen.blit(text,(10,window_size+10)) | |
| pygame.display.flip() | |
| clock.tick(60) | |
| def process_drawing(screen): | |
| surface = pygame.surfarray.array3d(screen) | |
| gray = np.dot(surface[...,:3],[0.2989,0.5870,0.1140]) | |
| gray = np.transpose(gray,(1,0)) | |
| gray = cv2.resize(gray,(28,28),interpolation=cv2.INTER_AREA) | |
| gray = gray.astype(np.float32)/255.0 | |
| gray = (gray - 0.5)/0.5 | |
| tensor = torch.tensor(gray,dtype=torch.float32).unsqueeze(0).unsqueeze(0) | |
| return tensor | |
| def predict_digit(screen): | |
| image = process_drawing(screen) | |
| if image is None: | |
| return None | |
| model.eval() | |
| with torch.no_grad(): | |
| output = model(image) | |
| _, predicted = torch.max(output.data,1) | |
| return predicted.item() | |
| draw_digi() |