From a0c3443eaa67c70f69eca0abec37798f7d991235 Mon Sep 17 00:00:00 2001 From: Phil Wang Date: Fri, 8 Jul 2022 17:43:41 -0700 Subject: [PATCH] optimizer should be saved and loaded --- denoising_diffusion_pytorch/denoising_diffusion_pytorch.py | 6 ++++++ setup.py | 2 +- 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py index 78dbcb1..c0dce55 100644 --- a/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +++ b/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py @@ -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']) diff --git a/setup.py b/setup.py index 9fbf615..7691cef 100644 --- a/setup.py +++ b/setup.py @@ -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',