|
|
|
@@ -37,6 +37,9 @@ def default(val, d):
|
|
|
|
|
return val
|
|
|
|
|
return d() if callable(d) else d
|
|
|
|
|
|
|
|
|
|
def identity(t, *args, **kwargs):
|
|
|
|
|
return t
|
|
|
|
|
|
|
|
|
|
def cycle(dl):
|
|
|
|
|
while True:
|
|
|
|
|
for data in dl:
|
|
|
|
@@ -53,7 +56,7 @@ def num_to_groups(num, divisor):
|
|
|
|
|
arr.append(remainder)
|
|
|
|
|
return arr
|
|
|
|
|
|
|
|
|
|
def convert_image_to(img_type, image):
|
|
|
|
|
def convert_image_to_fn(img_type, image):
|
|
|
|
|
if image.mode != img_type:
|
|
|
|
|
return image.convert(img_type)
|
|
|
|
|
return image
|
|
|
|
@@ -427,6 +430,7 @@ class GaussianDiffusion(nn.Module):
|
|
|
|
|
):
|
|
|
|
|
super().__init__()
|
|
|
|
|
assert not (type(self) == GaussianDiffusion and model.channels != model.out_dim)
|
|
|
|
|
assert not model.learned_sinusoidal_cond
|
|
|
|
|
|
|
|
|
|
self.model = model
|
|
|
|
|
self.channels = self.model.channels
|
|
|
|
@@ -446,7 +450,7 @@ class GaussianDiffusion(nn.Module):
|
|
|
|
|
raise ValueError(f'unknown beta schedule {beta_schedule}')
|
|
|
|
|
|
|
|
|
|
alphas = 1. - betas
|
|
|
|
|
alphas_cumprod = torch.cumprod(alphas, axis=0)
|
|
|
|
|
alphas_cumprod = torch.cumprod(alphas, dim=0)
|
|
|
|
|
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value = 1.)
|
|
|
|
|
|
|
|
|
|
timesteps, = betas.shape
|
|
|
|
@@ -516,16 +520,19 @@ class GaussianDiffusion(nn.Module):
|
|
|
|
|
posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, t, x_t.shape)
|
|
|
|
|
return posterior_mean, posterior_variance, posterior_log_variance_clipped
|
|
|
|
|
|
|
|
|
|
def model_predictions(self, x, t, x_self_cond = None):
|
|
|
|
|
def model_predictions(self, x, t, x_self_cond = None, clip_x_start = False):
|
|
|
|
|
model_output = self.model(x, t, x_self_cond)
|
|
|
|
|
maybe_clip = partial(torch.clamp, min = -1., max = 1.) if clip_x_start else identity
|
|
|
|
|
|
|
|
|
|
if self.objective == 'pred_noise':
|
|
|
|
|
pred_noise = model_output
|
|
|
|
|
x_start = self.predict_start_from_noise(x, t, model_output)
|
|
|
|
|
x_start = self.predict_start_from_noise(x, t, pred_noise)
|
|
|
|
|
x_start = maybe_clip(x_start)
|
|
|
|
|
|
|
|
|
|
elif self.objective == 'pred_x0':
|
|
|
|
|
pred_noise = self.predict_noise_from_start(x, t, model_output)
|
|
|
|
|
x_start = model_output
|
|
|
|
|
x_start = maybe_clip(x_start)
|
|
|
|
|
pred_noise = self.predict_noise_from_start(x, t, x_start)
|
|
|
|
|
|
|
|
|
|
return ModelPrediction(pred_noise, x_start)
|
|
|
|
|
|
|
|
|
@@ -567,31 +574,30 @@ class GaussianDiffusion(nn.Module):
|
|
|
|
|
def ddim_sample(self, shape, clip_denoised = True):
|
|
|
|
|
batch, device, total_timesteps, sampling_timesteps, eta, objective = shape[0], self.betas.device, self.num_timesteps, self.sampling_timesteps, self.ddim_sampling_eta, self.objective
|
|
|
|
|
|
|
|
|
|
times = torch.linspace(0., total_timesteps, steps = sampling_timesteps + 2)[:-1]
|
|
|
|
|
times = torch.linspace(-1, total_timesteps - 1, steps=sampling_timesteps + 1) # [-1, 0, 1, 2, ..., T-1] when sampling_timesteps == total_timesteps
|
|
|
|
|
times = list(reversed(times.int().tolist()))
|
|
|
|
|
time_pairs = list(zip(times[:-1], times[1:]))
|
|
|
|
|
time_pairs = list(zip(times[:-1], times[1:])) # [(T-1, T-2), (T-2, T-3), ..., (1, 0), (0, -1)]
|
|
|
|
|
|
|
|
|
|
img = torch.randn(shape, device = device)
|
|
|
|
|
|
|
|
|
|
x_start = None
|
|
|
|
|
|
|
|
|
|
for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'):
|
|
|
|
|
alpha = self.alphas_cumprod_prev[time]
|
|
|
|
|
alpha_next = self.alphas_cumprod_prev[time_next]
|
|
|
|
|
|
|
|
|
|
time_cond = torch.full((batch,), time, device = device, dtype = torch.long)
|
|
|
|
|
|
|
|
|
|
time_cond = torch.full((batch,), time, device=device, dtype=torch.long)
|
|
|
|
|
self_cond = x_start if self.self_condition else None
|
|
|
|
|
pred_noise, x_start, *_ = self.model_predictions(img, time_cond, self_cond, clip_x_start = clip_denoised)
|
|
|
|
|
|
|
|
|
|
pred_noise, x_start, *_ = self.model_predictions(img, time_cond, self_cond)
|
|
|
|
|
if time_next < 0:
|
|
|
|
|
img = x_start
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
if clip_denoised:
|
|
|
|
|
x_start.clamp_(-1., 1.)
|
|
|
|
|
alpha = self.alphas_cumprod[time]
|
|
|
|
|
alpha_next = self.alphas_cumprod[time_next]
|
|
|
|
|
|
|
|
|
|
sigma = eta * ((1 - alpha / alpha_next) * (1 - alpha_next) / (1 - alpha)).sqrt()
|
|
|
|
|
c = ((1 - alpha_next) - sigma ** 2).sqrt()
|
|
|
|
|
c = (1 - alpha_next - sigma ** 2).sqrt()
|
|
|
|
|
|
|
|
|
|
noise = torch.randn_like(img) if time_next > 0 else 0.
|
|
|
|
|
noise = torch.randn_like(img)
|
|
|
|
|
|
|
|
|
|
img = x_start * alpha_next.sqrt() + \
|
|
|
|
|
c * pred_noise + \
|
|
|
|
@@ -698,7 +704,7 @@ class Dataset(Dataset):
|
|
|
|
|
self.image_size = image_size
|
|
|
|
|
self.paths = [p for ext in exts for p in Path(f'{folder}').glob(f'**/*.{ext}')]
|
|
|
|
|
|
|
|
|
|
maybe_convert_fn = partial(convert_image_to, convert_image_to) if exists(convert_image_to) else nn.Identity()
|
|
|
|
|
maybe_convert_fn = partial(convert_image_to_fn, convert_image_to) if exists(convert_image_to) else nn.Identity()
|
|
|
|
|
|
|
|
|
|
self.transform = T.Compose([
|
|
|
|
|
T.Lambda(maybe_convert_fn),
|
|
|
|
@@ -804,7 +810,10 @@ class Trainer(object):
|
|
|
|
|
torch.save(data, str(self.results_folder / f'model-{milestone}.pt'))
|
|
|
|
|
|
|
|
|
|
def load(self, milestone):
|
|
|
|
|
data = torch.load(str(self.results_folder / f'model-{milestone}.pt'))
|
|
|
|
|
accelerator = self.accelerator
|
|
|
|
|
device = accelerator.device
|
|
|
|
|
|
|
|
|
|
data = torch.load(str(self.results_folder / f'model-{milestone}.pt'), map_location=device)
|
|
|
|
|
|
|
|
|
|
model = self.accelerator.unwrap_model(self.model)
|
|
|
|
|
model.load_state_dict(data['model'])
|
|
|
|
@@ -836,6 +845,7 @@ class Trainer(object):
|
|
|
|
|
|
|
|
|
|
self.accelerator.backward(loss)
|
|
|
|
|
|
|
|
|
|
accelerator.clip_grad_norm_(self.model.parameters(), 1.0)
|
|
|
|
|
pbar.set_description(f'loss: {total_loss:.4f}')
|
|
|
|
|
|
|
|
|
|
accelerator.wait_for_everyone()
|
|
|
|
@@ -845,6 +855,7 @@ class Trainer(object):
|
|
|
|
|
|
|
|
|
|
accelerator.wait_for_everyone()
|
|
|
|
|
|
|
|
|
|
self.step += 1
|
|
|
|
|
if accelerator.is_main_process:
|
|
|
|
|
self.ema.to(device)
|
|
|
|
|
self.ema.update()
|
|
|
|
@@ -861,7 +872,6 @@ class Trainer(object):
|
|
|
|
|
utils.save_image(all_images, str(self.results_folder / f'sample-{milestone}.png'), nrow = int(math.sqrt(self.num_samples)))
|
|
|
|
|
self.save(milestone)
|
|
|
|
|
|
|
|
|
|
self.step += 1
|
|
|
|
|
pbar.update(1)
|
|
|
|
|
|
|
|
|
|
accelerator.print('training complete')
|
|
|
|
|