Compare commits

...
4 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
2 changed files with 4 additions and 4 deletions
@@ -56,7 +56,7 @@ def num_to_groups(num, divisor):
arr.append(remainder) arr.append(remainder)
return arr return arr
def convert_image_to(img_type, image): def convert_image_to_fn(img_type, image):
if image.mode != img_type: if image.mode != img_type:
return image.convert(img_type) return image.convert(img_type)
return image return image
@@ -450,7 +450,7 @@ class GaussianDiffusion(nn.Module):
raise ValueError(f'unknown beta schedule {beta_schedule}') raise ValueError(f'unknown beta schedule {beta_schedule}')
alphas = 1. - betas 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.) alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value = 1.)
timesteps, = betas.shape timesteps, = betas.shape
@@ -704,7 +704,7 @@ class Dataset(Dataset):
self.image_size = image_size self.image_size = image_size
self.paths = [p for ext in exts for p in Path(f'{folder}').glob(f'**/*.{ext}')] 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([ self.transform = T.Compose([
T.Lambda(maybe_convert_fn), T.Lambda(maybe_convert_fn),
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup( setup(
name = 'denoising-diffusion-pytorch', name = 'denoising-diffusion-pytorch',
packages = find_packages(), packages = find_packages(),
version = '0.27.9', version = '0.27.11',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',