(self, args, train_platform, model, diffusion, data)
| 36 | |
| 37 | class TrainLoop: |
| 38 | def __init__(self, args, train_platform, model, diffusion, data): |
| 39 | self.args = args |
| 40 | self.dataset = args.dataset |
| 41 | self.train_platform = train_platform |
| 42 | self.model = model |
| 43 | self.model_avg = None |
| 44 | if self.args.use_ema: |
| 45 | self.model_avg = copy.deepcopy(self.model) |
| 46 | self.model_for_eval = self.model_avg if self.args.use_ema else self.model |
| 47 | if args.gen_guidance_param != 1: |
| 48 | self.model_for_eval = ClassifierFreeSampleModel(self.model_for_eval) # wrapping model with the classifier-free sampler |
| 49 | self.diffusion = diffusion |
| 50 | self.cond_mode = model.cond_mode |
| 51 | self.data = data |
| 52 | self.batch_size = args.batch_size |
| 53 | self.microbatch = args.batch_size # deprecating this option |
| 54 | self.lr = args.lr |
| 55 | self.log_interval = args.log_interval |
| 56 | self.save_interval = args.save_interval |
| 57 | self.resume_checkpoint = args.resume_checkpoint |
| 58 | self.use_fp16 = False # deprecating this option |
| 59 | self.fp16_scale_growth = 1e-3 # deprecating this option |
| 60 | self.weight_decay = args.weight_decay |
| 61 | self.lr_anneal_steps = args.lr_anneal_steps |
| 62 | |
| 63 | self.step = 0 |
| 64 | self.resume_step = 0 |
| 65 | self.global_batch = self.batch_size # * dist.get_world_size() |
| 66 | self.num_steps = args.num_steps |
| 67 | self.num_epochs = self.num_steps // len(self.data) + 1 |
| 68 | |
| 69 | self.sync_cuda = torch.cuda.is_available() |
| 70 | |
| 71 | self._load_and_sync_parameters() |
| 72 | self.mp_trainer = MixedPrecisionTrainer( |
| 73 | model=self.model, |
| 74 | use_fp16=self.use_fp16, |
| 75 | fp16_scale_growth=self.fp16_scale_growth, |
| 76 | ) |
| 77 | |
| 78 | self.save_dir = args.save_dir |
| 79 | self.overwrite = args.overwrite |
| 80 | |
| 81 | if self.args.use_ema: |
| 82 | self.opt = AdamW( |
| 83 | # with amp, we don't need to use the mp_trainer's master_params |
| 84 | (self.model.parameters() |
| 85 | if self.use_fp16 else self.mp_trainer.master_params), |
| 86 | lr=self.lr, |
| 87 | weight_decay=self.weight_decay, |
| 88 | betas=(0.9, self.args.adam_beta2), |
| 89 | ) |
| 90 | else: |
| 91 | self.opt = AdamW( |
| 92 | self.mp_trainer.master_params, lr=self.lr, weight_decay=self.weight_decay |
| 93 | ) |
| 94 | |
| 95 | if self.resume_step: |
nothing calls this directly
no test coverage detected