mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-10 12:01:08 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
662172851b | ||
|
|
0248b5e4d3 | ||
|
|
6b56af08a2 |
@@ -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,13 +6,12 @@ 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
|
||||||
@@ -21,6 +20,8 @@ from einops.layers.torch import Rearrange
|
|||||||
from tqdm.auto import tqdm
|
from tqdm.auto import tqdm
|
||||||
from ema_pytorch import EMA
|
from ema_pytorch import EMA
|
||||||
|
|
||||||
|
from accelerate import Accelerator
|
||||||
|
|
||||||
# helpers functions
|
# helpers functions
|
||||||
|
|
||||||
def exists(x):
|
def exists(x):
|
||||||
@@ -36,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
|
||||||
@@ -44,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
|
||||||
|
|
||||||
@@ -562,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):
|
||||||
@@ -592,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')
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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.4',
|
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',
|
||||||
|
|||||||
Reference in New Issue
Block a user