CoolFace
Apppublic

emilios/codeformer-face-restorization

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
codeformer_model.py333 linesDownload Raw Back to models
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