forked from Jamie725/SRCNN-Pytorch
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsrcnn_training.py
More file actions
42 lines (36 loc) · 1.07 KB
/
Copy pathsrcnn_training.py
File metadata and controls
42 lines (36 loc) · 1.07 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
from srcnn_module import *
import torchvision.transforms as transforms
import torchvision.datasets as datasets
def train(epoch, trainSet):
epoch_loss = 0
for itr, data in enumerate(trainSet):
imgs, label = data
imgLR, imgHR = imgs
imgLR.unsqueeze_(0)
imgHR.unsqueeze_(0)
if use_gpu :
imgLR = imgLR.cuda()
imgHR = imgLR.cuda()
optimizer.zero_grad()
out_model = srcnn(imgLR)
loss = loss_func(out_model, imgHR)
epoch_loss += loss.item()
loss.backward()
optimizer.step()
#print("===> Epoch[{}]({}/{}): Loss: {:.4f}".format(epoch, itr, len(trainSet), loss.item()))
print("===> Epoch {} Complete: Avg. Loss: {:.4f}".format(epoch, epoch_loss / len(trainSet)))
def test(testSet):
sum_psnr = 0
for itr, data in enumerate(testSet):
imgs, label = data
imgLR, imgHR = imgs
imgLR.unsqueeze_(0)
imgHR.unsqueeze_(0)
if use_gpu :
imgLR = imgLR.cuda()
imgHR = imgLR.cuda()
sr_result = srcnn(imgLR)
MSE = loss_func(sr_result, imgHR)
psnr = 10*log10(1/MSE.item())
sum_psnr += psnr
print("**Average PSNR: {} dB".format(sum_psnr/len(testSet)))