CoolFace
Apppublic

GoodWin/Deep-Multi-scale

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
base_model.py174 linesDownload Raw Back to models
1import os2import torch3from collections import OrderedDict4from . import networks5 6 7class BaseModel():8 9    # modify parser to add command line options,10    # and also change the default values if needed11    @staticmethod12    def modify_commandline_options(parser, is_train):13        return parser14    15    def name(self):16        return 'BaseModel'17 18    def initialize(self, opt):19        self.opt = opt20        self.gpu_ids = opt.gpu_ids21        self.isTrain = opt.isTrain22        self.device = torch.device('cuda:{}'.format(self.gpu_ids[0])) if self.gpu_ids else torch.device('cpu')23        self.save_dir = os.path.join(opt.checkpoints_dir, opt.name)24        if opt.resize_or_crop != 'scale_width':25            torch.backends.cudnn.benchmark = True26        self.loss_names = []27        self.model_names = []28        self.visual_names = []29        self.image_paths = []30        # self.optimizers = []31 32    def set_input(self, input):33        self.input = input34 35    def forward(self):36        pass37 38    # load and print networks; create schedulers39    def setup(self, opt, parser=None):40        if self.isTrain:41            self.schedulers = [networks.get_scheduler(optimizer, opt) for optimizer in self.optimizers]42        if not self.isTrain or opt.continue_train:43            self.load_networks(opt.which_epoch)44        self.print_networks(opt.verbose)45 46    # make models eval mode during test time47    def eval(self):48        for name in self.model_names:49            if isinstance(name, str):50                net = getattr(self, 'net' + name)51                net.eval()52 53    # used in test time, wrapping `forward` in no_grad() so we don't save54    # intermediate steps for backprop55    def test(self):56        with torch.no_grad():57            self.forward()58 59    # get image paths60    def get_image_paths(self):61        return self.image_paths62 63    def optimize_parameters(self):64        pass65 66    # update learning rate (called once every epoch)67    def update_learning_rate(self):68        for scheduler in self.schedulers:69            scheduler.step()70        lr = self.optimizers[0].param_groups[0]['lr']71        print('learning rate = %.7f' % lr)72 73    # return visualization images. train.py will display these images, and save the images to a html74    def get_current_visuals(self):75        visual_ret = OrderedDict()76        for name in self.visual_names:77            if isinstance(name, str):78                visual_ret[name] = getattr(self, name)79        return visual_ret80 81    # return traning losses/errors. train.py will print out these errors as debugging information82    def get_current_losses(self):83        errors_ret = OrderedDict()84        for name in self.loss_names:85            if isinstance(name, str):86                # float(...) works for both scalar tensor and float number87                errors_ret[name] = float(getattr(self, 'loss_' + name))88        return errors_ret89 90    # save models to the disk91    def save_networks(self, which_epoch):92        for name in self.model_names:93            if isinstance(name, str):94                save_filename = '%s_net_%s.pth' % (which_epoch, name)95                save_path = os.path.join(self.save_dir, save_filename)96                net = getattr(self, 'net' + name)97 98                if len(self.gpu_ids) > 0 and torch.cuda.is_available():99                    torch.save(net.module.cpu().state_dict(), save_path)100                    net.cuda(self.gpu_ids[0])101                else:102                    torch.save(net.cpu().state_dict(), save_path)103 104    def __patch_instance_norm_state_dict(self, state_dict, module, keys, i=0):105        key = keys[i]106        if i + 1 == len(keys):  # at the end, pointing to a parameter/buffer107            if module.__class__.__name__.startswith('InstanceNorm') and \108                    (key == 'running_mean' or key == 'running_var'):109                if getattr(module, key) is None:110                    state_dict.pop('.'.join(keys))111            if module.__class__.__name__.startswith('InstanceNorm') and \112               (key == 'num_batches_tracked'):113                state_dict.pop('.'.join(keys))114        else:115            self.__patch_instance_norm_state_dict(state_dict, getattr(module, key), keys, i + 1)116 117    # load models from the disk118    def load_networks(self, which_epoch):119        for name in self.model_names:120            if isinstance(name, str):121                load_filename = '%s_net_%s.pth' % (which_epoch, name)122                load_path = os.path.join(self.save_dir, load_filename)123                net = getattr(self, 'net' + name)124                if isinstance(net, torch.nn.DataParallel):125                    net = net.module126                # print('loading the model from %s' % load_path)127                # if you are using PyTorch newer than 0.4 (e.g., built from128                # GitHub source), you can remove str() on self.device129                if not os.path.exists(load_path):130                    continue131                state_dict = torch.load(load_path, map_location=str(self.device))132                if hasattr(state_dict, '_metadata'):133                    del state_dict._metadata134 135                # patch InstanceNorm checkpoints prior to 0.4136                # for key in list(state_dict.keys()):  # need to copy keys here because we mutate in loop137                #     self.__patch_instance_norm_state_dict(state_dict, net, key.split('.'))138                model_dict = net.state_dict()139                # new_dict = {k: v for k, v in state_dict.items() if k in model_dict.keys()}140                new_dict = {}141                for k, v in state_dict.items():142                    if k in model_dict.keys():143                        # print(k)144                        # if k == 'sff_branch.0.sff0.MaskModel.0.weight' or k =='sff_branch.0.sff1.MaskModel.0.weight' or k == 'sff_branch.1.sff0.MaskModel.0.weight' or k =='sff_branch.1.sff1.MaskModel.0.weight'  or k == 'sff_branch.2.sff0.MaskModel.0.weight' or k =='sff_branch.2.sff1.MaskModel.0.weight'  or k == 'sff_branch.3.sff0.MaskModel.0.weight' or k =='sff_branch.3.sff1.MaskModel.0.weight'  or k == 'sff_branch.4.MaskModel.0.weight' :145                        #     continue146                        # if 'Mask_CModel.model' in k:147                        #     continue148                        new_dict[k] = v149                model_dict.update(new_dict)150                net.load_state_dict(model_dict)151 152    # print network information153    def print_networks(self, verbose):154        # print('---------- Networks initialized -------------')155        for name in self.model_names:156            if isinstance(name, str):157                net = getattr(self, 'net' + name)158                num_params = 0159                for param in net.parameters():160                    num_params += param.numel()161                # if verbose:162                #     print(net)163                # print('[Network %s] Total number of parameters : %.3f M' % (name, num_params / 1e6))164        # print('-----------------------------------------------')165 166    # set requies_grad=Fasle to avoid computation167    def set_requires_grad(self, nets, requires_grad=False):168        if not isinstance(nets, list):169            nets = [nets]170        for net in nets:171            if net is not None:172                for param in net.parameters():173                    param.requires_grad = requires_grad174