-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
113 lines (78 loc) · 3.19 KB
/
Copy pathmain.py
File metadata and controls
113 lines (78 loc) · 3.19 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
101
102
103
104
105
106
107
108
109
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Created on Mon Feb 11 17:41:22 2019
@author: dogaykamar
"""
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import utils
from dataset import FaceDataset
from network import SuperRes, SmallTestSampler
import numpy as np
import time
BATCHSIZE = 32
MAX_EPOCHS = 250
start_epoch = 0
learning_rate = 1e-3
weightDecay = 0
history = {"train_loss": [], "val_loss": []}
device = torch.device("cuda:0")
image_file = 'datasets/celebahq/'
print("Start Data Load")
train_data = FaceDataset(image_file, upscale_factor = 3, mode = 'train')
val_data = FaceDataset(image_file, upscale_factor = 3, mode = 'val')
sampler = SmallTestSampler()
train_loader = DataLoader(train_data, batch_size = BATCHSIZE, shuffle=True, pin_memory = True)
val_loader = DataLoader(val_data, batch_size = BATCHSIZE, shuffle=True, pin_memory = True)
model = SuperRes(upscale_factor = 3)
model.to(device)
criteria = nn.MSELoss()
optimizer = optim.Adam(model.parameters(), lr = learning_rate)
start = time.time()
print("Begin!")
for epoch in range(start_epoch, MAX_EPOCHS):
avg_loss = 0
print("Epoch: ",epoch+1,"/",MAX_EPOCHS)
print("Training start")
model.train()
num_batch = len(train_loader)
for i, batch in enumerate(train_loader):
input_img = batch[0].cuda(device)
target_img = batch[1].cuda(device)
prediction = model(input_img)
loss = criteria.forward(prediction, target_img)
optimizer.zero_grad()
loss.backward()
optimizer.step()
print("MiniBatch: ", i+1, "/", num_batch, " Loss: ", loss.item(), end="\r")
avg_loss = (avg_loss*i)+loss.item()
avg_loss = avg_loss/(i+1)
del input_img, target_img, prediction, loss
time_since_start = time.time()-start
print("\nTrain Loss: ", avg_loss, " found in {:.0f}m {:.0f}s".format(time_since_start// 60, time_since_start % 60))
history["train_loss"].append(avg_loss)
avg_loss = 0
print("Validation start")
model.eval()
num_batch = len(val_loader)
for i, batch in enumerate(val_loader):
input_img = batch[0].cuda(device)
target_img = batch[1].cuda(device)
prediction = model(input_img)
loss = criteria.forward(prediction, target_img)
print("MiniBatch: ", i+1, "/", num_batch, " Loss: ", loss.item(), end="\r")
avg_loss = (avg_loss*i)+loss.item()
avg_loss = avg_loss/(i+1)
if (i == 0):
save_image = torch.cat([target_img, prediction], dim=2).detach().cpu()
utils.save_image(save_image, 'model_out.png')
del input_img, target_img, prediction, loss
time_since_start = time.time()-start
print("\nValidation Loss: ",avg_loss, " found in {:.0f}m {:.0f}s".format(time_since_start// 60, time_since_start % 60))
history["val_loss"].append(avg_loss)
state = {"model": model.state_dict(), "optimizer": optimizer.state_dict()}
torch.save(state, "models/model_{}.pth".format(epoch+1))
np.save("models/history.npy", history)