CoolFace
Apppublic

emilios/codeformer-face-restorization

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