mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-10 12:01:08 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f762d33c17 | ||
|
|
dfbafee555 | ||
|
|
40dd8ba1de | ||
|
|
2ac3f94a80 | ||
|
|
98f2eeac35 | ||
|
|
6e8a0f2082 | ||
|
|
8c36559295 | ||
|
|
f74f536339 | ||
|
|
d85b8bbe2e | ||
|
|
e0a1bed31a | ||
|
|
82b67fc00a | ||
|
|
7c0cd05c27 | ||
|
|
6dda508ff6 | ||
|
|
9ec8d27217 | ||
|
|
4b4ebab7c3 | ||
|
|
e4a4e4acaa | ||
|
|
cd8329cdd7 | ||
|
|
aec2a26984 | ||
|
|
c78709f887 | ||
|
|
4436128a0b | ||
|
|
e46a89e2bc | ||
|
|
42158d6248 | ||
|
|
44f95e2e9d | ||
|
|
d9275a744c |
@@ -8,8 +8,12 @@ This implementation was transcribed from the official Tensorflow version <a href
|
|||||||
|
|
||||||
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>
|
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>
|
<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>
|
<img src="./images/sample.png" width="500px"><img>
|
||||||
|
|
||||||
[](https://badge.fury.io/py/denoising-diffusion-pytorch)
|
[](https://badge.fury.io/py/denoising-diffusion-pytorch)
|
||||||
@@ -181,3 +185,13 @@ $ accelerate launch train.py
|
|||||||
primaryClass = {cs.CV}
|
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}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|||||||
@@ -37,6 +37,9 @@ def default(val, d):
|
|||||||
return val
|
return val
|
||||||
return d() if callable(d) else d
|
return d() if callable(d) else d
|
||||||
|
|
||||||
|
def identity(t, *args, **kwargs):
|
||||||
|
return t
|
||||||
|
|
||||||
def cycle(dl):
|
def cycle(dl):
|
||||||
while True:
|
while True:
|
||||||
for data in dl:
|
for data in dl:
|
||||||
@@ -53,14 +56,11 @@ def num_to_groups(num, divisor):
|
|||||||
arr.append(remainder)
|
arr.append(remainder)
|
||||||
return arr
|
return arr
|
||||||
|
|
||||||
def convert_image_to(img_type, image):
|
def convert_image_to_fn(img_type, image):
|
||||||
if image.mode != img_type:
|
if image.mode != img_type:
|
||||||
return image.convert(img_type)
|
return image.convert(img_type)
|
||||||
return image
|
return image
|
||||||
|
|
||||||
def l2norm(t):
|
|
||||||
return F.normalize(t, dim = -1)
|
|
||||||
|
|
||||||
# normalization functions
|
# normalization functions
|
||||||
|
|
||||||
def normalize_to_neg_one_to_one(img):
|
def normalize_to_neg_one_to_one(img):
|
||||||
@@ -88,6 +88,21 @@ def Upsample(dim, dim_out = None):
|
|||||||
def Downsample(dim, dim_out = None):
|
def Downsample(dim, dim_out = None):
|
||||||
return nn.Conv2d(dim, default(dim_out, dim), 4, 2, 1)
|
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):
|
class LayerNorm(nn.Module):
|
||||||
def __init__(self, dim):
|
def __init__(self, dim):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -147,7 +162,7 @@ class LearnedSinusoidalPosEmb(nn.Module):
|
|||||||
class Block(nn.Module):
|
class Block(nn.Module):
|
||||||
def __init__(self, dim, dim_out, groups = 8):
|
def __init__(self, dim, dim_out, groups = 8):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.proj = nn.Conv2d(dim, dim_out, 3, padding = 1)
|
self.proj = WeightStandardizedConv2d(dim, dim_out, 3, padding = 1)
|
||||||
self.norm = nn.GroupNorm(groups, dim_out)
|
self.norm = nn.GroupNorm(groups, dim_out)
|
||||||
self.act = nn.SiLU()
|
self.act = nn.SiLU()
|
||||||
|
|
||||||
@@ -219,11 +234,12 @@ class LinearAttention(nn.Module):
|
|||||||
return self.to_out(out)
|
return self.to_out(out)
|
||||||
|
|
||||||
class Attention(nn.Module):
|
class Attention(nn.Module):
|
||||||
def __init__(self, dim, heads = 4, dim_head = 32, scale = 16):
|
def __init__(self, dim, heads = 4, dim_head = 32, scale = 10):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.scale = scale
|
self.scale = dim_head ** -0.5
|
||||||
self.heads = heads
|
self.heads = heads
|
||||||
hidden_dim = dim_head * heads
|
hidden_dim = dim_head * heads
|
||||||
|
|
||||||
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
|
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
|
||||||
self.to_out = nn.Conv2d(hidden_dim, dim, 1)
|
self.to_out = nn.Conv2d(hidden_dim, dim, 1)
|
||||||
|
|
||||||
@@ -232,12 +248,12 @@ class Attention(nn.Module):
|
|||||||
qkv = self.to_qkv(x).chunk(3, dim = 1)
|
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, k, v = map(lambda t: rearrange(t, 'b (h c) x y -> b h c (x y)', h = self.heads), qkv)
|
||||||
|
|
||||||
q, k = map(l2norm, (q, k))
|
q = q * self.scale
|
||||||
|
|
||||||
sim = einsum('b h d i, b h d j -> b h i j', q, k) * self.scale
|
sim = einsum('b h d i, b h d j -> b h i j', q, k)
|
||||||
attn = sim.softmax(dim = -1)
|
attn = sim.softmax(dim = -1)
|
||||||
|
|
||||||
out = einsum('b h i j, b h d j -> b h i d', attn, v)
|
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)
|
out = rearrange(out, 'b h (x y) d -> b (h d) x y', x = h, y = w)
|
||||||
return self.to_out(out)
|
return self.to_out(out)
|
||||||
|
|
||||||
@@ -413,6 +429,7 @@ class GaussianDiffusion(nn.Module):
|
|||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim)
|
assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim)
|
||||||
|
assert not model.learned_sinusoidal_cond
|
||||||
|
|
||||||
self.model = model
|
self.model = model
|
||||||
self.channels = self.model.channels
|
self.channels = self.model.channels
|
||||||
@@ -432,7 +449,7 @@ class GaussianDiffusion(nn.Module):
|
|||||||
raise ValueError(f'unknown beta schedule {beta_schedule}')
|
raise ValueError(f'unknown beta schedule {beta_schedule}')
|
||||||
|
|
||||||
alphas = 1. - betas
|
alphas = 1. - betas
|
||||||
alphas_cumprod = torch.cumprod(alphas, axis=0)
|
alphas_cumprod = torch.cumprod(alphas, dim=0)
|
||||||
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value = 1.)
|
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value = 1.)
|
||||||
|
|
||||||
timesteps, = betas.shape
|
timesteps, = betas.shape
|
||||||
@@ -502,16 +519,19 @@ class GaussianDiffusion(nn.Module):
|
|||||||
posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, 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
|
return posterior_mean, posterior_variance, posterior_log_variance_clipped
|
||||||
|
|
||||||
def model_predictions(self, x, t, x_self_cond = None):
|
def model_predictions(self, x, t, x_self_cond = None, clip_x_start = False):
|
||||||
model_output = self.model(x, t, x_self_cond)
|
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':
|
if self.objective == 'pred_noise':
|
||||||
pred_noise = model_output
|
pred_noise = model_output
|
||||||
x_start = self.predict_start_from_noise(x, t, model_output)
|
x_start = self.predict_start_from_noise(x, t, pred_noise)
|
||||||
|
x_start = maybe_clip(x_start)
|
||||||
|
|
||||||
elif self.objective == 'pred_x0':
|
elif self.objective == 'pred_x0':
|
||||||
pred_noise = self.predict_noise_from_start(x, t, model_output)
|
|
||||||
x_start = model_output
|
x_start = model_output
|
||||||
|
x_start = maybe_clip(x_start)
|
||||||
|
pred_noise = self.predict_noise_from_start(x, t, x_start)
|
||||||
|
|
||||||
return ModelPrediction(pred_noise, x_start)
|
return ModelPrediction(pred_noise, x_start)
|
||||||
|
|
||||||
@@ -553,31 +573,30 @@ class GaussianDiffusion(nn.Module):
|
|||||||
def ddim_sample(self, shape, clip_denoised = True):
|
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
|
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(0., total_timesteps, steps = sampling_timesteps + 2)[:-1]
|
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()))
|
times = list(reversed(times.int().tolist()))
|
||||||
time_pairs = list(zip(times[:-1], times[1:]))
|
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)
|
img = torch.randn(shape, device = device)
|
||||||
|
|
||||||
x_start = None
|
x_start = None
|
||||||
|
|
||||||
for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'):
|
for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'):
|
||||||
alpha = self.alphas_cumprod_prev[time]
|
time_cond = torch.full((batch,), time, device=device, dtype=torch.long)
|
||||||
alpha_next = self.alphas_cumprod_prev[time_next]
|
|
||||||
|
|
||||||
time_cond = torch.full((batch,), time, device = device, dtype = torch.long)
|
|
||||||
|
|
||||||
self_cond = x_start if self.self_condition else None
|
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)
|
||||||
|
|
||||||
pred_noise, x_start, *_ = self.model_predictions(img, time_cond, self_cond)
|
if time_next < 0:
|
||||||
|
img = x_start
|
||||||
|
continue
|
||||||
|
|
||||||
if clip_denoised:
|
alpha = self.alphas_cumprod[time]
|
||||||
x_start.clamp_(-1., 1.)
|
alpha_next = self.alphas_cumprod[time_next]
|
||||||
|
|
||||||
sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
|
sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
|
||||||
c = ((1 - alpha_next) - sigma ** 2).sqrt()
|
c = (1 - alpha_next - sigma ** 2).sqrt()
|
||||||
|
|
||||||
noise = torch.randn_like(img) if time_next > 0 else 0.
|
noise = torch.randn_like(img)
|
||||||
|
|
||||||
img = x_start * alpha_next.sqrt() + \
|
img = x_start * alpha_next.sqrt() + \
|
||||||
c * pred_noise + \
|
c * pred_noise + \
|
||||||
@@ -684,7 +703,7 @@ class Dataset(Dataset):
|
|||||||
self.image_size = image_size
|
self.image_size = image_size
|
||||||
self.paths = [p for ext in exts for p in Path(f'{folder}').glob(f'**/*.{ext}')]
|
self.paths = [p for ext in exts for p in Path(f'{folder}').glob(f'**/*.{ext}')]
|
||||||
|
|
||||||
maybe_convert_fn = partial(convert_image_to, convert_image_to) if exists(convert_image_to) else nn.Identity()
|
maybe_convert_fn = partial(convert_image_to_fn, convert_image_to) if exists(convert_image_to) else nn.Identity()
|
||||||
|
|
||||||
self.transform = T.Compose([
|
self.transform = T.Compose([
|
||||||
T.Lambda(maybe_convert_fn),
|
T.Lambda(maybe_convert_fn),
|
||||||
@@ -790,7 +809,10 @@ class Trainer(object):
|
|||||||
torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
|
torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
|
||||||
|
|
||||||
def load(self, milestone):
|
def load(self, milestone):
|
||||||
data = torch.load(str(self.results_folder / f'model-{milestone}.pt'))
|
accelerator = self.accelerator
|
||||||
|
device = accelerator.device
|
||||||
|
|
||||||
|
data = torch.load(str(self.results_folder / f'model-{milestone}.pt'), map_location=device)
|
||||||
|
|
||||||
model = self.accelerator.unwrap_model(self.model)
|
model = self.accelerator.unwrap_model(self.model)
|
||||||
model.load_state_dict(data['model'])
|
model.load_state_dict(data['model'])
|
||||||
@@ -822,6 +844,7 @@ class Trainer(object):
|
|||||||
|
|
||||||
self.accelerator.backward(loss)
|
self.accelerator.backward(loss)
|
||||||
|
|
||||||
|
accelerator.clip_grad_norm_(self.model.parameters(), 1.0)
|
||||||
pbar.set_description(f'loss: {total_loss:.4f}')
|
pbar.set_description(f'loss: {total_loss:.4f}')
|
||||||
|
|
||||||
accelerator.wait_for_everyone()
|
accelerator.wait_for_everyone()
|
||||||
@@ -831,6 +854,7 @@ class Trainer(object):
|
|||||||
|
|
||||||
accelerator.wait_for_everyone()
|
accelerator.wait_for_everyone()
|
||||||
|
|
||||||
|
self.step += 1
|
||||||
if accelerator.is_main_process:
|
if accelerator.is_main_process:
|
||||||
self.ema.to(device)
|
self.ema.to(device)
|
||||||
self.ema.update()
|
self.ema.update()
|
||||||
@@ -847,7 +871,6 @@ class Trainer(object):
|
|||||||
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = int(math.sqrt(self.num_samples)))
|
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = int(math.sqrt(self.num_samples)))
|
||||||
self.save(milestone)
|
self.save(milestone)
|
||||||
|
|
||||||
self.step += 1
|
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
|
|
||||||
accelerator.print('training complete')
|
accelerator.print('training complete')
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
|
|||||||
setup(
|
setup(
|
||||||
name = 'denoising-diffusion-pytorch',
|
name = 'denoising-diffusion-pytorch',
|
||||||
packages = find_packages(),
|
packages = find_packages(),
|
||||||
version = '0.27.2',
|
version = '0.28.0',
|
||||||
license='MIT',
|
license='MIT',
|
||||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||||
author = 'Phil Wang',
|
author = 'Phil Wang',
|
||||||
@@ -30,4 +30,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