Compare commits

...
6 Commits
5 changed files with 164 additions and 85 deletions
+16
View File
@@ -80,6 +80,22 @@ trainer.train()
Samples and model checkpoints will be logged to `./results` periodically 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
```
## Citations ## Citations
```bibtex ```bibtex
@@ -6,25 +6,21 @@ import torch.nn.functional as F
from inspect import isfunction from inspect import isfunction
from functools import partial from functools import partial
from torch.utils import data from torch.utils.data import Dataset, DataLoader
from multiprocessing import cpu_count from multiprocessing import cpu_count
from torch.cuda.amp import autocast, GradScaler
from pathlib import Path from pathlib import Path
from torch.optim import Adam from torch.optim import Adam
from torchvision import transforms, utils from torchvision import transforms as T, utils
from PIL import Image from PIL import Image
from einops import rearrange, reduce from einops import rearrange, reduce
from einops.layers.torch import Rearrange from einops.layers.torch import Rearrange
from tqdm.auto import tqdm
from ema_pytorch import EMA from ema_pytorch import EMA
import sys from accelerate import Accelerator
if 'ipykernel' in sys.modules:
from tqdm.notebook import tqdm
else:
from tqdm import tqdm
# helpers functions # helpers functions
@@ -41,6 +37,9 @@ def cycle(dl):
for data in dl: for data in dl:
yield data yield data
def has_int_squareroot(num):
return (math.sqrt(num) ** 2) == num
def num_to_groups(num, divisor): def num_to_groups(num, divisor):
groups = num // divisor groups = num // divisor
remainder = num % divisor remainder = num % divisor
@@ -49,6 +48,13 @@ def num_to_groups(num, divisor):
arr.append(remainder) arr.append(remainder)
return arr return arr
def convert_image_to(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): def normalize_to_neg_one_to_one(img):
return img * 2 - 1 return img * 2 - 1
@@ -567,18 +573,28 @@ class GaussianDiffusion(nn.Module):
# dataset classes # dataset classes
class Dataset(data.Dataset): class Dataset(Dataset):
def __init__(self, folder, image_size, exts = ['jpg', 'jpeg', 'png'], augment_horizontal_flip = False): def __init__(
self,
folder,
image_size,
exts = ['jpg', 'jpeg', 'png', 'tiff'],
augment_horizontal_flip = False,
convert_image_to = None
):
super().__init__() super().__init__()
self.folder = folder self.folder = folder
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}')]
self.transform = transforms.Compose([ maybe_convert_fn = partial(convert_image_to, convert_image_to) if exists(convert_image_to) else nn.Identity()
transforms.Resize(image_size),
transforms.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(), self.transform = T.Compose([
transforms.CenterCrop(image_size), T.Lambda(maybe_convert_fn),
transforms.ToTensor() T.Resize(image_size),
T.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(),
T.CenterCrop(image_size),
T.ToTensor()
]) ])
def __len__(self): def __len__(self):
@@ -597,92 +613,135 @@ class Trainer(object):
diffusion_model, diffusion_model,
folder, folder,
*, *,
ema_decay = 0.995, train_batch_size = 16,
train_batch_size = 32, gradient_accumulate_every = 1,
augment_horizontal_flip = True,
train_lr = 1e-4, train_lr = 1e-4,
train_num_steps = 100000, train_num_steps = 100000,
gradient_accumulate_every = 2,
amp = False,
step_start_ema = 2000,
ema_update_every = 10, ema_update_every = 10,
ema_decay = 0.995,
save_and_sample_every = 1000, save_and_sample_every = 1000,
num_samples = 25,
results_folder = './results', results_folder = './results',
augment_horizontal_flip = True amp = False,
fp16 = False,
split_batches = True,
convert_image_to = None
): ):
super().__init__() super().__init__()
self.image_size = diffusion_model.image_size
self.accelerator = Accelerator(
split_batches = split_batches,
mixed_precision = 'fp16' if fp16 else 'no'
)
self.accelerator.native_amp = amp
self.model = diffusion_model self.model = diffusion_model
self.ema = EMA(diffusion_model, beta = ema_decay, update_every = ema_update_every)
self.step_start_ema = step_start_ema assert has_int_squareroot(num_samples), 'number of samples must have an integer square root'
self.num_samples = num_samples
self.save_and_sample_every = save_and_sample_every self.save_and_sample_every = save_and_sample_every
self.batch_size = train_batch_size self.batch_size = train_batch_size
self.image_size = diffusion_model.image_size
self.gradient_accumulate_every = gradient_accumulate_every self.gradient_accumulate_every = gradient_accumulate_every
self.train_num_steps = train_num_steps
self.ds = Dataset(folder, self.image_size, augment_horizontal_flip = augment_horizontal_flip) self.train_num_steps = train_num_steps
self.dl = cycle(data.DataLoader(self.ds, batch_size = train_batch_size, shuffle = True, pin_memory = True, num_workers = cpu_count())) self.image_size = diffusion_model.image_size
# dataset and dataloader
self.ds = Dataset(folder, self.image_size, augment_horizontal_flip = augment_horizontal_flip, convert_image_to = convert_image_to)
dl = DataLoader(self.ds, batch_size = train_batch_size, shuffle = True, pin_memory = True, num_workers = cpu_count())
self.dl = cycle(dl)
# optimizer
self.opt = Adam(diffusion_model.parameters(), lr = train_lr) self.opt = Adam(diffusion_model.parameters(), lr = train_lr)
# for logging results in a folder periodically
if self.accelerator.is_main_process:
self.ema = EMA(diffusion_model, beta = ema_decay, update_every = ema_update_every)
self.results_folder = Path(results_folder)
self.results_folder.mkdir(exist_ok = True)
# step counter state
self.step = 0 self.step = 0
self.amp = amp # prepare model, dataloader, optimizer with accelerator
self.scaler = GradScaler(enabled = amp)
self.results_folder = Path(results_folder) self.model, self.dl, self.opt = self.accelerator.prepare(self.model, self.dl, self.opt)
self.results_folder.mkdir(exist_ok = True)
def save(self, milestone): def save(self, milestone):
if not self.accelerator.is_main_process:
return
data = { data = {
'step': self.step, 'step': self.step,
'model': self.model.state_dict(), 'model': self.accelerator.get_state_dict(self.model),
'ema': self.ema.state_dict(), 'ema': self.ema.state_dict(),
'scaler': self.scaler.state_dict() 'scaler': self.accelerator.scaler.state_dict() if exists(self.accelerator.scaler) else None
} }
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')) data = torch.load(str(self.results_folder / f'model-{milestone}.pt'))
model = self.accelerator.unwrap_model(self.model)
model.load_state_dict(data['model'])
self.step = data['step'] self.step = data['step']
self.model.load_state_dict(data['model'])
self.ema.load_state_dict(data['ema']) self.ema.load_state_dict(data['ema'])
self.scaler.load_state_dict(data['scaler'])
if exists(self.accelerator.scaler) and exists(data['scaler']):
self.accelerator.scaler.load_state_dict(data['scaler'])
def train(self): def train(self):
with tqdm(initial = self.step, total = self.train_num_steps) as pbar: accelerator = self.accelerator
device = accelerator.device
with tqdm(initial = self.step, total = self.train_num_steps, disable = not accelerator.is_main_process) as pbar:
while self.step < self.train_num_steps: while self.step < self.train_num_steps:
for i in range(self.gradient_accumulate_every):
data = next(self.dl).cuda()
with autocast(enabled = self.amp): for _ in range(self.gradient_accumulate_every):
data = next(self.dl).to(device)
with self.accelerator.autocast():
loss = self.model(data) loss = self.model(data)
self.scaler.scale(loss / self.gradient_accumulate_every).backward() self.accelerator.backward(loss / self.gradient_accumulate_every)
pbar.set_description(f'loss: {loss.item():.4f}') pbar.set_description(f'loss: {loss.item():.4f}')
self.scaler.step(self.opt) accelerator.wait_for_everyone()
self.scaler.update()
self.opt.step()
self.opt.zero_grad() self.opt.zero_grad()
self.ema.update() accelerator.wait_for_everyone()
if self.step != 0 and self.step % self.save_and_sample_every == 0: if accelerator.is_main_process:
self.ema.ema_model.eval() self.ema.to(device)
with torch.no_grad(): self.ema.update()
milestone = self.step // self.save_and_sample_every
batches = num_to_groups(36, self.batch_size)
all_images_list = list(map(lambda n: self.ema.ema_model.sample(batch_size=n), batches))
all_images = torch.cat(all_images_list, dim=0) if self.step != 0 and self.step % self.save_and_sample_every == 0:
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6) self.ema.ema_model.eval()
self.save(milestone)
with torch.no_grad():
milestone = self.step // self.save_and_sample_every
batches = num_to_groups(self.num_samples, self.batch_size)
all_images_list = list(map(lambda n: self.ema.ema_model.sample(batch_size=n), batches))
all_images = torch.cat(all_images_list, dim = 0)
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = int(math.sqrt(self.num_samples)))
self.save(milestone)
self.step += 1 self.step += 1
pbar.update(1) pbar.update(1)
print('training complete') accelerator.print('training complete')
@@ -96,13 +96,30 @@ class ElucidatedDiffusion(nn.Module):
def c_noise(self, sigma): def c_noise(self, sigma):
return log(sigma) * 0.25 return log(sigma) * 0.25
# noise distribution # preconditioned network output
# equation (7) in the paper
def noise_distribution(self, batch_size): def preconditioned_network_forward(self, noised_images, sigma, clamp = False):
return (self.P_mean + self.P_std * torch.randn((batch_size,), device = self.device)).exp() batch, device = noised_images.shape[0], noised_images.device
def loss_weight(self, sigma): if isinstance(sigma, float):
return (sigma ** 2 + self.sigma_data ** 2) * (sigma * self.sigma_data) ** -2 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)
)
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 # sample schedule
# equation (5) in the paper # equation (5) in the paper
@@ -119,28 +136,8 @@ class ElucidatedDiffusion(nn.Module):
sigmas = F.pad(sigmas, (0, 1), value = 0.) # last step is sigma value of 0. sigmas = F.pad(sigmas, (0, 1), value = 0.) # last step is sigma value of 0.
return sigmas return sigmas
# preconditioned network output
# equation (7) in the paper
def preconditioned_network_forward(self, noised_images, sigma):
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)
)
return self.c_skip(padded_sigma) * noised_images + self.c_out(padded_sigma) * net_out
# sampling
@torch.no_grad() @torch.no_grad()
def sample(self, batch_size = 16, num_sample_steps = None): def sample(self, batch_size = 16, num_sample_steps = None, clamp = True):
num_sample_steps = default(num_sample_steps, self.num_sample_steps) num_sample_steps = default(num_sample_steps, self.num_sample_steps)
shape = (batch_size, self.channels, self.image_size, self.image_size) shape = (batch_size, self.channels, self.image_size, self.image_size)
@@ -173,7 +170,7 @@ class ElucidatedDiffusion(nn.Module):
sigma_hat = sigma + gamma * sigma sigma_hat = sigma + gamma * sigma
images_hat = images + sqrt(sigma_hat ** 2 - sigma ** 2) * eps images_hat = images + sqrt(sigma_hat ** 2 - sigma ** 2) * eps
model_output = self.preconditioned_network_forward(images_hat, sigma_hat) model_output = self.preconditioned_network_forward(images_hat, sigma_hat, clamp = clamp)
denoised_over_sigma = (images_hat - model_output) / sigma_hat denoised_over_sigma = (images_hat - model_output) / sigma_hat
images_next = images_hat + (sigma_next - sigma_hat) * denoised_over_sigma images_next = images_hat + (sigma_next - sigma_hat) * denoised_over_sigma
@@ -181,7 +178,7 @@ class ElucidatedDiffusion(nn.Module):
# second order correction, if not the last timestep # second order correction, if not the last timestep
if sigma_next != 0: if sigma_next != 0:
model_output_next = self.preconditioned_network_forward(images_next, sigma_next) model_output_next = self.preconditioned_network_forward(images_next, sigma_next, clamp = clamp)
denoised_prime_over_sigma = (images_next - model_output_next) / sigma_next 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_next = images_hat + 0.5 * (sigma_next - sigma_hat) * (denoised_over_sigma + denoised_prime_over_sigma)
@@ -192,6 +189,12 @@ class ElucidatedDiffusion(nn.Module):
# training # 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): def forward(self, images):
batch_size, c, h, w, device, image_size, channels = *images.shape, images.device, self.image_size, self.channels batch_size, c, h, w, device, image_size, channels = *images.shape, images.device, self.image_size, self.channels
@@ -22,7 +22,7 @@ def default(val, d):
# tensor helpers # tensor helpers
def log(t, eps = 1e-12): def log(t, eps = 1e-15):
return torch.log(t.clamp(min = eps)) return torch.log(t.clamp(min = eps))
def meanflat(x): def meanflat(x):
+2 -1
View File
@@ -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.23.2', version = '0.24.2',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',
@@ -15,6 +15,7 @@ setup(
'generative models' 'generative models'
], ],
install_requires=[ install_requires=[
'accelerate',
'einops', 'einops',
'ema-pytorch', 'ema-pytorch',
'pillow', 'pillow',