sneedium/captcha_pixelplanet
1
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 