MCPcopy Create free account
hub / github.com/cientgu/VQ-Diffusion / __init__

Method __init__

image_synthesis/engine/solver.py:37–110  ·  view source on GitHub ↗
(self, config, args, model, dataloader, logger)

Source from the content-addressed store, hash-verified

35
36class 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:

Callers

nothing calls this directly

Calls 6

instantiate_from_configFunction · 0.90
EMAClass · 0.90
log_infoMethod · 0.80
getMethod · 0.45

Tested by

no test coverage detected