CoolFace
Apppublic

faisalhr1997/codeformer

sourceHugging Faceupdated 4y agoView on Hugging Face
0likes
train.py226 linesDownload Raw Back to basicsr
1import argparse2import datetime3import logging4import math5import copy6import random7import time8import torch9from os import path as osp10 11from basicsr.data import build_dataloader, build_dataset12from basicsr.data.data_sampler import EnlargedSampler13from basicsr.data.prefetch_dataloader import CPUPrefetcher, CUDAPrefetcher14from basicsr.models import build_model15from basicsr.utils import (MessageLogger, check_resume, get_env_info, get_root_logger, init_tb_logger,16                           init_wandb_logger, make_exp_dirs, mkdir_and_rename, set_random_seed)17from basicsr.utils.dist_util import get_dist_info, init_dist18from basicsr.utils.options import dict2str, parse19 20import warnings21# ignore UserWarning: Detected call of `lr_scheduler.step()` before `optimizer.step()`.22warnings.filterwarnings("ignore", category=UserWarning)23 24def parse_options(root_path, is_train=True):25    parser = argparse.ArgumentParser()26    parser.add_argument('-opt', type=str, required=True, help='Path to option YAML file.')27    parser.add_argument('--launcher', choices=['none', 'pytorch', 'slurm'], default='none', help='job launcher')28    parser.add_argument('--local_rank', type=int, default=0)29    args = parser.parse_args()30    opt = parse(args.opt, root_path, is_train=is_train)31 32    # distributed settings33    if args.launcher == 'none':34        opt['dist'] = False35        print('Disable distributed.', flush=True)36    else:37        opt['dist'] = True38        if args.launcher == 'slurm' and 'dist_params' in opt:39            init_dist(args.launcher, **opt['dist_params'])40        else:41            init_dist(args.launcher)42 43    opt['rank'], opt['world_size'] = get_dist_info()44 45    # random seed46    seed = opt.get('manual_seed')47    if seed is None:48        seed = random.randint(1, 10000)49        opt['manual_seed'] = seed50    set_random_seed(seed + opt['rank'])51 52    return opt53 54 55def init_loggers(opt):56    log_file = osp.join(opt['path']['log'], f"train_{opt['name']}.log")57    logger = get_root_logger(logger_name='basicsr', log_level=logging.INFO, log_file=log_file)58    logger.info(get_env_info())59    logger.info(dict2str(opt))60 61    # initialize wandb logger before tensorboard logger to allow proper sync:62    if (opt['logger'].get('wandb') is not None) and (opt['logger']['wandb'].get('project') is not None):63        assert opt['logger'].get('use_tb_logger') is True, ('should turn on tensorboard when using wandb')64        init_wandb_logger(opt)65    tb_logger = None66    if opt['logger'].get('use_tb_logger'):67        tb_logger = init_tb_logger(log_dir=osp.join('tb_logger', opt['name']))68    return logger, tb_logger69 70 71def create_train_val_dataloader(opt, logger):72    # create train and val dataloaders73    train_loader, val_loader = None, None74    for phase, dataset_opt in opt['datasets'].items():75        if phase == 'train':76            dataset_enlarge_ratio = dataset_opt.get('dataset_enlarge_ratio', 1)77            train_set = build_dataset(dataset_opt)78            train_sampler = EnlargedSampler(train_set, opt['world_size'], opt['rank'], dataset_enlarge_ratio)79            train_loader = build_dataloader(80                train_set,81                dataset_opt,82                num_gpu=opt['num_gpu'],83                dist=opt['dist'],84                sampler=train_sampler,85                seed=opt['manual_seed'])86 87            num_iter_per_epoch = math.ceil(88                len(train_set) * dataset_enlarge_ratio / (dataset_opt['batch_size_per_gpu'] * opt['world_size']))89            total_iters = int(opt['train']['total_iter'])90            total_epochs = math.ceil(total_iters / (num_iter_per_epoch))91            logger.info('Training statistics:'92                        f'\n\tNumber of train images: {len(train_set)}'93                        f'\n\tDataset enlarge ratio: {dataset_enlarge_ratio}'94                        f'\n\tBatch size per gpu: {dataset_opt["batch_size_per_gpu"]}'95                        f'\n\tWorld size (gpu number): {opt["world_size"]}'96                        f'\n\tRequire iter number per epoch: {num_iter_per_epoch}'97                        f'\n\tTotal epochs: {total_epochs}; iters: {total_iters}.')98 99        elif phase == 'val':100            val_set = build_dataset(dataset_opt)101            val_loader = build_dataloader(102                val_set, dataset_opt, num_gpu=opt['num_gpu'], dist=opt['dist'], sampler=None, seed=opt['manual_seed'])103            logger.info(f'Number of val images/folders in {dataset_opt["name"]}: ' f'{len(val_set)}')104        else:105            raise ValueError(f'Dataset phase {phase} is not recognized.')106 107    return train_loader, train_sampler, val_loader, total_epochs, total_iters108 109 110def train_pipeline(root_path):111    # parse options, set distributed setting, set ramdom seed112    opt = parse_options(root_path, is_train=True)113 114    torch.backends.cudnn.benchmark = True115    # torch.backends.cudnn.deterministic = True116 117    # load resume states if necessary118    if opt['path'].get('resume_state'):119        device_id = torch.cuda.current_device()120        resume_state = torch.load(121            opt['path']['resume_state'], map_location=lambda storage, loc: storage.cuda(device_id))122    else:123        resume_state = None124 125    # mkdir for experiments and logger126    if resume_state is None:127        make_exp_dirs(opt)128        if opt['logger'].get('use_tb_logger') and opt['rank'] == 0:129            mkdir_and_rename(osp.join('tb_logger', opt['name']))130 131    # initialize loggers132    logger, tb_logger = init_loggers(opt)133    134    # create train and validation dataloaders135    result = create_train_val_dataloader(opt, logger)136    train_loader, train_sampler, val_loader, total_epochs, total_iters = result137 138    # create model139    if resume_state:  # resume training140        check_resume(opt, resume_state['iter'])141        model = build_model(opt)142        model.resume_training(resume_state)  # handle optimizers and schedulers143        logger.info(f"Resuming training from epoch: {resume_state['epoch']}, " f"iter: {resume_state['iter']}.")144        start_epoch = resume_state['epoch']145        current_iter = resume_state['iter']146    else:147        model = build_model(opt)148        start_epoch = 0149        current_iter = 0150 151    # create message logger (formatted outputs)152    msg_logger = MessageLogger(opt, current_iter, tb_logger)153 154    # dataloader prefetcher155    prefetch_mode = opt['datasets']['train'].get('prefetch_mode')156    if prefetch_mode is None or prefetch_mode == 'cpu':157        prefetcher = CPUPrefetcher(train_loader)158    elif prefetch_mode == 'cuda':159        prefetcher = CUDAPrefetcher(train_loader, opt)160        logger.info(f'Use {prefetch_mode} prefetch dataloader')161        if opt['datasets']['train'].get('pin_memory') is not True:162            raise ValueError('Please set pin_memory=True for CUDAPrefetcher.')163    else:164        raise ValueError(f'Wrong prefetch_mode {prefetch_mode}.' "Supported ones are: None, 'cuda', 'cpu'.")165 166    # training167    logger.info(f'Start training from epoch: {start_epoch}, iter: {current_iter+1}')168    data_time, iter_time = time.time(), time.time()169    start_time = time.time()170 171    for epoch in range(start_epoch, total_epochs + 1):172        train_sampler.set_epoch(epoch)173        prefetcher.reset()174        train_data = prefetcher.next()175 176        while train_data is not None:177            data_time = time.time() - data_time178 179            current_iter += 1180            if current_iter > total_iters:181                break182            # update learning rate183            model.update_learning_rate(current_iter, warmup_iter=opt['train'].get('warmup_iter', -1))184            # training185            model.feed_data(train_data)186            model.optimize_parameters(current_iter)187            iter_time = time.time() - iter_time188            # log189            if current_iter % opt['logger']['print_freq'] == 0:190                log_vars = {'epoch': epoch, 'iter': current_iter}191                log_vars.update({'lrs': model.get_current_learning_rate()})192                log_vars.update({'time': iter_time, 'data_time': data_time})193                log_vars.update(model.get_current_log())194                msg_logger(log_vars)195 196            # save models and training states197            if current_iter % opt['logger']['save_checkpoint_freq'] == 0:198                logger.info('Saving models and training states.')199                model.save(epoch, current_iter)200 201            # validation202            if opt.get('val') is not None and opt['datasets'].get('val') is not None \203                and (current_iter % opt['val']['val_freq'] == 0):204                model.validation(val_loader, current_iter, tb_logger, opt['val']['save_img'])205 206            data_time = time.time()207            iter_time = time.time()208            train_data = prefetcher.next()209        # end of iter210 211    # end of epoch212 213    consumed_time = str(datetime.timedelta(seconds=int(time.time() - start_time)))214    logger.info(f'End of training. Time consumed: {consumed_time}')215    logger.info('Save the latest model.')216    model.save(epoch=-1, current_iter=-1)  # -1 stands for the latest217    if opt.get('val') is not None and opt['datasets'].get('val'):218        model.validation(val_loader, current_iter, tb_logger, opt['val']['save_img'])219    if tb_logger:220        tb_logger.close()221 222 223if __name__ == '__main__':224    root_path = osp.abspath(osp.join(__file__, osp.pardir, osp.pardir))225    train_pipeline(root_path)226