ReflectionEraser/ReflectionEraserApp
0
1import os
2import torch
3import util.util as util
4from tools import mutils
5
6
7class BaseModel:
8 def name(self):
9 return self.__class__.__name__.lower()
10
11 def initialize(self, opt):
12 self.opt = opt
13 self.gpu_ids = opt.gpu_ids
14 self.isTrain = opt.isTrain
15 self.Tensor = torch.cuda.FloatTensor if self.gpu_ids else torch.Tensor
16 last_split = opt.checkpoints_dir.split('/')[-1]
17 if opt.resume and last_split != 'checkpoints' and (last_split != opt.name or opt.supp_eval):
18
19 self.save_dir = opt.checkpoints_dir
20 self.model_save_dir = os.path.join(opt.checkpoints_dir.replace(opt.checkpoints_dir.split('/')[-1], ''),
21 opt.name, 'weights', mutils.get_formatted_time())
22 else:
23 self.save_dir = os.path.join(opt.checkpoints_dir, opt.name)
24 self.model_save_dir = os.path.join(opt.checkpoints_dir, opt.name, 'weights', mutils.get_formatted_time())
25 self._count = 0
26
27 def set_input(self, input):
28 self.input = input
29
30 def forward(self, mode='train'):
31 pass
32
33 # used in test time, no backprop
34 def test(self):
35 pass
36
37 def get_image_paths(self):
38 pass
39
40 def optimize_parameters(self):
41 pass
42
43 def get_current_visuals(self):
44 return self.input
45
46 def get_current_errors(self):
47 return {}
48
49 def print_optimizer_param(self):
50 print(self.optimizers[-1])
51
52 def save(self, label=None):
53 epoch = self.epoch
54 iterations = self.iterations
55
56 os.makedirs(self.model_save_dir, exist_ok=True)
57 if label is None:
58 model_name = os.path.join(self.model_save_dir, self.opt.name + '_%03d_%08d.pt' % ((epoch), (iterations)))
59 else:
60 model_name = os.path.join(self.model_save_dir, self.opt.name + '_' + label + '.pt')
61
62 torch.save(self.state_dict(), model_name)
63
64 def save_eval(self, label=None):
65 model_name = os.path.join(self.model_save_dir, label + '.pt')
66
67 torch.save(self.state_dict_eval(), model_name)
68
69 def _init_optimizer(self, optimizers):
70 self.optimizers = optimizers
71 for optimizer in self.optimizers:
72 util.set_opt_param(optimizer, 'initial_lr', self.opt.lr)
73 util.set_opt_param(optimizer, 'weight_decay', self.opt.wd)
74 