Download train.py from ByteJoseph/realtime-digit-draw: direct link, hf CLI and curl.
- Browser
- Download file 1.55 kB
-
https://huggingface.co/ByteJoseph/realtime-digit-draw/resolve/main/train.py
- Command line
-
hf download hf://ByteJoseph/realtime-digit-draw/train.py
-
curl -L -o train.py https://huggingface.co/ByteJoseph/realtime-digit-draw/resolve/main/train.py
1.55 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 | |
| from model import ANN | |
| 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}%") | |
| torch.save(model.state_dict(), "weights.pth") | |
| print("weights saved") | |