mirror of
https://github.com/wassname/denoising-diffusion-pytorch.git
synced 2026-09-11 12:11:43 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
60128257c5 |
@@ -603,34 +603,37 @@ class Trainer(object):
|
|||||||
self.scaler.load_state_dict(data['scaler'])
|
self.scaler.load_state_dict(data['scaler'])
|
||||||
|
|
||||||
def train(self):
|
def train(self):
|
||||||
while self.step < self.train_num_steps:
|
with tqdm(initial = self.step, total = self.train_num_steps) as pbar:
|
||||||
for i in range(self.gradient_accumulate_every):
|
|
||||||
data = next(self.dl).cuda()
|
|
||||||
|
|
||||||
with autocast(enabled = self.amp):
|
while self.step < self.train_num_steps:
|
||||||
loss = self.model(data)
|
for i in range(self.gradient_accumulate_every):
|
||||||
self.scaler.scale(loss / self.gradient_accumulate_every).backward()
|
data = next(self.dl).cuda()
|
||||||
|
|
||||||
print(f'{self.step}: {loss.item()}')
|
with autocast(enabled = self.amp):
|
||||||
|
loss = self.model(data)
|
||||||
|
self.scaler.scale(loss / self.gradient_accumulate_every).backward()
|
||||||
|
|
||||||
self.scaler.step(self.opt)
|
pbar.set_description(f'loss: {loss.item():.4f}')
|
||||||
self.scaler.update()
|
|
||||||
self.opt.zero_grad()
|
|
||||||
|
|
||||||
if self.step % self.update_ema_every == 0:
|
self.scaler.step(self.opt)
|
||||||
self.step_ema()
|
self.scaler.update()
|
||||||
|
self.opt.zero_grad()
|
||||||
|
|
||||||
if self.step != 0 and self.step % self.save_and_sample_every == 0:
|
if self.step % self.update_ema_every == 0:
|
||||||
self.ema_model.eval()
|
self.step_ema()
|
||||||
|
|
||||||
milestone = self.step // self.save_and_sample_every
|
if self.step != 0 and self.step % self.save_and_sample_every == 0:
|
||||||
batches = num_to_groups(36, self.batch_size)
|
self.ema_model.eval()
|
||||||
all_images_list = list(map(lambda n: self.ema_model.sample(batch_size=n), batches))
|
|
||||||
all_images = torch.cat(all_images_list, dim=0)
|
|
||||||
all_images = unnormalize_to_zero_to_one(all_images)
|
|
||||||
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6)
|
|
||||||
self.save(milestone)
|
|
||||||
|
|
||||||
self.step += 1
|
milestone = self.step // self.save_and_sample_every
|
||||||
|
batches = num_to_groups(36, self.batch_size)
|
||||||
|
all_images_list = list(map(lambda n: self.ema_model.sample(batch_size=n), batches))
|
||||||
|
all_images = torch.cat(all_images_list, dim=0)
|
||||||
|
all_images = unnormalize_to_zero_to_one(all_images)
|
||||||
|
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = 6)
|
||||||
|
self.save(milestone)
|
||||||
|
|
||||||
print('training completed')
|
self.step += 1
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
print('training complete')
|
||||||
|
|||||||
@@ -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.15.2',
|
version = '0.15.3',
|
||||||
license='MIT',
|
license='MIT',
|
||||||
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
description = 'Denoising Diffusion Probabilistic Models - Pytorch',
|
||||||
author = 'Phil Wang',
|
author = 'Phil Wang',
|
||||||
|
|||||||
Reference in New Issue
Block a user