Compare commits

...
2 Commits
3 changed files with 19 additions and 10 deletions
+1 -1
View File
@@ -69,7 +69,7 @@ trainer = Trainer(
diffusion,
'path/to/your/images',
train_batch_size = 32,
train_lr = 1e-4,
train_lr = 8e-5,
train_num_steps = 700000, # total training steps
gradient_accumulate_every = 2, # gradient accumulation steps
ema_decay = 0.995, # exponential moving average decay
@@ -58,6 +58,9 @@ def convert_image_to(img_type, image):
return image.convert(img_type)
return image
def l2norm(t):
return F.normalize(t, dim = -1)
# normalization functions
def normalize_to_neg_one_to_one(img):
@@ -215,9 +218,9 @@ class LinearAttention(nn.Module):
return self.to_out(out)
class Attention(nn.Module):
def __init__(self, dim, heads = 4, dim_head = 32):
def __init__(self, dim, heads = 4, dim_head = 32, scale = 16):
super().__init__()
self.scale = dim_head ** -0.5
self.scale = scale
self.heads = heads
hidden_dim = dim_head * heads
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias = False)
@@ -227,10 +230,10 @@ class Attention(nn.Module):
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)
sim = sim - sim.amax(dim = -1, keepdim = True).detach()
q, k = map(l2norm, (q, k))
sim = einsum('b h d i, b h d j -> b h i j', q, k) * self.scale
attn = sim.softmax(dim = -1)
out = einsum('b h i j, b h d j -> b h i d', attn, v)
@@ -681,6 +684,7 @@ class Trainer(object):
train_num_steps = 100000,
ema_update_every = 10,
ema_decay = 0.995,
adam_betas = (0.9, 0.99),
save_and_sample_every = 1000,
num_samples = 25,
results_folder = './results',
@@ -719,7 +723,7 @@ class Trainer(object):
# optimizer
self.opt = Adam(diffusion_model.parameters(), lr = train_lr)
self.opt = Adam(diffusion_model.parameters(), lr = train_lr, betas = adam_betas)
# for logging results in a folder periodically
@@ -772,14 +776,19 @@ class Trainer(object):
while self.step < self.train_num_steps:
total_loss = 0.
for _ in range(self.gradient_accumulate_every):
data = next(self.dl).to(device)
with self.accelerator.autocast():
loss = self.model(data)
self.accelerator.backward(loss / self.gradient_accumulate_every)
loss = loss / self.gradient_accumulate_every
total_loss += loss.item()
pbar.set_description(f'loss: {loss.item():.4f}')
self.accelerator.backward(loss)
pbar.set_description(f'loss: {total_loss:.4f}')
accelerator.wait_for_everyone()
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup(
name = 'denoising-diffusion-pytorch',
packages = find_packages(),
version = '0.25.2',
version = '0.26.0',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',