-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
103 lines (85 loc) · 3.51 KB
/
Copy pathtrain.py
File metadata and controls
103 lines (85 loc) · 3.51 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
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
import numpy as np
from PIL import Image
import argparse
import matplotlib.pyplot as plt
import fmodel
import torch
from torch import nn, optim
import torchvision
import futility
import torchvision.transforms as transforms
import torchvision.models as models
from torch.autograd import Variable
from collections import OrderedDict
import json
import torchvision.datasets as datasets
parser = argparse.ArgumentParser(description='Parser for train.py')
parser.add_argument('data_dir', action="store", default="./flowers/")
parser.add_argument('--save_dir', action="store", default="./checkpoint.pth")
parser.add_argument('--arch', action="store", default="vgg16")
parser.add_argument('--learning_rate', action="store", type=float, default=0.001)
parser.add_argument('--hidden_units', action="store", type=int, default=512)
parser.add_argument('--epochs', action="store", type=int, default=3)
parser.add_argument('--dropout', action="store", type=float, default=0.2)
parser.add_argument('--gpu', action="store_true", default=False)
args = parser.parse_args()
where = args.data_dir
path = args.save_dir
lr = args.learning_rate
struct = args.arch
hidden_units = args.hidden_units
power = args.gpu
epochs = args.epochs
dropout = args.dropout
device = torch.device("cuda:0" if torch.cuda.is_available() and power else "cpu")
def main():
trainloader, validloader, testloader, train_data = futility.load_data(where)
model, criterion, optimizer = fmodel.setup_network(struct, dropout, hidden_units, lr, power)
# Train Model
steps = 0
running_loss = 0
print_every = 5
print("--Training starting--")
for epoch in range(epochs):
for inputs, labels in trainloader:
steps += 1
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
log_ps = model.forward(inputs)
loss = criterion(log_ps, labels)
loss.backward()
optimizer.step()
running_loss += loss.item()
if steps % print_every == 0:
valid_loss = 0
accuracy = 0
model.eval()
with torch.no_grad():
for inputs, labels in validloader:
inputs, labels = inputs.to(device), labels.to(device)
log_ps = model.forward(inputs)
batch_loss = criterion(log_ps, labels)
valid_loss += batch_loss.item()
# Calculate accuracy
ps = torch.exp(log_ps)
top_p, top_class = ps.topk(1, dim=1)
equals = top_class == labels.view(*top_class.shape)
accuracy += torch.mean(equals.type(torch.FloatTensor)).item()
print(f"Epoch {epoch+1}/{epochs}.. "
f"Loss: {running_loss/print_every:.3f}.. "
f"Validation Loss: {valid_loss/len(validloader):.3f}.. "
f"Accuracy: {accuracy/len(validloader):.3f}")
running_loss = 0
model.train()
model.class_to_idx = train_data.class_to_idx
torch.save({'structure': struct,
'hidden_units': hidden_units,
'dropout': dropout,
'learning_rate': lr,
'no_of_epochs': epochs,
'state_dict': model.state_dict(),
'class_to_idx': model.class_to_idx},
path)
print("Saved checkpoint!")
if __name__ == "__main__":
main