ReflectionEraser/ReflectionEraserApp
0
1import torch
2import util.util as util
3from models import make_model
4import time
5import os
6import sys
7from os.path import join
8from util.visualizer import Visualizer
9
10
11class Engine(object):
12 def __init__(self, opt):
13 self.opt = opt
14 self.writer = None
15 self.visualizer = None
16 self.model = None
17 self.best_val_loss = 1e6
18
19 self.__setup()
20
21 def __setup(self):#private methods prefixed with __
22 self.basedir = join('checkpoints', self.opt.name) #creating directory path for checkpoint files
23 os.makedirs(self.basedir, exist_ok=True)
24
25 opt = self.opt
26
27 """(mingcv)Model"""
28 self.model = make_model(self.opt.model)() #(mingcv)models.__dict__[self.opt.model]()
29 self.model.initialize(opt)
30 if not opt.no_log: # setting up logging and visualisation
31 self.writer = util.get_summary_writer(os.path.join(self.basedir, 'logs'))
32 self.visualizer = Visualizer(opt)
33
34 def train(self, train_loader, **kwargs):
35 print('\nEpoch: %d' % self.epoch)
36 avg_meters = util.AverageMeters()#generating running averga eof metrics
37 opt = self.opt
38 model = self.model
39 epoch = self.epoch
40
41 epoch_start_time = time.time()
42 for i, data in enumerate(train_loader):
43 iter_start_time = time.time()
44 iterations = self.iterations
45
46 model.set_input(data, mode='train')
47 model.optimize_parameters(**kwargs)
48
49 errors = model.get_current_errors()
50 avg_meters.update(errors)
51 util.progress_bar(i, len(train_loader), str(avg_meters))
52
53 if not opt.no_log:#update logging and visualisation every iteration of training
54 util.write_loss(self.writer, 'train', avg_meters, iterations)
55
56 if iterations % opt.display_freq == 0 and opt.display_id != 0:
57 save_result = iterations % opt.update_html_freq == 0
58 self.visualizer.display_current_results(model.get_current_visuals(), epoch, save_result)
59
60 if iterations % opt.print_freq == 0 and opt.display_id != 0:
61 t = (time.time() - iter_start_time)
62
63 self.iterations += 1
64
65 self.epoch += 1
66
67 if not self.opt.no_log:#
68 if self.epoch % opt.save_epoch_freq == 0:
69 print('saving the model at epoch %d, iters %d' %
70 (self.epoch, self.iterations))
71 model.save()
72
73 print('saving the latest model at the end of epoch %d, iters %d' %
74 (self.epoch, self.iterations))
75 model.save(label='latest')
76
77 print('Time Taken: %d sec' %
78 (time.time() - epoch_start_time))
79
80 # model.update_learning_rate()
81 try:
82 train_loader.reset()
83 except:
84 pass
85
86 def eval(self, val_loader, dataset_name, savedir='./tmp', loss_key=None, max_save_size=None, **kwargs):
87 #(mingcv)print(dataset_name)
88 if savedir is not None:# create a directory for saving if not exist
89 os.makedirs(savedir, exist_ok=True)
90 self.f = open(os.path.join(savedir, 'metrics.txt'), 'w+')
91 self.f.write(dataset_name + '\n')
92 avg_meters = util.AverageMeters()#get running average metrics
93 model = self.model
94 opt = self.opt
95 with torch.no_grad():#disable gradient calculation to save memory
96 for i, data in enumerate(val_loader):
97 if opt.selected and data['fn'][0].split('.')[0] not in opt.selected:
98 continue
99 if max_save_size is not None and i > max_save_size:
100 index = model.eval(data, savedir=None, **kwargs)
101 else:
102 index = model.eval(data, savedir=savedir, **kwargs)
103
104 #(mingcv) print(data['fn'][0], index)
105 if savedir is not None:
106 self.f.write(f"{data['fn'][0]} {index['PSNR']} {index['SSIM']}\n")
107 avg_meters.update(index)
108 util.progress_bar(i, len(val_loader), str(avg_meters))
109
110 if not opt.no_log:
111 util.write_loss(self.writer, join('eval', dataset_name), avg_meters, self.epoch)
112
113 if loss_key is not None:
114 val_loss = avg_meters[loss_key]
115 if val_loss < self.best_val_loss:
116 self.best_val_loss = val_loss
117 print('saving the best model at the end of epoch %d, iters %d' %
118 (self.epoch, self.iterations))
119 model.save(label='best_{}_{}'.format(loss_key, dataset_name))
120
121 return avg_meters
122
123 def test(self, test_loader, savedir=None, **kwargs):
124 model = self.model
125 opt = self.opt
126 with torch.no_grad():
127 for i, data in enumerate(test_loader):
128 model.test(data, savedir=savedir, **kwargs)
129 util.progress_bar(i, len(test_loader))
130
131 def save_model(self):
132 self.model.save()
133
134 def save_eval(self, label):
135 self.model.save_eval(label)
136
137 @property
138 def iterations(self):
139 return self.model.iterations
140
141 @iterations.setter
142 def iterations(self, i):
143 self.model.iterations = i
144
145 @property
146 def epoch(self):
147 return self.model.epoch
148
149 @epoch.setter
150 def epoch(self, e):
151 self.model.epoch = e
152 