mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-10 12:01:08 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
88922929cb | ||
|
|
cf4f44909f | ||
|
|
a7f2d670bb | ||
|
|
6ab29d5cea | ||
|
|
ddc31bc489 | ||
|
|
e872ec3618 | ||
|
|
60f760ead5 | ||
|
|
6205e20e5e | ||
|
|
aece9c2a2a | ||
|
|
12079cadee | ||
|
|
0ffff59ca0 | ||
|
|
aadaa7d288 | ||
|
|
23fd887a5f | ||
|
|
dfbafee555 | ||
|
|
40dd8ba1de | ||
|
|
2ac3f94a80 | ||
|
|
98f2eeac35 | ||
|
|
6e8a0f2082 | ||
|
|
8c36559295 | ||
|
|
f74f536339 | ||
|
|
d85b8bbe2e | ||
|
|
e0a1bed31a | ||
|
|
82b67fc00a | ||
|
|
7c0cd05c27 | ||
|
|
6dda508ff6 | ||
|
|
9ec8d27217 | ||
|
|
4b4ebab7c3 | ||
|
|
e4a4e4acaa | ||
|
|
cd8329cdd7 | ||
|
|
aec2a26984 | ||
|
|
c78709f887 | ||
|
|
4436128a0b | ||
|
|
e46a89e2bc | ||
|
|
42158d6248 | ||
|
|
44f95e2e9d | ||
|
|
d9275a744c | ||
|
|
beb2f2d8dd | ||
|
|
f0d59acdfd | ||
|
|
689593a579 | ||
|
|
eba44498d1 | ||
|
|
12f95b33d8 | ||
|
|
6b504c4ae9 | ||
|
|
37334ae824 | ||
|
|
6eba6cdd50 | ||
|
|
555566c188 | ||
|
|
2b742dd2cc | ||
|
|
1345a8a41d | ||
|
|
931a5af2c3 | ||
|
|
a0c3443eaa | ||
|
|
662172851b | ||
|
|
0248b5e4d3 | ||
|
|
6b56af08a2 | ||
|
|
d4420248f1 | ||
|
|
a536e5bee9 | ||
|
|
1b85379d3a | ||
|
|
8859864f63 | ||
|
|
32657f035f | ||
|
|
d97bc0278c | ||
|
|
8408775cfc | ||
|
|
86fcb6785b | ||
|
|
c535d31fc5 | ||
|
|
5db64fec4b | ||
|
|
b87ea27781 | ||
|
|
a8403b83fe | ||
|
|
f4b1d7a67c | ||
|
|
618493714f | ||
|
|
76b79aa847 | ||
|
|
be2bd8d320 | ||
|
|
c3d1607019 | ||
|
|
06b2e52645 | ||
|
|
09b8a1c805 | ||
|
|
d26acbcae6 | ||
|
|
9939a48139 | ||
|
|
75ea49a7ef | ||
|
|
8c3609a6e3 | ||
|
|
1586d1a8a0 | ||
|
|
b4fb8804d2 | ||
|
|
9fd05f1b1f | ||
|
|
ec2397f0ba | ||
|
|
844e557dfb | ||
|
|
8b30be8042 | ||
|
|
f2f3994b92 | ||
|
|
8ec4ea56a5 | ||
|
|
99cf9b5b96 | ||
|
|
ecc6f30901 | ||
|
|
f900f40f14 | ||
|
|
479f60c178 | ||
|
|
96bb2ff310 | ||
|
|
582bfe275b | ||
|
|
4284c8840d | ||
|
|
d4ffa3fced | ||
|
|
c44d3ea01d | ||
|
|
c4991f576f | ||
|
|
a19331aa59 | ||
|
|
94eabaca1a | ||
|
|
eaf9d9fdc4 | ||
|
|
3bf5e768c2 | ||
|
|
532178a6a3 | ||
|
|
3bbb6ebf16 | ||
|
|
6b93fa48f6 | ||
|
|
a291da5098 | ||
|
|
e5a18bb25c | ||
|
|
fc8e4547aa | ||
|
|
cae9f4a71f | ||
|
|
91f03fb88b | ||
|
|
60128257c5 | ||
|
|
cf6db71985 | ||
|
|
84ebb9ad13 | ||
|
|
caa5af170d | ||
|
|
55c658b967 | ||
|
|
e0f26677d6 | ||
|
|
e147839d74 | ||
|
|
62e8490385 | ||
|
|
d412d8816b | ||
|
|
402b7c26df | ||
|
|
09613a40f3 | ||
|
|
c6966ae95a | ||
|
|
73591cf1ad | ||
|
|
989f0fcb8e | ||
|
|
84731bb03d | ||
|
|
c6ecca555b | ||
|
|
1f5c233072 | ||
|
|
de378158e5 | ||
|
|
e274fb305a | ||
|
|
f39b3b1d3f | ||
|
|
782c904d3b | ||
|
|
71953ebd22 | ||
|
|
0b8cdb4c8b | ||
|
|
e504e0e554 | ||
|
|
bd1e3b676e | ||
|
|
f4615599bc | ||
|
|
eb6e1b508e | ||
|
|
91cff45939 | ||
|
|
7b51e30da7 | ||
|
|
dadbf20154 | ||
|
|
7706bdfc6f | ||
|
|
183e5f3cc5 | ||
|
|
16c9ae7bb3 | ||
|
|
f5916111f8 | ||
|
|
ad9e303ff3 | ||
|
|
ae42f48f6a | ||
|
|
5989f4c77e | ||
|
|
2082046888 | ||
|
|
3c5b7e2d56 | ||
|
|
d4ce9f6c38 | ||
|
|
ff451f697e | ||
|
|
3d96532c60 | ||
|
|
ef2ca0b625 | ||
|
|
9f95a03c07 | ||
|
|
a4c68d3569 | ||
|
|
b33a48e342 | ||
|
|
8e5fb17063 | ||
|
|
4bf28914bc | ||
|
|
88f83d0ff2 | ||
|
|
26b5cab6c8 | ||
|
|
1307b3115d | ||
|
|
81fb2a0386 | ||
|
|
698227ae13 | ||
|
|
c479adf960 | ||
|
|
d70fb08f8a | ||
|
|
e700a7c6de |
@@ -1,3 +1,6 @@
|
|||||||
|
# Generation results
|
||||||
|
results/
|
||||||
|
|
||||||
# Byte-compiled / optimized / DLL files
|
# Byte-compiled / optimized / DLL files
|
||||||
__pycache__/
|
__pycache__/
|
||||||
*.py[cod]
|
*.py[cod]
|
||||||
|
|||||||
@@ -1,8 +1,22 @@
|
|||||||
<img src="./denoising-diffusion.png" width="500px"></img>
|
<img src="./images/denoising-diffusion.png" width="500px"></img>
|
||||||
|
|
||||||
## Denoising Diffusion Probabilistic Model, in Pytorch
|
## Denoising Diffusion Probabilistic Model, in Pytorch
|
||||||
|
|
||||||
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>.
|
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>
|
||||||
|
|
||||||
|
Youtube AI Educators - <a href="https://www.youtube.com/watch?v=W-O7AZNzbzQ">Yannic Kilcher</a> | <a href="https://www.youtube.com/watch?v=344w5h24-h8">AI Coffeebreak with Letitia</a> | <a href="https://www.youtube.com/watch?v=HoKDTa5jHvg">Outlier</a>
|
||||||
|
|
||||||
|
<a href="https://github.com/yiyixuxu/denoising-diffusion-flax">Flax implementation</a> from <a href="https://github.com/yiyixuxu">YiYi Xu</a>
|
||||||
|
|
||||||
|
<a href="https://huggingface.co/blog/annotated-diffusion">Annotated code</a> by Research Scientists / Engineers from <a href="https://huggingface.co/">🤗 Huggingface</a>
|
||||||
|
|
||||||
|
Update: Turns out none of the technicalities really matters at all | <a href="https://arxiv.org/abs/2208.09392">"Cold Diffusion" paper</a>
|
||||||
|
|
||||||
|
<img src="./images/sample.png" width="500px"><img>
|
||||||
|
|
||||||
|
[](https://badge.fury.io/py/denoising-diffusion-pytorch)
|
||||||
|
|
||||||
## Install
|
## Install
|
||||||
|
|
||||||
@@ -23,19 +37,18 @@ model = Unet(
|
|||||||
|
|
||||||
diffusion = GaussianDiffusion(
|
diffusion = GaussianDiffusion(
|
||||||
model,
|
model,
|
||||||
beta_start = 0.0001,
|
image_size = 128,
|
||||||
beta_end = 0.02,
|
timesteps = 1000, # number of steps
|
||||||
num_diffusion_timesteps = 1000, # number of steps
|
loss_type = 'l1' # L1 or L2
|
||||||
loss_type = 'l1' # L1 or L2 (wavegrad paper claims l1 is better?)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
training_images = torch.randn(8, 3, 128, 128)
|
training_images = torch.rand(8, 3, 128, 128) # images are normalized from 0 to 1
|
||||||
loss = diffusion(training_images)
|
loss = diffusion(training_images)
|
||||||
loss.backward()
|
loss.backward()
|
||||||
# after a lot of training
|
# after a lot of training
|
||||||
|
|
||||||
sampled_images = diffusion.p_sample_loop((1, 3, 128, 128))
|
sampled_images = diffusion.sample(batch_size = 4)
|
||||||
sampled_images.shape # (1, 3, 128, 128)
|
sampled_images.shape # (4, 3, 128, 128)
|
||||||
```
|
```
|
||||||
|
|
||||||
Or, if you simply want to pass in a folder name and the desired image dimensions, you can use the `Trainer` class to easily train a model.
|
Or, if you simply want to pass in a folder name and the desired image dimensions, you can use the `Trainer` class to easily train a model.
|
||||||
@@ -50,36 +63,198 @@ model = Unet(
|
|||||||
|
|
||||||
diffusion = GaussianDiffusion(
|
diffusion = GaussianDiffusion(
|
||||||
model,
|
model,
|
||||||
beta_start = 0.0001,
|
image_size = 128,
|
||||||
beta_end = 0.02,
|
timesteps = 1000, # number of steps
|
||||||
num_diffusion_timesteps = 1000, # number of steps
|
sampling_timesteps = 250, # number of sampling timesteps (using ddim for faster inference [see citation for ddim paper])
|
||||||
loss_type = 'l1' # L1 or L2
|
loss_type = 'l1' # L1 or L2
|
||||||
).cuda()
|
).cuda()
|
||||||
|
|
||||||
trainer = Trainer(
|
trainer = Trainer(
|
||||||
diffusion,
|
diffusion,
|
||||||
'path/to/your/images',
|
'path/to/your/images',
|
||||||
image_size = 128,
|
|
||||||
train_batch_size = 32,
|
train_batch_size = 32,
|
||||||
train_lr = 2e-5,
|
train_lr = 8e-5,
|
||||||
train_num_steps = 100000,
|
train_num_steps = 700000, # total training steps
|
||||||
gradient_accumulate_every = 2
|
gradient_accumulate_every = 2, # gradient accumulation steps
|
||||||
|
ema_decay = 0.995, # exponential moving average decay
|
||||||
|
amp = True # turn on mixed precision
|
||||||
)
|
)
|
||||||
|
|
||||||
trainer.train()
|
trainer.train()
|
||||||
```
|
```
|
||||||
|
|
||||||
Todo: Command line tool for one-line training
|
Samples and model checkpoints will be logged to `./results` periodically
|
||||||
|
|
||||||
|
## Multi-GPU Training
|
||||||
|
|
||||||
|
The `Trainer` class is now equipped with <a href="https://huggingface.co/docs/accelerate/accelerator">🤗 Accelerator</a>. You can easily do multi-gpu training in two steps using their `accelerate` CLI
|
||||||
|
|
||||||
|
At the project root directory, where the training script is, run
|
||||||
|
|
||||||
|
```python
|
||||||
|
$ accelerate config
|
||||||
|
```
|
||||||
|
|
||||||
|
Then, in the same directory
|
||||||
|
|
||||||
|
```python
|
||||||
|
$ accelerate launch train.py
|
||||||
|
```
|
||||||
|
|
||||||
|
## Miscellaneous
|
||||||
|
|
||||||
|
### 1D Sequence
|
||||||
|
|
||||||
|
By popular request, a 1D Unet + Gaussian Diffusion implementation. You will have to do the training code yourself
|
||||||
|
|
||||||
|
```python
|
||||||
|
import torch
|
||||||
|
from denoising_diffusion_pytorch import Unet1D, GaussianDiffusion1D
|
||||||
|
|
||||||
|
model = Unet1D(
|
||||||
|
dim = 64,
|
||||||
|
dim_mults = (1, 2, 4, 8),
|
||||||
|
channels = 32
|
||||||
|
)
|
||||||
|
|
||||||
|
diffusion = GaussianDiffusion1D(
|
||||||
|
model,
|
||||||
|
seq_length = 128,
|
||||||
|
timesteps = 1000,
|
||||||
|
objective = 'pred_v'
|
||||||
|
)
|
||||||
|
|
||||||
|
training_seq = torch.rand(8, 32, 128) # features are normalized from 0 to 1
|
||||||
|
loss = diffusion(training_seq)
|
||||||
|
loss.backward()
|
||||||
|
|
||||||
|
# after a lot of training
|
||||||
|
|
||||||
|
sampled_seq = diffusion.sample(batch_size = 4)
|
||||||
|
sampled_seq.shape # (4, 32, 128)
|
||||||
|
```
|
||||||
|
|
||||||
## 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
|
||||||
|
@InProceedings{pmlr-v139-nichol21a,
|
||||||
|
title = {Improved Denoising Diffusion Probabilistic Models},
|
||||||
|
author = {Nichol, Alexander Quinn and Dhariwal, Prafulla},
|
||||||
|
booktitle = {Proceedings of the 38th International Conference on Machine Learning},
|
||||||
|
pages = {8162--8171},
|
||||||
|
year = {2021},
|
||||||
|
editor = {Meila, Marina and Zhang, Tong},
|
||||||
|
volume = {139},
|
||||||
|
series = {Proceedings of Machine Learning Research},
|
||||||
|
month = {18--24 Jul},
|
||||||
|
publisher = {PMLR},
|
||||||
|
pdf = {http://proceedings.mlr.press/v139/nichol21a/nichol21a.pdf},
|
||||||
|
url = {https://proceedings.mlr.press/v139/nichol21a.html},
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@inproceedings{kingma2021on,
|
||||||
|
title = {On Density Estimation with Diffusion Models},
|
||||||
|
author = {Diederik P Kingma and Tim Salimans and Ben Poole and Jonathan Ho},
|
||||||
|
booktitle = {Advances in Neural Information Processing Systems},
|
||||||
|
editor = {A. Beygelzimer and Y. Dauphin and P. Liang and J. Wortman Vaughan},
|
||||||
|
year = {2021},
|
||||||
|
url = {https://openreview.net/forum?id=2LdBqxc1Yv}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@article{Choi2022PerceptionPT,
|
||||||
|
title = {Perception Prioritized Training of Diffusion Models},
|
||||||
|
author = {Jooyoung Choi and Jungbeom Lee and Chaehun Shin and Sungwon Kim and Hyunwoo J. Kim and Sung-Hoon Yoon},
|
||||||
|
journal = {ArXiv},
|
||||||
|
year = {2022},
|
||||||
|
volume = {abs/2204.00227}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@article{Karras2022ElucidatingTD,
|
||||||
|
title = {Elucidating the Design Space of Diffusion-Based Generative Models},
|
||||||
|
author = {Tero Karras and Miika Aittala and Timo Aila and Samuli Laine},
|
||||||
|
journal = {ArXiv},
|
||||||
|
year = {2022},
|
||||||
|
volume = {abs/2206.00364}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@article{Song2021DenoisingDI,
|
||||||
|
title = {Denoising Diffusion Implicit Models},
|
||||||
|
author = {Jiaming Song and Chenlin Meng and Stefano Ermon},
|
||||||
|
journal = {ArXiv},
|
||||||
|
year = {2021},
|
||||||
|
volume = {abs/2010.02502}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@misc{chen2022analog,
|
||||||
|
title = {Analog Bits: Generating Discrete Data using Diffusion Models with Self-Conditioning},
|
||||||
|
author = {Ting Chen and Ruixiang Zhang and Geoffrey Hinton},
|
||||||
|
year = {2022},
|
||||||
|
eprint = {2208.04202},
|
||||||
|
archivePrefix = {arXiv},
|
||||||
|
primaryClass = {cs.CV}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@article{Qiao2019WeightS,
|
||||||
|
title = {Weight Standardization},
|
||||||
|
author = {Siyuan Qiao and Huiyu Wang and Chenxi Liu and Wei Shen and Alan Loddon Yuille},
|
||||||
|
journal = {ArXiv},
|
||||||
|
year = {2019},
|
||||||
|
volume = {abs/1903.10520}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@article{Salimans2022ProgressiveDF,
|
||||||
|
title = {Progressive Distillation for Fast Sampling of Diffusion Models},
|
||||||
|
author = {Tim Salimans and Jonathan Ho},
|
||||||
|
journal = {ArXiv},
|
||||||
|
year = {2022},
|
||||||
|
volume = {abs/2202.00512}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@article{Ho2022ClassifierFreeDG,
|
||||||
|
title = {Classifier-Free Diffusion Guidance},
|
||||||
|
author = {Jonathan Ho},
|
||||||
|
journal = {ArXiv},
|
||||||
|
year = {2022},
|
||||||
|
volume = {abs/2207.12598}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@article{Sunkara2022NoMS,
|
||||||
|
title = {No More Strided Convolutions or Pooling: A New CNN Building Block for Low-Resolution Images and Small Objects},
|
||||||
|
author = {Raja Sunkara and Tie Luo},
|
||||||
|
journal = {ArXiv},
|
||||||
|
year = {2022},
|
||||||
|
volume = {abs/2208.03641}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -1 +1,10 @@
|
|||||||
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, Unet, Trainer
|
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, Unet, Trainer
|
||||||
|
|
||||||
|
from denoising_diffusion_pytorch.learned_gaussian_diffusion import LearnedGaussianDiffusion
|
||||||
|
from denoising_diffusion_pytorch.continuous_time_gaussian_diffusion import ContinuousTimeGaussianDiffusion
|
||||||
|
from denoising_diffusion_pytorch.weighted_objective_gaussian_diffusion import WeightedObjectiveGaussianDiffusion
|
||||||
|
from denoising_diffusion_pytorch.elucidated_diffusion import ElucidatedDiffusion
|
||||||
|
from denoising_diffusion_pytorch.v_param_continuous_time_gaussian_diffusion import VParamContinuousTimeGaussianDiffusion
|
||||||
|
|
||||||
|
from denoising_diffusion_pytorch.denoising_diffusion_pytorch_1d import GaussianDiffusion1D, Unet1D
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,792 @@
|
|||||||
|
import math
|
||||||
|
import copy
|
||||||
|
from pathlib import Path
|
||||||
|
from random import random
|
||||||
|
from functools import partial
|
||||||
|
from collections import namedtuple
|
||||||
|
from multiprocessing import cpu_count
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn, einsum
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from einops import rearrange, reduce, repeat
|
||||||
|
from einops.layers.torch import Rearrange
|
||||||
|
|
||||||
|
from tqdm.auto import tqdm
|
||||||
|
|
||||||
|
# constants
|
||||||
|
|
||||||
|
ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start'])
|
||||||
|
|
||||||
|
# helpers functions
|
||||||
|
|
||||||
|
def exists(x):
|
||||||
|
return x is not None
|
||||||
|
|
||||||
|
def default(val, d):
|
||||||
|
if exists(val):
|
||||||
|
return val
|
||||||
|
return d() if callable(d) else d
|
||||||
|
|
||||||
|
def identity(t, *args, **kwargs):
|
||||||
|
return t
|
||||||
|
|
||||||
|
def cycle(dl):
|
||||||
|
while True:
|
||||||
|
for data in dl:
|
||||||
|
yield data
|
||||||
|
|
||||||
|
def has_int_squareroot(num):
|
||||||
|
return (math.sqrt(num) ** 2) == num
|
||||||
|
|
||||||
|
def num_to_groups(num, divisor):
|
||||||
|
groups = num // divisor
|
||||||
|
remainder = num % divisor
|
||||||
|
arr = [divisor] * groups
|
||||||
|
if remainder > 0:
|
||||||
|
arr.append(remainder)
|
||||||
|
return arr
|
||||||
|
|
||||||
|
def convert_image_to_fn(img_type, image):
|
||||||
|
if image.mode != img_type:
|
||||||
|
return image.convert(img_type)
|
||||||
|
return image
|
||||||
|
|
||||||
|
# normalization functions
|
||||||
|
|
||||||
|
def normalize_to_neg_one_to_one(img):
|
||||||
|
return img * 2 - 1
|
||||||
|
|
||||||
|
def unnormalize_to_zero_to_one(t):
|
||||||
|
return (t + 1) * 0.5
|
||||||
|
|
||||||
|
# classifier free guidance functions
|
||||||
|
|
||||||
|
def uniform(shape, device):
|
||||||
|
return torch.zeros(shape, device = device).float().uniform_(0, 1)
|
||||||
|
|
||||||
|
def prob_mask_like(shape, prob, device):
|
||||||
|
if prob == 1:
|
||||||
|
return torch.ones(shape, device = device, dtype = torch.bool)
|
||||||
|
elif prob == 0:
|
||||||
|
return torch.zeros(shape, device = device, dtype = torch.bool)
|
||||||
|
else:
|
||||||
|
return torch.zeros(shape, device = device).float().uniform_(0, 1) < prob
|
||||||
|
|
||||||
|
# small helper modules
|
||||||
|
|
||||||
|
class Residual(nn.Module):
|
||||||
|
def __init__(self, fn):
|
||||||
|
super().__init__()
|
||||||
|
self.fn = fn
|
||||||
|
|
||||||
|
def forward(self, x, *args, **kwargs):
|
||||||
|
return self.fn(x, *args, **kwargs) + x
|
||||||
|
|
||||||
|
def Upsample(dim, dim_out = None):
|
||||||
|
return nn.Sequential(
|
||||||
|
nn.Upsample(scale_factor = 2, mode = 'nearest'),
|
||||||
|
nn.Conv2d(dim, default(dim_out, dim), 3, padding = 1)
|
||||||
|
)
|
||||||
|
|
||||||
|
def Downsample(dim, dim_out = None):
|
||||||
|
return nn.Conv2d(dim, default(dim_out, dim), 4, 2, 1)
|
||||||
|
|
||||||
|
class WeightStandardizedConv2d(nn.Conv2d):
|
||||||
|
"""
|
||||||
|
https://arxiv.org/abs/1903.10520
|
||||||
|
weight standardization purportedly works synergistically with group normalization
|
||||||
|
"""
|
||||||
|
def forward(self, x):
|
||||||
|
eps = 1e-5 if x.dtype == torch.float32 else 1e-3
|
||||||
|
|
||||||
|
weight = self.weight
|
||||||
|
mean = reduce(weight, 'o ... -> o 1 1 1', 'mean')
|
||||||
|
var = reduce(weight, 'o ... -> o 1 1 1', partial(torch.var, unbiased = False))
|
||||||
|
normalized_weight = (weight - mean) * (var + eps).rsqrt()
|
||||||
|
|
||||||
|
return F.conv2d(x, normalized_weight, self.bias, self.stride, self.padding, self.dilation, self.groups)
|
||||||
|
|
||||||
|
class LayerNorm(nn.Module):
|
||||||
|
def __init__(self, dim):
|
||||||
|
super().__init__()
|
||||||
|
self.g = nn.Parameter(torch.ones(1, dim, 1, 1))
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
eps = 1e-5 if x.dtype == torch.float32 else 1e-3
|
||||||
|
var = torch.var(x, dim = 1, unbiased = False, keepdim = True)
|
||||||
|
mean = torch.mean(x, dim = 1, keepdim = True)
|
||||||
|
return (x - mean) * (var + eps).rsqrt() * self.g
|
||||||
|
|
||||||
|
class PreNorm(nn.Module):
|
||||||
|
def __init__(self, dim, fn):
|
||||||
|
super().__init__()
|
||||||
|
self.fn = fn
|
||||||
|
self.norm = LayerNorm(dim)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
x = self.norm(x)
|
||||||
|
return self.fn(x)
|
||||||
|
|
||||||
|
# sinusoidal positional embeds
|
||||||
|
|
||||||
|
class SinusoidalPosEmb(nn.Module):
|
||||||
|
def __init__(self, dim):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
device = x.device
|
||||||
|
half_dim = self.dim // 2
|
||||||
|
emb = math.log(10000) / (half_dim - 1)
|
||||||
|
emb = torch.exp(torch.arange(half_dim, device=device) * -emb)
|
||||||
|
emb = x[:, None] * emb[None, :]
|
||||||
|
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
|
||||||
|
return emb
|
||||||
|
|
||||||
|
class RandomOrLearnedSinusoidalPosEmb(nn.Module):
|
||||||
|
""" following @crowsonkb 's lead with random (learned optional) sinusoidal pos emb """
|
||||||
|
""" https://github.com/crowsonkb/v-diffusion-jax/blob/master/diffusion/models/danbooru_128.py#L8 """
|
||||||
|
|
||||||
|
def __init__(self, dim, is_random = False):
|
||||||
|
super().__init__()
|
||||||
|
assert (dim % 2) == 0
|
||||||
|
half_dim = dim // 2
|
||||||
|
self.weights = nn.Parameter(torch.randn(half_dim), requires_grad = not is_random)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
x = rearrange(x, 'b -> b 1')
|
||||||
|
freqs = x * rearrange(self.weights, 'd -> 1 d') * 2 * math.pi
|
||||||
|
fouriered = torch.cat((freqs.sin(), freqs.cos()), dim = -1)
|
||||||
|
fouriered = torch.cat((x, fouriered), dim = -1)
|
||||||
|
return fouriered
|
||||||
|
|
||||||
|
# building block modules
|
||||||
|
|
||||||
|
class Block(nn.Module):
|
||||||
|
def __init__(self, dim, dim_out, groups = 8):
|
||||||
|
super().__init__()
|
||||||
|
self.proj = WeightStandardizedConv2d(dim, dim_out, 3, padding = 1)
|
||||||
|
self.norm = nn.GroupNorm(groups, dim_out)
|
||||||
|
self.act = nn.SiLU()
|
||||||
|
|
||||||
|
def forward(self, x, scale_shift = None):
|
||||||
|
x = self.proj(x)
|
||||||
|
x = self.norm(x)
|
||||||
|
|
||||||
|
if exists(scale_shift):
|
||||||
|
scale, shift = scale_shift
|
||||||
|
x = x * (scale + 1) + shift
|
||||||
|
|
||||||
|
x = self.act(x)
|
||||||
|
return x
|
||||||
|
|
||||||
|
class ResnetBlock(nn.Module):
|
||||||
|
def __init__(self, dim, dim_out, *, time_emb_dim = None, classes_emb_dim = None, groups = 8):
|
||||||
|
super().__init__()
|
||||||
|
self.mlp = nn.Sequential(
|
||||||
|
nn.SiLU(),
|
||||||
|
nn.Linear(int(time_emb_dim) + int(classes_emb_dim), dim_out * 2)
|
||||||
|
) if exists(time_emb_dim) or exists(classes_emb_dim) else None
|
||||||
|
|
||||||
|
self.block1 = Block(dim, dim_out, groups = groups)
|
||||||
|
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()
|
||||||
|
|
||||||
|
def forward(self, x, time_emb = None, class_emb = None):
|
||||||
|
|
||||||
|
scale_shift = None
|
||||||
|
if exists(self.mlp) and (exists(time_emb) or exists(class_emb)):
|
||||||
|
cond_emb = tuple(filter(exists, (time_emb, class_emb)))
|
||||||
|
cond_emb = torch.cat(cond_emb, dim = -1)
|
||||||
|
cond_emb = self.mlp(cond_emb)
|
||||||
|
cond_emb = rearrange(cond_emb, 'b c -> b c 1 1')
|
||||||
|
scale_shift = cond_emb.chunk(2, dim = 1)
|
||||||
|
|
||||||
|
h = self.block1(x, scale_shift = scale_shift)
|
||||||
|
|
||||||
|
h = self.block2(h)
|
||||||
|
|
||||||
|
return h + self.res_conv(x)
|
||||||
|
|
||||||
|
class LinearAttention(nn.Module):
|
||||||
|
def __init__(self, dim, heads = 4, dim_head = 32):
|
||||||
|
super().__init__()
|
||||||
|
self.scale = dim_head ** -0.5
|
||||||
|
self.heads = heads
|
||||||
|
hidden_dim = dim_head * heads
|
||||||
|
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
|
||||||
|
|
||||||
|
self.to_out = nn.Sequential(
|
||||||
|
nn.Conv2d(hidden_dim, dim, 1),
|
||||||
|
LayerNorm(dim)
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
b, c, h, w = x.shape
|
||||||
|
qkv = self.to_qkv(x).chunk(3, dim = 1)
|
||||||
|
q, k, v = map(lambda t: rearrange(t, 'b (h c) x y -> b h c (x y)', h = self.heads), qkv)
|
||||||
|
|
||||||
|
q = q.softmax(dim = -2)
|
||||||
|
k = k.softmax(dim = -1)
|
||||||
|
|
||||||
|
q = q * self.scale
|
||||||
|
v = v / (h * w)
|
||||||
|
|
||||||
|
context = torch.einsum('b h d n, b h e n -> b h d e', k, v)
|
||||||
|
|
||||||
|
out = torch.einsum('b h d e, b h d n -> b h e n', context, q)
|
||||||
|
out = rearrange(out, 'b h c (x y) -> b (h c) x y', h = self.heads, x = h, y = w)
|
||||||
|
return self.to_out(out)
|
||||||
|
|
||||||
|
class Attention(nn.Module):
|
||||||
|
def __init__(self, dim, heads = 4, dim_head = 32):
|
||||||
|
super().__init__()
|
||||||
|
self.scale = dim_head ** -0.5
|
||||||
|
self.heads = heads
|
||||||
|
hidden_dim = dim_head * heads
|
||||||
|
|
||||||
|
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
|
||||||
|
self.to_out = nn.Conv2d(hidden_dim, dim, 1)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
b, c, h, w = x.shape
|
||||||
|
qkv = self.to_qkv(x).chunk(3, dim = 1)
|
||||||
|
q, k, v = map(lambda t: rearrange(t, 'b (h c) x y -> b h c (x y)', h = self.heads), qkv)
|
||||||
|
|
||||||
|
q = q * self.scale
|
||||||
|
|
||||||
|
sim = einsum('b h d i, b h d j -> b h i j', q, k)
|
||||||
|
attn = sim.softmax(dim = -1)
|
||||||
|
out = einsum('b h i j, b h d j -> b h i d', attn, v)
|
||||||
|
|
||||||
|
out = rearrange(out, 'b h (x y) d -> b (h d) x y', x = h, y = w)
|
||||||
|
return self.to_out(out)
|
||||||
|
|
||||||
|
# model
|
||||||
|
|
||||||
|
class Unet(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim,
|
||||||
|
num_classes,
|
||||||
|
cond_drop_prob = 0.5,
|
||||||
|
init_dim = None,
|
||||||
|
out_dim = None,
|
||||||
|
dim_mults=(1, 2, 4, 8),
|
||||||
|
channels = 3,
|
||||||
|
resnet_block_groups = 8,
|
||||||
|
learned_variance = False,
|
||||||
|
learned_sinusoidal_cond = False,
|
||||||
|
random_fourier_features = False,
|
||||||
|
learned_sinusoidal_dim = 16,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
# classifier free guidance stuff
|
||||||
|
|
||||||
|
self.cond_drop_prob = cond_drop_prob
|
||||||
|
|
||||||
|
# determine dimensions
|
||||||
|
|
||||||
|
self.channels = channels
|
||||||
|
input_channels = channels
|
||||||
|
|
||||||
|
init_dim = default(init_dim, dim)
|
||||||
|
self.init_conv = nn.Conv2d(input_channels, init_dim, 7, padding = 3)
|
||||||
|
|
||||||
|
dims = [init_dim, *map(lambda m: dim * m, dim_mults)]
|
||||||
|
in_out = list(zip(dims[:-1], dims[1:]))
|
||||||
|
|
||||||
|
block_klass = partial(ResnetBlock, groups = resnet_block_groups)
|
||||||
|
|
||||||
|
# time embeddings
|
||||||
|
|
||||||
|
time_dim = dim * 4
|
||||||
|
|
||||||
|
self.random_or_learned_sinusoidal_cond = learned_sinusoidal_cond or random_fourier_features
|
||||||
|
|
||||||
|
if self.random_or_learned_sinusoidal_cond:
|
||||||
|
sinu_pos_emb = RandomOrLearnedSinusoidalPosEmb(learned_sinusoidal_dim, random_fourier_features)
|
||||||
|
fourier_dim = learned_sinusoidal_dim + 1
|
||||||
|
else:
|
||||||
|
sinu_pos_emb = SinusoidalPosEmb(dim)
|
||||||
|
fourier_dim = dim
|
||||||
|
|
||||||
|
self.time_mlp = nn.Sequential(
|
||||||
|
sinu_pos_emb,
|
||||||
|
nn.Linear(fourier_dim, time_dim),
|
||||||
|
nn.GELU(),
|
||||||
|
nn.Linear(time_dim, time_dim)
|
||||||
|
)
|
||||||
|
|
||||||
|
# class embeddings
|
||||||
|
|
||||||
|
self.classes_emb = nn.Embedding(num_classes, dim)
|
||||||
|
self.null_classes_emb = nn.Parameter(torch.randn(dim))
|
||||||
|
|
||||||
|
classes_dim = dim * 4
|
||||||
|
|
||||||
|
self.classes_mlp = nn.Sequential(
|
||||||
|
nn.Linear(dim, classes_dim),
|
||||||
|
nn.GELU(),
|
||||||
|
nn.Linear(classes_dim, classes_dim)
|
||||||
|
)
|
||||||
|
|
||||||
|
# layers
|
||||||
|
|
||||||
|
self.downs = nn.ModuleList([])
|
||||||
|
self.ups = nn.ModuleList([])
|
||||||
|
num_resolutions = len(in_out)
|
||||||
|
|
||||||
|
for ind, (dim_in, dim_out) in enumerate(in_out):
|
||||||
|
is_last = ind >= (num_resolutions - 1)
|
||||||
|
|
||||||
|
self.downs.append(nn.ModuleList([
|
||||||
|
block_klass(dim_in, dim_in, time_emb_dim = time_dim, classes_emb_dim = classes_dim),
|
||||||
|
block_klass(dim_in, dim_in, time_emb_dim = time_dim, classes_emb_dim = classes_dim),
|
||||||
|
Residual(PreNorm(dim_in, LinearAttention(dim_in))),
|
||||||
|
Downsample(dim_in, dim_out) if not is_last else nn.Conv2d(dim_in, dim_out, 3, padding = 1)
|
||||||
|
]))
|
||||||
|
|
||||||
|
mid_dim = dims[-1]
|
||||||
|
self.mid_block1 = block_klass(mid_dim, mid_dim, time_emb_dim = time_dim, classes_emb_dim = classes_dim)
|
||||||
|
self.mid_attn = Residual(PreNorm(mid_dim, Attention(mid_dim)))
|
||||||
|
self.mid_block2 = block_klass(mid_dim, mid_dim, time_emb_dim = time_dim, classes_emb_dim = classes_dim)
|
||||||
|
|
||||||
|
for ind, (dim_in, dim_out) in enumerate(reversed(in_out)):
|
||||||
|
is_last = ind == (len(in_out) - 1)
|
||||||
|
|
||||||
|
self.ups.append(nn.ModuleList([
|
||||||
|
block_klass(dim_out + dim_in, dim_out, time_emb_dim = time_dim, classes_emb_dim = classes_dim),
|
||||||
|
block_klass(dim_out + dim_in, dim_out, time_emb_dim = time_dim, classes_emb_dim = classes_dim),
|
||||||
|
Residual(PreNorm(dim_out, LinearAttention(dim_out))),
|
||||||
|
Upsample(dim_out, dim_in) if not is_last else nn.Conv2d(dim_out, dim_in, 3, padding = 1)
|
||||||
|
]))
|
||||||
|
|
||||||
|
default_out_dim = channels * (1 if not learned_variance else 2)
|
||||||
|
self.out_dim = default(out_dim, default_out_dim)
|
||||||
|
|
||||||
|
self.final_res_block = block_klass(dim * 2, dim, time_emb_dim = time_dim, classes_emb_dim = classes_dim)
|
||||||
|
self.final_conv = nn.Conv2d(dim, self.out_dim, 1)
|
||||||
|
|
||||||
|
def forward_with_cond_scale(
|
||||||
|
self,
|
||||||
|
*args,
|
||||||
|
cond_scale = 1.,
|
||||||
|
**kwargs
|
||||||
|
):
|
||||||
|
logits = self.forward(*args, **kwargs)
|
||||||
|
|
||||||
|
if cond_scale == 1:
|
||||||
|
return logits
|
||||||
|
|
||||||
|
null_logits = self.forward(*args, cond_drop_prob = 1., **kwargs)
|
||||||
|
return null_logits + (logits - null_logits) * cond_scale
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x,
|
||||||
|
time,
|
||||||
|
classes,
|
||||||
|
cond_drop_prob = None
|
||||||
|
):
|
||||||
|
batch, device = x.shape[0], x.device
|
||||||
|
|
||||||
|
cond_drop_prob = default(cond_drop_prob, self.cond_drop_prob)
|
||||||
|
|
||||||
|
# derive condition, with condition dropout for classifier free guidance
|
||||||
|
|
||||||
|
classes_emb = self.classes_emb(classes)
|
||||||
|
|
||||||
|
if cond_drop_prob > 0:
|
||||||
|
keep_mask = prob_mask_like((batch,), 1 - cond_drop_prob, device = device)
|
||||||
|
null_classes_emb = repeat(self.null_classes_emb, 'd -> b d', b = batch)
|
||||||
|
|
||||||
|
classes_emb = torch.where(
|
||||||
|
rearrange(keep_mask, 'b -> b 1'),
|
||||||
|
classes_emb,
|
||||||
|
null_classes_emb
|
||||||
|
)
|
||||||
|
|
||||||
|
c = self.classes_mlp(classes_emb)
|
||||||
|
|
||||||
|
# unet
|
||||||
|
|
||||||
|
x = self.init_conv(x)
|
||||||
|
r = x.clone()
|
||||||
|
|
||||||
|
t = self.time_mlp(time)
|
||||||
|
|
||||||
|
h = []
|
||||||
|
|
||||||
|
for block1, block2, attn, downsample in self.downs:
|
||||||
|
x = block1(x, t, c)
|
||||||
|
h.append(x)
|
||||||
|
|
||||||
|
x = block2(x, t, c)
|
||||||
|
x = attn(x)
|
||||||
|
h.append(x)
|
||||||
|
|
||||||
|
x = downsample(x)
|
||||||
|
|
||||||
|
x = self.mid_block1(x, t, c)
|
||||||
|
x = self.mid_attn(x)
|
||||||
|
x = self.mid_block2(x, t, c)
|
||||||
|
|
||||||
|
for block1, block2, attn, upsample in self.ups:
|
||||||
|
x = torch.cat((x, h.pop()), dim = 1)
|
||||||
|
x = block1(x, t, c)
|
||||||
|
|
||||||
|
x = torch.cat((x, h.pop()), dim = 1)
|
||||||
|
x = block2(x, t, c)
|
||||||
|
x = attn(x)
|
||||||
|
|
||||||
|
x = upsample(x)
|
||||||
|
|
||||||
|
x = torch.cat((x, r), dim = 1)
|
||||||
|
|
||||||
|
x = self.final_res_block(x, t, c)
|
||||||
|
return self.final_conv(x)
|
||||||
|
|
||||||
|
# gaussian diffusion trainer class
|
||||||
|
|
||||||
|
def extract(a, t, x_shape):
|
||||||
|
b, *_ = t.shape
|
||||||
|
out = a.gather(-1, t)
|
||||||
|
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
||||||
|
|
||||||
|
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
|
||||||
|
as proposed in https://openreview.net/forum?id=-NEXDKk8gZ
|
||||||
|
"""
|
||||||
|
steps = timesteps + 1
|
||||||
|
x = torch.linspace(0, timesteps, steps, dtype = torch.float64)
|
||||||
|
alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * math.pi * 0.5) ** 2
|
||||||
|
alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
|
||||||
|
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
|
||||||
|
return torch.clip(betas, 0, 0.999)
|
||||||
|
|
||||||
|
class GaussianDiffusion(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model,
|
||||||
|
*,
|
||||||
|
image_size,
|
||||||
|
timesteps = 1000,
|
||||||
|
sampling_timesteps = None,
|
||||||
|
loss_type = 'l1',
|
||||||
|
objective = 'pred_noise',
|
||||||
|
beta_schedule = 'cosine',
|
||||||
|
p2_loss_weight_gamma = 0., # p2 loss weight, from https://arxiv.org/abs/2204.00227 - 0 is equivalent to weight of 1 across time - 1. is recommended
|
||||||
|
p2_loss_weight_k = 1,
|
||||||
|
ddim_sampling_eta = 0.
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim)
|
||||||
|
assert not model.random_or_learned_sinusoidal_cond
|
||||||
|
|
||||||
|
self.model = model
|
||||||
|
self.channels = self.model.channels
|
||||||
|
|
||||||
|
self.image_size = image_size
|
||||||
|
|
||||||
|
self.objective = objective
|
||||||
|
|
||||||
|
assert objective in {'pred_noise', 'pred_x0', 'pred_v'}, 'objective must be either pred_noise (predict noise) or pred_x0 (predict image start) or pred_v (predict v [v-parameterization as defined in appendix D of progressive distillation paper, used in imagen-video successfully])'
|
||||||
|
|
||||||
|
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, dim=0)
|
||||||
|
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value = 1.)
|
||||||
|
|
||||||
|
timesteps, = betas.shape
|
||||||
|
self.num_timesteps = int(timesteps)
|
||||||
|
self.loss_type = loss_type
|
||||||
|
|
||||||
|
# sampling related parameters
|
||||||
|
|
||||||
|
self.sampling_timesteps = default(sampling_timesteps, timesteps) # default num sampling timesteps to number of timesteps at training
|
||||||
|
|
||||||
|
assert self.sampling_timesteps <= timesteps
|
||||||
|
self.is_ddim_sampling = self.sampling_timesteps < timesteps
|
||||||
|
self.ddim_sampling_eta = ddim_sampling_eta
|
||||||
|
|
||||||
|
# helper function to register buffer from float64 to float32
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
register_buffer('sqrt_alphas_cumprod', torch.sqrt(alphas_cumprod))
|
||||||
|
register_buffer('sqrt_one_minus_alphas_cumprod', torch.sqrt(1. - alphas_cumprod))
|
||||||
|
register_buffer('log_one_minus_alphas_cumprod', torch.log(1. - alphas_cumprod))
|
||||||
|
register_buffer('sqrt_recip_alphas_cumprod', torch.sqrt(1. / alphas_cumprod))
|
||||||
|
register_buffer('sqrt_recipm1_alphas_cumprod', torch.sqrt(1. / alphas_cumprod - 1))
|
||||||
|
|
||||||
|
# calculations for posterior q(x_{t-1} | x_t, x_0)
|
||||||
|
|
||||||
|
posterior_variance = betas * (1. - alphas_cumprod_prev) / (1. - alphas_cumprod)
|
||||||
|
|
||||||
|
# above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t)
|
||||||
|
|
||||||
|
register_buffer('posterior_variance', posterior_variance)
|
||||||
|
|
||||||
|
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
|
||||||
|
|
||||||
|
register_buffer('posterior_log_variance_clipped', torch.log(posterior_variance.clamp(min =1e-20)))
|
||||||
|
register_buffer('posterior_mean_coef1', betas * torch.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod))
|
||||||
|
register_buffer('posterior_mean_coef2', (1. - alphas_cumprod_prev) * torch.sqrt(alphas) / (1. - alphas_cumprod))
|
||||||
|
|
||||||
|
# calculate p2 reweighting
|
||||||
|
|
||||||
|
register_buffer('p2_loss_weight', (p2_loss_weight_k + alphas_cumprod / (1 - alphas_cumprod)) ** -p2_loss_weight_gamma)
|
||||||
|
|
||||||
|
def predict_start_from_noise(self, x_t, t, noise):
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t -
|
||||||
|
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * noise
|
||||||
|
)
|
||||||
|
|
||||||
|
def predict_noise_from_start(self, x_t, t, x0):
|
||||||
|
return (
|
||||||
|
(extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - x0) / \
|
||||||
|
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
|
||||||
|
)
|
||||||
|
|
||||||
|
def predict_v(self, x_start, t, noise):
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_alphas_cumprod, t, x_start.shape) * noise -
|
||||||
|
extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * x_start
|
||||||
|
)
|
||||||
|
|
||||||
|
def predict_start_from_v(self, x_t, t, v):
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_alphas_cumprod, t, x_t.shape) * x_t -
|
||||||
|
extract(self.sqrt_one_minus_alphas_cumprod, t, x_t.shape) * v
|
||||||
|
)
|
||||||
|
|
||||||
|
def q_posterior(self, x_start, x_t, t):
|
||||||
|
posterior_mean = (
|
||||||
|
extract(self.posterior_mean_coef1, t, x_t.shape) * x_start +
|
||||||
|
extract(self.posterior_mean_coef2, t, x_t.shape) * x_t
|
||||||
|
)
|
||||||
|
posterior_variance = extract(self.posterior_variance, t, x_t.shape)
|
||||||
|
posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, t, x_t.shape)
|
||||||
|
return posterior_mean, posterior_variance, posterior_log_variance_clipped
|
||||||
|
|
||||||
|
def model_predictions(self, x, t, classes, cond_scale = 3., clip_x_start = False):
|
||||||
|
model_output = self.model.forward_with_cond_scale(x, t, classes, cond_scale = cond_scale)
|
||||||
|
maybe_clip = partial(torch.clamp, min = -1., max = 1.) if clip_x_start else identity
|
||||||
|
|
||||||
|
if self.objective == 'pred_noise':
|
||||||
|
pred_noise = model_output
|
||||||
|
x_start = self.predict_start_from_noise(x, t, pred_noise)
|
||||||
|
x_start = maybe_clip(x_start)
|
||||||
|
|
||||||
|
elif self.objective == 'pred_x0':
|
||||||
|
x_start = model_output
|
||||||
|
x_start = maybe_clip(x_start)
|
||||||
|
pred_noise = self.predict_noise_from_start(x, t, x_start)
|
||||||
|
|
||||||
|
elif self.objective == 'pred_v':
|
||||||
|
v = model_output
|
||||||
|
x_start = self.predict_start_from_v(x, t, v)
|
||||||
|
x_start = maybe_clip(x_start)
|
||||||
|
pred_noise = self.predict_noise_from_start(x, t, x_start)
|
||||||
|
|
||||||
|
return ModelPrediction(pred_noise, x_start)
|
||||||
|
|
||||||
|
def p_mean_variance(self, x, t, classes, cond_scale, clip_denoised = True):
|
||||||
|
preds = self.model_predictions(x, t, classes, cond_scale)
|
||||||
|
x_start = preds.pred_x_start
|
||||||
|
|
||||||
|
if clip_denoised:
|
||||||
|
x_start.clamp_(-1., 1.)
|
||||||
|
|
||||||
|
model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start = x_start, x_t = x, t = t)
|
||||||
|
return model_mean, posterior_variance, posterior_log_variance, x_start
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def p_sample(self, x, t: int, classes, cond_scale = 3., clip_denoised = True):
|
||||||
|
b, *_, device = *x.shape, x.device
|
||||||
|
batched_times = torch.full((x.shape[0],), t, device = x.device, dtype = torch.long)
|
||||||
|
model_mean, _, model_log_variance, x_start = self.p_mean_variance(x = x, t = batched_times, classes = classes, cond_scale = cond_scale, clip_denoised = clip_denoised)
|
||||||
|
noise = torch.randn_like(x) if t > 0 else 0. # no noise if t == 0
|
||||||
|
pred_img = model_mean + (0.5 * model_log_variance).exp() * noise
|
||||||
|
return pred_img, x_start
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def p_sample_loop(self, classes, shape, cond_scale = 3.):
|
||||||
|
batch, device = shape[0], self.betas.device
|
||||||
|
|
||||||
|
img = torch.randn(shape, device=device)
|
||||||
|
|
||||||
|
x_start = None
|
||||||
|
|
||||||
|
for t in tqdm(reversed(range(0, self.num_timesteps)), desc = 'sampling loop time step', total = self.num_timesteps):
|
||||||
|
img, x_start = self.p_sample(img, t, classes, cond_scale)
|
||||||
|
|
||||||
|
img = unnormalize_to_zero_to_one(img)
|
||||||
|
return img
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def ddim_sample(self, classes, shape, cond_scale = 3., clip_denoised = True):
|
||||||
|
batch, device, total_timesteps, sampling_timesteps, eta, objective = shape[0], self.betas.device, self.num_timesteps, self.sampling_timesteps, self.ddim_sampling_eta, self.objective
|
||||||
|
|
||||||
|
times = torch.linspace(-1, total_timesteps - 1, steps=sampling_timesteps + 1) # [-1, 0, 1, 2, ..., T-1] when sampling_timesteps == total_timesteps
|
||||||
|
times = list(reversed(times.int().tolist()))
|
||||||
|
time_pairs = list(zip(times[:-1], times[1:])) # [(T-1, T-2), (T-2, T-3), ..., (1, 0), (0, -1)]
|
||||||
|
|
||||||
|
img = torch.randn(shape, device = device)
|
||||||
|
|
||||||
|
x_start = None
|
||||||
|
|
||||||
|
for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'):
|
||||||
|
time_cond = torch.full((batch,), time, device=device, dtype=torch.long)
|
||||||
|
pred_noise, x_start, *_ = self.model_predictions(img, time_cond, classes, cond_scale = cond_scale, clip_x_start = clip_denoised)
|
||||||
|
|
||||||
|
if time_next < 0:
|
||||||
|
img = x_start
|
||||||
|
continue
|
||||||
|
|
||||||
|
alpha = self.alphas_cumprod[time]
|
||||||
|
alpha_next = self.alphas_cumprod[time_next]
|
||||||
|
|
||||||
|
sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
|
||||||
|
c = (1 - alpha_next - sigma ** 2).sqrt()
|
||||||
|
|
||||||
|
noise = torch.randn_like(img)
|
||||||
|
|
||||||
|
img = x_start * alpha_next.sqrt() + \
|
||||||
|
c * pred_noise + \
|
||||||
|
sigma * noise
|
||||||
|
|
||||||
|
img = unnormalize_to_zero_to_one(img)
|
||||||
|
return img
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def sample(self, classes, cond_scale = 3.):
|
||||||
|
batch_size, image_size, channels = classes.shape[0], self.image_size, self.channels
|
||||||
|
sample_fn = self.p_sample_loop if not self.is_ddim_sampling else self.ddim_sample
|
||||||
|
return sample_fn(classes, (batch_size, channels, image_size, image_size), cond_scale)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def interpolate(self, x1, x2, t = None, lam = 0.5):
|
||||||
|
b, *_, device = *x1.shape, x1.device
|
||||||
|
t = default(t, self.num_timesteps - 1)
|
||||||
|
|
||||||
|
assert x1.shape == x2.shape
|
||||||
|
|
||||||
|
t_batched = torch.stack([torch.tensor(t, device = device)] * b)
|
||||||
|
xt1, xt2 = map(lambda x: self.q_sample(x, t = t_batched), (x1, x2))
|
||||||
|
|
||||||
|
img = (1 - lam) * xt1 + lam * xt2
|
||||||
|
for i in tqdm(reversed(range(0, t)), desc = 'interpolation sample time step', total = t):
|
||||||
|
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long))
|
||||||
|
|
||||||
|
return img
|
||||||
|
|
||||||
|
def q_sample(self, x_start, t, noise=None):
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start +
|
||||||
|
extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def loss_fn(self):
|
||||||
|
if self.loss_type == 'l1':
|
||||||
|
return F.l1_loss
|
||||||
|
elif self.loss_type == 'l2':
|
||||||
|
return F.mse_loss
|
||||||
|
else:
|
||||||
|
raise ValueError(f'invalid loss type {self.loss_type}')
|
||||||
|
|
||||||
|
def p_losses(self, x_start, t, *, classes, noise = None):
|
||||||
|
b, c, h, w = x_start.shape
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
|
||||||
|
# noise sample
|
||||||
|
|
||||||
|
x = self.q_sample(x_start = x_start, t = t, noise = noise)
|
||||||
|
|
||||||
|
# predict and take gradient step
|
||||||
|
|
||||||
|
model_out = self.model(x, t, classes)
|
||||||
|
|
||||||
|
if self.objective == 'pred_noise':
|
||||||
|
target = noise
|
||||||
|
elif self.objective == 'pred_x0':
|
||||||
|
target = x_start
|
||||||
|
elif self.objective == 'pred_v':
|
||||||
|
v = self.predict_v(x_start, t, noise)
|
||||||
|
target = v
|
||||||
|
else:
|
||||||
|
raise ValueError(f'unknown objective {self.objective}')
|
||||||
|
|
||||||
|
loss = self.loss_fn(model_out, target, reduction = 'none')
|
||||||
|
loss = reduce(loss, 'b ... -> b (...)', 'mean')
|
||||||
|
|
||||||
|
loss = loss * extract(self.p2_loss_weight, t, loss.shape)
|
||||||
|
return loss.mean()
|
||||||
|
|
||||||
|
def forward(self, img, *args, **kwargs):
|
||||||
|
b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
|
||||||
|
assert h == img_size and w == img_size, f'height and width of image must be {img_size}'
|
||||||
|
t = torch.randint(0, self.num_timesteps, (b,), device=device).long()
|
||||||
|
|
||||||
|
img = normalize_to_neg_one_to_one(img)
|
||||||
|
return self.p_losses(img, t, *args, **kwargs)
|
||||||
|
|
||||||
|
# example
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
num_classes = 10
|
||||||
|
|
||||||
|
model = Unet(
|
||||||
|
dim = 64,
|
||||||
|
dim_mults = (1, 2, 4, 8),
|
||||||
|
num_classes = num_classes,
|
||||||
|
cond_drop_prob = 0.5
|
||||||
|
)
|
||||||
|
|
||||||
|
diffusion = GaussianDiffusion(
|
||||||
|
model,
|
||||||
|
image_size = 128,
|
||||||
|
timesteps = 1000
|
||||||
|
).cuda()
|
||||||
|
|
||||||
|
training_images = torch.randn(8, 3, 128, 128).cuda() # images are normalized from 0 to 1
|
||||||
|
image_classes = torch.randint(0, num_classes, (8,)).cuda() # say 10 classes
|
||||||
|
|
||||||
|
loss = diffusion(training_images, classes = image_classes)
|
||||||
|
loss.backward()
|
||||||
|
|
||||||
|
# do above for many steps
|
||||||
|
|
||||||
|
sampled_images = diffusion.sample(
|
||||||
|
classes = image_classes,
|
||||||
|
cond_scale = 3. # condition scaling, anything greater than 1 strengthens the classifier free guidance. reportedly 3-8 is good empirically
|
||||||
|
)
|
||||||
|
|
||||||
|
sampled_images.shape # (8, 3, 128, 128)
|
||||||
@@ -0,0 +1,288 @@
|
|||||||
|
import math
|
||||||
|
import torch
|
||||||
|
from torch import sqrt
|
||||||
|
from torch import nn, einsum
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch.special import expm1
|
||||||
|
|
||||||
|
from tqdm import tqdm
|
||||||
|
from einops import rearrange, repeat, reduce
|
||||||
|
from einops.layers.torch import Rearrange
|
||||||
|
|
||||||
|
# helpers
|
||||||
|
|
||||||
|
def exists(val):
|
||||||
|
return val is not None
|
||||||
|
|
||||||
|
def default(val, d):
|
||||||
|
if exists(val):
|
||||||
|
return val
|
||||||
|
return d() if callable(d) else d
|
||||||
|
|
||||||
|
# normalization functions
|
||||||
|
|
||||||
|
def normalize_to_neg_one_to_one(img):
|
||||||
|
return img * 2 - 1
|
||||||
|
|
||||||
|
def unnormalize_to_zero_to_one(t):
|
||||||
|
return (t + 1) * 0.5
|
||||||
|
|
||||||
|
# diffusion helpers
|
||||||
|
|
||||||
|
def right_pad_dims_to(x, t):
|
||||||
|
padding_dims = x.ndim - t.ndim
|
||||||
|
if padding_dims <= 0:
|
||||||
|
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
|
||||||
|
# @crowsonkb Katherine's repository also helped here https://github.com/crowsonkb/v-diffusion-jax/blob/master/diffusion/utils.py
|
||||||
|
|
||||||
|
# log(snr) that approximates the original linear schedule
|
||||||
|
|
||||||
|
def log(t, eps = 1e-20):
|
||||||
|
return torch.log(t.clamp(min = eps))
|
||||||
|
|
||||||
|
def beta_linear_log_snr(t):
|
||||||
|
return -log(expm1(1e-4 + 10 * (t ** 2)))
|
||||||
|
|
||||||
|
def alpha_cosine_log_snr(t, s = 0.008):
|
||||||
|
return -log((torch.cos((t + s) / (1 + s) * math.pi * 0.5) ** -2) - 1, eps = 1e-5)
|
||||||
|
|
||||||
|
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.
|
||||||
|
):
|
||||||
|
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)
|
||||||
|
|
||||||
|
class ContinuousTimeGaussianDiffusion(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model,
|
||||||
|
*,
|
||||||
|
image_size,
|
||||||
|
channels = 3,
|
||||||
|
loss_type = 'l1',
|
||||||
|
noise_schedule = 'linear',
|
||||||
|
num_sample_steps = 500,
|
||||||
|
clip_sample_denoised = True,
|
||||||
|
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
|
||||||
|
p2_loss_weight_gamma = 0., # p2 loss weight, from https://arxiv.org/abs/2204.00227 - 0 is equivalent to weight of 1 across time
|
||||||
|
p2_loss_weight_k = 1
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
assert model.random_or_learned_sinusoidal_cond
|
||||||
|
assert not model.self_condition, 'not supported yet'
|
||||||
|
|
||||||
|
self.model = model
|
||||||
|
|
||||||
|
# image dimensions
|
||||||
|
|
||||||
|
self.channels = channels
|
||||||
|
self.image_size = image_size
|
||||||
|
|
||||||
|
# continuous noise schedule related stuff
|
||||||
|
|
||||||
|
self.loss_type = loss_type
|
||||||
|
|
||||||
|
if noise_schedule == 'linear':
|
||||||
|
self.log_snr = beta_linear_log_snr
|
||||||
|
elif noise_schedule == 'cosine':
|
||||||
|
self.log_snr = alpha_cosine_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}')
|
||||||
|
|
||||||
|
# sampling
|
||||||
|
|
||||||
|
self.num_sample_steps = num_sample_steps
|
||||||
|
self.clip_sample_denoised = clip_sample_denoised
|
||||||
|
|
||||||
|
# p2 loss weight
|
||||||
|
# proposed https://arxiv.org/abs/2204.00227
|
||||||
|
|
||||||
|
assert p2_loss_weight_gamma <= 2, 'in paper, they noticed any gamma greater than 2 is harmful'
|
||||||
|
|
||||||
|
self.p2_loss_weight_gamma = p2_loss_weight_gamma # recommended to be 0.5 or 1
|
||||||
|
self.p2_loss_weight_k = p2_loss_weight_k
|
||||||
|
|
||||||
|
@property
|
||||||
|
def device(self):
|
||||||
|
return next(self.model.parameters()).device
|
||||||
|
|
||||||
|
@property
|
||||||
|
def loss_fn(self):
|
||||||
|
if self.loss_type == 'l1':
|
||||||
|
return F.l1_loss
|
||||||
|
elif self.loss_type == 'l2':
|
||||||
|
return F.mse_loss
|
||||||
|
else:
|
||||||
|
raise ValueError(f'invalid loss type {self.loss_type}')
|
||||||
|
|
||||||
|
def p_mean_variance(self, x, time, time_next):
|
||||||
|
# reviewer found an error in the equation in the paper (missing sigma)
|
||||||
|
# following - https://openreview.net/forum?id=2LdBqxc1Yv¬eId=rIQgH0zKsRt
|
||||||
|
|
||||||
|
log_snr = self.log_snr(time)
|
||||||
|
log_snr_next = self.log_snr(time_next)
|
||||||
|
c = -expm1(log_snr - log_snr_next)
|
||||||
|
|
||||||
|
squared_alpha, squared_alpha_next = log_snr.sigmoid(), log_snr_next.sigmoid()
|
||||||
|
squared_sigma, squared_sigma_next = (-log_snr).sigmoid(), (-log_snr_next).sigmoid()
|
||||||
|
|
||||||
|
alpha, sigma, alpha_next = map(sqrt, (squared_alpha, squared_sigma, squared_alpha_next))
|
||||||
|
|
||||||
|
batch_log_snr = repeat(log_snr, ' -> b', b = x.shape[0])
|
||||||
|
pred_noise = self.model(x, batch_log_snr)
|
||||||
|
|
||||||
|
if self.clip_sample_denoised:
|
||||||
|
x_start = (x - sigma * pred_noise) / alpha
|
||||||
|
|
||||||
|
# in Imagen, this was changed to dynamic thresholding
|
||||||
|
x_start.clamp_(-1., 1.)
|
||||||
|
|
||||||
|
model_mean = alpha_next * (x * (1 - c) / alpha + c * x_start)
|
||||||
|
else:
|
||||||
|
model_mean = alpha_next / alpha * (x - c * sigma * pred_noise)
|
||||||
|
|
||||||
|
posterior_variance = squared_sigma_next * c
|
||||||
|
|
||||||
|
return model_mean, posterior_variance
|
||||||
|
|
||||||
|
# sampling related functions
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def p_sample(self, x, time, time_next):
|
||||||
|
batch, *_, device = *x.shape, x.device
|
||||||
|
|
||||||
|
model_mean, model_variance = self.p_mean_variance(x = x, time = time, time_next = time_next)
|
||||||
|
|
||||||
|
if time_next == 0:
|
||||||
|
return model_mean
|
||||||
|
|
||||||
|
noise = torch.randn_like(x)
|
||||||
|
return model_mean + sqrt(model_variance) * noise
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def p_sample_loop(self, shape):
|
||||||
|
batch = shape[0]
|
||||||
|
|
||||||
|
img = torch.randn(shape, device = self.device)
|
||||||
|
steps = torch.linspace(1., 0., self.num_sample_steps + 1, device = self.device)
|
||||||
|
|
||||||
|
for i in tqdm(range(self.num_sample_steps), desc = 'sampling loop time step', total = self.num_sample_steps):
|
||||||
|
times = steps[i]
|
||||||
|
times_next = steps[i + 1]
|
||||||
|
img = self.p_sample(img, times, times_next)
|
||||||
|
|
||||||
|
img.clamp_(-1., 1.)
|
||||||
|
img = unnormalize_to_zero_to_one(img)
|
||||||
|
return img
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def sample(self, batch_size = 16):
|
||||||
|
return self.p_sample_loop((batch_size, self.channels, self.image_size, self.image_size))
|
||||||
|
|
||||||
|
# training related functions - noise prediction
|
||||||
|
|
||||||
|
def q_sample(self, x_start, times, noise = None):
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
|
||||||
|
log_snr = self.log_snr(times)
|
||||||
|
|
||||||
|
log_snr_padded = right_pad_dims_to(x_start, log_snr)
|
||||||
|
alpha, sigma = sqrt(log_snr_padded.sigmoid()), sqrt((-log_snr_padded).sigmoid())
|
||||||
|
x_noised = x_start * alpha + noise * sigma
|
||||||
|
|
||||||
|
return x_noised, log_snr
|
||||||
|
|
||||||
|
def random_times(self, batch_size):
|
||||||
|
# times are now uniform from 0 to 1
|
||||||
|
return torch.zeros((batch_size,), device = self.device).float().uniform_(0, 1)
|
||||||
|
|
||||||
|
def p_losses(self, x_start, times, noise = None):
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
|
||||||
|
x, log_snr = self.q_sample(x_start = x_start, times = times, noise = noise)
|
||||||
|
model_out = self.model(x, log_snr)
|
||||||
|
|
||||||
|
losses = self.loss_fn(model_out, noise, reduction = 'none')
|
||||||
|
losses = reduce(losses, 'b ... -> b', 'mean')
|
||||||
|
|
||||||
|
if self.p2_loss_weight_gamma >= 0:
|
||||||
|
# following eq 8. in https://arxiv.org/abs/2204.00227
|
||||||
|
loss_weight = (self.p2_loss_weight_k + log_snr.exp()) ** -self.p2_loss_weight_gamma
|
||||||
|
losses = losses * loss_weight
|
||||||
|
|
||||||
|
return losses.mean()
|
||||||
|
|
||||||
|
def forward(self, img, *args, **kwargs):
|
||||||
|
b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
|
||||||
|
assert h == img_size and w == img_size, f'height and width of image must be {img_size}'
|
||||||
|
|
||||||
|
times = self.random_times(b)
|
||||||
|
img = normalize_to_neg_one_to_one(img)
|
||||||
|
return self.p_losses(img, times, *args, **kwargs)
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,695 @@
|
|||||||
|
import math
|
||||||
|
from random import random
|
||||||
|
from functools import partial
|
||||||
|
from collections import namedtuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn, einsum
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from einops import rearrange, reduce
|
||||||
|
from einops.layers.torch import Rearrange
|
||||||
|
|
||||||
|
from tqdm.auto import tqdm
|
||||||
|
|
||||||
|
# constants
|
||||||
|
|
||||||
|
ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start'])
|
||||||
|
|
||||||
|
# helpers functions
|
||||||
|
|
||||||
|
def exists(x):
|
||||||
|
return x is not None
|
||||||
|
|
||||||
|
def default(val, d):
|
||||||
|
if exists(val):
|
||||||
|
return val
|
||||||
|
return d() if callable(d) else d
|
||||||
|
|
||||||
|
def identity(t, *args, **kwargs):
|
||||||
|
return t
|
||||||
|
|
||||||
|
def cycle(dl):
|
||||||
|
while True:
|
||||||
|
for data in dl:
|
||||||
|
yield data
|
||||||
|
|
||||||
|
def has_int_squareroot(num):
|
||||||
|
return (math.sqrt(num) ** 2) == num
|
||||||
|
|
||||||
|
def num_to_groups(num, divisor):
|
||||||
|
groups = num // divisor
|
||||||
|
remainder = num % divisor
|
||||||
|
arr = [divisor] * groups
|
||||||
|
if remainder > 0:
|
||||||
|
arr.append(remainder)
|
||||||
|
return arr
|
||||||
|
|
||||||
|
def convert_image_to_fn(img_type, image):
|
||||||
|
if image.mode != img_type:
|
||||||
|
return image.convert(img_type)
|
||||||
|
return image
|
||||||
|
|
||||||
|
# normalization functions
|
||||||
|
|
||||||
|
def normalize_to_neg_one_to_one(img):
|
||||||
|
return img * 2 - 1
|
||||||
|
|
||||||
|
def unnormalize_to_zero_to_one(t):
|
||||||
|
return (t + 1) * 0.5
|
||||||
|
|
||||||
|
# small helper modules
|
||||||
|
|
||||||
|
class Residual(nn.Module):
|
||||||
|
def __init__(self, fn):
|
||||||
|
super().__init__()
|
||||||
|
self.fn = fn
|
||||||
|
|
||||||
|
def forward(self, x, *args, **kwargs):
|
||||||
|
return self.fn(x, *args, **kwargs) + x
|
||||||
|
|
||||||
|
def Upsample(dim, dim_out = None):
|
||||||
|
return nn.Sequential(
|
||||||
|
nn.Upsample(scale_factor = 2, mode = 'nearest'),
|
||||||
|
nn.Conv1d(dim, default(dim_out, dim), 3, padding = 1)
|
||||||
|
)
|
||||||
|
|
||||||
|
def Downsample(dim, dim_out = None):
|
||||||
|
return nn.Conv1d(dim, default(dim_out, dim), 4, 2, 1)
|
||||||
|
|
||||||
|
class WeightStandardizedConv2d(nn.Conv1d):
|
||||||
|
"""
|
||||||
|
https://arxiv.org/abs/1903.10520
|
||||||
|
weight standardization purportedly works synergistically with group normalization
|
||||||
|
"""
|
||||||
|
def forward(self, x):
|
||||||
|
eps = 1e-5 if x.dtype == torch.float32 else 1e-3
|
||||||
|
|
||||||
|
weight = self.weight
|
||||||
|
mean = reduce(weight, 'o ... -> o 1 1', 'mean')
|
||||||
|
var = reduce(weight, 'o ... -> o 1 1', partial(torch.var, unbiased = False))
|
||||||
|
normalized_weight = (weight - mean) * (var + eps).rsqrt()
|
||||||
|
|
||||||
|
return F.conv1d(x, normalized_weight, self.bias, self.stride, self.padding, self.dilation, self.groups)
|
||||||
|
|
||||||
|
class LayerNorm(nn.Module):
|
||||||
|
def __init__(self, dim):
|
||||||
|
super().__init__()
|
||||||
|
self.g = nn.Parameter(torch.ones(1, dim, 1))
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
eps = 1e-5 if x.dtype == torch.float32 else 1e-3
|
||||||
|
var = torch.var(x, dim = 1, unbiased = False, keepdim = True)
|
||||||
|
mean = torch.mean(x, dim = 1, keepdim = True)
|
||||||
|
return (x - mean) * (var + eps).rsqrt() * self.g
|
||||||
|
|
||||||
|
class PreNorm(nn.Module):
|
||||||
|
def __init__(self, dim, fn):
|
||||||
|
super().__init__()
|
||||||
|
self.fn = fn
|
||||||
|
self.norm = LayerNorm(dim)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
x = self.norm(x)
|
||||||
|
return self.fn(x)
|
||||||
|
|
||||||
|
# sinusoidal positional embeds
|
||||||
|
|
||||||
|
class SinusoidalPosEmb(nn.Module):
|
||||||
|
def __init__(self, dim):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
device = x.device
|
||||||
|
half_dim = self.dim // 2
|
||||||
|
emb = math.log(10000) / (half_dim - 1)
|
||||||
|
emb = torch.exp(torch.arange(half_dim, device=device) * -emb)
|
||||||
|
emb = x[:, None] * emb[None, :]
|
||||||
|
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
|
||||||
|
return emb
|
||||||
|
|
||||||
|
class RandomOrLearnedSinusoidalPosEmb(nn.Module):
|
||||||
|
""" following @crowsonkb 's lead with random (learned optional) sinusoidal pos emb """
|
||||||
|
""" https://github.com/crowsonkb/v-diffusion-jax/blob/master/diffusion/models/danbooru_128.py#L8 """
|
||||||
|
|
||||||
|
def __init__(self, dim, is_random = False):
|
||||||
|
super().__init__()
|
||||||
|
assert (dim % 2) == 0
|
||||||
|
half_dim = dim // 2
|
||||||
|
self.weights = nn.Parameter(torch.randn(half_dim), requires_grad = not is_random)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
x = rearrange(x, 'b -> b 1')
|
||||||
|
freqs = x * rearrange(self.weights, 'd -> 1 d') * 2 * math.pi
|
||||||
|
fouriered = torch.cat((freqs.sin(), freqs.cos()), dim = -1)
|
||||||
|
fouriered = torch.cat((x, fouriered), dim = -1)
|
||||||
|
return fouriered
|
||||||
|
|
||||||
|
# building block modules
|
||||||
|
|
||||||
|
class Block(nn.Module):
|
||||||
|
def __init__(self, dim, dim_out, groups = 8):
|
||||||
|
super().__init__()
|
||||||
|
self.proj = WeightStandardizedConv2d(dim, dim_out, 3, padding = 1)
|
||||||
|
self.norm = nn.GroupNorm(groups, dim_out)
|
||||||
|
self.act = nn.SiLU()
|
||||||
|
|
||||||
|
def forward(self, x, scale_shift = None):
|
||||||
|
x = self.proj(x)
|
||||||
|
x = self.norm(x)
|
||||||
|
|
||||||
|
if exists(scale_shift):
|
||||||
|
scale, shift = scale_shift
|
||||||
|
x = x * (scale + 1) + shift
|
||||||
|
|
||||||
|
x = self.act(x)
|
||||||
|
return x
|
||||||
|
|
||||||
|
class ResnetBlock(nn.Module):
|
||||||
|
def __init__(self, dim, dim_out, *, time_emb_dim = None, groups = 8):
|
||||||
|
super().__init__()
|
||||||
|
self.mlp = nn.Sequential(
|
||||||
|
nn.SiLU(),
|
||||||
|
nn.Linear(time_emb_dim, dim_out * 2)
|
||||||
|
) if exists(time_emb_dim) else None
|
||||||
|
|
||||||
|
self.block1 = Block(dim, dim_out, groups = groups)
|
||||||
|
self.block2 = Block(dim_out, dim_out, groups = groups)
|
||||||
|
self.res_conv = nn.Conv1d(dim, dim_out, 1) if dim != dim_out else nn.Identity()
|
||||||
|
|
||||||
|
def forward(self, x, time_emb = None):
|
||||||
|
|
||||||
|
scale_shift = None
|
||||||
|
if exists(self.mlp) and exists(time_emb):
|
||||||
|
time_emb = self.mlp(time_emb)
|
||||||
|
time_emb = rearrange(time_emb, 'b c -> b c 1')
|
||||||
|
scale_shift = time_emb.chunk(2, dim = 1)
|
||||||
|
|
||||||
|
h = self.block1(x, scale_shift = scale_shift)
|
||||||
|
|
||||||
|
h = self.block2(h)
|
||||||
|
|
||||||
|
return h + self.res_conv(x)
|
||||||
|
|
||||||
|
class LinearAttention(nn.Module):
|
||||||
|
def __init__(self, dim, heads = 4, dim_head = 32):
|
||||||
|
super().__init__()
|
||||||
|
self.scale = dim_head ** -0.5
|
||||||
|
self.heads = heads
|
||||||
|
hidden_dim = dim_head * heads
|
||||||
|
self.to_qkv = nn.Conv1d(dim, hidden_dim * 3, 1, bias = False)
|
||||||
|
|
||||||
|
self.to_out = nn.Sequential(
|
||||||
|
nn.Conv1d(hidden_dim, dim, 1),
|
||||||
|
LayerNorm(dim)
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
b, c, n = x.shape
|
||||||
|
qkv = self.to_qkv(x).chunk(3, dim = 1)
|
||||||
|
q, k, v = map(lambda t: rearrange(t, 'b (h c) n -> b h c n', h = self.heads), qkv)
|
||||||
|
|
||||||
|
q = q.softmax(dim = -2)
|
||||||
|
k = k.softmax(dim = -1)
|
||||||
|
|
||||||
|
q = q * self.scale
|
||||||
|
|
||||||
|
context = torch.einsum('b h d n, b h e n -> b h d e', k, v)
|
||||||
|
|
||||||
|
out = torch.einsum('b h d e, b h d n -> b h e n', context, q)
|
||||||
|
out = rearrange(out, 'b h c n -> b (h c) n', h = self.heads)
|
||||||
|
return self.to_out(out)
|
||||||
|
|
||||||
|
class Attention(nn.Module):
|
||||||
|
def __init__(self, dim, heads = 4, dim_head = 32):
|
||||||
|
super().__init__()
|
||||||
|
self.scale = dim_head ** -0.5
|
||||||
|
self.heads = heads
|
||||||
|
hidden_dim = dim_head * heads
|
||||||
|
|
||||||
|
self.to_qkv = nn.Conv1d(dim, hidden_dim * 3, 1, bias = False)
|
||||||
|
self.to_out = nn.Conv1d(hidden_dim, dim, 1)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
b, c, n = x.shape
|
||||||
|
qkv = self.to_qkv(x).chunk(3, dim = 1)
|
||||||
|
q, k, v = map(lambda t: rearrange(t, 'b (h c) n -> b h c n', h = self.heads), qkv)
|
||||||
|
|
||||||
|
q = q * self.scale
|
||||||
|
|
||||||
|
sim = einsum('b h d i, b h d j -> b h i j', q, k)
|
||||||
|
attn = sim.softmax(dim = -1)
|
||||||
|
out = einsum('b h i j, b h d j -> b h i d', attn, v)
|
||||||
|
|
||||||
|
out = rearrange(out, 'b h n d -> b (h d) n')
|
||||||
|
return self.to_out(out)
|
||||||
|
|
||||||
|
# model
|
||||||
|
|
||||||
|
class Unet1D(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim,
|
||||||
|
init_dim = None,
|
||||||
|
out_dim = None,
|
||||||
|
dim_mults=(1, 2, 4, 8),
|
||||||
|
channels = 3,
|
||||||
|
self_condition = False,
|
||||||
|
resnet_block_groups = 8,
|
||||||
|
learned_variance = False,
|
||||||
|
learned_sinusoidal_cond = False,
|
||||||
|
random_fourier_features = False,
|
||||||
|
learned_sinusoidal_dim = 16
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
# determine dimensions
|
||||||
|
|
||||||
|
self.channels = channels
|
||||||
|
self.self_condition = self_condition
|
||||||
|
input_channels = channels * (2 if self_condition else 1)
|
||||||
|
|
||||||
|
init_dim = default(init_dim, dim)
|
||||||
|
self.init_conv = nn.Conv1d(input_channels, init_dim, 7, padding = 3)
|
||||||
|
|
||||||
|
dims = [init_dim, *map(lambda m: dim * m, dim_mults)]
|
||||||
|
in_out = list(zip(dims[:-1], dims[1:]))
|
||||||
|
|
||||||
|
block_klass = partial(ResnetBlock, groups = resnet_block_groups)
|
||||||
|
|
||||||
|
# time embeddings
|
||||||
|
|
||||||
|
time_dim = dim * 4
|
||||||
|
|
||||||
|
self.random_or_learned_sinusoidal_cond = learned_sinusoidal_cond or random_fourier_features
|
||||||
|
|
||||||
|
if self.random_or_learned_sinusoidal_cond:
|
||||||
|
sinu_pos_emb = RandomOrLearnedSinusoidalPosEmb(learned_sinusoidal_dim, random_fourier_features)
|
||||||
|
fourier_dim = learned_sinusoidal_dim + 1
|
||||||
|
else:
|
||||||
|
sinu_pos_emb = SinusoidalPosEmb(dim)
|
||||||
|
fourier_dim = dim
|
||||||
|
|
||||||
|
self.time_mlp = nn.Sequential(
|
||||||
|
sinu_pos_emb,
|
||||||
|
nn.Linear(fourier_dim, time_dim),
|
||||||
|
nn.GELU(),
|
||||||
|
nn.Linear(time_dim, time_dim)
|
||||||
|
)
|
||||||
|
|
||||||
|
# layers
|
||||||
|
|
||||||
|
self.downs = nn.ModuleList([])
|
||||||
|
self.ups = nn.ModuleList([])
|
||||||
|
num_resolutions = len(in_out)
|
||||||
|
|
||||||
|
for ind, (dim_in, dim_out) in enumerate(in_out):
|
||||||
|
is_last = ind >= (num_resolutions - 1)
|
||||||
|
|
||||||
|
self.downs.append(nn.ModuleList([
|
||||||
|
block_klass(dim_in, dim_in, time_emb_dim = time_dim),
|
||||||
|
block_klass(dim_in, dim_in, time_emb_dim = time_dim),
|
||||||
|
Residual(PreNorm(dim_in, LinearAttention(dim_in))),
|
||||||
|
Downsample(dim_in, dim_out) if not is_last else nn.Conv1d(dim_in, dim_out, 3, padding = 1)
|
||||||
|
]))
|
||||||
|
|
||||||
|
mid_dim = dims[-1]
|
||||||
|
self.mid_block1 = block_klass(mid_dim, mid_dim, time_emb_dim = time_dim)
|
||||||
|
self.mid_attn = Residual(PreNorm(mid_dim, Attention(mid_dim)))
|
||||||
|
self.mid_block2 = block_klass(mid_dim, mid_dim, time_emb_dim = time_dim)
|
||||||
|
|
||||||
|
for ind, (dim_in, dim_out) in enumerate(reversed(in_out)):
|
||||||
|
is_last = ind == (len(in_out) - 1)
|
||||||
|
|
||||||
|
self.ups.append(nn.ModuleList([
|
||||||
|
block_klass(dim_out + dim_in, dim_out, time_emb_dim = time_dim),
|
||||||
|
block_klass(dim_out + dim_in, dim_out, time_emb_dim = time_dim),
|
||||||
|
Residual(PreNorm(dim_out, LinearAttention(dim_out))),
|
||||||
|
Upsample(dim_out, dim_in) if not is_last else nn.Conv1d(dim_out, dim_in, 3, padding = 1)
|
||||||
|
]))
|
||||||
|
|
||||||
|
default_out_dim = channels * (1 if not learned_variance else 2)
|
||||||
|
self.out_dim = default(out_dim, default_out_dim)
|
||||||
|
|
||||||
|
self.final_res_block = block_klass(dim * 2, dim, time_emb_dim = time_dim)
|
||||||
|
self.final_conv = nn.Conv1d(dim, self.out_dim, 1)
|
||||||
|
|
||||||
|
def forward(self, x, time, x_self_cond = None):
|
||||||
|
if self.self_condition:
|
||||||
|
x_self_cond = default(x_self_cond, lambda: torch.zeros_like(x))
|
||||||
|
x = torch.cat((x_self_cond, x), dim = 1)
|
||||||
|
|
||||||
|
x = self.init_conv(x)
|
||||||
|
r = x.clone()
|
||||||
|
|
||||||
|
t = self.time_mlp(time)
|
||||||
|
|
||||||
|
h = []
|
||||||
|
|
||||||
|
for block1, block2, attn, downsample in self.downs:
|
||||||
|
x = block1(x, t)
|
||||||
|
h.append(x)
|
||||||
|
|
||||||
|
x = block2(x, t)
|
||||||
|
x = attn(x)
|
||||||
|
h.append(x)
|
||||||
|
|
||||||
|
x = downsample(x)
|
||||||
|
|
||||||
|
x = self.mid_block1(x, t)
|
||||||
|
x = self.mid_attn(x)
|
||||||
|
x = self.mid_block2(x, t)
|
||||||
|
|
||||||
|
for block1, block2, attn, upsample in self.ups:
|
||||||
|
x = torch.cat((x, h.pop()), dim = 1)
|
||||||
|
x = block1(x, t)
|
||||||
|
|
||||||
|
x = torch.cat((x, h.pop()), dim = 1)
|
||||||
|
x = block2(x, t)
|
||||||
|
x = attn(x)
|
||||||
|
|
||||||
|
x = upsample(x)
|
||||||
|
|
||||||
|
x = torch.cat((x, r), dim = 1)
|
||||||
|
|
||||||
|
x = self.final_res_block(x, t)
|
||||||
|
return self.final_conv(x)
|
||||||
|
|
||||||
|
# gaussian diffusion trainer class
|
||||||
|
|
||||||
|
def extract(a, t, x_shape):
|
||||||
|
b, *_ = t.shape
|
||||||
|
out = a.gather(-1, t)
|
||||||
|
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
||||||
|
|
||||||
|
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
|
||||||
|
as proposed in https://openreview.net/forum?id=-NEXDKk8gZ
|
||||||
|
"""
|
||||||
|
steps = timesteps + 1
|
||||||
|
x = torch.linspace(0, timesteps, steps, dtype = torch.float64)
|
||||||
|
alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * math.pi * 0.5) ** 2
|
||||||
|
alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
|
||||||
|
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
|
||||||
|
return torch.clip(betas, 0, 0.999)
|
||||||
|
|
||||||
|
class GaussianDiffusion1D(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model,
|
||||||
|
*,
|
||||||
|
seq_length,
|
||||||
|
timesteps = 1000,
|
||||||
|
sampling_timesteps = None,
|
||||||
|
loss_type = 'l1',
|
||||||
|
objective = 'pred_noise',
|
||||||
|
beta_schedule = 'cosine',
|
||||||
|
p2_loss_weight_gamma = 0.,
|
||||||
|
p2_loss_weight_k = 1,
|
||||||
|
ddim_sampling_eta = 1.
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.model = model
|
||||||
|
self.channels = self.model.channels
|
||||||
|
self.self_condition = self.model.self_condition
|
||||||
|
|
||||||
|
self.seq_length = seq_length
|
||||||
|
|
||||||
|
self.objective = objective
|
||||||
|
|
||||||
|
assert objective in {'pred_noise', 'pred_x0', 'pred_v'}, 'objective must be either pred_noise (predict noise) or pred_x0 (predict image start) or pred_v (predict v [v-parameterization as defined in appendix D of progressive distillation paper, used in imagen-video successfully])'
|
||||||
|
|
||||||
|
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, dim=0)
|
||||||
|
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value = 1.)
|
||||||
|
|
||||||
|
timesteps, = betas.shape
|
||||||
|
self.num_timesteps = int(timesteps)
|
||||||
|
self.loss_type = loss_type
|
||||||
|
|
||||||
|
# sampling related parameters
|
||||||
|
|
||||||
|
self.sampling_timesteps = default(sampling_timesteps, timesteps) # default num sampling timesteps to number of timesteps at training
|
||||||
|
|
||||||
|
assert self.sampling_timesteps <= timesteps
|
||||||
|
self.is_ddim_sampling = self.sampling_timesteps < timesteps
|
||||||
|
self.ddim_sampling_eta = ddim_sampling_eta
|
||||||
|
|
||||||
|
# helper function to register buffer from float64 to float32
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
register_buffer('sqrt_alphas_cumprod', torch.sqrt(alphas_cumprod))
|
||||||
|
register_buffer('sqrt_one_minus_alphas_cumprod', torch.sqrt(1. - alphas_cumprod))
|
||||||
|
register_buffer('log_one_minus_alphas_cumprod', torch.log(1. - alphas_cumprod))
|
||||||
|
register_buffer('sqrt_recip_alphas_cumprod', torch.sqrt(1. / alphas_cumprod))
|
||||||
|
register_buffer('sqrt_recipm1_alphas_cumprod', torch.sqrt(1. / alphas_cumprod - 1))
|
||||||
|
|
||||||
|
# calculations for posterior q(x_{t-1} | x_t, x_0)
|
||||||
|
|
||||||
|
posterior_variance = betas * (1. - alphas_cumprod_prev) / (1. - alphas_cumprod)
|
||||||
|
|
||||||
|
# above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t)
|
||||||
|
|
||||||
|
register_buffer('posterior_variance', posterior_variance)
|
||||||
|
|
||||||
|
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
|
||||||
|
|
||||||
|
register_buffer('posterior_log_variance_clipped', torch.log(posterior_variance.clamp(min =1e-20)))
|
||||||
|
register_buffer('posterior_mean_coef1', betas * torch.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod))
|
||||||
|
register_buffer('posterior_mean_coef2', (1. - alphas_cumprod_prev) * torch.sqrt(alphas) / (1. - alphas_cumprod))
|
||||||
|
|
||||||
|
# calculate p2 reweighting
|
||||||
|
|
||||||
|
register_buffer('p2_loss_weight', (p2_loss_weight_k + alphas_cumprod / (1 - alphas_cumprod)) ** -p2_loss_weight_gamma)
|
||||||
|
|
||||||
|
def predict_start_from_noise(self, x_t, t, noise):
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t -
|
||||||
|
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * noise
|
||||||
|
)
|
||||||
|
|
||||||
|
def predict_noise_from_start(self, x_t, t, x0):
|
||||||
|
return (
|
||||||
|
(extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - x0) / \
|
||||||
|
extract(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape)
|
||||||
|
)
|
||||||
|
|
||||||
|
def predict_v(self, x_start, t, noise):
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_alphas_cumprod, t, x_start.shape) * noise -
|
||||||
|
extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * x_start
|
||||||
|
)
|
||||||
|
|
||||||
|
def predict_start_from_v(self, x_t, t, v):
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_alphas_cumprod, t, x_t.shape) * x_t -
|
||||||
|
extract(self.sqrt_one_minus_alphas_cumprod, t, x_t.shape) * v
|
||||||
|
)
|
||||||
|
|
||||||
|
def q_posterior(self, x_start, x_t, t):
|
||||||
|
posterior_mean = (
|
||||||
|
extract(self.posterior_mean_coef1, t, x_t.shape) * x_start +
|
||||||
|
extract(self.posterior_mean_coef2, t, x_t.shape) * x_t
|
||||||
|
)
|
||||||
|
posterior_variance = extract(self.posterior_variance, t, x_t.shape)
|
||||||
|
posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, t, x_t.shape)
|
||||||
|
return posterior_mean, posterior_variance, posterior_log_variance_clipped
|
||||||
|
|
||||||
|
def model_predictions(self, x, t, x_self_cond = None, clip_x_start = False):
|
||||||
|
model_output = self.model(x, t, x_self_cond)
|
||||||
|
maybe_clip = partial(torch.clamp, min = -1., max = 1.) if clip_x_start else identity
|
||||||
|
|
||||||
|
if self.objective == 'pred_noise':
|
||||||
|
pred_noise = model_output
|
||||||
|
x_start = self.predict_start_from_noise(x, t, pred_noise)
|
||||||
|
x_start = maybe_clip(x_start)
|
||||||
|
|
||||||
|
elif self.objective == 'pred_x0':
|
||||||
|
x_start = model_output
|
||||||
|
x_start = maybe_clip(x_start)
|
||||||
|
pred_noise = self.predict_noise_from_start(x, t, x_start)
|
||||||
|
|
||||||
|
elif self.objective == 'pred_v':
|
||||||
|
v = model_output
|
||||||
|
x_start = self.predict_start_from_v(x, t, v)
|
||||||
|
x_start = maybe_clip(x_start)
|
||||||
|
pred_noise = self.predict_noise_from_start(x, t, x_start)
|
||||||
|
|
||||||
|
return ModelPrediction(pred_noise, x_start)
|
||||||
|
|
||||||
|
def p_mean_variance(self, x, t, x_self_cond = None, clip_denoised = True):
|
||||||
|
preds = self.model_predictions(x, t, x_self_cond)
|
||||||
|
x_start = preds.pred_x_start
|
||||||
|
|
||||||
|
if clip_denoised:
|
||||||
|
x_start.clamp_(-1., 1.)
|
||||||
|
|
||||||
|
model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start = x_start, x_t = x, t = t)
|
||||||
|
return model_mean, posterior_variance, posterior_log_variance, x_start
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def p_sample(self, x, t: int, x_self_cond = None, clip_denoised = True):
|
||||||
|
b, *_, device = *x.shape, x.device
|
||||||
|
batched_times = torch.full((x.shape[0],), t, device = x.device, dtype = torch.long)
|
||||||
|
model_mean, _, model_log_variance, x_start = self.p_mean_variance(x = x, t = batched_times, x_self_cond = x_self_cond, clip_denoised = clip_denoised)
|
||||||
|
noise = torch.randn_like(x) if t > 0 else 0. # no noise if t == 0
|
||||||
|
pred_img = model_mean + (0.5 * model_log_variance).exp() * noise
|
||||||
|
return pred_img, x_start
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def p_sample_loop(self, shape):
|
||||||
|
batch, device = shape[0], self.betas.device
|
||||||
|
|
||||||
|
img = torch.randn(shape, device=device)
|
||||||
|
|
||||||
|
x_start = None
|
||||||
|
|
||||||
|
for t in tqdm(reversed(range(0, self.num_timesteps)), desc = 'sampling loop time step', total = self.num_timesteps):
|
||||||
|
self_cond = x_start if self.self_condition else None
|
||||||
|
img, x_start = self.p_sample(img, t, self_cond)
|
||||||
|
|
||||||
|
img = unnormalize_to_zero_to_one(img)
|
||||||
|
return img
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def ddim_sample(self, shape, clip_denoised = True):
|
||||||
|
batch, device, total_timesteps, sampling_timesteps, eta, objective = shape[0], self.betas.device, self.num_timesteps, self.sampling_timesteps, self.ddim_sampling_eta, self.objective
|
||||||
|
|
||||||
|
times = torch.linspace(-1, total_timesteps - 1, steps=sampling_timesteps + 1) # [-1, 0, 1, 2, ..., T-1] when sampling_timesteps == total_timesteps
|
||||||
|
times = list(reversed(times.int().tolist()))
|
||||||
|
time_pairs = list(zip(times[:-1], times[1:])) # [(T-1, T-2), (T-2, T-3), ..., (1, 0), (0, -1)]
|
||||||
|
|
||||||
|
img = torch.randn(shape, device = device)
|
||||||
|
|
||||||
|
x_start = None
|
||||||
|
|
||||||
|
for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'):
|
||||||
|
time_cond = torch.full((batch,), time, device=device, dtype=torch.long)
|
||||||
|
self_cond = x_start if self.self_condition else None
|
||||||
|
pred_noise, x_start, *_ = self.model_predictions(img, time_cond, self_cond, clip_x_start = clip_denoised)
|
||||||
|
|
||||||
|
if time_next < 0:
|
||||||
|
img = x_start
|
||||||
|
continue
|
||||||
|
|
||||||
|
alpha = self.alphas_cumprod[time]
|
||||||
|
alpha_next = self.alphas_cumprod[time_next]
|
||||||
|
|
||||||
|
sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
|
||||||
|
c = (1 - alpha_next - sigma ** 2).sqrt()
|
||||||
|
|
||||||
|
noise = torch.randn_like(img)
|
||||||
|
|
||||||
|
img = x_start * alpha_next.sqrt() + \
|
||||||
|
c * pred_noise + \
|
||||||
|
sigma * noise
|
||||||
|
|
||||||
|
img = unnormalize_to_zero_to_one(img)
|
||||||
|
return img
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def sample(self, batch_size = 16):
|
||||||
|
seq_length, channels = self.seq_length, self.channels
|
||||||
|
sample_fn = self.p_sample_loop if not self.is_ddim_sampling else self.ddim_sample
|
||||||
|
return sample_fn((batch_size, channels, seq_length))
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def interpolate(self, x1, x2, t = None, lam = 0.5):
|
||||||
|
b, *_, device = *x1.shape, x1.device
|
||||||
|
t = default(t, self.num_timesteps - 1)
|
||||||
|
|
||||||
|
assert x1.shape == x2.shape
|
||||||
|
|
||||||
|
t_batched = torch.stack([torch.tensor(t, device = device)] * b)
|
||||||
|
xt1, xt2 = map(lambda x: self.q_sample(x, t = t_batched), (x1, x2))
|
||||||
|
|
||||||
|
img = (1 - lam) * xt1 + lam * xt2
|
||||||
|
for i in tqdm(reversed(range(0, t)), desc = 'interpolation sample time step', total = t):
|
||||||
|
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long))
|
||||||
|
|
||||||
|
return img
|
||||||
|
|
||||||
|
def q_sample(self, x_start, t, noise=None):
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
|
||||||
|
return (
|
||||||
|
extract(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start +
|
||||||
|
extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def loss_fn(self):
|
||||||
|
if self.loss_type == 'l1':
|
||||||
|
return F.l1_loss
|
||||||
|
elif self.loss_type == 'l2':
|
||||||
|
return F.mse_loss
|
||||||
|
else:
|
||||||
|
raise ValueError(f'invalid loss type {self.loss_type}')
|
||||||
|
|
||||||
|
def p_losses(self, x_start, t, noise = None):
|
||||||
|
b, c, n = x_start.shape
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
|
||||||
|
# noise sample
|
||||||
|
|
||||||
|
x = self.q_sample(x_start = x_start, t = t, noise = noise)
|
||||||
|
|
||||||
|
# if doing self-conditioning, 50% of the time, predict x_start from current set of times
|
||||||
|
# and condition with unet with that
|
||||||
|
# this technique will slow down training by 25%, but seems to lower FID significantly
|
||||||
|
|
||||||
|
x_self_cond = None
|
||||||
|
if self.self_condition and random() < 0.5:
|
||||||
|
with torch.no_grad():
|
||||||
|
x_self_cond = self.model_predictions(x, t).pred_x_start
|
||||||
|
x_self_cond.detach_()
|
||||||
|
|
||||||
|
# predict and take gradient step
|
||||||
|
|
||||||
|
model_out = self.model(x, t, x_self_cond)
|
||||||
|
|
||||||
|
if self.objective == 'pred_noise':
|
||||||
|
target = noise
|
||||||
|
elif self.objective == 'pred_x0':
|
||||||
|
target = x_start
|
||||||
|
elif self.objective == 'pred_v':
|
||||||
|
v = self.predict_v(x_start, t, noise)
|
||||||
|
target = v
|
||||||
|
else:
|
||||||
|
raise ValueError(f'unknown objective {self.objective}')
|
||||||
|
|
||||||
|
loss = self.loss_fn(model_out, target, reduction = 'none')
|
||||||
|
loss = reduce(loss, 'b ... -> b (...)', 'mean')
|
||||||
|
|
||||||
|
loss = loss * extract(self.p2_loss_weight, t, loss.shape)
|
||||||
|
return loss.mean()
|
||||||
|
|
||||||
|
def forward(self, img, *args, **kwargs):
|
||||||
|
b, c, n, device, seq_length, = *img.shape, img.device, self.seq_length
|
||||||
|
assert n == seq_length, f'seq length must be {seq_length}'
|
||||||
|
t = torch.randint(0, self.num_timesteps, (b,), device=device).long()
|
||||||
|
|
||||||
|
img = normalize_to_neg_one_to_one(img)
|
||||||
|
return self.p_losses(img, t, *args, **kwargs)
|
||||||
@@ -0,0 +1,240 @@
|
|||||||
|
from math import sqrt
|
||||||
|
from random import random
|
||||||
|
import torch
|
||||||
|
from torch import nn, einsum
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from tqdm import tqdm
|
||||||
|
from einops import rearrange, repeat, reduce
|
||||||
|
|
||||||
|
# helpers
|
||||||
|
|
||||||
|
def exists(val):
|
||||||
|
return val is not None
|
||||||
|
|
||||||
|
def default(val, d):
|
||||||
|
if exists(val):
|
||||||
|
return val
|
||||||
|
return d() if callable(d) else d
|
||||||
|
|
||||||
|
# tensor helpers
|
||||||
|
|
||||||
|
def log(t, eps = 1e-20):
|
||||||
|
return torch.log(t.clamp(min = eps))
|
||||||
|
|
||||||
|
# normalization functions
|
||||||
|
|
||||||
|
def normalize_to_neg_one_to_one(img):
|
||||||
|
return img * 2 - 1
|
||||||
|
|
||||||
|
def unnormalize_to_zero_to_one(t):
|
||||||
|
return (t + 1) * 0.5
|
||||||
|
|
||||||
|
# main class
|
||||||
|
|
||||||
|
class ElucidatedDiffusion(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
net,
|
||||||
|
*,
|
||||||
|
image_size,
|
||||||
|
channels = 3,
|
||||||
|
num_sample_steps = 32, # number of sampling steps
|
||||||
|
sigma_min = 0.002, # min noise level
|
||||||
|
sigma_max = 80, # max noise level
|
||||||
|
sigma_data = 0.5, # standard deviation of data distribution
|
||||||
|
rho = 7, # controls the sampling schedule
|
||||||
|
P_mean = -1.2, # mean of log-normal distribution from which noise is drawn for training
|
||||||
|
P_std = 1.2, # standard deviation of log-normal distribution from which noise is drawn for training
|
||||||
|
S_churn = 80, # parameters for stochastic sampling - depends on dataset, Table 5 in apper
|
||||||
|
S_tmin = 0.05,
|
||||||
|
S_tmax = 50,
|
||||||
|
S_noise = 1.003,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
assert net.random_or_learned_sinusoidal_cond
|
||||||
|
self.self_condition = net.self_condition
|
||||||
|
|
||||||
|
self.net = net
|
||||||
|
|
||||||
|
# image dimensions
|
||||||
|
|
||||||
|
self.channels = channels
|
||||||
|
self.image_size = image_size
|
||||||
|
|
||||||
|
# parameters
|
||||||
|
|
||||||
|
self.sigma_min = sigma_min
|
||||||
|
self.sigma_max = sigma_max
|
||||||
|
self.sigma_data = sigma_data
|
||||||
|
|
||||||
|
self.rho = rho
|
||||||
|
|
||||||
|
self.P_mean = P_mean
|
||||||
|
self.P_std = P_std
|
||||||
|
|
||||||
|
self.num_sample_steps = num_sample_steps # otherwise known as N in the paper
|
||||||
|
|
||||||
|
self.S_churn = S_churn
|
||||||
|
self.S_tmin = S_tmin
|
||||||
|
self.S_tmax = S_tmax
|
||||||
|
self.S_noise = S_noise
|
||||||
|
|
||||||
|
@property
|
||||||
|
def device(self):
|
||||||
|
return next(self.net.parameters()).device
|
||||||
|
|
||||||
|
# derived preconditioning params - Table 1
|
||||||
|
|
||||||
|
def c_skip(self, sigma):
|
||||||
|
return (self.sigma_data ** 2) / (sigma ** 2 + self.sigma_data ** 2)
|
||||||
|
|
||||||
|
def c_out(self, sigma):
|
||||||
|
return sigma * self.sigma_data * (self.sigma_data ** 2 + sigma ** 2) ** -0.5
|
||||||
|
|
||||||
|
def c_in(self, sigma):
|
||||||
|
return 1 * (sigma ** 2 + self.sigma_data ** 2) ** -0.5
|
||||||
|
|
||||||
|
def c_noise(self, sigma):
|
||||||
|
return log(sigma) * 0.25
|
||||||
|
|
||||||
|
# preconditioned network output
|
||||||
|
# equation (7) in the paper
|
||||||
|
|
||||||
|
def preconditioned_network_forward(self, noised_images, sigma, self_cond = None, clamp = False):
|
||||||
|
batch, device = noised_images.shape[0], noised_images.device
|
||||||
|
|
||||||
|
if isinstance(sigma, float):
|
||||||
|
sigma = torch.full((batch,), sigma, device = device)
|
||||||
|
|
||||||
|
padded_sigma = rearrange(sigma, 'b -> b 1 1 1')
|
||||||
|
|
||||||
|
net_out = self.net(
|
||||||
|
self.c_in(padded_sigma) * noised_images,
|
||||||
|
self.c_noise(sigma),
|
||||||
|
self_cond
|
||||||
|
)
|
||||||
|
|
||||||
|
out = self.c_skip(padded_sigma) * noised_images + self.c_out(padded_sigma) * net_out
|
||||||
|
|
||||||
|
if clamp:
|
||||||
|
out = out.clamp(-1., 1.)
|
||||||
|
|
||||||
|
return out
|
||||||
|
|
||||||
|
# sampling
|
||||||
|
|
||||||
|
# sample schedule
|
||||||
|
# equation (5) in the paper
|
||||||
|
|
||||||
|
def sample_schedule(self, num_sample_steps = None):
|
||||||
|
num_sample_steps = default(num_sample_steps, self.num_sample_steps)
|
||||||
|
|
||||||
|
N = num_sample_steps
|
||||||
|
inv_rho = 1 / self.rho
|
||||||
|
|
||||||
|
steps = torch.arange(num_sample_steps, device = self.device, dtype = torch.float32)
|
||||||
|
sigmas = (self.sigma_max ** inv_rho + steps / (N - 1) * (self.sigma_min ** inv_rho - self.sigma_max ** inv_rho)) ** self.rho
|
||||||
|
|
||||||
|
sigmas = F.pad(sigmas, (0, 1), value = 0.) # last step is sigma value of 0.
|
||||||
|
return sigmas
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def sample(self, batch_size = 16, num_sample_steps = None, clamp = True):
|
||||||
|
num_sample_steps = default(num_sample_steps, self.num_sample_steps)
|
||||||
|
|
||||||
|
shape = (batch_size, self.channels, self.image_size, self.image_size)
|
||||||
|
|
||||||
|
# get the schedule, which is returned as (sigma, gamma) tuple, and pair up with the next sigma and gamma
|
||||||
|
|
||||||
|
sigmas = self.sample_schedule(num_sample_steps)
|
||||||
|
|
||||||
|
gammas = torch.where(
|
||||||
|
(sigmas >= self.S_tmin) & (sigmas <= self.S_tmax),
|
||||||
|
min(self.S_churn / num_sample_steps, sqrt(2) - 1),
|
||||||
|
0.
|
||||||
|
)
|
||||||
|
|
||||||
|
sigmas_and_gammas = list(zip(sigmas[:-1], sigmas[1:], gammas[:-1]))
|
||||||
|
|
||||||
|
# images is noise at the beginning
|
||||||
|
|
||||||
|
init_sigma = sigmas[0]
|
||||||
|
|
||||||
|
images = init_sigma * torch.randn(shape, device = self.device)
|
||||||
|
|
||||||
|
# for self conditioning
|
||||||
|
|
||||||
|
x_start = None
|
||||||
|
|
||||||
|
# gradually denoise
|
||||||
|
|
||||||
|
for sigma, sigma_next, gamma in tqdm(sigmas_and_gammas, desc = 'sampling time step'):
|
||||||
|
sigma, sigma_next, gamma = map(lambda t: t.item(), (sigma, sigma_next, gamma))
|
||||||
|
|
||||||
|
eps = self.S_noise * torch.randn(shape, device = self.device) # stochastic sampling
|
||||||
|
|
||||||
|
sigma_hat = sigma + gamma * sigma
|
||||||
|
images_hat = images + sqrt(sigma_hat ** 2 - sigma ** 2) * eps
|
||||||
|
|
||||||
|
self_cond = x_start if self.self_condition else None
|
||||||
|
|
||||||
|
model_output = self.preconditioned_network_forward(images_hat, sigma_hat, self_cond, clamp = clamp)
|
||||||
|
denoised_over_sigma = (images_hat - model_output) / sigma_hat
|
||||||
|
|
||||||
|
images_next = images_hat + (sigma_next - sigma_hat) * denoised_over_sigma
|
||||||
|
|
||||||
|
# second order correction, if not the last timestep
|
||||||
|
|
||||||
|
if sigma_next != 0:
|
||||||
|
self_cond = model_output if self.self_condition else None
|
||||||
|
|
||||||
|
model_output_next = self.preconditioned_network_forward(images_next, sigma_next, self_cond, clamp = clamp)
|
||||||
|
denoised_prime_over_sigma = (images_next - model_output_next) / sigma_next
|
||||||
|
images_next = images_hat + 0.5 * (sigma_next - sigma_hat) * (denoised_over_sigma + denoised_prime_over_sigma)
|
||||||
|
|
||||||
|
images = images_next
|
||||||
|
x_start = model_output
|
||||||
|
|
||||||
|
images = images.clamp(-1., 1.)
|
||||||
|
return unnormalize_to_zero_to_one(images)
|
||||||
|
|
||||||
|
# training
|
||||||
|
|
||||||
|
def loss_weight(self, sigma):
|
||||||
|
return (sigma ** 2 + self.sigma_data ** 2) * (sigma * self.sigma_data) ** -2
|
||||||
|
|
||||||
|
def noise_distribution(self, batch_size):
|
||||||
|
return (self.P_mean + self.P_std * torch.randn((batch_size,), device = self.device)).exp()
|
||||||
|
|
||||||
|
def forward(self, images):
|
||||||
|
batch_size, c, h, w, device, image_size, channels = *images.shape, images.device, self.image_size, self.channels
|
||||||
|
|
||||||
|
assert h == image_size and w == image_size, f'height and width of image must be {image_size}'
|
||||||
|
assert c == channels, 'mismatch of image channels'
|
||||||
|
|
||||||
|
images = normalize_to_neg_one_to_one(images)
|
||||||
|
|
||||||
|
sigmas = self.noise_distribution(batch_size)
|
||||||
|
padded_sigmas = rearrange(sigmas, 'b -> b 1 1 1')
|
||||||
|
|
||||||
|
noise = torch.randn_like(images)
|
||||||
|
|
||||||
|
noised_images = images + padded_sigmas * noise # alphas are 1. in the paper
|
||||||
|
|
||||||
|
self_cond = None
|
||||||
|
|
||||||
|
if self.self_condition and random() < 0.5:
|
||||||
|
# from hinton's group's bit diffusion paper
|
||||||
|
with torch.no_grad():
|
||||||
|
self_cond = self.preconditioned_network_forward(noised_images, sigmas)
|
||||||
|
self_cond.detach_()
|
||||||
|
|
||||||
|
denoised = self.preconditioned_network_forward(noised_images, sigmas, self_cond)
|
||||||
|
|
||||||
|
losses = F.mse_loss(denoised, images, reduction = 'none')
|
||||||
|
losses = reduce(losses, 'b ... -> b', 'mean')
|
||||||
|
|
||||||
|
losses = losses * self.loss_weight(sigmas)
|
||||||
|
|
||||||
|
return losses.mean()
|
||||||
@@ -0,0 +1,151 @@
|
|||||||
|
import torch
|
||||||
|
from collections import namedtuple
|
||||||
|
from math import pi, sqrt, log as ln
|
||||||
|
from inspect import isfunction
|
||||||
|
from torch import nn, einsum
|
||||||
|
from einops import rearrange
|
||||||
|
|
||||||
|
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, extract, unnormalize_to_zero_to_one
|
||||||
|
|
||||||
|
# constants
|
||||||
|
|
||||||
|
NAT = 1. / ln(2)
|
||||||
|
|
||||||
|
ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start', 'pred_variance'])
|
||||||
|
|
||||||
|
# helper functions
|
||||||
|
|
||||||
|
def exists(x):
|
||||||
|
return x is not None
|
||||||
|
|
||||||
|
def default(val, d):
|
||||||
|
if exists(val):
|
||||||
|
return val
|
||||||
|
return d() if isfunction(d) else d
|
||||||
|
|
||||||
|
# tensor helpers
|
||||||
|
|
||||||
|
def log(t, eps = 1e-15):
|
||||||
|
return torch.log(t.clamp(min = eps))
|
||||||
|
|
||||||
|
def meanflat(x):
|
||||||
|
return x.mean(dim = tuple(range(1, len(x.shape))))
|
||||||
|
|
||||||
|
def normal_kl(mean1, logvar1, mean2, logvar2):
|
||||||
|
"""
|
||||||
|
KL divergence between normal distributions parameterized by mean and log-variance.
|
||||||
|
"""
|
||||||
|
return 0.5 * (-1.0 + logvar2 - logvar1 + torch.exp(logvar1 - logvar2) + ((mean1 - mean2) ** 2) * torch.exp(-logvar2))
|
||||||
|
|
||||||
|
def approx_standard_normal_cdf(x):
|
||||||
|
return 0.5 * (1.0 + torch.tanh(sqrt(2.0 / pi) * (x + 0.044715 * (x ** 3))))
|
||||||
|
|
||||||
|
def discretized_gaussian_log_likelihood(x, *, means, log_scales, thres = 0.999):
|
||||||
|
assert x.shape == means.shape == log_scales.shape
|
||||||
|
|
||||||
|
centered_x = x - means
|
||||||
|
inv_stdv = torch.exp(-log_scales)
|
||||||
|
plus_in = inv_stdv * (centered_x + 1. / 255.)
|
||||||
|
cdf_plus = approx_standard_normal_cdf(plus_in)
|
||||||
|
min_in = inv_stdv * (centered_x - 1. / 255.)
|
||||||
|
cdf_min = approx_standard_normal_cdf(min_in)
|
||||||
|
log_cdf_plus = log(cdf_plus)
|
||||||
|
log_one_minus_cdf_min = log(1. - cdf_min)
|
||||||
|
cdf_delta = cdf_plus - cdf_min
|
||||||
|
|
||||||
|
log_probs = torch.where(x < -thres,
|
||||||
|
log_cdf_plus,
|
||||||
|
torch.where(x > thres,
|
||||||
|
log_one_minus_cdf_min,
|
||||||
|
log(cdf_delta)))
|
||||||
|
|
||||||
|
return log_probs
|
||||||
|
|
||||||
|
# https://arxiv.org/abs/2102.09672
|
||||||
|
|
||||||
|
# i thought the results were questionable, if one were to focus only on FID
|
||||||
|
# but may as well get this in here for others to try, as GLIDE is using it (and DALL-E2 first stage of cascade)
|
||||||
|
# gaussian diffusion for learned variance + hybrid eps simple + vb loss
|
||||||
|
|
||||||
|
class LearnedGaussianDiffusion(GaussianDiffusion):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model,
|
||||||
|
vb_loss_weight = 0.001, # lambda was 0.001 in the paper
|
||||||
|
*args,
|
||||||
|
**kwargs
|
||||||
|
):
|
||||||
|
super().__init__(model, *args, **kwargs)
|
||||||
|
assert model.out_dim == (model.channels * 2), 'dimension out of unet must be twice the number of channels for learned variance - you can also set the `learned_variance` keyword argument on the Unet to be `True`'
|
||||||
|
assert not model.self_condition, 'not supported yet'
|
||||||
|
|
||||||
|
self.vb_loss_weight = vb_loss_weight
|
||||||
|
|
||||||
|
def model_predictions(self, x, t):
|
||||||
|
model_output = self.model(x, t)
|
||||||
|
model_output, pred_variance = model_output.chunk(2, dim = 1)
|
||||||
|
|
||||||
|
if self.objective == 'pred_noise':
|
||||||
|
pred_noise = model_output
|
||||||
|
x_start = self.predict_start_from_noise(x, t, model_output)
|
||||||
|
|
||||||
|
elif self.objective == 'pred_x0':
|
||||||
|
pred_noise = self.predict_noise_from_start(x, t, model_output)
|
||||||
|
x_start = model_output
|
||||||
|
|
||||||
|
return ModelPrediction(pred_noise, x_start, pred_variance)
|
||||||
|
|
||||||
|
def p_mean_variance(self, *, x, t, clip_denoised, model_output = None):
|
||||||
|
model_output = default(model_output, lambda: self.model(x, t))
|
||||||
|
pred_noise, var_interp_frac_unnormalized = model_output.chunk(2, dim = 1)
|
||||||
|
|
||||||
|
min_log = extract(self.posterior_log_variance_clipped, t, x.shape)
|
||||||
|
max_log = extract(torch.log(self.betas), t, x.shape)
|
||||||
|
var_interp_frac = unnormalize_to_zero_to_one(var_interp_frac_unnormalized)
|
||||||
|
|
||||||
|
model_log_variance = var_interp_frac * max_log + (1 - var_interp_frac) * min_log
|
||||||
|
model_variance = model_log_variance.exp()
|
||||||
|
|
||||||
|
x_start = self.predict_start_from_noise(x, t, pred_noise)
|
||||||
|
|
||||||
|
if clip_denoised:
|
||||||
|
x_start.clamp_(-1., 1.)
|
||||||
|
|
||||||
|
model_mean, _, _ = self.q_posterior(x_start, x, t)
|
||||||
|
|
||||||
|
return model_mean, model_variance, model_log_variance
|
||||||
|
|
||||||
|
def p_losses(self, x_start, t, noise = None, clip_denoised = False):
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
x_t = self.q_sample(x_start = x_start, t = t, noise = noise)
|
||||||
|
|
||||||
|
# model output
|
||||||
|
|
||||||
|
model_output = self.model(x_t, t)
|
||||||
|
|
||||||
|
# calculating kl loss for learned variance (interpolation)
|
||||||
|
|
||||||
|
true_mean, _, true_log_variance_clipped = self.q_posterior(x_start = x_start, x_t = x_t, t = t)
|
||||||
|
model_mean, _, model_log_variance = self.p_mean_variance(x = x_t, t = t, clip_denoised = clip_denoised, model_output = model_output)
|
||||||
|
|
||||||
|
# kl loss with detached model predicted mean, for stability reasons as in paper
|
||||||
|
|
||||||
|
detached_model_mean = model_mean.detach()
|
||||||
|
|
||||||
|
kl = normal_kl(true_mean, true_log_variance_clipped, detached_model_mean, model_log_variance)
|
||||||
|
kl = meanflat(kl) * NAT
|
||||||
|
|
||||||
|
decoder_nll = -discretized_gaussian_log_likelihood(x_start, means = detached_model_mean, log_scales = 0.5 * model_log_variance)
|
||||||
|
decoder_nll = meanflat(decoder_nll) * NAT
|
||||||
|
|
||||||
|
# at the first timestep return the decoder NLL, otherwise return KL(q(x_{t-1}|x_t,x_0) || p(x_{t-1}|x_t))
|
||||||
|
|
||||||
|
vb_losses = torch.where(t == 0, decoder_nll, kl)
|
||||||
|
|
||||||
|
# simple loss - predicting noise, x0, or x_prev
|
||||||
|
|
||||||
|
pred_noise, _ = model_output.chunk(2, dim = 1)
|
||||||
|
|
||||||
|
simple_losses = self.loss_fn(pred_noise, noise)
|
||||||
|
|
||||||
|
return simple_losses + vb_losses.mean() * self.vb_loss_weight
|
||||||
@@ -0,0 +1,184 @@
|
|||||||
|
import math
|
||||||
|
import torch
|
||||||
|
from torch import sqrt
|
||||||
|
from torch import nn, einsum
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch.special import expm1
|
||||||
|
|
||||||
|
from tqdm import tqdm
|
||||||
|
from einops import rearrange, repeat, reduce
|
||||||
|
from einops.layers.torch import Rearrange
|
||||||
|
|
||||||
|
# helpers
|
||||||
|
|
||||||
|
def exists(val):
|
||||||
|
return val is not None
|
||||||
|
|
||||||
|
def default(val, d):
|
||||||
|
if exists(val):
|
||||||
|
return val
|
||||||
|
return d() if callable(d) else d
|
||||||
|
|
||||||
|
# normalization functions
|
||||||
|
|
||||||
|
def normalize_to_neg_one_to_one(img):
|
||||||
|
return img * 2 - 1
|
||||||
|
|
||||||
|
def unnormalize_to_zero_to_one(t):
|
||||||
|
return (t + 1) * 0.5
|
||||||
|
|
||||||
|
# diffusion helpers
|
||||||
|
|
||||||
|
def right_pad_dims_to(x, t):
|
||||||
|
padding_dims = x.ndim - t.ndim
|
||||||
|
if padding_dims <= 0:
|
||||||
|
return t
|
||||||
|
return t.view(*t.shape, *((1,) * padding_dims))
|
||||||
|
|
||||||
|
# continuous schedules
|
||||||
|
# log(snr) that approximates the original linear schedule
|
||||||
|
|
||||||
|
def log(t, eps = 1e-20):
|
||||||
|
return torch.log(t.clamp(min = eps))
|
||||||
|
|
||||||
|
def alpha_cosine_log_snr(t, s = 0.008):
|
||||||
|
return -log((torch.cos((t + s) / (1 + s) * math.pi * 0.5) ** -2) - 1, eps = 1e-5)
|
||||||
|
|
||||||
|
class VParamContinuousTimeGaussianDiffusion(nn.Module):
|
||||||
|
"""
|
||||||
|
a new type of parameterization in v-space proposed in https://arxiv.org/abs/2202.00512 that
|
||||||
|
(1) allows for improved distillation over noise prediction objective and
|
||||||
|
(2) noted in imagen-video to improve upsampling unets by removing the color shifting artifacts
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model,
|
||||||
|
*,
|
||||||
|
image_size,
|
||||||
|
channels = 3,
|
||||||
|
num_sample_steps = 500,
|
||||||
|
clip_sample_denoised = True,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
assert model.random_or_learned_sinusoidal_cond
|
||||||
|
assert not model.self_condition, 'not supported yet'
|
||||||
|
|
||||||
|
self.model = model
|
||||||
|
|
||||||
|
# image dimensions
|
||||||
|
|
||||||
|
self.channels = channels
|
||||||
|
self.image_size = image_size
|
||||||
|
|
||||||
|
# continuous noise schedule related stuff
|
||||||
|
|
||||||
|
self.log_snr = alpha_cosine_log_snr
|
||||||
|
|
||||||
|
# sampling
|
||||||
|
|
||||||
|
self.num_sample_steps = num_sample_steps
|
||||||
|
self.clip_sample_denoised = clip_sample_denoised
|
||||||
|
|
||||||
|
@property
|
||||||
|
def device(self):
|
||||||
|
return next(self.model.parameters()).device
|
||||||
|
|
||||||
|
def p_mean_variance(self, x, time, time_next):
|
||||||
|
# reviewer found an error in the equation in the paper (missing sigma)
|
||||||
|
# following - https://openreview.net/forum?id=2LdBqxc1Yv¬eId=rIQgH0zKsRt
|
||||||
|
|
||||||
|
log_snr = self.log_snr(time)
|
||||||
|
log_snr_next = self.log_snr(time_next)
|
||||||
|
c = -expm1(log_snr - log_snr_next)
|
||||||
|
|
||||||
|
squared_alpha, squared_alpha_next = log_snr.sigmoid(), log_snr_next.sigmoid()
|
||||||
|
squared_sigma, squared_sigma_next = (-log_snr).sigmoid(), (-log_snr_next).sigmoid()
|
||||||
|
|
||||||
|
alpha, sigma, alpha_next = map(sqrt, (squared_alpha, squared_sigma, squared_alpha_next))
|
||||||
|
|
||||||
|
batch_log_snr = repeat(log_snr, ' -> b', b = x.shape[0])
|
||||||
|
|
||||||
|
pred_v = self.model(x, batch_log_snr)
|
||||||
|
|
||||||
|
# shown in Appendix D in the paper
|
||||||
|
x_start = alpha * x - sigma * pred_v
|
||||||
|
|
||||||
|
if self.clip_sample_denoised:
|
||||||
|
x_start.clamp_(-1., 1.)
|
||||||
|
|
||||||
|
model_mean = alpha_next * (x * (1 - c) / alpha + c * x_start)
|
||||||
|
|
||||||
|
posterior_variance = squared_sigma_next * c
|
||||||
|
|
||||||
|
return model_mean, posterior_variance
|
||||||
|
|
||||||
|
# sampling related functions
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def p_sample(self, x, time, time_next):
|
||||||
|
batch, *_, device = *x.shape, x.device
|
||||||
|
|
||||||
|
model_mean, model_variance = self.p_mean_variance(x = x, time = time, time_next = time_next)
|
||||||
|
|
||||||
|
if time_next == 0:
|
||||||
|
return model_mean
|
||||||
|
|
||||||
|
noise = torch.randn_like(x)
|
||||||
|
return model_mean + sqrt(model_variance) * noise
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def p_sample_loop(self, shape):
|
||||||
|
batch = shape[0]
|
||||||
|
|
||||||
|
img = torch.randn(shape, device = self.device)
|
||||||
|
steps = torch.linspace(1., 0., self.num_sample_steps + 1, device = self.device)
|
||||||
|
|
||||||
|
for i in tqdm(range(self.num_sample_steps), desc = 'sampling loop time step', total = self.num_sample_steps):
|
||||||
|
times = steps[i]
|
||||||
|
times_next = steps[i + 1]
|
||||||
|
img = self.p_sample(img, times, times_next)
|
||||||
|
|
||||||
|
img.clamp_(-1., 1.)
|
||||||
|
img = unnormalize_to_zero_to_one(img)
|
||||||
|
return img
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def sample(self, batch_size = 16):
|
||||||
|
return self.p_sample_loop((batch_size, self.channels, self.image_size, self.image_size))
|
||||||
|
|
||||||
|
# training related functions - noise prediction
|
||||||
|
|
||||||
|
def q_sample(self, x_start, times, noise = None):
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
|
||||||
|
log_snr = self.log_snr(times)
|
||||||
|
|
||||||
|
log_snr_padded = right_pad_dims_to(x_start, log_snr)
|
||||||
|
alpha, sigma = sqrt(log_snr_padded.sigmoid()), sqrt((-log_snr_padded).sigmoid())
|
||||||
|
x_noised = x_start * alpha + noise * sigma
|
||||||
|
|
||||||
|
return x_noised, log_snr, alpha, sigma
|
||||||
|
|
||||||
|
def random_times(self, batch_size):
|
||||||
|
return torch.zeros((batch_size,), device = self.device).float().uniform_(0, 1)
|
||||||
|
|
||||||
|
def p_losses(self, x_start, times, noise = None):
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
|
||||||
|
x, log_snr, alpha, sigma = self.q_sample(x_start = x_start, times = times, noise = noise)
|
||||||
|
|
||||||
|
# described in section 4 as the prediction objective, with derivation in Appendix D
|
||||||
|
v = alpha * noise - sigma * x_start
|
||||||
|
|
||||||
|
model_out = self.model(x, log_snr)
|
||||||
|
|
||||||
|
return F.mse_loss(model_out, v)
|
||||||
|
|
||||||
|
def forward(self, img, *args, **kwargs):
|
||||||
|
b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
|
||||||
|
assert h == img_size and w == img_size, f'height and width of image must be {img_size}'
|
||||||
|
|
||||||
|
times = self.random_times(b)
|
||||||
|
img = normalize_to_neg_one_to_one(img)
|
||||||
|
return self.p_losses(img, times, *args, **kwargs)
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
__version__ = '0.1.5'
|
||||||
@@ -0,0 +1,82 @@
|
|||||||
|
import torch
|
||||||
|
from inspect import isfunction
|
||||||
|
from torch import nn, einsum
|
||||||
|
from einops import rearrange
|
||||||
|
|
||||||
|
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion
|
||||||
|
|
||||||
|
# helper functions
|
||||||
|
|
||||||
|
def exists(x):
|
||||||
|
return x is not None
|
||||||
|
|
||||||
|
def default(val, d):
|
||||||
|
if exists(val):
|
||||||
|
return val
|
||||||
|
return d() if isfunction(d) else d
|
||||||
|
|
||||||
|
# some improvisation on my end
|
||||||
|
# where i have the model learn to both predict noise and x0
|
||||||
|
# and learn the weighted sum for each depending on time step
|
||||||
|
|
||||||
|
class WeightedObjectiveGaussianDiffusion(GaussianDiffusion):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model,
|
||||||
|
*args,
|
||||||
|
pred_noise_loss_weight = 0.1,
|
||||||
|
pred_x_start_loss_weight = 0.1,
|
||||||
|
**kwargs
|
||||||
|
):
|
||||||
|
super().__init__(model, *args, **kwargs)
|
||||||
|
channels = model.channels
|
||||||
|
assert model.out_dim == (channels * 2 + 2), 'dimension out (out_dim) of unet must be twice the number of channels + 2 (for the softmax weighted sum) - for channels of 3, this should be (3 * 2) + 2 = 8'
|
||||||
|
assert not model.self_condition, 'not supported yet'
|
||||||
|
assert not self.is_ddim_sampling, 'ddim sampling cannot be used'
|
||||||
|
|
||||||
|
self.split_dims = (channels, channels, 2)
|
||||||
|
self.pred_noise_loss_weight = pred_noise_loss_weight
|
||||||
|
self.pred_x_start_loss_weight = pred_x_start_loss_weight
|
||||||
|
|
||||||
|
def p_mean_variance(self, *, x, t, clip_denoised, model_output = None):
|
||||||
|
model_output = self.model(x, t)
|
||||||
|
|
||||||
|
pred_noise, pred_x_start, weights = model_output.split(self.split_dims, dim = 1)
|
||||||
|
normalized_weights = weights.softmax(dim = 1)
|
||||||
|
|
||||||
|
x_start_from_noise = self.predict_start_from_noise(x, t = t, noise = pred_noise)
|
||||||
|
|
||||||
|
x_starts = torch.stack((x_start_from_noise, pred_x_start), dim = 1)
|
||||||
|
weighted_x_start = einsum('b j h w, b j c h w -> b c h w', normalized_weights, x_starts)
|
||||||
|
|
||||||
|
if clip_denoised:
|
||||||
|
weighted_x_start.clamp_(-1., 1.)
|
||||||
|
|
||||||
|
model_mean, model_variance, model_log_variance = self.q_posterior(weighted_x_start, x, t)
|
||||||
|
|
||||||
|
return model_mean, model_variance, model_log_variance
|
||||||
|
|
||||||
|
def p_losses(self, x_start, t, noise = None, clip_denoised = False):
|
||||||
|
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||||
|
x_t = self.q_sample(x_start = x_start, t = t, noise = noise)
|
||||||
|
|
||||||
|
model_output = self.model(x_t, t)
|
||||||
|
pred_noise, pred_x_start, weights = model_output.split(self.split_dims, dim = 1)
|
||||||
|
|
||||||
|
# get loss for predicted noise and x_start
|
||||||
|
# with the loss weight given at initialization
|
||||||
|
|
||||||
|
noise_loss = self.loss_fn(noise, pred_noise) * self.pred_noise_loss_weight
|
||||||
|
x_start_loss = self.loss_fn(x_start, pred_x_start) * self.pred_x_start_loss_weight
|
||||||
|
|
||||||
|
# calculate x_start from predicted noise
|
||||||
|
# then do a weighted sum of the x_start prediction, weights also predicted by the model (softmax normalized)
|
||||||
|
|
||||||
|
x_start_from_pred_noise = self.predict_start_from_noise(x_t, t, pred_noise)
|
||||||
|
x_start_from_pred_noise = x_start_from_pred_noise.clamp(-2., 2.)
|
||||||
|
weighted_x_start = einsum('b j h w, b j c h w -> b c h w', weights.softmax(dim = 1), torch.stack((x_start_from_pred_noise, pred_x_start), dim = 1))
|
||||||
|
|
||||||
|
# main loss to x_start with the weighted one
|
||||||
|
|
||||||
|
weighted_x_start_loss = self.loss_fn(x_start, weighted_x_start)
|
||||||
|
return weighted_x_start_loss + x_start_loss + noise_loss
|
||||||
Binary file not shown.
|
After Width: | Height: | Size: 40 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 842 KiB |
@@ -1,21 +1,25 @@
|
|||||||
from setuptools import setup, find_packages
|
from setuptools import setup, find_packages
|
||||||
|
|
||||||
|
exec(open('denoising_diffusion_pytorch/version.py').read())
|
||||||
|
|
||||||
setup(
|
setup(
|
||||||
name = 'denoising-diffusion-pytorch',
|
name = 'denoising-diffusion-pytorch',
|
||||||
packages = find_packages(),
|
packages = find_packages(),
|
||||||
version = '0.1.2',
|
version = __version__,
|
||||||
license='MIT',
|
license='MIT',
|
||||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||||
author = 'Phil Wang',
|
author = 'Phil Wang',
|
||||||
author_email = 'lucidrains@gmail.com',
|
author_email = 'lucidrains@gmail.com',
|
||||||
url = 'https://github.com/lucidrains/denoising-diffusion-pytorch',
|
url = 'https://github.com/lucidrains/denoising-diffusion-pytorch',
|
||||||
|
long_description_content_type = 'text/markdown',
|
||||||
keywords = [
|
keywords = [
|
||||||
'artificial intelligence',
|
'artificial intelligence',
|
||||||
'generative models'
|
'generative models'
|
||||||
],
|
],
|
||||||
install_requires=[
|
install_requires=[
|
||||||
|
'accelerate',
|
||||||
'einops',
|
'einops',
|
||||||
'numpy',
|
'ema-pytorch',
|
||||||
'pillow',
|
'pillow',
|
||||||
'torch',
|
'torch',
|
||||||
'torchvision',
|
'torchvision',
|
||||||
@@ -28,4 +32,4 @@ setup(
|
|||||||
'License :: OSI Approved :: MIT License',
|
'License :: OSI Approved :: MIT License',
|
||||||
'Programming Language :: Python :: 3.6',
|
'Programming Language :: Python :: 3.6',
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user