emilios/codeformer-face-restorization
0
1import logging2import os3import torch4from collections import OrderedDict5from copy import deepcopy6from torch.nn.parallel import DataParallel, DistributedDataParallel7 8from basicsr.models import lr_scheduler as lr_scheduler9from basicsr.utils.dist_util import master_only10 11logger = logging.getLogger('basicsr')12 13 14class BaseModel():15 """Base model."""16 17 def __init__(self, opt):18 self.opt = opt19 self.device = torch.device('cuda' if opt['num_gpu'] != 0 else 'cpu')20 self.is_train = opt['is_train']21 self.schedulers = []22 self.optimizers = []23 24 def feed_data(self, data):25 pass26 27 def optimize_parameters(self):28 pass29 30 def get_current_visuals(self):31 pass32 33 def save(self, epoch, current_iter):34 """Save networks and training state."""35 pass36 37 def validation(self, dataloader, current_iter, tb_logger, save_img=False):38 """Validation function.39 40 Args:41 dataloader (torch.utils.data.DataLoader): Validation dataloader.42 current_iter (int): Current iteration.43 tb_logger (tensorboard logger): Tensorboard logger.44 save_img (bool): Whether to save images. Default: False.45 """46 if self.opt['dist']:47 self.dist_validation(dataloader, current_iter, tb_logger, save_img)48 else:49 self.nondist_validation(dataloader, current_iter, tb_logger, save_img)50 51 def model_ema(self, decay=0.999):52 net_g = self.get_bare_model(self.net_g)53 54 net_g_params = dict(net_g.named_parameters())55 net_g_ema_params = dict(self.net_g_ema.named_parameters())56 57 for k in net_g_ema_params.keys():58 net_g_ema_params[k].data.mul_(decay).add_(net_g_params[k].data, alpha=1 - decay)59 60 def get_current_log(self):61 return self.log_dict62 63 def model_to_device(self, net):64 """Model to device. It also warps models with DistributedDataParallel65 or DataParallel.66 67 Args:68 net (nn.Module)69 """70 net = net.to(self.device)71 if self.opt['dist']:72 find_unused_parameters = self.opt.get('find_unused_parameters', False)73 net = DistributedDataParallel(74 net, device_ids=[torch.cuda.current_device()], find_unused_parameters=find_unused_parameters)75 elif self.opt['num_gpu'] > 1:76 net = DataParallel(net)77 return net78 79 def get_optimizer(self, optim_type, params, lr, **kwargs):80 if optim_type == 'Adam':81 optimizer = torch.optim.Adam(params, lr, **kwargs)82 else:83 raise NotImplementedError(f'optimizer {optim_type} is not supperted yet.')84 return optimizer85 86 def setup_schedulers(self):87 """Set up schedulers."""88 train_opt = self.opt['train']89 scheduler_type = train_opt['scheduler'].pop('type')90 if scheduler_type in ['MultiStepLR', 'MultiStepRestartLR']:91 for optimizer in self.optimizers:92 self.schedulers.append(lr_scheduler.MultiStepRestartLR(optimizer, **train_opt['scheduler']))93 elif scheduler_type == 'CosineAnnealingRestartLR':94 for optimizer in self.optimizers:95 self.schedulers.append(lr_scheduler.CosineAnnealingRestartLR(optimizer, **train_opt['scheduler']))96 else:97 raise NotImplementedError(f'Scheduler {scheduler_type} is not implemented yet.')98 99 def get_bare_model(self, net):100 """Get bare model, especially under wrapping with101 DistributedDataParallel or DataParallel.102 """103 if isinstance(net, (DataParallel, DistributedDataParallel)):104 net = net.module105 return net106 107 @master_only108 def print_network(self, net):109 """Print the str and parameter number of a network.110 111 Args:112 net (nn.Module)113 """114 if isinstance(net, (DataParallel, DistributedDataParallel)):115 net_cls_str = (f'{net.__class__.__name__} - ' f'{net.module.__class__.__name__}')116 else:117 net_cls_str = f'{net.__class__.__name__}'118 119 net = self.get_bare_model(net)120 net_str = str(net)121 net_params = sum(map(lambda x: x.numel(), net.parameters()))122 123 logger.info(f'Network: {net_cls_str}, with parameters: {net_params:,d}')124 logger.info(net_str)125 126 def _set_lr(self, lr_groups_l):127 """Set learning rate for warmup.128 129 Args:130 lr_groups_l (list): List for lr_groups, each for an optimizer.131 """132 for optimizer, lr_groups in zip(self.optimizers, lr_groups_l):133 for param_group, lr in zip(optimizer.param_groups, lr_groups):134 param_group['lr'] = lr135 136 def _get_init_lr(self):137 """Get the initial lr, which is set by the scheduler.138 """139 init_lr_groups_l = []140 for optimizer in self.optimizers:141 init_lr_groups_l.append([v['initial_lr'] for v in optimizer.param_groups])142 return init_lr_groups_l143 144 def update_learning_rate(self, current_iter, warmup_iter=-1):145 """Update learning rate.146 147 Args:148 current_iter (int): Current iteration.149 warmup_iter (int): Warmup iter numbers. -1 for no warmup.150 Default: -1.151 """152 if current_iter > 1:153 for scheduler in self.schedulers:154 scheduler.step()155 # set up warm-up learning rate156 if current_iter < warmup_iter:157 # get initial lr for each group158 init_lr_g_l = self._get_init_lr()159 # modify warming-up learning rates160 # currently only support linearly warm up161 warm_up_lr_l = []162 for init_lr_g in init_lr_g_l:163 warm_up_lr_l.append([v / warmup_iter * current_iter for v in init_lr_g])164 # set learning rate165 self._set_lr(warm_up_lr_l)166 167 def get_current_learning_rate(self):168 return [param_group['lr'] for param_group in self.optimizers[0].param_groups]169 170 @master_only171 def save_network(self, net, net_label, current_iter, param_key='params'):172 """Save networks.173 174 Args:175 net (nn.Module | list[nn.Module]): Network(s) to be saved.176 net_label (str): Network label.177 current_iter (int): Current iter number.178 param_key (str | list[str]): The parameter key(s) to save network.179 Default: 'params'.180 """181 if current_iter == -1:182 current_iter = 'latest'183 save_filename = f'{net_label}_{current_iter}.pth'184 save_path = os.path.join(self.opt['path']['models'], save_filename)185 186 net = net if isinstance(net, list) else [net]187 param_key = param_key if isinstance(param_key, list) else [param_key]188 assert len(net) == len(param_key), 'The lengths of net and param_key should be the same.'189 190 save_dict = {}191 for net_, param_key_ in zip(net, param_key):192 net_ = self.get_bare_model(net_)193 state_dict = net_.state_dict()194 for key, param in state_dict.items():195 if key.startswith('module.'): # remove unnecessary 'module.'196 key = key[7:]197 state_dict[key] = param.cpu()198 save_dict[param_key_] = state_dict199 200 torch.save(save_dict, save_path)201 202 def _print_different_keys_loading(self, crt_net, load_net, strict=True):203 """Print keys with differnet name or different size when loading models.204 205 1. Print keys with differnet names.206 2. If strict=False, print the same key but with different tensor size.207 It also ignore these keys with different sizes (not load).208 209 Args:210 crt_net (torch model): Current network.211 load_net (dict): Loaded network.212 strict (bool): Whether strictly loaded. Default: True.213 """214 crt_net = self.get_bare_model(crt_net)215 crt_net = crt_net.state_dict()216 crt_net_keys = set(crt_net.keys())217 load_net_keys = set(load_net.keys())218 219 if crt_net_keys != load_net_keys:220 logger.warning('Current net - loaded net:')221 for v in sorted(list(crt_net_keys - load_net_keys)):222 logger.warning(f' {v}')223 logger.warning('Loaded net - current net:')224 for v in sorted(list(load_net_keys - crt_net_keys)):225 logger.warning(f' {v}')226 227 # check the size for the same keys228 if not strict:229 common_keys = crt_net_keys & load_net_keys230 for k in common_keys:231 if crt_net[k].size() != load_net[k].size():232 logger.warning(f'Size different, ignore [{k}]: crt_net: '233 f'{crt_net[k].shape}; load_net: {load_net[k].shape}')234 load_net[k + '.ignore'] = load_net.pop(k)235 236 def load_network(self, net, load_path, strict=True, param_key='params'):237 """Load network.238 239 Args:240 load_path (str): The path of networks to be loaded.241 net (nn.Module): Network.242 strict (bool): Whether strictly loaded.243 param_key (str): The parameter key of loaded network. If set to244 None, use the root 'path'.245 Default: 'params'.246 """247 net = self.get_bare_model(net)248 logger.info(f'Loading {net.__class__.__name__} model from {load_path}.')249 load_net = torch.load(load_path, map_location=lambda storage, loc: storage)250 if param_key is not None:251 if param_key not in load_net and 'params' in load_net:252 param_key = 'params'253 logger.info('Loading: params_ema does not exist, use params.')254 load_net = load_net[param_key]255 # remove unnecessary 'module.'256 for k, v in deepcopy(load_net).items():257 if k.startswith('module.'):258 load_net[k[7:]] = v259 load_net.pop(k)260 self._print_different_keys_loading(net, load_net, strict)261 net.load_state_dict(load_net, strict=strict)262 263 @master_only264 def save_training_state(self, epoch, current_iter):265 """Save training states during training, which will be used for266 resuming.267 268 Args:269 epoch (int): Current epoch.270 current_iter (int): Current iteration.271 """272 if current_iter != -1:273 state = {'epoch': epoch, 'iter': current_iter, 'optimizers': [], 'schedulers': []}274 for o in self.optimizers:275 state['optimizers'].append(o.state_dict())276 for s in self.schedulers:277 state['schedulers'].append(s.state_dict())278 save_filename = f'{current_iter}.state'279 save_path = os.path.join(self.opt['path']['training_states'], save_filename)280 torch.save(state, save_path)281 282 def resume_training(self, resume_state):283 """Reload the optimizers and schedulers for resumed training.284 285 Args:286 resume_state (dict): Resume state.287 """288 resume_optimizers = resume_state['optimizers']289 resume_schedulers = resume_state['schedulers']290 assert len(resume_optimizers) == len(self.optimizers), 'Wrong lengths of optimizers'291 assert len(resume_schedulers) == len(self.schedulers), 'Wrong lengths of schedulers'292 for i, o in enumerate(resume_optimizers):293 self.optimizers[i].load_state_dict(o)294 for i, s in enumerate(resume_schedulers):295 self.schedulers[i].load_state_dict(s)296 297 def reduce_loss_dict(self, loss_dict):298 """reduce loss dict.299 300 In distributed training, it averages the losses among different GPUs .301 302 Args:303 loss_dict (OrderedDict): Loss dict.304 """305 with torch.no_grad():306 if self.opt['dist']:307 keys = []308 losses = []309 for name, value in loss_dict.items():310 keys.append(name)311 losses.append(value)312 losses = torch.stack(losses, 0)313 torch.distributed.reduce(losses, dst=0)314 if self.opt['rank'] == 0:315 losses /= self.opt['world_size']316 loss_dict = {key: loss for key, loss in zip(keys, losses)}317 318 log_dict = OrderedDict()319 for name, value in loss_dict.items():320 log_dict[name] = value.mean().item()321 322 return log_dict323 