mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-11 12:11:43 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
662172851b | ||
|
|
0248b5e4d3 |
@@ -579,15 +579,18 @@ class Dataset(Dataset):
|
|||||||
folder,
|
folder,
|
||||||
image_size,
|
image_size,
|
||||||
exts = ['jpg', 'jpeg', 'png', 'tiff'],
|
exts = ['jpg', 'jpeg', 'png', 'tiff'],
|
||||||
augment_horizontal_flip = False
|
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}')]
|
||||||
|
|
||||||
|
maybe_convert_fn = partial(convert_image_to, convert_image_to) if exists(convert_image_to) else nn.Identity()
|
||||||
|
|
||||||
self.transform = T.Compose([
|
self.transform = T.Compose([
|
||||||
T.Lambda(partial(convert_image_to, 'RGB')),
|
T.Lambda(maybe_convert_fn),
|
||||||
T.Resize(image_size),
|
T.Resize(image_size),
|
||||||
T.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(),
|
T.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(),
|
||||||
T.CenterCrop(image_size),
|
T.CenterCrop(image_size),
|
||||||
@@ -622,7 +625,8 @@ class Trainer(object):
|
|||||||
results_folder = './results',
|
results_folder = './results',
|
||||||
amp = False,
|
amp = False,
|
||||||
fp16 = False,
|
fp16 = False,
|
||||||
split_batches = True
|
split_batches = True,
|
||||||
|
convert_image_to = None
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
@@ -647,7 +651,7 @@ class Trainer(object):
|
|||||||
|
|
||||||
# dataset and dataloader
|
# dataset and dataloader
|
||||||
|
|
||||||
self.ds = Dataset(folder, self.image_size, augment_horizontal_flip = augment_horizontal_flip)
|
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())
|
dl = DataLoader(self.ds, batch_size = train_batch_size, shuffle = True, pin_memory = True, num_workers = cpu_count())
|
||||||
|
|
||||||
self.dl = cycle(dl)
|
self.dl = cycle(dl)
|
||||||
@@ -694,7 +698,7 @@ class Trainer(object):
|
|||||||
self.step = data['step']
|
self.step = data['step']
|
||||||
self.ema.load_state_dict(data['ema'])
|
self.ema.load_state_dict(data['ema'])
|
||||||
|
|
||||||
if exists(self.accelerator.scaler):
|
if exists(self.accelerator.scaler) and exists(data['scaler']):
|
||||||
self.accelerator.scaler.load_state_dict(data['scaler'])
|
self.accelerator.scaler.load_state_dict(data['scaler'])
|
||||||
|
|
||||||
def train(self):
|
def train(self):
|
||||||
|
|||||||
@@ -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.24.0',
|
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',
|
||||||
|
|||||||
Reference in New Issue
Block a user