Team Ai
Apppublic

ReflectionEraser/ReflectionEraserApp

sourceHugging Faceotherupdated 2y agoView on Hugging Face
0likes
engine.py152 linesDownload Raw Back to DSRNet
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