CoolFace
Apppublic

emilios/codeformer-face-restorization

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
sr_model.py210 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_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