Compare commits

..
1 Commits
3 changed files with 8 additions and 72 deletions
-2
View File
@@ -6,8 +6,6 @@ Implementation of <a href="https://arxiv.org/abs/2006.11239">Denoising Diffusion
This implementation was transcribed from the official Tensorflow version <a href="https://github.com/hojonathanho/diffusion">here</a>
<a href="https://huggingface.co/blog/annotated-diffusion">Annotated code</a> by Research Scientists / Engineers from <a href="https://huggingface.co/">🤗 Huggingface</a>
<img src="./sample.png" width="500px"><img>
[![PyPI version](https://badge.fury.io/py/denoising-diffusion-pytorch.svg)](https://badge.fury.io/py/denoising-diffusion-pytorch)
@@ -6,7 +6,6 @@ from torch.special import expm1
from tqdm import tqdm
from einops import rearrange, repeat
from einops.layers.torch import Rearrange
# helpers
@@ -34,24 +33,6 @@ def right_pad_dims_to(x, t):
return t
return t.view(*t.shape, *((1,) * padding_dims))
# neural net helpers
class Residual(nn.Module):
def __init__(self, fn):
super().__init__()
self.fn = fn
def forward(self, x):
return x + self.fn(x)
class MonotonicLinear(nn.Module):
def __init__(self, *args, **kwargs):
super().__init__()
self.net = nn.Linear(*args, **kwargs)
def forward(self, x):
return F.linear(x, self.net.weight.abs(), self.net.bias.abs())
# continuous schedules
# equations are taken from https://openreview.net/attachment?id=2LdBqxc1Yv&name=supplementary_material
@@ -66,44 +47,10 @@ def alpha_cosine_log_snr(t):
raise NotImplementedError
class learned_noise_schedule(nn.Module):
""" described in section H and then I.2 of the supplementary material for variational ddpm paper """
def __init__(
self,
*,
log_snr_max,
log_snr_min,
hidden_dim = 1024,
frac_gradient = 1.
):
def __init__(self):
super().__init__()
self.slope = log_snr_min - log_snr_max
self.intercept = log_snr_max
self.net = nn.Sequential(
Rearrange('... -> ... 1'),
MonotonicLinear(1, 1),
Residual(nn.Sequential(
MonotonicLinear(1, hidden_dim),
nn.Sigmoid(),
MonotonicLinear(hidden_dim, 1)
)),
Rearrange('... 1 -> ...'),
)
self.frac_gradient = frac_gradient
def forward(self, x):
frac_gradient = self.frac_gradient
device = x.device
out_zero = self.net(torch.zeros_like(x))
out_one = self.net(torch.ones_like(x))
x = self.net(x)
normed = self.slope * ((x - out_zero) / (out_one - out_zero)) + self.intercept
return normed * frac_gradient + normed.detach() * (1 - frac_gradient)
raise NotImplementedError
# learned noise schedule, using learned monotonic MLP (weights kept positive) in the paper
class ContinuousTimeGaussianDiffusion(nn.Module):
def __init__(
@@ -112,11 +59,10 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
*,
image_size,
channels = 3,
cond_scale = 500,
loss_type = 'l1',
noise_schedule = 'linear',
num_sample_steps = 500,
learned_schedule_net_hidden_dim = 1024,
learned_noise_schedule_frac_gradient = 1. # between 0 and 1, determines what percentage of gradients go back, so one can update the learned noise schedule more slowly
num_sample_steps = 500
):
super().__init__()
assert not denoise_fn.sinusoidal_cond_mlp
@@ -130,19 +76,11 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
# continuous noise schedule related stuff
self.cond_scale = cond_scale # the log(snr) will be scaled by this value
self.loss_type = loss_type
if noise_schedule == 'linear':
self.log_snr = beta_linear_log_snr
elif noise_schedule == 'learned':
log_snr_max, log_snr_min = [beta_linear_log_snr(torch.tensor([time])).item() for time in (0., 1.)]
self.log_snr = learned_noise_schedule(
log_snr_max = log_snr_max,
log_snr_min = log_snr_min,
hidden_dim = learned_schedule_net_hidden_dim,
frac_gradient = learned_noise_schedule_frac_gradient
)
else:
raise ValueError(f'unknown noise schedule {noise_schedule}')
@@ -241,7 +179,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
x, log_snr = self.q_sample(x_start = x_start, times = times, noise = noise)
model_out = self.denoise_fn(x, log_snr)
model_out = self.denoise_fn(x, log_snr * self.cond_scale)
return self.loss_fn(model_out, noise)
def forward(self, img, *args, **kwargs):
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup(
name = 'denoising-diffusion-pytorch',
packages = find_packages(),
version = '0.17.4',
version = '0.16.6',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',