microsoft/AI-For-Beginners

Public

mirrored from https://github.com/microsoft/AI-For-BeginnersAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
b5aaeeb689986fdcf3ca48f00343ef6e2e31fbea

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

lessons/4-ComputerVision/07-ConvNets/pytorchcv.py

169lines · modecode

1# Script file to hide implementation details for PyTorch computer vision module
2
3import builtins
4import torch
5import torch.nn as nn
6from torch.utils import data
7import torchvision
8from torchvision.transforms import ToTensor
9import matplotlib.pyplot as plt
10import numpy as np
11from PIL import Image
12import glob
13import os
14import zipfile
15
16default_device = 'cuda' if torch.cuda.is_available() else 'cpu'
17
18def load_mnist(batch_size=64):
19 builtins.data_train = torchvision.datasets.MNIST('./data',
20 download=True,train=True,transform=ToTensor())
21 builtins.data_test = torchvision.datasets.MNIST('./data',
22 download=True,train=False,transform=ToTensor())
23 builtins.train_loader = torch.utils.data.DataLoader(data_train,batch_size=batch_size)
24 builtins.test_loader = torch.utils.data.DataLoader(data_test,batch_size=batch_size)
25
26def train_epoch(net,dataloader,lr=0.01,optimizer=None,loss_fn = nn.NLLLoss()):
27 optimizer = optimizer or torch.optim.Adam(net.parameters(),lr=lr)
28 net.train()
29 total_loss,acc,count = 0,0,0
30 for features,labels in dataloader:
31 optimizer.zero_grad()
32 lbls = labels.to(default_device)
33 out = net(features.to(default_device))
34 loss = loss_fn(out,lbls) #cross_entropy(out,labels)
35 loss.backward()
36 optimizer.step()
37 total_loss+=loss
38 _,predicted = torch.max(out,1)
39 acc+=(predicted==lbls).sum()
40 count+=len(labels)
41 return total_loss.item()/count, acc.item()/count
42
43def validate(net, dataloader,loss_fn=nn.NLLLoss()):
44 net.eval()
45 count,acc,loss = 0,0,0
46 with torch.no_grad():
47 for features,labels in dataloader:
48 lbls = labels.to(default_device)
49 out = net(features.to(default_device))
50 loss += loss_fn(out,lbls)
51 pred = torch.max(out,1)[1]
52 acc += (pred==lbls).sum()
53 count += len(labels)
54 return loss.item()/count, acc.item()/count
55
56def train(net,train_loader,test_loader,optimizer=None,lr=0.01,epochs=10,loss_fn=nn.NLLLoss()):
57 optimizer = optimizer or torch.optim.Adam(net.parameters(),lr=lr)
58 res = { 'train_loss' : [], 'train_acc': [], 'val_loss': [], 'val_acc': []}
59 for ep in range(epochs):
60 tl,ta = train_epoch(net,train_loader,optimizer=optimizer,lr=lr,loss_fn=loss_fn)
61 vl,va = validate(net,test_loader,loss_fn=loss_fn)
62 print(f"Epoch {ep:2}, Train acc={ta:.3f}, Val acc={va:.3f}, Train loss={tl:.3f}, Val loss={vl:.3f}")
63 res['train_loss'].append(tl)
64 res['train_acc'].append(ta)
65 res['val_loss'].append(vl)
66 res['val_acc'].append(va)
67 return res
68
69def train_long(net,train_loader,test_loader,epochs=5,lr=0.01,optimizer=None,loss_fn = nn.NLLLoss(),print_freq=10):
70 optimizer = optimizer or torch.optim.Adam(net.parameters(),lr=lr)
71 for epoch in range(epochs):
72 net.train()
73 total_loss,acc,count = 0,0,0
74 for i, (features,labels) in enumerate(train_loader):
75 lbls = labels.to(default_device)
76 optimizer.zero_grad()
77 out = net(features.to(default_device))
78 loss = loss_fn(out,lbls)
79 loss.backward()
80 optimizer.step()
81 total_loss+=loss
82 _,predicted = torch.max(out,1)
83 acc+=(predicted==lbls).sum()
84 count+=len(labels)
85 if i%print_freq==0:
86 print("Epoch {}, minibatch {}: train acc = {}, train loss = {}".format(epoch,i,acc.item()/count,total_loss.item()/count))
87 vl,va = validate(net,test_loader,loss_fn)
88 print("Epoch {} done, validation acc = {}, validation loss = {}".format(epoch,va,vl))
89
90
91def plot_results(hist):
92 plt.figure(figsize=(15,5))
93 plt.subplot(121)
94 plt.plot(hist['train_acc'], label='Training acc')
95 plt.plot(hist['val_acc'], label='Validation acc')
96 plt.legend()
97 plt.subplot(122)
98 plt.plot(hist['train_loss'], label='Training loss')
99 plt.plot(hist['val_loss'], label='Validation loss')
100 plt.legend()
101
102def plot_convolution(t,title=''):
103 with torch.no_grad():
104 c = nn.Conv2d(kernel_size=(3,3),out_channels=1,in_channels=1)
105 c.weight.copy_(t)
106 fig, ax = plt.subplots(2,6,figsize=(8,3))
107 fig.suptitle(title,fontsize=16)
108 for i in range(5):
109 im = data_train[i][0]
110 ax[0][i].imshow(im[0])
111 ax[1][i].imshow(c(im.unsqueeze(0))[0][0])
112 ax[0][i].axis('off')
113 ax[1][i].axis('off')
114 ax[0,5].imshow(t)
115 ax[0,5].axis('off')
116 ax[1,5].axis('off')
117 #plt.tight_layout()
118 plt.show()
119
120def display_dataset(dataset, n=10,classes=None):
121 fig,ax = plt.subplots(1,n,figsize=(15,3))
122 mn = min([dataset[i][0].min() for i in range(n)])
123 mx = max([dataset[i][0].max() for i in range(n)])
124 for i in range(n):
125 ax[i].imshow(np.transpose((dataset[i][0]-mn)/(mx-mn),(1,2,0)))
126 ax[i].axis('off')
127 if classes:
128 ax[i].set_title(classes[dataset[i][1]])
129
130
131def check_image(fn):
132 try:
133 im = Image.open(fn)
134 im.verify()
135 return True
136 except:
137 return False
138
139def check_image_dir(path):
140 for fn in glob.glob(path):
141 if not check_image(fn):
142 print("Corrupt image: {}".format(fn))
143 os.remove(fn)
144
145
146def common_transform():
147 std_normalize = torchvision.transforms.Normalize(mean=[0.485, 0.456, 0.406],
148 std=[0.229, 0.224, 0.225])
149 trans = torchvision.transforms.Compose([
150 torchvision.transforms.Resize(256),
151 torchvision.transforms.CenterCrop(224),
152 torchvision.transforms.ToTensor(),
153 std_normalize])
154 return trans
155
156def load_cats_dogs_dataset():
157 if not os.path.exists('data/PetImages'):
158 with zipfile.ZipFile('data/kagglecatsanddogs_3367a.zip', 'r') as zip_ref:
159 zip_ref.extractall('data')
160
161 check_image_dir('data/PetImages/Cat/*.jpg')
162 check_image_dir('data/PetImages/Dog/*.jpg')
163
164 dataset = torchvision.datasets.ImageFolder('data/PetImages',transform=common_transform())
165 trainset, testset = torch.utils.data.random_split(dataset,[20000,len(dataset)-20000])
166 trainloader = torch.utils.data.DataLoader(trainset,batch_size=32)
167 testloader = torch.utils.data.DataLoader(trainset,batch_size=32)
168 return dataset, trainloader, testloader