remove outdated apex in favor of native pytorch AMP

This commit is contained in:
Phil Wang
2022-04-13 08:59:18 -07:00
parent e504e0e554
commit 0b8cdb4c8b
3 changed files with 18 additions and 28 deletions
+1 -1
View File
@@ -68,7 +68,7 @@ trainer = Trainer(
train_num_steps = 700000, # total training steps
gradient_accumulate_every = 2, # gradient accumulation steps
ema_decay = 0.995, # exponential moving average decay
fp16 = True # turn on mixed precision training with apex
amp = True # turn on mixed precision
)
trainer.train()
@@ -7,6 +7,8 @@ from inspect import isfunction
from functools import partial
from torch.utils import data
from torch.cuda.amp import autocast, GradScaler
from pathlib import Path
from torch.optim import Adam
from torchvision import transforms, utils
@@ -15,12 +17,6 @@ from PIL import Image
from tqdm import tqdm
from einops import rearrange
try:
from apex import amp
APEX_AVAILABLE = True
except:
APEX_AVAILABLE = False
# helpers functions
def exists(x):
@@ -44,13 +40,6 @@ def num_to_groups(num, divisor):
arr.append(remainder)
return arr
def loss_backwards(fp16, loss, optimizer, **kwargs):
if fp16:
with amp.scale_loss(loss, optimizer) as scaled_loss:
scaled_loss.backward(**kwargs)
else:
loss.backward(**kwargs)
# small helper modules
class EMA():
@@ -502,7 +491,7 @@ class Trainer(object):
train_lr = 2e-5,
train_num_steps = 100000,
gradient_accumulate_every = 2,
fp16 = False,
amp = False,
step_start_ema = 2000,
update_ema_every = 10,
save_and_sample_every = 1000,
@@ -528,11 +517,8 @@ class Trainer(object):
self.step = 0
assert not fp16 or fp16 and APEX_AVAILABLE, 'Apex must be installed in order for mixed precision training to be turned on'
self.fp16 = fp16
if fp16:
(self.model, self.ema_model), self.opt = amp.initialize([self.model, self.ema_model], self.opt, opt_level='O1')
self.amp = amp
self.scaler = GradScaler(enabled = amp)
self.results_folder = Path(results_folder)
self.results_folder.mkdir(exist_ok = True)
@@ -552,7 +538,8 @@ class Trainer(object):
data = {
'step': self.step,
'model': self.model.state_dict(),
'ema': self.ema_model.state_dict()
'ema': self.ema_model.state_dict(),
'scaler': self.scaler.state_dict()
}
torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
@@ -562,18 +549,21 @@ class Trainer(object):
self.step = data['step']
self.model.load_state_dict(data['model'])
self.ema_model.load_state_dict(data['ema'])
self.scaler.load_state_dict(data['scaler'])
def train(self):
backwards = partial(loss_backwards, self.fp16)
while self.step < self.train_num_steps:
for i in range(self.gradient_accumulate_every):
data = next(self.dl).cuda()
loss = self.model(data)
print(f'{self.step}: {loss.item()}')
backwards(loss / self.gradient_accumulate_every, self.opt)
self.opt.step()
with autocast(enabled = self.amp):
loss = self.model(data)
self.scaler.scale(loss / self.gradient_accumulate_every).backward()
print(f'{self.step}: {loss.item()}')
self.scaler.step(self.opt)
self.scaler.update()
self.opt.zero_grad()
if self.step % self.update_ema_every == 0:
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup(
name = 'denoising-diffusion-pytorch',
packages = find_packages(),
version = '0.8.1',
version = '0.9.0',
license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang',