CoolFace
Apppublic

emilios/codeformer-face-restorization

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
base_model.py323 linesDownload Raw Back to models
1import logging2import os3import torch4from collections import OrderedDict5from copy import deepcopy6from torch.nn.parallel import DataParallel, DistributedDataParallel7 8from basicsr.models import lr_scheduler as lr_scheduler9from basicsr.utils.dist_util import master_only10 11logger = logging.getLogger('basicsr')12 13 14class BaseModel():15    """Base model."""16 17    def __init__(self, opt):18        self.opt = opt19        self.device = torch.device('cuda' if opt['num_gpu'] != 0 else 'cpu')20        self.is_train = opt['is_train']21        self.schedulers = []22        self.optimizers = []23 24    def feed_data(self, data):25        pass26 27    def optimize_parameters(self):28        pass29 30    def get_current_visuals(self):31        pass32 33    def save(self, epoch, current_iter):34        """Save networks and training state."""35        pass36 37    def validation(self, dataloader, current_iter, tb_logger, save_img=False):38        """Validation function.39 40        Args:41            dataloader (torch.utils.data.DataLoader): Validation dataloader.42            current_iter (int): Current iteration.43            tb_logger (tensorboard logger): Tensorboard logger.44            save_img (bool): Whether to save images. Default: False.45        """46        if self.opt['dist']:47            self.dist_validation(dataloader, current_iter, tb_logger, save_img)48        else:49            self.nondist_validation(dataloader, current_iter, tb_logger, save_img)50 51    def model_ema(self, decay=0.999):52        net_g = self.get_bare_model(self.net_g)53 54        net_g_params = dict(net_g.named_parameters())55        net_g_ema_params = dict(self.net_g_ema.named_parameters())56 57        for k in net_g_ema_params.keys():58            net_g_ema_params[k].data.mul_(decay).add_(net_g_params[k].data, alpha=1 - decay)59 60    def get_current_log(self):61        return self.log_dict62 63    def model_to_device(self, net):64        """Model to device. It also warps models with DistributedDataParallel65        or DataParallel.66 67        Args:68            net (nn.Module)69        """70        net = net.to(self.device)71        if self.opt['dist']:72            find_unused_parameters = self.opt.get('find_unused_parameters', False)73            net = DistributedDataParallel(74                net, device_ids=[torch.cuda.current_device()], find_unused_parameters=find_unused_parameters)75        elif self.opt['num_gpu'] > 1:76            net = DataParallel(net)77        return net78 79    def get_optimizer(self, optim_type, params, lr, **kwargs):80        if optim_type == 'Adam':81            optimizer = torch.optim.Adam(params, lr, **kwargs)82        else:83            raise NotImplementedError(f'optimizer {optim_type} is not supperted yet.')84        return optimizer85 86    def setup_schedulers(self):87        """Set up schedulers."""88        train_opt = self.opt['train']89        scheduler_type = train_opt['scheduler'].pop('type')90        if scheduler_type in ['MultiStepLR', 'MultiStepRestartLR']:91            for optimizer in self.optimizers:92                self.schedulers.append(lr_scheduler.MultiStepRestartLR(optimizer, **train_opt['scheduler']))93        elif scheduler_type == 'CosineAnnealingRestartLR':94            for optimizer in self.optimizers:95                self.schedulers.append(lr_scheduler.CosineAnnealingRestartLR(optimizer, **train_opt['scheduler']))96        else:97            raise NotImplementedError(f'Scheduler {scheduler_type} is not implemented yet.')98 99    def get_bare_model(self, net):100        """Get bare model, especially under wrapping with101        DistributedDataParallel or DataParallel.102        """103        if isinstance(net, (DataParallel, DistributedDataParallel)):104            net = net.module105        return net106 107    @master_only108    def print_network(self, net):109        """Print the str and parameter number of a network.110 111        Args:112            net (nn.Module)113        """114        if isinstance(net, (DataParallel, DistributedDataParallel)):115            net_cls_str = (f'{net.__class__.__name__} - ' f'{net.module.__class__.__name__}')116        else:117            net_cls_str = f'{net.__class__.__name__}'118 119        net = self.get_bare_model(net)120        net_str = str(net)121        net_params = sum(map(lambda x: x.numel(), net.parameters()))122 123        logger.info(f'Network: {net_cls_str}, with parameters: {net_params:,d}')124        logger.info(net_str)125 126    def _set_lr(self, lr_groups_l):127        """Set learning rate for warmup.128 129        Args:130            lr_groups_l (list): List for lr_groups, each for an optimizer.131        """132        for optimizer, lr_groups in zip(self.optimizers, lr_groups_l):133            for param_group, lr in zip(optimizer.param_groups, lr_groups):134                param_group['lr'] = lr135 136    def _get_init_lr(self):137        """Get the initial lr, which is set by the scheduler.138        """139        init_lr_groups_l = []140        for optimizer in self.optimizers:141            init_lr_groups_l.append([v['initial_lr'] for v in optimizer.param_groups])142        return init_lr_groups_l143 144    def update_learning_rate(self, current_iter, warmup_iter=-1):145        """Update learning rate.146 147        Args:148            current_iter (int): Current iteration.149            warmup_iter (int): Warmup iter numbers. -1 for no warmup.150                Default: -1.151        """152        if current_iter > 1:153            for scheduler in self.schedulers:154                scheduler.step()155        # set up warm-up learning rate156        if current_iter < warmup_iter:157            # get initial lr for each group158            init_lr_g_l = self._get_init_lr()159            # modify warming-up learning rates160            # currently only support linearly warm up161            warm_up_lr_l = []162            for init_lr_g in init_lr_g_l:163                warm_up_lr_l.append([v / warmup_iter * current_iter for v in init_lr_g])164            # set learning rate165            self._set_lr(warm_up_lr_l)166 167    def get_current_learning_rate(self):168        return [param_group['lr'] for param_group in self.optimizers[0].param_groups]169 170    @master_only171    def save_network(self, net, net_label, current_iter, param_key='params'):172        """Save networks.173 174        Args:175            net (nn.Module | list[nn.Module]): Network(s) to be saved.176            net_label (str): Network label.177            current_iter (int): Current iter number.178            param_key (str | list[str]): The parameter key(s) to save network.179                Default: 'params'.180        """181        if current_iter == -1:182            current_iter = 'latest'183        save_filename = f'{net_label}_{current_iter}.pth'184        save_path = os.path.join(self.opt['path']['models'], save_filename)185 186        net = net if isinstance(net, list) else [net]187        param_key = param_key if isinstance(param_key, list) else [param_key]188        assert len(net) == len(param_key), 'The lengths of net and param_key should be the same.'189 190        save_dict = {}191        for net_, param_key_ in zip(net, param_key):192            net_ = self.get_bare_model(net_)193            state_dict = net_.state_dict()194            for key, param in state_dict.items():195                if key.startswith('module.'):  # remove unnecessary 'module.'196                    key = key[7:]197                state_dict[key] = param.cpu()198            save_dict[param_key_] = state_dict199 200        torch.save(save_dict, save_path)201 202    def _print_different_keys_loading(self, crt_net, load_net, strict=True):203        """Print keys with differnet name or different size when loading models.204 205        1. Print keys with differnet names.206        2. If strict=False, print the same key but with different tensor size.207            It also ignore these keys with different sizes (not load).208 209        Args:210            crt_net (torch model): Current network.211            load_net (dict): Loaded network.212            strict (bool): Whether strictly loaded. Default: True.213        """214        crt_net = self.get_bare_model(crt_net)215        crt_net = crt_net.state_dict()216        crt_net_keys = set(crt_net.keys())217        load_net_keys = set(load_net.keys())218 219        if crt_net_keys != load_net_keys:220            logger.warning('Current net - loaded net:')221            for v in sorted(list(crt_net_keys - load_net_keys)):222                logger.warning(f'  {v}')223            logger.warning('Loaded net - current net:')224            for v in sorted(list(load_net_keys - crt_net_keys)):225                logger.warning(f'  {v}')226 227        # check the size for the same keys228        if not strict:229            common_keys = crt_net_keys & load_net_keys230            for k in common_keys:231                if crt_net[k].size() != load_net[k].size():232                    logger.warning(f'Size different, ignore [{k}]: crt_net: '233                                   f'{crt_net[k].shape}; load_net: {load_net[k].shape}')234                    load_net[k + '.ignore'] = load_net.pop(k)235 236    def load_network(self, net, load_path, strict=True, param_key='params'):237        """Load network.238 239        Args:240            load_path (str): The path of networks to be loaded.241            net (nn.Module): Network.242            strict (bool): Whether strictly loaded.243            param_key (str): The parameter key of loaded network. If set to244                None, use the root 'path'.245                Default: 'params'.246        """247        net = self.get_bare_model(net)248        logger.info(f'Loading {net.__class__.__name__} model from {load_path}.')249        load_net = torch.load(load_path, map_location=lambda storage, loc: storage)250        if param_key is not None:251            if param_key not in load_net and 'params' in load_net:252                param_key = 'params'253                logger.info('Loading: params_ema does not exist, use params.')254            load_net = load_net[param_key]255        # remove unnecessary 'module.'256        for k, v in deepcopy(load_net).items():257            if k.startswith('module.'):258                load_net[k[7:]] = v259                load_net.pop(k)260        self._print_different_keys_loading(net, load_net, strict)261        net.load_state_dict(load_net, strict=strict)262 263    @master_only264    def save_training_state(self, epoch, current_iter):265        """Save training states during training, which will be used for266        resuming.267 268        Args:269            epoch (int): Current epoch.270            current_iter (int): Current iteration.271        """272        if current_iter != -1:273            state = {'epoch': epoch, 'iter': current_iter, 'optimizers': [], 'schedulers': []}274            for o in self.optimizers:275                state['optimizers'].append(o.state_dict())276            for s in self.schedulers:277                state['schedulers'].append(s.state_dict())278            save_filename = f'{current_iter}.state'279            save_path = os.path.join(self.opt['path']['training_states'], save_filename)280            torch.save(state, save_path)281 282    def resume_training(self, resume_state):283        """Reload the optimizers and schedulers for resumed training.284 285        Args:286            resume_state (dict): Resume state.287        """288        resume_optimizers = resume_state['optimizers']289        resume_schedulers = resume_state['schedulers']290        assert len(resume_optimizers) == len(self.optimizers), 'Wrong lengths of optimizers'291        assert len(resume_schedulers) == len(self.schedulers), 'Wrong lengths of schedulers'292        for i, o in enumerate(resume_optimizers):293            self.optimizers[i].load_state_dict(o)294        for i, s in enumerate(resume_schedulers):295            self.schedulers[i].load_state_dict(s)296 297    def reduce_loss_dict(self, loss_dict):298        """reduce loss dict.299 300        In distributed training, it averages the losses among different GPUs .301 302        Args:303            loss_dict (OrderedDict): Loss dict.304        """305        with torch.no_grad():306            if self.opt['dist']:307                keys = []308                losses = []309                for name, value in loss_dict.items():310                    keys.append(name)311                    losses.append(value)312                losses = torch.stack(losses, 0)313                torch.distributed.reduce(losses, dst=0)314                if self.opt['rank'] == 0:315                    losses /= self.opt['world_size']316                loss_dict = {key: loss for key, loss in zip(keys, losses)}317 318            log_dict = OrderedDict()319            for name, value in loss_dict.items():320                log_dict[name] = value.mean().item()321 322            return log_dict323