From 5ad56dda25f545f01e9dfbf4c2deb2970ceeb625 Mon Sep 17 00:00:00 2001 From: Phil Wang Date: Fri, 9 Oct 2020 21:39:46 -0700 Subject: [PATCH] save samples and models to ./results path --- .../denoising_diffusion_pytorch.py | 9 ++++++--- setup.py | 2 +- 2 files changed, 7 insertions(+), 4 deletions(-) diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index e24d213..83f0caf 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -28,6 +28,9 @@ SAVE_AND_SAMPLE_EVERY = 1000 UPDATE_EMA_EVERY = 10 EXTS = ['jpg', 'jpeg', 'png'] +RESULTS_FOLDER = Path('./results') +RESULTS_FOLDER.mkdir(exist_ok = True) + # helpers functions def exists(x): @@ -497,10 +500,10 @@ class Trainer(object): 'model': self.model.state_dict(), 'ema': self.ema_model.state_dict() } - torch.save(data, f'./model-{milestone}.pt') + torch.save(data, str(RESULTS_FOLDER / f'model-{milestone}.pt')) def load(self, milestone): - data = torch.load(f'./model-{milestone}.pt') + data = torch.load(str(RESULTS_FOLDER / 'model-{milestone}.pt')) self.step = data['step'] self.model.load_state_dict(data['model']) @@ -527,7 +530,7 @@ class Trainer(object): batches = num_to_groups(36, self.batch_size) all_images_list = list(map(lambda n: self.ema_model.sample(self.image_size, batch_size=n), batches)) all_images = torch.cat(all_images_list, dim=0) - utils.save_image(all_images, f'./sample-{milestone}.png', nrow=6) + utils.save_image(all_images, str(RESULTS_FOLDER / 'sample-{milestone}.png'), nrow=6) self.save(milestone) self.step += 1 diff --git a/setup.py b/setup.py index 71860ed..2af3250 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.5.0', + version = '0.5.1', license='MIT', description = 'Denoising Diffusion Probabilistic Models - Pytorch', author = 'Phil Wang',