Team Ai
Apppublic

ReflectionEraser/ReflectionEraserApp

sourceHugging Faceotherupdated 2y agoView on Hugging Face
0likes
base_model.py74 linesDownload Raw Back to models
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