emilios/codeformer-face-restorization
0
1import torch2from collections import OrderedDict3from os import path as osp4from tqdm import tqdm5 6from basicsr.archs import build_network7from basicsr.losses import build_loss8from basicsr.metrics import calculate_metric9from basicsr.utils import get_root_logger, imwrite, tensor2img10from basicsr.utils.registry import MODEL_REGISTRY11from .base_model import BaseModel12 13@MODEL_REGISTRY.register()14class SRModel(BaseModel):15 """Base SR model for single image super-resolution."""16 17 def __init__(self, opt):18 super(SRModel, self).__init__(opt)19 20 # define network21 self.net_g = build_network(opt['network_g'])22 self.net_g = self.model_to_device(self.net_g)23 self.print_network(self.net_g)24 25 # load pretrained models26 load_path = self.opt['path'].get('pretrain_network_g', None)27 if load_path is not None:28 param_key = self.opt['path'].get('param_key_g', 'params')29 self.load_network(self.net_g, load_path, self.opt['path'].get('strict_load_g', True), param_key)30 31 if self.is_train:32 self.init_training_settings()33 34 def init_training_settings(self):35 self.net_g.train()36 train_opt = self.opt['train']37 38 self.ema_decay = train_opt.get('ema_decay', 0)39 if self.ema_decay > 0:40 logger = get_root_logger()41 logger.info(f'Use Exponential Moving Average with decay: {self.ema_decay}')42 # define network net_g with Exponential Moving Average (EMA)43 # net_g_ema is used only for testing on one GPU and saving44 # There is no need to wrap with DistributedDataParallel45 self.net_g_ema = build_network(self.opt['network_g']).to(self.device)46 # load pretrained model47 load_path = self.opt['path'].get('pretrain_network_g', None)48 if load_path is not None:49 self.load_network(self.net_g_ema, load_path, self.opt['path'].get('strict_load_g', True), 'params_ema')50 else:51 self.model_ema(0) # copy net_g weight52 self.net_g_ema.eval()53 54 # define losses55 if train_opt.get('pixel_opt'):56 self.cri_pix = build_loss(train_opt['pixel_opt']).to(self.device)57 else:58 self.cri_pix = None59 60 if train_opt.get('perceptual_opt'):61 self.cri_perceptual = build_loss(train_opt['perceptual_opt']).to(self.device)62 else:63 self.cri_perceptual = None64 65 if self.cri_pix is None and self.cri_perceptual is None:66 raise ValueError('Both pixel and perceptual losses are None.')67 68 # set up optimizers and schedulers69 self.setup_optimizers()70 self.setup_schedulers()71 72 def setup_optimizers(self):73 train_opt = self.opt['train']74 optim_params = []75 for k, v in self.net_g.named_parameters():76 if v.requires_grad:77 optim_params.append(v)78 else:79 logger = get_root_logger()80 logger.warning(f'Params {k} will not be optimized.')81 82 optim_type = train_opt['optim_g'].pop('type')83 self.optimizer_g = self.get_optimizer(optim_type, optim_params, **train_opt['optim_g'])84 self.optimizers.append(self.optimizer_g)85 86 def feed_data(self, data):87 self.lq = data['lq'].to(self.device)88 if 'gt' in data:89 self.gt = data['gt'].to(self.device)90 91 def optimize_parameters(self, current_iter):92 self.optimizer_g.zero_grad()93 self.output = self.net_g(self.lq)94 95 l_total = 096 loss_dict = OrderedDict()97 # pixel loss98 if self.cri_pix:99 l_pix = self.cri_pix(self.output, self.gt)100 l_total += l_pix101 loss_dict['l_pix'] = l_pix102 # perceptual loss103 if self.cri_perceptual:104 l_percep, l_style = self.cri_perceptual(self.output, self.gt)105 if l_percep is not None:106 l_total += l_percep107 loss_dict['l_percep'] = l_percep108 if l_style is not None:109 l_total += l_style110 loss_dict['l_style'] = l_style111 112 l_total.backward()113 self.optimizer_g.step()114 115 self.log_dict = self.reduce_loss_dict(loss_dict)116 117 if self.ema_decay > 0:118 self.model_ema(decay=self.ema_decay)119 120 def test(self):121 if hasattr(self, 'ema_decay'):122 self.net_g_ema.eval()123 with torch.no_grad():124 self.output = self.net_g_ema(self.lq)125 else:126 self.net_g.eval()127 with torch.no_grad():128 self.output = self.net_g(self.lq)129 self.net_g.train()130 131 def dist_validation(self, dataloader, current_iter, tb_logger, save_img):132 if self.opt['rank'] == 0:133 self.nondist_validation(dataloader, current_iter, tb_logger, save_img)134 135 def nondist_validation(self, dataloader, current_iter, tb_logger, save_img):136 dataset_name = dataloader.dataset.opt['name']137 with_metrics = self.opt['val'].get('metrics') is not None138 if with_metrics:139 self.metric_results = {metric: 0 for metric in self.opt['val']['metrics'].keys()}140 pbar = tqdm(total=len(dataloader), unit='image')141 142 for idx, val_data in enumerate(dataloader):143 img_name = osp.splitext(osp.basename(val_data['lq_path'][0]))[0]144 self.feed_data(val_data)145 self.test()146 147 visuals = self.get_current_visuals()148 sr_img = tensor2img([visuals['result']])149 if 'gt' in visuals:150 gt_img = tensor2img([visuals['gt']])151 del self.gt152 153 # tentative for out of GPU memory154 del self.lq155 del self.output156 torch.cuda.empty_cache()157 158 if save_img:159 if self.opt['is_train']:160 save_img_path = osp.join(self.opt['path']['visualization'], img_name,161 f'{img_name}_{current_iter}.png')162 else:163 if self.opt['val']['suffix']:164 save_img_path = osp.join(self.opt['path']['visualization'], dataset_name,165 f'{img_name}_{self.opt["val"]["suffix"]}.png')166 else:167 save_img_path = osp.join(self.opt['path']['visualization'], dataset_name,168 f'{img_name}_{self.opt["name"]}.png')169 imwrite(sr_img, save_img_path)170 171 if with_metrics:172 # calculate metrics173 for name, opt_ in self.opt['val']['metrics'].items():174 metric_data = dict(img1=sr_img, img2=gt_img)175 self.metric_results[name] += calculate_metric(metric_data, opt_)176 pbar.update(1)177 pbar.set_description(f'Test {img_name}')178 pbar.close()179 180 if with_metrics:181 for metric in self.metric_results.keys():182 self.metric_results[metric] /= (idx + 1)183 184 self._log_validation_metric_values(current_iter, dataset_name, tb_logger)185 186 def _log_validation_metric_values(self, current_iter, dataset_name, tb_logger):187 log_str = f'Validation {dataset_name}\n'188 for metric, value in self.metric_results.items():189 log_str += f'\t # {metric}: {value:.4f}\n'190 logger = get_root_logger()191 logger.info(log_str)192 if tb_logger:193 for metric, value in self.metric_results.items():194 tb_logger.add_scalar(f'metrics/{metric}', value, current_iter)195 196 def get_current_visuals(self):197 out_dict = OrderedDict()198 out_dict['lq'] = self.lq.detach().cpu()199 out_dict['result'] = self.output.detach().cpu()200 if hasattr(self, 'gt'):201 out_dict['gt'] = self.gt.detach().cpu()202 return out_dict203 204 def save(self, epoch, current_iter):205 if hasattr(self, 'ema_decay'):206 self.save_network([self.net_g, self.net_g_ema], 'net_g', current_iter, param_key=['params', 'params_ema'])207 else:208 self.save_network(self.net_g, 'net_g', current_iter)209 self.save_training_state(epoch, current_iter)210 