(self, config, args, model, dataloader, logger)
| 35 | |
| 36 | class Solver(object): |
| 37 | def __init__(self, config, args, model, dataloader, logger): |
| 38 | self.config = config |
| 39 | self.args = args |
| 40 | self.model = model |
| 41 | self.dataloader = dataloader |
| 42 | self.logger = logger |
| 43 | |
| 44 | self.max_epochs = config['solver']['max_epochs'] |
| 45 | self.save_epochs = config['solver']['save_epochs'] |
| 46 | self.save_iterations = config['solver'].get('save_iterations', -1) |
| 47 | self.sample_iterations = config['solver']['sample_iterations'] |
| 48 | if self.sample_iterations == 'epoch': |
| 49 | self.sample_iterations = self.dataloader['train_iterations'] |
| 50 | self.validation_epochs = config['solver'].get('validation_epochs', 2) |
| 51 | assert isinstance(self.save_epochs, (int, list)) |
| 52 | assert isinstance(self.validation_epochs, (int, list)) |
| 53 | self.debug = config['solver'].get('debug', False) |
| 54 | |
| 55 | self.last_epoch = -1 |
| 56 | self.last_iter = -1 |
| 57 | self.ckpt_dir = os.path.join(args.save_dir, 'checkpoint') |
| 58 | self.image_dir = os.path.join(args.save_dir, 'images') |
| 59 | os.makedirs(self.ckpt_dir, exist_ok=True) |
| 60 | os.makedirs(self.image_dir, exist_ok=True) |
| 61 | |
| 62 | # get grad_clipper |
| 63 | if 'clip_grad_norm' in config['solver']: |
| 64 | self.clip_grad_norm = instantiate_from_config(config['solver']['clip_grad_norm']) |
| 65 | else: |
| 66 | self.clip_grad_norm = None |
| 67 | |
| 68 | # get lr |
| 69 | adjust_lr = config['solver'].get('adjust_lr', 'sqrt') |
| 70 | base_lr = config['solver'].get('base_lr', 1.0e-4) |
| 71 | if adjust_lr == 'none': |
| 72 | self.lr = base_lr |
| 73 | elif adjust_lr == 'sqrt': |
| 74 | self.lr = base_lr * math.sqrt(args.world_size * config['dataloader']['batch_size']) |
| 75 | elif adjust_lr == 'linear': |
| 76 | self.lr = base_lr * args.world_size * config['dataloader']['batch_size'] |
| 77 | else: |
| 78 | raise NotImplementedError('Unknown type of adjust lr {}!'.format(adjust_lr)) |
| 79 | self.logger.log_info('Get lr {} from base lr {} with {}'.format(self.lr, base_lr, adjust_lr)) |
| 80 | |
| 81 | if hasattr(model, 'get_optimizer_and_scheduler') and callable(getattr(model, 'get_optimizer_and_scheduler')): |
| 82 | optimizer_and_scheduler = model.get_optimizer_and_scheduler(config['solver']['optimizers_and_schedulers']) |
| 83 | else: |
| 84 | optimizer_and_scheduler = self._get_optimizer_and_scheduler(config['solver']['optimizers_and_schedulers']) |
| 85 | |
| 86 | assert type(optimizer_and_scheduler) == type({}), 'optimizer and schduler should be a dict!' |
| 87 | self.optimizer_and_scheduler = optimizer_and_scheduler |
| 88 | |
| 89 | # configre for ema |
| 90 | if 'ema' in config['solver'] and args.local_rank == 0: |
| 91 | ema_args = config['solver']['ema'] |
| 92 | ema_args['model'] = self.model |
| 93 | self.ema = EMA(**ema_args) |
| 94 | else: |
nothing calls this directly
no test coverage detected