GoodWin/Deep-Multi-scale
0
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 