-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmnist_model.py
More file actions
77 lines (65 loc) · 2.31 KB
/
Copy pathmnist_model.py
File metadata and controls
77 lines (65 loc) · 2.31 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
import torch
from torch import nn
from torch.optim import Adam
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
class SimpleNet(torch.nn.Module):
def __init__(self):
super(SimpleNet, self).__init__()
self.conv1 = nn.Conv2d(1, 32, kernel_size=3)
self.conv2 = nn.Conv2d(32, 32, kernel_size=3)
self.conv3 = nn.Conv2d(32, 64, kernel_size=3)
self.conv4 = nn.Conv2d(64, 64, kernel_size=3)
self.fc1 = nn.Linear(1024, 200)
self.fc2 = nn.Linear(200, 200)
self.fc3 = nn.Linear(200, 10)
def forward(self, x):
x = torch.relu(self.conv1(x))
x = torch.relu(self.conv2(x))
x = torch.max_pool2d(x, 2)
x = torch.relu(self.conv3(x))
x = torch.relu(self.conv4(x))
x = torch.max_pool2d(x, 2)
x = x.view(-1, 1024)
x = torch.relu(self.fc1(x))
x = torch.relu(self.fc2(x))
return self.fc3(x)
def train_on_mnist(self, name=None):
transform = self.get_transform()
mnist_dataset = datasets.MNIST(download=True, root='.', train=True, transform=transform)
train_loader = DataLoader(mnist_dataset, batch_size=256, shuffle=True)
optimizer = Adam(lr=0.01, params=self.parameters())
loss_fn = torch.nn.CrossEntropyLoss()
epochs = 2
self.train()
for e in range(epochs):
print(f'EPOCH {e}')
for i, (samples, labels) in enumerate(train_loader):
optimizer.zero_grad()
preds = self(samples)
loss = loss_fn(preds, labels)
loss.backward()
optimizer.step()
if i % 10 == 1:
print(f'Batch {i}, train loss: {loss}')
torch.save(self.state_dict(), name if name is not None else './mnist_net.pth')
return self
def load_pretrained_mnist(self, name=None):
state_dict = torch.load(name if name is not None else './mnist_net.pth')
self.load_state_dict(state_dict)
self.eval()
return self
def get_transform(self):
return transforms.ToTensor()
def test_mnist(self):
self.eval()
accuracy = 0
transform = self.get_transform()
mnist_test_dataset = datasets.MNIST(download=True, root='.', train=False, transform=transform)
test_loader = DataLoader(mnist_test_dataset, batch_size=256)
for tdata in test_loader:
samples, labels = tdata
pred_labels = self(samples).argmax(dim=1, keepdim=True)
accuracy += pred_labels.eq(labels.view_as(pred_labels)).sum().item()
accuracy = accuracy / len(test_loader.dataset)
print(f'Test accuracy: {accuracy}')