mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-11 12:11:43 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ecc6f30901 | ||
|
|
f900f40f14 |
@@ -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
|
||||||
|
|||||||
@@ -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',
|
||||||
|
|||||||
Reference in New Issue
Block a user