CoolFace
Apppublic

sneedium/captcha_pixelplanet

sourceHugging Facebsdupdated 4y agoView on Hugging Face
1likes
main.py247 linesDownload Raw Back to root
1import argparse2import logging3import os4import random5 6import torch7from fastai.callbacks.general_sched import GeneralScheduler, TrainingPhase8from fastai.distributed import *9from fastai.vision import *10from torch.backends import cudnn11 12from callbacks import DumpPrediction, IterationCallback, TextAccuracy, TopKTextAccuracy13from dataset import ImageDataset, TextDataset14from losses import MultiLosses15from utils import Config, Logger, MyDataParallel, MyConcatDataset16 17 18def _set_random_seed(seed):19    if seed is not None:20        random.seed(seed)21        torch.manual_seed(seed)22        cudnn.deterministic = True23        logging.warning('You have chosen to seed training. '24                        'This will slow down your training!')25 26def _get_training_phases(config, n):27    lr = np.array(config.optimizer_lr)28    periods = config.optimizer_scheduler_periods29    sigma = [config.optimizer_scheduler_gamma ** i for i in range(len(periods))]30    phases = [TrainingPhase(n * periods[i]).schedule_hp('lr', lr * sigma[i])31                for i in range(len(periods))]32    return phases33 34def _get_dataset(ds_type, paths, is_training, config, **kwargs):35    kwargs.update({36        'img_h': config.dataset_image_height,37        'img_w': config.dataset_image_width,38        'max_length': config.dataset_max_length,39        'case_sensitive': config.dataset_case_sensitive,40        'charset_path': config.dataset_charset_path,41        'data_aug': config.dataset_data_aug,42        'deteriorate_ratio': config.dataset_deteriorate_ratio,43        'is_training': is_training,44        'multiscales': config.dataset_multiscales,45        'one_hot_y': config.dataset_one_hot_y,46    })47    datasets = [ds_type(p, **kwargs) for p in paths]48    if len(datasets) > 1: return MyConcatDataset(datasets)49    else: return datasets[0]50 51 52def _get_language_databaunch(config):53    kwargs = {54        'max_length': config.dataset_max_length,55        'case_sensitive': config.dataset_case_sensitive,56        'charset_path': config.dataset_charset_path,57        'smooth_label': config.dataset_smooth_label,58        'smooth_factor': config.dataset_smooth_factor,59        'one_hot_y': config.dataset_one_hot_y,60        'use_sm': config.dataset_use_sm,61    }62    train_ds = TextDataset(config.dataset_train_roots[0], is_training=True, **kwargs)63    valid_ds = TextDataset(config.dataset_test_roots[0], is_training=False, **kwargs)64    data = DataBunch.create(65        path=train_ds.path,66        train_ds=train_ds,67        valid_ds=valid_ds,68        bs=config.dataset_train_batch_size,69        val_bs=config.dataset_test_batch_size,70        num_workers=config.dataset_num_workers,71        pin_memory=config.dataset_pin_memory)72    logging.info(f'{len(data.train_ds)} training items found.')73    if not data.empty_val:74        logging.info(f'{len(data.valid_ds)} valid items found.')75    return data76 77def _get_databaunch(config):78    # An awkward way to reduce loadding data time during test79    if config.global_phase == 'test': config.dataset_train_roots = config.dataset_test_roots80    train_ds = _get_dataset(ImageDataset, config.dataset_train_roots, True, config)81    valid_ds = _get_dataset(ImageDataset, config.dataset_test_roots, False, config)82    data = ImageDataBunch.create(83        train_ds=train_ds,84        valid_ds=valid_ds,85        bs=config.dataset_train_batch_size,86        val_bs=config.dataset_test_batch_size,87        num_workers=config.dataset_num_workers,88        pin_memory=config.dataset_pin_memory).normalize(imagenet_stats)89    ar_tfm = lambda x: ((x[0], x[1]), x[1])  # auto-regression only for dtd90    data.add_tfm(ar_tfm)91 92    logging.info(f'{len(data.train_ds)} training items found.')93    if not data.empty_val:94        logging.info(f'{len(data.valid_ds)} valid items found.')95    96    return data97 98def _get_model(config):99    import importlib100    names = config.model_name.split('.')101    module_name, class_name = '.'.join(names[:-1]), names[-1]102    cls = getattr(importlib.import_module(module_name), class_name)103    model = cls(config)104    logging.info(model)105    return model106 107 108def _get_learner(config, data, model, local_rank=None):109    strict = ifnone(config.model_strict, True)110    if config.global_stage == 'pretrain-language':111        metrics = [TopKTextAccuracy(112            k=ifnone(config.model_k, 5),113            charset_path=config.dataset_charset_path,114            max_length=config.dataset_max_length + 1,115            case_sensitive=config.dataset_eval_case_sensisitves,116            model_eval=config.model_eval)] 117    else:118        metrics = [TextAccuracy(119            charset_path=config.dataset_charset_path,120            max_length=config.dataset_max_length + 1,121            case_sensitive=config.dataset_eval_case_sensisitves,122            model_eval=config.model_eval)]123    opt_type = getattr(torch.optim, config.optimizer_type)124    learner = Learner(data, model, silent=True, model_dir='.',125        true_wd=config.optimizer_true_wd, 126        wd=config.optimizer_wd,127        bn_wd=config.optimizer_bn_wd,128        path=config.global_workdir,129        metrics=metrics,130        opt_func=partial(opt_type, **config.optimizer_args or dict()), 131        loss_func=MultiLosses(one_hot=config.dataset_one_hot_y))132    learner.split(lambda m: children(m))133 134    if config.global_phase == 'train':135        num_replicas = 1 if local_rank is None else torch.distributed.get_world_size()136        phases = _get_training_phases(config, len(learner.data.train_dl)//num_replicas)137        learner.callback_fns += [138            partial(GeneralScheduler, phases=phases),139            partial(GradientClipping, clip=config.optimizer_clip_grad),140            partial(IterationCallback, name=config.global_name,141                    show_iters=config.training_show_iters,142                    eval_iters=config.training_eval_iters,143                    save_iters=config.training_save_iters,144                    start_iters=config.training_start_iters,145                    stats_iters=config.training_stats_iters)]146    else:147        learner.callbacks += [148            DumpPrediction(learn=learner,149                    dataset='-'.join([Path(p).name for p in config.dataset_test_roots]),charset_path=config.dataset_charset_path,150                    model_eval=config.model_eval,151                    debug=config.global_debug,152                    image_only=config.global_image_only)]153 154    learner.rank = local_rank155    if local_rank is not None:156        logging.info(f'Set model to distributed with rank {local_rank}.')157        learner.model = torch.nn.SyncBatchNorm.convert_sync_batchnorm(learner.model)158        learner.model.to(local_rank)159        learner = learner.to_distributed(local_rank)160 161    if torch.cuda.device_count() > 1 and local_rank is None:162        logging.info(f'Use {torch.cuda.device_count()} GPUs.')163        learner.model = MyDataParallel(learner.model)164 165    if config.model_checkpoint:166        if Path(config.model_checkpoint).exists():167            with open(config.model_checkpoint, 'rb') as f:168                buffer = io.BytesIO(f.read())169            learner.load(buffer, strict=strict)170        else:171            from distutils.dir_util import copy_tree172            src = Path('/data/fangsc/model')/config.global_name173            trg = Path('/output')/config.global_name174            if src.exists(): copy_tree(str(src), str(trg))175            learner.load(config.model_checkpoint, strict=strict)176        logging.info(f'Read model from {config.model_checkpoint}')177    elif config.global_phase == 'test':178        learner.load(f'best-{config.global_name}', strict=strict)179        logging.info(f'Read model from best-{config.global_name}')180 181    if learner.opt_func.func.__name__ == 'Adadelta':    # fastai bug, fix after 1.0.60182        learner.fit(epochs=0, lr=config.optimizer_lr)183        learner.opt.mom = 0.184 185    return learner186 187def main():188    parser = argparse.ArgumentParser()189    parser.add_argument('--config', type=str, required=True,190                        help='path to config file')191    parser.add_argument('--phase', type=str, default=None, choices=['train', 'test'])192    parser.add_argument('--name', type=str, default=None)193    parser.add_argument('--checkpoint', type=str, default=None)194    parser.add_argument('--test_root', type=str, default=None)195    parser.add_argument("--local_rank", type=int, default=None)196    parser.add_argument('--debug', action='store_true', default=None)197    parser.add_argument('--image_only', action='store_true', default=None)198    parser.add_argument('--model_strict', action='store_false', default=None)199    parser.add_argument('--model_eval', type=str, default=None, 200                        choices=['alignment', 'vision', 'language'])201    args = parser.parse_args()202    config = Config(args.config)203    if args.name is not None: config.global_name = args.name204    if args.phase is not None: config.global_phase = args.phase205    if args.test_root is not None: config.dataset_test_roots = [args.test_root]206    if args.checkpoint is not None: config.model_checkpoint = args.checkpoint207    if args.debug is not None: config.global_debug = args.debug208    if args.image_only is not None: config.global_image_only = args.image_only209    if args.model_eval is not None: config.model_eval = args.model_eval210    if args.model_strict is not None: config.model_strict = args.model_strict211 212    Logger.init(config.global_workdir, config.global_name, config.global_phase)213    Logger.enable_file()214    _set_random_seed(config.global_seed)215    logging.info(config)216 217    if args.local_rank is not None:218        logging.info(f'Init distribution training at device {args.local_rank}.')219        torch.cuda.set_device(args.local_rank)220        torch.distributed.init_process_group(backend='nccl', init_method='env://')221 222    logging.info('Construct dataset.')223    if config.global_stage == 'pretrain-language': data = _get_language_databaunch(config)224    else: data = _get_databaunch(config)225 226    logging.info('Construct model.')227    model = _get_model(config)228 229    logging.info('Construct learner.')230    learner = _get_learner(config, data, model, args.local_rank)231 232    if config.global_phase == 'train':233        logging.info('Start training.')234        learner.fit(epochs=config.training_epochs,235                    lr=config.optimizer_lr)236    else:237        logging.info('Start validate')238        last_metrics = learner.validate()239        log_str = f'eval loss = {last_metrics[0]:6.3f},  ' \240                  f'ccr = {last_metrics[1]:6.3f},  cwr = {last_metrics[2]:6.3f},  ' \241                  f'ted = {last_metrics[3]:6.3f},  ned = {last_metrics[4]:6.0f},  ' \242                  f'ted/w = {last_metrics[5]:6.3f}, '243        logging.info(log_str)244 245if __name__ == '__main__':246    main()247