From a291da50986d1c3d5193bca869dd38e586ca00ac Mon Sep 17 00:00:00 2001 From: Phil Wang Date: Fri, 27 May 2022 19:13:05 -0700 Subject: [PATCH] bring back linear noise schedule, but default to cosine --- .../denoising_diffusion_pytorch.py | 16 ++++++++++++++-- setup.py | 2 +- 2 files changed, 15 insertions(+), 3 deletions(-) diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index 8477c3a..21202a3 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -329,6 +329,12 @@ def noise_like(shape, device, repeat=False): noise = lambda: torch.randn(shape, device=device) return repeat_noise() if repeat else noise() +def linear_beta_schedule(timesteps): + scale = 1000 / timesteps + beta_start = scale * 0.0001 + beta_end = scale * 0.02 + return torch.linspace(beta_start, beta_end, timesteps, dtype = torch.float64) + def cosine_beta_schedule(timesteps, s = 0.008): """ cosine schedule @@ -350,7 +356,8 @@ class GaussianDiffusion(nn.Module): channels = 3, timesteps = 1000, loss_type = 'l1', - objective = 'pred_noise' + objective = 'pred_noise', + beta_schedule = 'cosine' ): super().__init__() assert not (type(self) == GaussianDiffusion and denoise_fn.channels != denoise_fn.out_dim) @@ -360,7 +367,12 @@ class GaussianDiffusion(nn.Module): self.denoise_fn = denoise_fn self.objective = objective - betas = cosine_beta_schedule(timesteps) + if beta_schedule == 'linear': + betas = linear_beta_schedule(timesteps) + elif beta_schedule == 'cosine': + betas = cosine_beta_schedule(timesteps) + else: + raise ValueError(f'unknown beta schedule {beta_schedule}') alphas = 1. - betas alphas_cumprod = torch.cumprod(alphas, axis=0) diff --git a/setup.py b/setup.py index 8f4653d..926fb66 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.16.0', + version = '0.16.1', license='MIT', description = 'Denoising Diffusion Probabilistic Models - Pytorch', author = 'Phil Wang',