Compare commits

..
7 Commits
Author SHA1 Message Date
Phil Wang 6e8a0f2082 fix auto-conversion of images to mode in dataset 2022-09-20 19:29:35 -07:00
Phil Wang 8c36559295 0.27.10 2022-09-16 17:15:02 -07:00
Phil Wang f74f536339 Merge pull request #90 from kashif/patch-1
fix torch.cumprod
2022-09-16 17:14:48 -07:00
Kashif Rasul d85b8bbe2e fix torch.cumprod 2022-09-16 17:19:00 +02:00
Phil Wang e0a1bed31a 0.27.9 2022-09-05 02:43:11 -07:00
Phil Wang 82b67fc00a Merge pull request #85 from RyannDaGreat/main
Trainer.load can use GPU's other than cuda:0
2022-09-05 02:42:49 -07:00
Ryan Burgert 7c0cd05c27 Trainer.load can use GPU's other than cuda:0 2022-09-04 22:11:35 -04:00
2 changed files with 8 additions and 5 deletions
@@ -56,7 +56,7 @@ def num_to_groups(num, divisor):
arr.append(remainder)
return arr
def convert_image_to(img_type, image):
def convert_image_to_fn(img_type, image):
if image.mode != img_type:
return image.convert(img_type)
return image
@@ -450,7 +450,7 @@ class GaussianDiffusion(nn.Module):
raise ValueError(f'unknown beta schedule {beta_schedule}')
alphas = 1. - betas
alphas_cumprod = torch.cumprod(alphas, axis=0)
alphas_cumprod = torch.cumprod(alphas, dim=0)
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value = 1.)
timesteps, = betas.shape
@@ -704,7 +704,7 @@ class Dataset(Dataset):
self.image_size = image_size
self.paths = [p for ext in exts for p in Path(f'{folder}').glob(f'**/*.{ext}')]
maybe_convert_fn = partial(convert_image_to, convert_image_to) if exists(convert_image_to) else nn.Identity()
maybe_convert_fn = partial(convert_image_to_fn, convert_image_to) if exists(convert_image_to) else nn.Identity()
self.transform = T.Compose([
T.Lambda(maybe_convert_fn),
@@ -810,7 +810,10 @@ class Trainer(object):
torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
def load(self, milestone):
data = torch.load(str(self.results_folder / f'model-{milestone}.pt'))
accelerator = self.accelerator
device = accelerator.device
data = torch.load(str(self.results_folder / f'model-{milestone}.pt'), map_location=device)
model = self.accelerator.unwrap_model(self.model)
model.load_state_dict(data['model'])
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup(
name = 'denoising-diffusion-pytorch',
packages = find_packages(),
version = '0.27.8',
version = '0.27.11',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',