mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-11 12:11:43 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
402b7c26df | ||
|
|
09613a40f3 | ||
|
|
c6966ae95a | ||
|
|
73591cf1ad | ||
|
|
989f0fcb8e | ||
|
|
84731bb03d | ||
|
|
c6ecca555b |
@@ -4,7 +4,7 @@
|
|||||||
|
|
||||||
Implementation of <a href="https://arxiv.org/abs/2006.11239">Denoising Diffusion Probabilistic Model</a> in Pytorch. It is a new approach to generative modeling that may <a href="https://ajolicoeur.wordpress.com/the-new-contender-to-gans-score-matching-with-langevin-sampling/">have the potential</a> to rival GANs. It uses denoising score matching to estimate the gradient of the data distribution, followed by Langevin sampling to sample from the true distribution.
|
Implementation of <a href="https://arxiv.org/abs/2006.11239">Denoising Diffusion Probabilistic Model</a> in Pytorch. It is a new approach to generative modeling that may <a href="https://ajolicoeur.wordpress.com/the-new-contender-to-gans-score-matching-with-langevin-sampling/">have the potential</a> to rival GANs. It uses denoising score matching to estimate the gradient of the data distribution, followed by Langevin sampling to sample from the true distribution.
|
||||||
|
|
||||||
This implementation was transcribed from the official Tensorflow version <a href="https://github.com/hojonathanho/diffusion">here</a> and then modified to use <a href="https://arxiv.org/abs/2201.03545">ConvNext</a> blocks instead of Resnets.
|
This implementation was transcribed from the official Tensorflow version <a href="https://github.com/hojonathanho/diffusion">here</a>
|
||||||
|
|
||||||
<img src="./sample.png" width="500px"><img>
|
<img src="./sample.png" width="500px"><img>
|
||||||
|
|
||||||
@@ -79,34 +79,32 @@ Samples and model checkpoints will be logged to `./results` periodically
|
|||||||
## Citations
|
## Citations
|
||||||
|
|
||||||
```bibtex
|
```bibtex
|
||||||
@misc{ho2020denoising,
|
@inproceedings{NEURIPS2020_4c5bcfec,
|
||||||
title = {Denoising Diffusion Probabilistic Models},
|
author = {Ho, Jonathan and Jain, Ajay and Abbeel, Pieter},
|
||||||
author = {Jonathan Ho and Ajay Jain and Pieter Abbeel},
|
booktitle = {Advances in Neural Information Processing Systems},
|
||||||
year = {2020},
|
editor = {H. Larochelle and M. Ranzato and R. Hadsell and M.F. Balcan and H. Lin},
|
||||||
eprint = {2006.11239},
|
pages = {6840--6851},
|
||||||
archivePrefix = {arXiv},
|
publisher = {Curran Associates, Inc.},
|
||||||
primaryClass = {cs.LG}
|
title = {Denoising Diffusion Probabilistic Models},
|
||||||
|
url = {https://proceedings.neurips.cc/paper/2020/file/4c5bcfec8584af0d967f1ab10179ca4b-Paper.pdf},
|
||||||
|
volume = {33},
|
||||||
|
year = {2020}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
```bibtex
|
```bibtex
|
||||||
@inproceedings{anonymous2021improved,
|
@InProceedings{pmlr-v139-nichol21a,
|
||||||
title = {Improved Denoising Diffusion Probabilistic Models},
|
title = {Improved Denoising Diffusion Probabilistic Models},
|
||||||
author = {Anonymous},
|
author = {Nichol, Alexander Quinn and Dhariwal, Prafulla},
|
||||||
booktitle = {Submitted to International Conference on Learning Representations},
|
booktitle = {Proceedings of the 38th International Conference on Machine Learning},
|
||||||
year = {2021},
|
pages = {8162--8171},
|
||||||
url = {https://openreview.net/forum?id=-NEXDKk8gZ},
|
year = {2021},
|
||||||
note = {under review}
|
editor = {Meila, Marina and Zhang, Tong},
|
||||||
}
|
volume = {139},
|
||||||
```
|
series = {Proceedings of Machine Learning Research},
|
||||||
|
month = {18--24 Jul},
|
||||||
```bibtex
|
publisher = {PMLR},
|
||||||
@misc{liu2022convnet,
|
pdf = {http://proceedings.mlr.press/v139/nichol21a/nichol21a.pdf},
|
||||||
title = {A ConvNet for the 2020s},
|
url = {https://proceedings.mlr.press/v139/nichol21a.html},
|
||||||
author = {Zhuang Liu and Hanzi Mao and Chao-Yuan Wu and Christoph Feichtenhofer and Trevor Darrell and Saining Xie},
|
|
||||||
year = {2022},
|
|
||||||
eprint = {2201.03545},
|
|
||||||
archivePrefix = {arXiv},
|
|
||||||
primaryClass = {cs.CV}
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -128,8 +128,8 @@ class ResnetBlock(nn.Module):
|
|||||||
nn.Linear(time_emb_dim, dim_out)
|
nn.Linear(time_emb_dim, dim_out)
|
||||||
) if exists(time_emb_dim) else None
|
) if exists(time_emb_dim) else None
|
||||||
|
|
||||||
self.block1 = Block(dim, dim_out)
|
self.block1 = Block(dim, dim_out, groups = groups)
|
||||||
self.block2 = Block(dim_out, dim_out)
|
self.block2 = Block(dim_out, dim_out, groups = groups)
|
||||||
self.res_conv = nn.Conv2d(dim, dim_out, 1) if dim != dim_out else nn.Identity()
|
self.res_conv = nn.Conv2d(dim, dim_out, 1) if dim != dim_out else nn.Identity()
|
||||||
|
|
||||||
def forward(self, x, time_emb = None):
|
def forward(self, x, time_emb = None):
|
||||||
@@ -142,39 +142,6 @@ class ResnetBlock(nn.Module):
|
|||||||
h = self.block2(h)
|
h = self.block2(h)
|
||||||
return h + self.res_conv(x)
|
return h + self.res_conv(x)
|
||||||
|
|
||||||
class ConvNextBlock(nn.Module):
|
|
||||||
""" https://arxiv.org/abs/2201.03545 """
|
|
||||||
|
|
||||||
def __init__(self, dim, dim_out, *, time_emb_dim = None, mult = 2, norm = True):
|
|
||||||
super().__init__()
|
|
||||||
self.mlp = nn.Sequential(
|
|
||||||
nn.GELU(),
|
|
||||||
nn.Linear(time_emb_dim, dim)
|
|
||||||
) if exists(time_emb_dim) else None
|
|
||||||
|
|
||||||
self.ds_conv = nn.Conv2d(dim, dim, 7, padding = 3, groups = dim)
|
|
||||||
|
|
||||||
self.net = nn.Sequential(
|
|
||||||
LayerNorm(dim) if norm else nn.Identity(),
|
|
||||||
nn.Conv2d(dim, dim_out * mult, 3, padding = 1),
|
|
||||||
nn.GELU(),
|
|
||||||
LayerNorm(dim_out * mult),
|
|
||||||
nn.Conv2d(dim_out * mult, dim_out, 3, padding = 1)
|
|
||||||
)
|
|
||||||
|
|
||||||
self.res_conv = nn.Conv2d(dim, dim_out, 1) if dim != dim_out else nn.Identity()
|
|
||||||
|
|
||||||
def forward(self, x, time_emb = None):
|
|
||||||
h = self.ds_conv(x)
|
|
||||||
|
|
||||||
if exists(self.mlp) and exists(time_emb):
|
|
||||||
assert exists(time_emb), 'time emb must be passed in'
|
|
||||||
condition = self.mlp(time_emb)
|
|
||||||
h = h + rearrange(condition, 'b c -> b c 1 1')
|
|
||||||
|
|
||||||
h = self.net(h)
|
|
||||||
return h + self.res_conv(x)
|
|
||||||
|
|
||||||
class LinearAttention(nn.Module):
|
class LinearAttention(nn.Module):
|
||||||
def __init__(self, dim, heads = 4, dim_head = 32):
|
def __init__(self, dim, heads = 4, dim_head = 32):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -237,7 +204,6 @@ class Unet(nn.Module):
|
|||||||
dim_mults=(1, 2, 4, 8),
|
dim_mults=(1, 2, 4, 8),
|
||||||
channels = 3,
|
channels = 3,
|
||||||
with_time_emb = True,
|
with_time_emb = True,
|
||||||
use_convnext = False,
|
|
||||||
resnet_block_groups = 8
|
resnet_block_groups = 8
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -252,12 +218,7 @@ class Unet(nn.Module):
|
|||||||
dims = [init_dim, *map(lambda m: dim * m, dim_mults)]
|
dims = [init_dim, *map(lambda m: dim * m, dim_mults)]
|
||||||
in_out = list(zip(dims[:-1], dims[1:]))
|
in_out = list(zip(dims[:-1], dims[1:]))
|
||||||
|
|
||||||
# resnet or convnext
|
block_klass = partial(ResnetBlock, groups = resnet_block_groups)
|
||||||
|
|
||||||
if use_convnext:
|
|
||||||
block_klass = ConvNextBlock
|
|
||||||
else:
|
|
||||||
block_klass = partial(ResnetBlock, groups = resnet_block_groups)
|
|
||||||
|
|
||||||
# time embeddings
|
# time embeddings
|
||||||
|
|
||||||
@@ -355,11 +316,11 @@ def cosine_beta_schedule(timesteps, s = 0.008):
|
|||||||
as proposed in https://openreview.net/forum?id=-NEXDKk8gZ
|
as proposed in https://openreview.net/forum?id=-NEXDKk8gZ
|
||||||
"""
|
"""
|
||||||
steps = timesteps + 1
|
steps = timesteps + 1
|
||||||
x = torch.linspace(0, timesteps, steps)
|
x = torch.linspace(0, timesteps, steps, dtype = torch.float64)
|
||||||
alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * torch.pi * 0.5) ** 2
|
alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * torch.pi * 0.5) ** 2
|
||||||
alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
|
alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
|
||||||
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
|
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
|
||||||
return torch.clip(betas, 0, 0.999)
|
return torch.clip(betas, 0, 0.9999)
|
||||||
|
|
||||||
class GaussianDiffusion(nn.Module):
|
class GaussianDiffusion(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -386,17 +347,21 @@ class GaussianDiffusion(nn.Module):
|
|||||||
self.num_timesteps = int(timesteps)
|
self.num_timesteps = int(timesteps)
|
||||||
self.loss_type = loss_type
|
self.loss_type = loss_type
|
||||||
|
|
||||||
self.register_buffer('betas', betas)
|
# helper function to register buffer from float64 to float32
|
||||||
self.register_buffer('alphas_cumprod', alphas_cumprod)
|
|
||||||
self.register_buffer('alphas_cumprod_prev', alphas_cumprod_prev)
|
register_buffer = lambda name, val: self.register_buffer(name, val.to(torch.float32))
|
||||||
|
|
||||||
|
register_buffer('betas', betas)
|
||||||
|
register_buffer('alphas_cumprod', alphas_cumprod)
|
||||||
|
register_buffer('alphas_cumprod_prev', alphas_cumprod_prev)
|
||||||
|
|
||||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
# calculations for diffusion q(x_t | x_{t-1}) and others
|
||||||
|
|
||||||
self.register_buffer('sqrt_alphas_cumprod', torch.sqrt(alphas_cumprod))
|
register_buffer('sqrt_alphas_cumprod', torch.sqrt(alphas_cumprod))
|
||||||
self.register_buffer('sqrt_one_minus_alphas_cumprod', torch.sqrt(1. - alphas_cumprod))
|
register_buffer('sqrt_one_minus_alphas_cumprod', torch.sqrt(1. - alphas_cumprod))
|
||||||
self.register_buffer('log_one_minus_alphas_cumprod', torch.log(1. - alphas_cumprod))
|
register_buffer('log_one_minus_alphas_cumprod', torch.log(1. - alphas_cumprod))
|
||||||
self.register_buffer('sqrt_recip_alphas_cumprod', torch.sqrt(1. / alphas_cumprod))
|
register_buffer('sqrt_recip_alphas_cumprod', torch.sqrt(1. / alphas_cumprod))
|
||||||
self.register_buffer('sqrt_recipm1_alphas_cumprod', torch.sqrt(1. / alphas_cumprod - 1))
|
register_buffer('sqrt_recipm1_alphas_cumprod', torch.sqrt(1. / alphas_cumprod - 1))
|
||||||
|
|
||||||
# calculations for posterior q(x_{t-1} | x_t, x_0)
|
# calculations for posterior q(x_{t-1} | x_t, x_0)
|
||||||
|
|
||||||
@@ -404,13 +369,13 @@ class GaussianDiffusion(nn.Module):
|
|||||||
|
|
||||||
# above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t)
|
# above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t)
|
||||||
|
|
||||||
self.register_buffer('posterior_variance', posterior_variance)
|
register_buffer('posterior_variance', posterior_variance)
|
||||||
|
|
||||||
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
|
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
|
||||||
|
|
||||||
self.register_buffer('posterior_log_variance_clipped', torch.log(posterior_variance.clamp(min =1e-20)))
|
register_buffer('posterior_log_variance_clipped', torch.log(posterior_variance.clamp(min =1e-20)))
|
||||||
self.register_buffer('posterior_mean_coef1', betas * torch.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod))
|
register_buffer('posterior_mean_coef1', betas * torch.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod))
|
||||||
self.register_buffer('posterior_mean_coef2', (1. - alphas_cumprod_prev) * torch.sqrt(alphas) / (1. - alphas_cumprod))
|
register_buffer('posterior_mean_coef2', (1. - alphas_cumprod_prev) * torch.sqrt(alphas) / (1. - alphas_cumprod))
|
||||||
|
|
||||||
def q_mean_variance(self, x_start, t):
|
def q_mean_variance(self, x_start, t):
|
||||||
mean = extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
|
mean = extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start
|
||||||
|
|||||||
@@ -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.11.0',
|
version = '0.12.1',
|
||||||
license='MIT',
|
license='MIT',
|
||||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||||
author = 'Phil Wang',
|
author = 'Phil Wang',
|
||||||
|
|||||||
Reference in New Issue
Block a user