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 |
@@ -680,10 +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': self.opt.state_dict(),
|
||||
'opt': opt.state_dict(),
|
||||
'ema': self.ema.state_dict(),
|
||||
'scaler': self.accelerator.scaler.state_dict() if exists(self.accelerator.scaler) else None
|
||||
}
|
||||
@@ -694,10 +696,12 @@ 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']):
|
||||
|
||||
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
|
||||
setup(
|
||||
name = 'denoising-diffusion-pytorch',
|
||||
packages = find_packages(),
|
||||
version = '0.24.3',
|
||||
version = '0.24.4',
|
||||
license='MIT',
|
||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||
author = 'Phil Wang',
|
||||
|
||||
Reference in New Issue
Block a user