Compare commits

..
4 Commits
4 changed files with 14 additions and 5 deletions
@@ -488,7 +488,7 @@ class GaussianDiffusion(nn.Module):
beta_schedule = 'cosine', beta_schedule = 'cosine',
p2_loss_weight_gamma = 0., # p2 loss weight, from https://arxiv.org/abs/2204.00227 - 0 is equivalent to weight of 1 across time - 1. is recommended p2_loss_weight_gamma = 0., # p2 loss weight, from https://arxiv.org/abs/2204.00227 - 0 is equivalent to weight of 1 across time - 1. is recommended
p2_loss_weight_k = 1, p2_loss_weight_k = 1,
ddim_sampling_eta = 1. ddim_sampling_eta = 0.
): ):
super().__init__() super().__init__()
assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim) assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim)
@@ -23,6 +23,8 @@ from ema_pytorch import EMA
from accelerate import Accelerator from accelerate import Accelerator
from denoising_diffusion_pytorch.version import __version__
# constants # constants
ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start']) ModelPrediction = namedtuple('ModelPrediction', ['pred_noise', 'pred_x_start'])
@@ -808,8 +810,8 @@ class Trainer(object):
if self.accelerator.is_main_process: if self.accelerator.is_main_process:
self.ema = EMA(diffusion_model, beta = ema_decay, update_every = ema_update_every) self.ema = EMA(diffusion_model, beta = ema_decay, update_every = ema_update_every)
self.results_folder = Path(results_folder) self.results_folder = Path(results_folder)
self.results_folder.mkdir(exist_ok = True) self.results_folder.mkdir(exist_ok = True)
# step counter state # step counter state
@@ -828,7 +830,8 @@ class Trainer(object):
'model': self.accelerator.get_state_dict(self.model), 'model': self.accelerator.get_state_dict(self.model),
'opt': self.opt.state_dict(), 'opt': self.opt.state_dict(),
'ema': self.ema.state_dict(), 'ema': self.ema.state_dict(),
'scaler': self.accelerator.scaler.state_dict() if exists(self.accelerator.scaler) else None 'scaler': self.accelerator.scaler.state_dict() if exists(self.accelerator.scaler) else None,
'version': __version__
} }
torch.save(data, str(self.results_folder / f'model-{milestone}.pt')) torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
@@ -846,6 +849,9 @@ class Trainer(object):
self.opt.load_state_dict(data['opt']) self.opt.load_state_dict(data['opt'])
self.ema.load_state_dict(data['ema']) self.ema.load_state_dict(data['ema'])
if 'version' in data:
print(f"loading from version {data['version']}")
if exists(self.accelerator.scaler) and exists(data['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'])
+1
View File
@@ -0,0 +1 @@
__version__ = '0.1.5'
+3 -1
View File
@@ -1,9 +1,11 @@
from setuptools import setup, find_packages from setuptools import setup, find_packages
exec(open('denoising_diffusion_pytorch/version.py').read())
setup( setup(
name = 'denoising-diffusion-pytorch', name = 'denoising-diffusion-pytorch',
packages = find_packages(), packages = find_packages(),
version = '0.32.0', version = __version__,
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',