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_REGISTRY11import torch.nn.functional as F12from .sr_model import SRModel13 14 15@MODEL_REGISTRY.register()16class CodeFormerModel(SRModel):17 def feed_data(self, data):18 self.gt = data['gt'].to(self.device)19 self.input = data['in'].to(self.device)20 self.b = self.gt.shape[0]21 22 if 'latent_gt' in data:23 self.idx_gt = data['latent_gt'].to(self.device)24 self.idx_gt = self.idx_gt.view(self.b, -1)25 else:26 self.idx_gt = None27 28 def init_training_settings(self):29 logger = get_root_logger()30 train_opt = self.opt['train']31 32 self.ema_decay = train_opt.get('ema_decay', 0)33 if self.ema_decay > 0:34 logger.info(f'Use Exponential Moving Average with decay: {self.ema_decay}')35 # define network net_g with Exponential Moving Average (EMA)36 # net_g_ema is used only for testing on one GPU and saving37 # There is no need to wrap with DistributedDataParallel38 self.net_g_ema = build_network(self.opt['network_g']).to(self.device)39 # load pretrained model40 load_path = self.opt['path'].get('pretrain_network_g', None)41 if load_path is not None:42 self.load_network(self.net_g_ema, load_path, self.opt['path'].get('strict_load_g', True), 'params_ema')43 else:44 self.model_ema(0) # copy net_g weight45 self.net_g_ema.eval()46 47 if self.opt.get('network_vqgan', None) is not None and self.opt['datasets'].get('latent_gt_path') is None:48 self.hq_vqgan_fix = build_network(self.opt['network_vqgan']).to(self.device)49 self.hq_vqgan_fix.eval()50 self.generate_idx_gt = True51 for param in self.hq_vqgan_fix.parameters():52 param.requires_grad = False53 else:54 self.generate_idx_gt = False55 56 self.hq_feat_loss = train_opt.get('use_hq_feat_loss', True)57 self.feat_loss_weight = train_opt.get('feat_loss_weight', 1.0)58 self.cross_entropy_loss = train_opt.get('cross_entropy_loss', True)59 self.entropy_loss_weight = train_opt.get('entropy_loss_weight', 0.5)60 self.fidelity_weight = train_opt.get('fidelity_weight', 1.0)61 self.scale_adaptive_gan_weight = train_opt.get('scale_adaptive_gan_weight', 0.8)62 63 64 self.net_g.train()65 # define network net_d66 if self.fidelity_weight > 0:67 self.net_d = build_network(self.opt['network_d'])68 self.net_d = self.model_to_device(self.net_d)69 self.print_network(self.net_d)70 71 # load pretrained models72 load_path = self.opt['path'].get('pretrain_network_d', None)73 if load_path is not None:74 self.load_network(self.net_d, load_path, self.opt['path'].get('strict_load_d', True))75 76 self.net_d.train()77 78 # define losses79 if train_opt.get('pixel_opt'):80 self.cri_pix = build_loss(train_opt['pixel_opt']).to(self.device)81 else:82 self.cri_pix = None83 84 if train_opt.get('perceptual_opt'):85 self.cri_perceptual = build_loss(train_opt['perceptual_opt']).to(self.device)86 else:87 self.cri_perceptual = None88 89 if train_opt.get('gan_opt'):90 self.cri_gan = build_loss(train_opt['gan_opt']).to(self.device)91 92 93 self.fix_generator = train_opt.get('fix_generator', True)94 logger.info(f'fix_generator: {self.fix_generator}')95 96 self.net_g_start_iter = train_opt.get('net_g_start_iter', 0)97 self.net_d_iters = train_opt.get('net_d_iters', 1)98 self.net_d_start_iter = train_opt.get('net_d_start_iter', 0)99 100 # set up optimizers and schedulers101 self.setup_optimizers()102 self.setup_schedulers()103 104 def calculate_adaptive_weight(self, recon_loss, g_loss, last_layer, disc_weight_max):105 recon_grads = torch.autograd.grad(recon_loss, last_layer, retain_graph=True)[0]106 g_grads = torch.autograd.grad(g_loss, last_layer, retain_graph=True)[0]107 108 d_weight = torch.norm(recon_grads) / (torch.norm(g_grads) + 1e-4)109 d_weight = torch.clamp(d_weight, 0.0, disc_weight_max).detach()110 return d_weight111 112 def setup_optimizers(self):113 train_opt = self.opt['train']114 # optimizer g115 optim_params_g = []116 for k, v in self.net_g.named_parameters():117 if v.requires_grad:118 optim_params_g.append(v)119 else:120 logger = get_root_logger()121 logger.warning(f'Params {k} will not be optimized.')122 optim_type = train_opt['optim_g'].pop('type')123 self.optimizer_g = self.get_optimizer(optim_type, optim_params_g, **train_opt['optim_g'])124 self.optimizers.append(self.optimizer_g)125 # optimizer d126 if self.fidelity_weight > 0:127 optim_type = train_opt['optim_d'].pop('type')128 self.optimizer_d = self.get_optimizer(optim_type, self.net_d.parameters(), **train_opt['optim_d'])129 self.optimizers.append(self.optimizer_d)130 131 def gray_resize_for_identity(self, out, size=128):132 out_gray = (0.2989 * out[:, 0, :, :] + 0.5870 * out[:, 1, :, :] + 0.1140 * out[:, 2, :, :])133 out_gray = out_gray.unsqueeze(1)134 out_gray = F.interpolate(out_gray, (size, size), mode='bilinear', align_corners=False)135 return out_gray136 137 def optimize_parameters(self, current_iter):138 logger = get_root_logger()139 # optimize net_g140 for p in self.net_d.parameters():141 p.requires_grad = False142 143 self.optimizer_g.zero_grad()144 145 if self.generate_idx_gt:146 x = self.hq_vqgan_fix.encoder(self.gt)147 output, _, quant_stats = self.hq_vqgan_fix.quantize(x)148 min_encoding_indices = quant_stats['min_encoding_indices']149 self.idx_gt = min_encoding_indices.view(self.b, -1)150 151 if self.fidelity_weight > 0:152 self.output, logits, lq_feat = self.net_g(self.input, w=self.fidelity_weight, detach_16=True)153 else:154 logits, lq_feat = self.net_g(self.input, w=0, code_only=True)155 156 if self.hq_feat_loss:157 # quant_feats158 quant_feat_gt = self.net_g.module.quantize.get_codebook_feat(self.idx_gt, shape=[self.b,16,16,256])159 160 l_g_total = 0161 loss_dict = OrderedDict()162 if current_iter % self.net_d_iters == 0 and current_iter > self.net_g_start_iter:163 # hq_feat_loss164 if self.hq_feat_loss: # codebook loss 165 l_feat_encoder = torch.mean((quant_feat_gt.detach()-lq_feat)**2) * self.feat_loss_weight166 l_g_total += l_feat_encoder167 loss_dict['l_feat_encoder'] = l_feat_encoder168 169 # cross_entropy_loss170 if self.cross_entropy_loss:171 # b(hw)n -> bn(hw)172 cross_entropy_loss = F.cross_entropy(logits.permute(0, 2, 1), self.idx_gt) * self.entropy_loss_weight173 l_g_total += cross_entropy_loss174 loss_dict['cross_entropy_loss'] = cross_entropy_loss175 176 if self.fidelity_weight > 0: # when fidelity_weight == 0 don't need image-level loss177 # pixel loss178 if self.cri_pix:179 l_g_pix = self.cri_pix(self.output, self.gt)180 l_g_total += l_g_pix181 loss_dict['l_g_pix'] = l_g_pix182 183 # perceptual loss184 if self.cri_perceptual:185 l_g_percep = self.cri_perceptual(self.output, self.gt)186 l_g_total += l_g_percep187 loss_dict['l_g_percep'] = l_g_percep188 189 # gan loss190 if current_iter > self.net_d_start_iter:191 fake_g_pred = self.net_d(self.output)192 l_g_gan = self.cri_gan(fake_g_pred, True, is_disc=False)193 recon_loss = l_g_pix + l_g_percep194 if not self.fix_generator:195 last_layer = self.net_g.module.generator.blocks[-1].weight196 d_weight = self.calculate_adaptive_weight(recon_loss, l_g_gan, last_layer, disc_weight_max=1.0)197 else:198 largest_fuse_size = self.opt['network_g']['connect_list'][-1]199 last_layer = self.net_g.module.fuse_convs_dict[largest_fuse_size].shift[-1].weight200 d_weight = self.calculate_adaptive_weight(recon_loss, l_g_gan, last_layer, disc_weight_max=1.0)201 202 d_weight *= self.scale_adaptive_gan_weight # 0.8203 loss_dict['d_weight'] = d_weight204 l_g_total += d_weight * l_g_gan205 loss_dict['l_g_gan'] = d_weight * l_g_gan206 207 l_g_total.backward()208 self.optimizer_g.step()209 210 if self.ema_decay > 0:211 self.model_ema(decay=self.ema_decay)212 213 # optimize net_d214 if current_iter > self.net_d_start_iter and self.fidelity_weight > 0:215 for p in self.net_d.parameters():216 p.requires_grad = True217 218 self.optimizer_d.zero_grad()219 # real220 real_d_pred = self.net_d(self.gt)221 l_d_real = self.cri_gan(real_d_pred, True, is_disc=True)222 loss_dict['l_d_real'] = l_d_real223 loss_dict['out_d_real'] = torch.mean(real_d_pred.detach())224 l_d_real.backward()225 # fake226 fake_d_pred = self.net_d(self.output.detach())227 l_d_fake = self.cri_gan(fake_d_pred, False, is_disc=True)228 loss_dict['l_d_fake'] = l_d_fake229 loss_dict['out_d_fake'] = torch.mean(fake_d_pred.detach())230 l_d_fake.backward()231 232 self.optimizer_d.step()233 234 self.log_dict = self.reduce_loss_dict(loss_dict)235 236 237 def test(self):238 with torch.no_grad():239 if hasattr(self, 'net_g_ema'):240 self.net_g_ema.eval()241 self.output, _, _ = self.net_g_ema(self.input, w=self.fidelity_weight)242 else:243 logger = get_root_logger()244 logger.warning('Do not have self.net_g_ema, use self.net_g.')245 self.net_g.eval()246 self.output, _, _ = self.net_g(self.input, w=self.fidelity_weight)247 self.net_g.train()248 249 250 def dist_validation(self, dataloader, current_iter, tb_logger, save_img):251 if self.opt['rank'] == 0:252 self.nondist_validation(dataloader, current_iter, tb_logger, save_img)253 254 255 def nondist_validation(self, dataloader, current_iter, tb_logger, save_img):256 dataset_name = dataloader.dataset.opt['name']257 with_metrics = self.opt['val'].get('metrics') is not None258 if with_metrics:259 self.metric_results = {metric: 0 for metric in self.opt['val']['metrics'].keys()}260 pbar = tqdm(total=len(dataloader), unit='image')261 262 for idx, val_data in enumerate(dataloader):263 img_name = osp.splitext(osp.basename(val_data['lq_path'][0]))[0]264 self.feed_data(val_data)265 self.test()266 267 visuals = self.get_current_visuals()268 sr_img = tensor2img([visuals['result']])269 if 'gt' in visuals:270 gt_img = tensor2img([visuals['gt']])271 del self.gt272 273 # tentative for out of GPU memory274 del self.lq275 del self.output276 torch.cuda.empty_cache()277 278 if save_img:279 if self.opt['is_train']:280 save_img_path = osp.join(self.opt['path']['visualization'], img_name,281 f'{img_name}_{current_iter}.png')282 else:283 if self.opt['val']['suffix']:284 save_img_path = osp.join(self.opt['path']['visualization'], dataset_name,285 f'{img_name}_{self.opt["val"]["suffix"]}.png')286 else:287 save_img_path = osp.join(self.opt['path']['visualization'], dataset_name,288 f'{img_name}_{self.opt["name"]}.png')289 imwrite(sr_img, save_img_path)290 291 if with_metrics:292 # calculate metrics293 for name, opt_ in self.opt['val']['metrics'].items():294 metric_data = dict(img1=sr_img, img2=gt_img)295 self.metric_results[name] += calculate_metric(metric_data, opt_)296 pbar.update(1)297 pbar.set_description(f'Test {img_name}')298 pbar.close()299 300 if with_metrics:301 for metric in self.metric_results.keys():302 self.metric_results[metric] /= (idx + 1)303 304 self._log_validation_metric_values(current_iter, dataset_name, tb_logger)305 306 307 def _log_validation_metric_values(self, current_iter, dataset_name, tb_logger):308 log_str = f'Validation {dataset_name}\n'309 for metric, value in self.metric_results.items():310 log_str += f'\t # {metric}: {value:.4f}\n'311 logger = get_root_logger()312 logger.info(log_str)313 if tb_logger:314 for metric, value in self.metric_results.items():315 tb_logger.add_scalar(f'metrics/{metric}', value, current_iter)316 317 318 def get_current_visuals(self):319 out_dict = OrderedDict()320 out_dict['gt'] = self.gt.detach().cpu()321 out_dict['result'] = self.output.detach().cpu()322 return out_dict323 324 325 def save(self, epoch, current_iter):326 if self.ema_decay > 0:327 self.save_network([self.net_g, self.net_g_ema], 'net_g', current_iter, param_key=['params', 'params_ema'])328 else:329 self.save_network(self.net_g, 'net_g', current_iter)330 if self.fidelity_weight > 0:331 self.save_network(self.net_d, 'net_d', current_iter)332 self.save_training_state(epoch, current_iter)333 