emilios/codeformer-face-restorization
0
1import torch2from collections import OrderedDict3from os import path as osp4from tqdm import tqdm5 6 7from basicsr.archs import build_network8from basicsr.losses import build_loss9from basicsr.metrics import calculate_metric10from basicsr.utils import get_root_logger, imwrite, tensor2img11from basicsr.utils.registry import MODEL_REGISTRY12import torch.nn.functional as F13from .sr_model import SRModel14 15 16@MODEL_REGISTRY.register()17class CodeFormerJointModel(SRModel):18 def feed_data(self, data):19 self.gt = data['gt'].to(self.device)20 self.input = data['in'].to(self.device)21 self.input_large_de = data['in_large_de'].to(self.device)22 self.b = self.gt.shape[0]23 24 if 'latent_gt' in data:25 self.idx_gt = data['latent_gt'].to(self.device)26 self.idx_gt = self.idx_gt.view(self.b, -1)27 else:28 self.idx_gt = None29 30 def init_training_settings(self):31 logger = get_root_logger()32 train_opt = self.opt['train']33 34 self.ema_decay = train_opt.get('ema_decay', 0)35 if self.ema_decay > 0:36 logger.info(f'Use Exponential Moving Average with decay: {self.ema_decay}')37 # define network net_g with Exponential Moving Average (EMA)38 # net_g_ema is used only for testing on one GPU and saving39 # There is no need to wrap with DistributedDataParallel40 self.net_g_ema = build_network(self.opt['network_g']).to(self.device)41 # load pretrained model42 load_path = self.opt['path'].get('pretrain_network_g', None)43 if load_path is not None:44 self.load_network(self.net_g_ema, load_path, self.opt['path'].get('strict_load_g', True), 'params_ema')45 else:46 self.model_ema(0) # copy net_g weight47 self.net_g_ema.eval()48 49 if self.opt['datasets']['train'].get('latent_gt_path', None) is not None:50 self.generate_idx_gt = False51 elif self.opt.get('network_vqgan', None) is not None:52 self.hq_vqgan_fix = build_network(self.opt['network_vqgan']).to(self.device)53 self.hq_vqgan_fix.eval()54 self.generate_idx_gt = True55 for param in self.hq_vqgan_fix.parameters():56 param.requires_grad = False57 else:58 raise NotImplementedError(f'Shoule have network_vqgan config or pre-calculated latent code.') 59 60 logger.info(f'Need to generate latent GT code: {self.generate_idx_gt}')61 62 self.hq_feat_loss = train_opt.get('use_hq_feat_loss', True)63 self.feat_loss_weight = train_opt.get('feat_loss_weight', 1.0)64 self.cross_entropy_loss = train_opt.get('cross_entropy_loss', True)65 self.entropy_loss_weight = train_opt.get('entropy_loss_weight', 0.5)66 self.scale_adaptive_gan_weight = train_opt.get('scale_adaptive_gan_weight', 0.8)67 68 # define network net_d69 self.net_d = build_network(self.opt['network_d'])70 self.net_d = self.model_to_device(self.net_d)71 self.print_network(self.net_d)72 73 # load pretrained models74 load_path = self.opt['path'].get('pretrain_network_d', None)75 if load_path is not None:76 self.load_network(self.net_d, load_path, self.opt['path'].get('strict_load_d', True))77 78 self.net_g.train()79 self.net_d.train()80 81 # define losses82 if train_opt.get('pixel_opt'):83 self.cri_pix = build_loss(train_opt['pixel_opt']).to(self.device)84 else:85 self.cri_pix = None86 87 if train_opt.get('perceptual_opt'):88 self.cri_perceptual = build_loss(train_opt['perceptual_opt']).to(self.device)89 else:90 self.cri_perceptual = None91 92 if train_opt.get('gan_opt'):93 self.cri_gan = build_loss(train_opt['gan_opt']).to(self.device)94 95 96 self.fix_generator = train_opt.get('fix_generator', True)97 logger.info(f'fix_generator: {self.fix_generator}')98 99 self.net_g_start_iter = train_opt.get('net_g_start_iter', 0)100 self.net_d_iters = train_opt.get('net_d_iters', 1)101 self.net_d_start_iter = train_opt.get('net_d_start_iter', 0)102 103 # set up optimizers and schedulers104 self.setup_optimizers()105 self.setup_schedulers()106 107 def calculate_adaptive_weight(self, recon_loss, g_loss, last_layer, disc_weight_max):108 recon_grads = torch.autograd.grad(recon_loss, last_layer, retain_graph=True)[0]109 g_grads = torch.autograd.grad(g_loss, last_layer, retain_graph=True)[0]110 111 d_weight = torch.norm(recon_grads) / (torch.norm(g_grads) + 1e-4)112 d_weight = torch.clamp(d_weight, 0.0, disc_weight_max).detach()113 return d_weight114 115 def setup_optimizers(self):116 train_opt = self.opt['train']117 # optimizer g118 optim_params_g = []119 for k, v in self.net_g.named_parameters():120 if v.requires_grad:121 optim_params_g.append(v)122 else:123 logger = get_root_logger()124 logger.warning(f'Params {k} will not be optimized.')125 optim_type = train_opt['optim_g'].pop('type')126 self.optimizer_g = self.get_optimizer(optim_type, optim_params_g, **train_opt['optim_g'])127 self.optimizers.append(self.optimizer_g)128 # optimizer d129 optim_type = train_opt['optim_d'].pop('type')130 self.optimizer_d = self.get_optimizer(optim_type, self.net_d.parameters(), **train_opt['optim_d'])131 self.optimizers.append(self.optimizer_d)132 133 def gray_resize_for_identity(self, out, size=128):134 out_gray = (0.2989 * out[:, 0, :, :] + 0.5870 * out[:, 1, :, :] + 0.1140 * out[:, 2, :, :])135 out_gray = out_gray.unsqueeze(1)136 out_gray = F.interpolate(out_gray, (size, size), mode='bilinear', align_corners=False)137 return out_gray138 139 def optimize_parameters(self, current_iter):140 logger = get_root_logger()141 # optimize net_g142 for p in self.net_d.parameters():143 p.requires_grad = False144 145 self.optimizer_g.zero_grad()146 147 if self.generate_idx_gt:148 x = self.hq_vqgan_fix.encoder(self.gt)149 output, _, quant_stats = self.hq_vqgan_fix.quantize(x)150 min_encoding_indices = quant_stats['min_encoding_indices']151 self.idx_gt = min_encoding_indices.view(self.b, -1)152 153 if current_iter <= 40000: # small degradation154 small_per_n = 1155 w = 1156 elif current_iter <= 80000: # small degradation157 small_per_n = 1158 w = 1.3 159 elif current_iter <= 120000: # large degradation160 small_per_n = 120000161 w = 0 162 else: # mixed degradation163 small_per_n = 15164 w = 1.3165 166 if current_iter % small_per_n == 0:167 self.output, logits, lq_feat = self.net_g(self.input, w=w, detach_16=True)168 large_de = False169 else:170 logits, lq_feat = self.net_g(self.input_large_de, code_only=True)171 large_de = True172 173 if self.hq_feat_loss:174 # quant_feats175 quant_feat_gt = self.net_g.module.quantize.get_codebook_feat(self.idx_gt, shape=[self.b,16,16,256])176 177 l_g_total = 0178 loss_dict = OrderedDict()179 if current_iter % self.net_d_iters == 0 and current_iter > self.net_g_start_iter:180 # hq_feat_loss181 if not 'transformer' in self.opt['network_g']['fix_modules']:182 if self.hq_feat_loss: # codebook loss 183 l_feat_encoder = torch.mean((quant_feat_gt.detach()-lq_feat)**2) * self.feat_loss_weight184 l_g_total += l_feat_encoder185 loss_dict['l_feat_encoder'] = l_feat_encoder186 187 # cross_entropy_loss188 if self.cross_entropy_loss:189 # b(hw)n -> bn(hw)190 cross_entropy_loss = F.cross_entropy(logits.permute(0, 2, 1), self.idx_gt) * self.entropy_loss_weight191 l_g_total += cross_entropy_loss192 loss_dict['cross_entropy_loss'] = cross_entropy_loss193 194 # pixel loss 195 if not large_de: # when large degradation don't need image-level loss196 if self.cri_pix:197 l_g_pix = self.cri_pix(self.output, self.gt)198 l_g_total += l_g_pix199 loss_dict['l_g_pix'] = l_g_pix200 201 # perceptual loss202 if self.cri_perceptual:203 l_g_percep = self.cri_perceptual(self.output, self.gt)204 l_g_total += l_g_percep205 loss_dict['l_g_percep'] = l_g_percep206 207 # gan loss208 if current_iter > self.net_d_start_iter:209 fake_g_pred = self.net_d(self.output)210 l_g_gan = self.cri_gan(fake_g_pred, True, is_disc=False)211 recon_loss = l_g_pix + l_g_percep212 if not self.fix_generator:213 last_layer = self.net_g.module.generator.blocks[-1].weight214 d_weight = self.calculate_adaptive_weight(recon_loss, l_g_gan, last_layer, disc_weight_max=1.0)215 else:216 largest_fuse_size = self.opt['network_g']['connect_list'][-1]217 last_layer = self.net_g.module.fuse_convs_dict[largest_fuse_size].shift[-1].weight218 d_weight = self.calculate_adaptive_weight(recon_loss, l_g_gan, last_layer, disc_weight_max=1.0)219 220 d_weight *= self.scale_adaptive_gan_weight # 0.8221 loss_dict['d_weight'] = d_weight222 l_g_total += d_weight * l_g_gan223 loss_dict['l_g_gan'] = d_weight * l_g_gan224 225 l_g_total.backward()226 self.optimizer_g.step()227 228 if self.ema_decay > 0:229 self.model_ema(decay=self.ema_decay)230 231 # optimize net_d232 if not large_de:233 if current_iter > self.net_d_start_iter:234 for p in self.net_d.parameters():235 p.requires_grad = True236 237 self.optimizer_d.zero_grad()238 # real239 real_d_pred = self.net_d(self.gt)240 l_d_real = self.cri_gan(real_d_pred, True, is_disc=True)241 loss_dict['l_d_real'] = l_d_real242 loss_dict['out_d_real'] = torch.mean(real_d_pred.detach())243 l_d_real.backward()244 # fake245 fake_d_pred = self.net_d(self.output.detach())246 l_d_fake = self.cri_gan(fake_d_pred, False, is_disc=True)247 loss_dict['l_d_fake'] = l_d_fake248 loss_dict['out_d_fake'] = torch.mean(fake_d_pred.detach())249 l_d_fake.backward()250 251 self.optimizer_d.step()252 253 self.log_dict = self.reduce_loss_dict(loss_dict)254 255 256 def test(self):257 with torch.no_grad():258 if hasattr(self, 'net_g_ema'):259 self.net_g_ema.eval()260 self.output, _, _ = self.net_g_ema(self.input, w=1)261 else:262 logger = get_root_logger()263 logger.warning('Do not have self.net_g_ema, use self.net_g.')264 self.net_g.eval()265 self.output, _, _ = self.net_g(self.input, w=1)266 self.net_g.train()267 268 269 def dist_validation(self, dataloader, current_iter, tb_logger, save_img):270 if self.opt['rank'] == 0:271 self.nondist_validation(dataloader, current_iter, tb_logger, save_img)272 273 274 def nondist_validation(self, dataloader, current_iter, tb_logger, save_img):275 dataset_name = dataloader.dataset.opt['name']276 with_metrics = self.opt['val'].get('metrics') is not None277 if with_metrics:278 self.metric_results = {metric: 0 for metric in self.opt['val']['metrics'].keys()}279 pbar = tqdm(total=len(dataloader), unit='image')280 281 for idx, val_data in enumerate(dataloader):282 img_name = osp.splitext(osp.basename(val_data['lq_path'][0]))[0]283 self.feed_data(val_data)284 self.test()285 286 visuals = self.get_current_visuals()287 sr_img = tensor2img([visuals['result']])288 if 'gt' in visuals:289 gt_img = tensor2img([visuals['gt']])290 del self.gt291 292 # tentative for out of GPU memory293 del self.lq294 del self.output295 torch.cuda.empty_cache()296 297 if save_img:298 if self.opt['is_train']:299 save_img_path = osp.join(self.opt['path']['visualization'], img_name,300 f'{img_name}_{current_iter}.png')301 else:302 if self.opt['val']['suffix']:303 save_img_path = osp.join(self.opt['path']['visualization'], dataset_name,304 f'{img_name}_{self.opt["val"]["suffix"]}.png')305 else:306 save_img_path = osp.join(self.opt['path']['visualization'], dataset_name,307 f'{img_name}_{self.opt["name"]}.png')308 imwrite(sr_img, save_img_path)309 310 if with_metrics:311 # calculate metrics312 for name, opt_ in self.opt['val']['metrics'].items():313 metric_data = dict(img1=sr_img, img2=gt_img)314 self.metric_results[name] += calculate_metric(metric_data, opt_)315 pbar.update(1)316 pbar.set_description(f'Test {img_name}')317 pbar.close()318 319 if with_metrics:320 for metric in self.metric_results.keys():321 self.metric_results[metric] /= (idx + 1)322 323 self._log_validation_metric_values(current_iter, dataset_name, tb_logger)324 325 326 def _log_validation_metric_values(self, current_iter, dataset_name, tb_logger):327 log_str = f'Validation {dataset_name}\n'328 for metric, value in self.metric_results.items():329 log_str += f'\t # {metric}: {value:.4f}\n'330 logger = get_root_logger()331 logger.info(log_str)332 if tb_logger:333 for metric, value in self.metric_results.items():334 tb_logger.add_scalar(f'metrics/{metric}', value, current_iter)335 336 337 def get_current_visuals(self):338 out_dict = OrderedDict()339 out_dict['gt'] = self.gt.detach().cpu()340 out_dict['result'] = self.output.detach().cpu()341 return out_dict342 343 344 def save(self, epoch, current_iter):345 if self.ema_decay > 0:346 self.save_network([self.net_g, self.net_g_ema], 'net_g', current_iter, param_key=['params', 'params_ema'])347 else:348 self.save_network(self.net_g, 'net_g', current_iter)349 self.save_network(self.net_d, 'net_d', current_iter)350 self.save_training_state(epoch, current_iter)351 