Compare commits

..
2 Commits
3 changed files with 10 additions and 8 deletions
@@ -5,7 +5,7 @@ import torch.nn.functional as F
from torch.special import expm1 from torch.special import expm1
from tqdm import tqdm from tqdm import tqdm
from einops import rearrange, repeat from einops import rearrange, repeat, reduce
from einops.layers.torch import Rearrange from einops.layers.torch import Rearrange
# helpers # helpers
@@ -268,7 +268,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
model_out = self.denoise_fn(x, log_snr) model_out = self.denoise_fn(x, log_snr)
losses = self.loss_fn(model_out, noise, reduction = 'none') losses = self.loss_fn(model_out, noise, reduction = 'none')
losses = losses.mean(dim = tuple(range(1, losses.ndim))) losses = reduce(losses, 'b ... -> b', 'mean')
if self.p2_loss_weight_gamma >= 0: if self.p2_loss_weight_gamma >= 0:
# following eq 8. in https://arxiv.org/abs/2204.00227 # following eq 8. in https://arxiv.org/abs/2204.00227
@@ -7,6 +7,7 @@ from inspect import isfunction
from functools import partial from functools import partial
from torch.utils import data from torch.utils import data
from multiprocessing import cpu_count
from torch.cuda.amp import autocast, GradScaler from torch.cuda.amp import autocast, GradScaler
from pathlib import Path from pathlib import Path
@@ -541,7 +542,7 @@ class GaussianDiffusion(nn.Module):
# dataset classes # dataset classes
class Dataset(data.Dataset): class Dataset(data.Dataset):
def __init__(self, folder, image_size, exts = ['jpg', 'jpeg', 'png']): def __init__(self, folder, image_size, exts = ['jpg', 'jpeg', 'png'], augment_horizontal_flip = False):
super().__init__() super().__init__()
self.folder = folder self.folder = folder
self.image_size = image_size self.image_size = image_size
@@ -549,7 +550,7 @@ class Dataset(data.Dataset):
self.transform = transforms.Compose([ self.transform = transforms.Compose([
transforms.Resize(image_size), transforms.Resize(image_size),
transforms.RandomHorizontalFlip(), transforms.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(),
transforms.CenterCrop(image_size), transforms.CenterCrop(image_size),
transforms.ToTensor() transforms.ToTensor()
]) ])
@@ -580,7 +581,8 @@ class Trainer(object):
step_start_ema = 2000, step_start_ema = 2000,
update_ema_every = 10, update_ema_every = 10,
save_and_sample_every = 1000, save_and_sample_every = 1000,
results_folder = './results' results_folder = './results',
augment_horizontal_flip = True
): ):
super().__init__() super().__init__()
self.model = diffusion_model self.model = diffusion_model
@@ -596,8 +598,8 @@ class Trainer(object):
self.gradient_accumulate_every = gradient_accumulate_every self.gradient_accumulate_every = gradient_accumulate_every
self.train_num_steps = train_num_steps self.train_num_steps = train_num_steps
self.ds = Dataset(folder, image_size) self.ds = Dataset(folder, image_size, augment_horizontal_flip = augment_horizontal_flip)
self.dl = cycle(data.DataLoader(self.ds, batch_size = train_batch_size, shuffle=True, pin_memory=True)) self.dl = cycle(data.DataLoader(self.ds, batch_size = train_batch_size, shuffle = True, pin_memory = True, num_workers = cpu_count()))
self.opt = Adam(diffusion_model.parameters(), lr=train_lr) self.opt = Adam(diffusion_model.parameters(), lr=train_lr)
self.step = 0 self.step = 0
+1 -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.18.0', version = '0.18.3',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',