mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-10 12:01:08 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a0c3443eaa | ||
|
|
662172851b | ||
|
|
0248b5e4d3 |
@@ -579,15 +579,18 @@ class Dataset(Dataset):
|
||||
folder,
|
||||
image_size,
|
||||
exts = ['jpg', 'jpeg', 'png', 'tiff'],
|
||||
augment_horizontal_flip = False
|
||||
augment_horizontal_flip = False,
|
||||
convert_image_to = None
|
||||
):
|
||||
super().__init__()
|
||||
self.folder = folder
|
||||
self.image_size = image_size
|
||||
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([
|
||||
T.Lambda(partial(convert_image_to, 'RGB')),
|
||||
T.Lambda(maybe_convert_fn),
|
||||
T.Resize(image_size),
|
||||
T.RandomHorizontalFlip() if augment_horizontal_flip else nn.Identity(),
|
||||
T.CenterCrop(image_size),
|
||||
@@ -622,7 +625,8 @@ class Trainer(object):
|
||||
results_folder = './results',
|
||||
amp = False,
|
||||
fp16 = False,
|
||||
split_batches = True
|
||||
split_batches = True,
|
||||
convert_image_to = None
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -647,7 +651,7 @@ class Trainer(object):
|
||||
|
||||
# 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())
|
||||
|
||||
self.dl = cycle(dl)
|
||||
@@ -676,9 +680,12 @@ class Trainer(object):
|
||||
if not self.accelerator.is_main_process:
|
||||
return
|
||||
|
||||
opt = self.accelerator.unwrap_model(self.opt)
|
||||
|
||||
data = {
|
||||
'step': self.step,
|
||||
'model': self.accelerator.get_state_dict(self.model),
|
||||
'opt': opt.state_dict(),
|
||||
'ema': self.ema.state_dict(),
|
||||
'scaler': self.accelerator.scaler.state_dict() if exists(self.accelerator.scaler) else None
|
||||
}
|
||||
@@ -689,12 +696,15 @@ class Trainer(object):
|
||||
data = torch.load(str(self.results_folder / f'model-{milestone}.pt'))
|
||||
|
||||
model = self.accelerator.unwrap_model(self.model)
|
||||
opt = self.accelerator.unwrap_model(self.opt)
|
||||
|
||||
model.load_state_dict(data['model'])
|
||||
opt.load_state_dict(data['opt'])
|
||||
|
||||
self.step = data['step']
|
||||
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'])
|
||||
|
||||
def train(self):
|
||||
|
||||
@@ -22,7 +22,7 @@ def default(val, d):
|
||||
|
||||
# tensor helpers
|
||||
|
||||
def log(t, eps = 1e-12):
|
||||
def log(t, eps = 1e-15):
|
||||
return torch.log(t.clamp(min = eps))
|
||||
|
||||
def meanflat(x):
|
||||
|
||||
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
|
||||
setup(
|
||||
name = 'denoising-diffusion-pytorch',
|
||||
packages = find_packages(),
|
||||
version = '0.24.0',
|
||||
version = '0.24.4',
|
||||
license='MIT',
|
||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||
author = 'Phil Wang',
|
||||
|
||||
Reference in New Issue
Block a user