Compare commits

..
2 Commits
Author SHA1 Message Date
Phil Wang 0b8cdb4c8b remove outdated apex in favor of native pytorch AMP 2022-04-13 08:59:18 -07:00
Phil Wang e504e0e554 cleanup 2022-04-12 13:02:18 -07:00
3 changed files with 18 additions and 30 deletions
+1 -1
View File
@@ -68,7 +68,7 @@ trainer = Trainer(
train_num_steps = 700000, # total training steps train_num_steps = 700000, # total training steps
gradient_accumulate_every = 2, # gradient accumulation steps gradient_accumulate_every = 2, # gradient accumulation steps
ema_decay = 0.995, # exponential moving average decay 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() trainer.train()
@@ -7,6 +7,8 @@ from inspect import isfunction
from functools import partial from functools import partial
from torch.utils import data from torch.utils import data
from torch.cuda.amp import autocast, GradScaler
from pathlib import Path from pathlib import Path
from torch.optim import Adam from torch.optim import Adam
from torchvision import transforms, utils from torchvision import transforms, utils
@@ -15,12 +17,6 @@ from PIL import Image
from tqdm import tqdm from tqdm import tqdm
from einops import rearrange from einops import rearrange
try:
from apex import amp
APEX_AVAILABLE = True
except:
APEX_AVAILABLE = False
# helpers functions # helpers functions
def exists(x): def exists(x):
@@ -44,13 +40,6 @@ def num_to_groups(num, divisor):
arr.append(remainder) arr.append(remainder)
return arr 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 # small helper modules
class EMA(): class EMA():
@@ -335,8 +324,6 @@ class GaussianDiffusion(nn.Module):
self.num_timesteps = int(timesteps) self.num_timesteps = int(timesteps)
self.loss_type = loss_type self.loss_type = loss_type
to_torch = partial(torch.tensor, dtype=torch.float32)
self.register_buffer('betas', betas) self.register_buffer('betas', betas)
self.register_buffer('alphas_cumprod', alphas_cumprod) self.register_buffer('alphas_cumprod', alphas_cumprod)
self.register_buffer('alphas_cumprod_prev', alphas_cumprod_prev) self.register_buffer('alphas_cumprod_prev', alphas_cumprod_prev)
@@ -504,7 +491,7 @@ class Trainer(object):
train_lr = 2e-5, train_lr = 2e-5,
train_num_steps = 100000, train_num_steps = 100000,
gradient_accumulate_every = 2, gradient_accumulate_every = 2,
fp16 = False, amp = False,
step_start_ema = 2000, step_start_ema = 2000,
update_ema_every = 10, update_ema_every = 10,
save_and_sample_every = 1000, save_and_sample_every = 1000,
@@ -530,11 +517,8 @@ class Trainer(object):
self.step = 0 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.amp = amp
self.scaler = GradScaler(enabled = amp)
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.results_folder = Path(results_folder) self.results_folder = Path(results_folder)
self.results_folder.mkdir(exist_ok = True) self.results_folder.mkdir(exist_ok = True)
@@ -554,7 +538,8 @@ class Trainer(object):
data = { data = {
'step': self.step, 'step': self.step,
'model': self.model.state_dict(), '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')) torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
@@ -564,18 +549,21 @@ class Trainer(object):
self.step = data['step'] self.step = data['step']
self.model.load_state_dict(data['model']) self.model.load_state_dict(data['model'])
self.ema_model.load_state_dict(data['ema']) self.ema_model.load_state_dict(data['ema'])
self.scaler.load_state_dict(data['scaler'])
def train(self): def train(self):
backwards = partial(loss_backwards, self.fp16)
while self.step < self.train_num_steps: while self.step < self.train_num_steps:
for i in range(self.gradient_accumulate_every): for i in range(self.gradient_accumulate_every):
data = next(self.dl).cuda() 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() self.opt.zero_grad()
if self.step % self.update_ema_every == 0: if self.step % self.update_ema_every == 0:
+1 -1
View File
@@ -3,7 +3,7 @@ from setuptools import setup, find_packages
setup( setup(
name = 'denoising-diffusion-pytorch', name = 'denoising-diffusion-pytorch',
packages = find_packages(), packages = find_packages(),
version = '0.8.1', version = '0.9.0',
license='MIT', license='MIT',
description = 'Denoising Diffusion Probabilistic Models - Pytorch', description = 'Denoising Diffusion Probabilistic Models - Pytorch',
author = 'Phil Wang', author = 'Phil Wang',