Compare commits

..
1 Commits
Author SHA1 Message Date
Phil Wang a0c3443eaa optimizer should be saved and loaded 2022-07-08 17:43:41 -07:00
2 changed files with 7 additions and 1 deletions
@@ -680,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
}
@@ -693,7 +696,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.ema.load_state_dict(data['ema'])
+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.2',
version = '0.24.4',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',