diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index 8a49a69..e8c1e9a 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -603,34 +603,37 @@ class Trainer(object): self.scaler.load_state_dict(data['scaler']) def train(self): - while self.step < self.train_num_steps: - for i in range(self.gradient_accumulate_every): - data = next(self.dl).cuda() + with tqdm(initial = self.step, total = self.train_num_steps) as pbar: - with autocast(enabled = self.amp): - loss = self.model(data) - self.scaler.scale(loss / self.gradient_accumulate_every).backward() + while self.step < self.train_num_steps: + for i in range(self.gradient_accumulate_every): + data = next(self.dl).cuda() - print(f'{self.step}: {loss.item()}') + with autocast(enabled = self.amp): + loss = self.model(data) + self.scaler.scale(loss / self.gradient_accumulate_every).backward() - self.scaler.step(self.opt) - self.scaler.update() - self.opt.zero_grad() + pbar.set_description(f'loss: {loss.item():.4f}') - if self.step % self.update_ema_every == 0: - self.step_ema() + self.scaler.step(self.opt) + self.scaler.update() + self.opt.zero_grad() - if self.step != 0 and self.step % self.save_and_sample_every == 0: - self.ema_model.eval() + if self.step % self.update_ema_every == 0: + self.step_ema() - milestone = self.step // self.save_and_sample_every - batches = num_to_groups(36, self.batch_size) - all_images_list = list(map(lambda n: self.ema_model.sample(batch_size=n), batches)) - all_images = torch.cat(all_images_list, dim=0) - all_images = unnormalize_to_zero_to_one(all_images) - utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6) - self.save(milestone) + if self.step != 0 and self.step % self.save_and_sample_every == 0: + self.ema_model.eval() - self.step += 1 + milestone = self.step // self.save_and_sample_every + batches = num_to_groups(36, self.batch_size) + all_images_list = list(map(lambda n: self.ema_model.sample(batch_size=n), batches)) + all_images = torch.cat(all_images_list, dim=0) + all_images = unnormalize_to_zero_to_one(all_images) + utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6) + self.save(milestone) - print('training completed') + self.step += 1 + pbar.update(1) + + print('training complete') diff --git a/setup.py b/setup.py index a507209..f9605bf 100644 --- a/setup.py +++ b/setup.py @@ -3,7 +3,7 @@ from setuptools import setup, find_packages setup( name = 'denoising-diffusion-pytorch', packages = find_packages(), - version = '0.15.2', + version = '0.15.3', license='MIT', description = 'Denoising Diffusion Probabilistic Models - Pytorch', author = 'Phil Wang',