mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-10 12:01:08 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
27853452a2 |
@@ -64,7 +64,7 @@ trainer = Trainer(
|
||||
diffusion,
|
||||
'path/to/your/images',
|
||||
train_batch_size = 32,
|
||||
train_lr = 1e-4,
|
||||
train_lr = 2e-5,
|
||||
train_num_steps = 700000, # total training steps
|
||||
gradient_accumulate_every = 2, # gradient accumulation steps
|
||||
ema_decay = 0.995, # exponential moving average decay
|
||||
|
||||
@@ -439,8 +439,6 @@ class GaussianDiffusion(nn.Module):
|
||||
|
||||
for i in tqdm(reversed(range(0, self.num_timesteps)), desc='sampling loop time step', total=self.num_timesteps):
|
||||
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long))
|
||||
|
||||
img = unnormalize_to_zero_to_one(img)
|
||||
return img
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -499,13 +497,11 @@ class GaussianDiffusion(nn.Module):
|
||||
loss = self.loss_fn(model_out, target)
|
||||
return loss
|
||||
|
||||
def forward(self, img, *args, **kwargs):
|
||||
b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
|
||||
def forward(self, x, *args, **kwargs):
|
||||
b, c, h, w, device, img_size, = *x.shape, x.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)
|
||||
return self.p_losses(x, t, *args, **kwargs)
|
||||
|
||||
# dataset classes
|
||||
|
||||
@@ -520,7 +516,8 @@ class Dataset(data.Dataset):
|
||||
transforms.Resize(image_size),
|
||||
transforms.RandomHorizontalFlip(),
|
||||
transforms.CenterCrop(image_size),
|
||||
transforms.ToTensor()
|
||||
transforms.ToTensor(),
|
||||
transforms.Lambda(normalize_to_neg_one_to_one)
|
||||
])
|
||||
|
||||
def __len__(self):
|
||||
@@ -542,7 +539,7 @@ class Trainer(object):
|
||||
ema_decay = 0.995,
|
||||
image_size = 128,
|
||||
train_batch_size = 32,
|
||||
train_lr = 1e-4,
|
||||
train_lr = 2e-5,
|
||||
train_num_steps = 100000,
|
||||
gradient_accumulate_every = 2,
|
||||
amp = False,
|
||||
@@ -606,36 +603,34 @@ class Trainer(object):
|
||||
self.scaler.load_state_dict(data['scaler'])
|
||||
|
||||
def train(self):
|
||||
with tqdm(initial = self.step, total = self.train_num_steps) as pbar:
|
||||
while self.step < self.train_num_steps:
|
||||
for i in range(self.gradient_accumulate_every):
|
||||
data = next(self.dl).cuda()
|
||||
|
||||
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):
|
||||
loss = self.model(data)
|
||||
self.scaler.scale(loss / self.gradient_accumulate_every).backward()
|
||||
|
||||
with autocast(enabled = self.amp):
|
||||
loss = self.model(data)
|
||||
self.scaler.scale(loss / self.gradient_accumulate_every).backward()
|
||||
print(f'{self.step}: {loss.item()}')
|
||||
|
||||
pbar.set_description(f'loss: {loss.item():.4f}')
|
||||
self.scaler.step(self.opt)
|
||||
self.scaler.update()
|
||||
self.opt.zero_grad()
|
||||
|
||||
self.scaler.step(self.opt)
|
||||
self.scaler.update()
|
||||
self.opt.zero_grad()
|
||||
if self.step % self.update_ema_every == 0:
|
||||
self.step_ema()
|
||||
|
||||
if self.step % self.update_ema_every == 0:
|
||||
self.step_ema()
|
||||
if self.step != 0 and self.step % self.save_and_sample_every == 0:
|
||||
self.ema_model.eval()
|
||||
|
||||
if self.step != 0 and self.step % self.save_and_sample_every == 0:
|
||||
self.ema_model.eval()
|
||||
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_model.sample(batch_size=n), batches))
|
||||
all_images = torch.cat(all_images_list, dim=0)
|
||||
all_images = unnormalize_to_zero_to_one(all_images)
|
||||
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6)
|
||||
self.save(milestone)
|
||||
|
||||
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_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 = 6)
|
||||
self.save(milestone)
|
||||
self.step += 1
|
||||
|
||||
self.step += 1
|
||||
pbar.update(1)
|
||||
|
||||
print('training complete')
|
||||
print('training completed')
|
||||
|
||||
@@ -3,7 +3,7 @@ from inspect import isfunction
|
||||
from torch import nn, einsum
|
||||
from einops import rearrange
|
||||
|
||||
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion
|
||||
from denoising_diffusion_pytorch.denoising_diffusion_pytorch import GaussianDiffusion, extract, unnormalize_to_zero_to_one
|
||||
|
||||
# helper functions
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
|
||||
setup(
|
||||
name = 'denoising-diffusion-pytorch',
|
||||
packages = find_packages(),
|
||||
version = '0.15.7',
|
||||
version = '0.15.1',
|
||||
license='MIT',
|
||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||
author = 'Phil Wang',
|
||||
|
||||
Reference in New Issue
Block a user