Compare commits

..
1 Commits
Author SHA1 Message Date
Phil Wang cee30ce304 optimizer should be saved and loaded 2022-07-08 17:07:32 -07:00
2 changed files with 3 additions and 7 deletions
@@ -680,12 +680,10 @@ 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(),
'opt': self.opt.state_dict(),
'ema': self.ema.state_dict(),
'scaler': self.accelerator.scaler.state_dict() if exists(self.accelerator.scaler) else None
}
@@ -696,12 +694,10 @@ 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.opt.load_state_dict(data['opt'])
self.ema.load_state_dict(data['ema'])
if exists(self.accelerator.scaler) and exists(data['scaler']):
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup(
name = 'denoising-diffusion-pytorch',
packages = find_packages(),
version = '0.24.4',
version = '0.24.3',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',