denoising-diffusion-pytorch 2.2.0__tar.gz → 2.2.2__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/PKG-INFO +1 -1
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/denoising_diffusion_pytorch_1d.py +34 -16
- denoising_diffusion_pytorch-2.2.2/denoising_diffusion_pytorch/version.py +1 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/PKG-INFO +1 -1
- denoising_diffusion_pytorch-2.2.0/denoising_diffusion_pytorch/version.py +0 -1
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/LICENSE +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/README.md +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/__init__.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/attend.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/classifier_free_guidance.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/elucidated_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/fid_evaluation.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/guided_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/karras_unet.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/karras_unet_1d.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/karras_unet_3d.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/learned_gaussian_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/repaint.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/simple_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/v_param_continuous_time_gaussian_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/weighted_objective_gaussian_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/SOURCES.txt +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/dependency_links.txt +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/requires.txt +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/top_level.txt +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/setup.cfg +0 -0
- {denoising_diffusion_pytorch-2.2.0 → denoising_diffusion_pytorch-2.2.2}/setup.py +0 -0
|
@@ -545,8 +545,12 @@ class GaussianDiffusion1D(Module):
|
|
|
545
545
|
posterior_log_variance_clipped = extract(self.posterior_log_variance_clipped, t, x_t.shape)
|
|
546
546
|
return posterior_mean, posterior_variance, posterior_log_variance_clipped
|
|
547
547
|
|
|
548
|
-
def model_predictions(self, x, t, x_self_cond = None, clip_x_start = False, rederive_pred_noise = False):
|
|
549
|
-
|
|
548
|
+
def model_predictions(self, x, t, x_self_cond = None, clip_x_start = False, rederive_pred_noise = False, model_forward_kwargs: dict = dict()):
|
|
549
|
+
|
|
550
|
+
if exists(x_self_cond):
|
|
551
|
+
model_forward_kwargs = {**model_forward_kwargs, 'self_cond': x_self_cond}
|
|
552
|
+
|
|
553
|
+
model_output = self.model(x, t, **model_forward_kwargs)
|
|
550
554
|
maybe_clip = partial(torch.clamp, min = -1., max = 1.) if clip_x_start else identity
|
|
551
555
|
|
|
552
556
|
if self.objective == 'pred_noise':
|
|
@@ -570,8 +574,12 @@ class GaussianDiffusion1D(Module):
|
|
|
570
574
|
|
|
571
575
|
return ModelPrediction(pred_noise, x_start)
|
|
572
576
|
|
|
573
|
-
def p_mean_variance(self, x, t, x_self_cond = None, clip_denoised = True):
|
|
574
|
-
|
|
577
|
+
def p_mean_variance(self, x, t, x_self_cond = None, clip_denoised = True, model_forward_kwargs: dict = dict()):
|
|
578
|
+
|
|
579
|
+
if exists(x_self_cond):
|
|
580
|
+
model_forward_kwargs = {**model_forward_kwargs, 'self_cond': x_self_cond}
|
|
581
|
+
|
|
582
|
+
preds = self.model_predictions(x, t, **model_forward_kwargs)
|
|
575
583
|
x_start = preds.pred_x_start
|
|
576
584
|
|
|
577
585
|
if clip_denoised:
|
|
@@ -581,45 +589,51 @@ class GaussianDiffusion1D(Module):
|
|
|
581
589
|
return model_mean, posterior_variance, posterior_log_variance, x_start
|
|
582
590
|
|
|
583
591
|
@torch.no_grad()
|
|
584
|
-
def p_sample(self, x, t: int, x_self_cond = None, clip_denoised = True):
|
|
592
|
+
def p_sample(self, x, t: int, x_self_cond = None, clip_denoised = True, model_forward_kwargs: dict = dict()):
|
|
585
593
|
b, *_, device = *x.shape, x.device
|
|
586
594
|
batched_times = torch.full((b,), t, device = x.device, dtype = torch.long)
|
|
587
|
-
model_mean, _, model_log_variance, x_start = self.p_mean_variance(x = x, t = batched_times, x_self_cond = x_self_cond, clip_denoised = clip_denoised)
|
|
595
|
+
model_mean, _, model_log_variance, x_start = self.p_mean_variance(x = x, t = batched_times, x_self_cond = x_self_cond, clip_denoised = clip_denoised, model_forward_kwargs = model_forward_kwargs)
|
|
588
596
|
noise = torch.randn_like(x) if t > 0 else 0. # no noise if t == 0
|
|
589
597
|
pred_img = model_mean + (0.5 * model_log_variance).exp() * noise
|
|
590
598
|
return pred_img, x_start
|
|
591
599
|
|
|
592
600
|
@torch.no_grad()
|
|
593
|
-
def p_sample_loop(self, shape):
|
|
601
|
+
def p_sample_loop(self, shape, return_noise = False, model_forward_kwargs: dict = dict()):
|
|
594
602
|
batch, device = shape[0], self.betas.device
|
|
595
603
|
|
|
596
|
-
|
|
604
|
+
noise = torch.randn(shape, device=device)
|
|
605
|
+
img = noise
|
|
597
606
|
|
|
598
607
|
x_start = None
|
|
599
608
|
|
|
600
609
|
for t in tqdm(reversed(range(0, self.num_timesteps)), desc = 'sampling loop time step', total = self.num_timesteps):
|
|
601
610
|
self_cond = x_start if self.self_condition else None
|
|
602
|
-
img, x_start = self.p_sample(img, t, self_cond)
|
|
611
|
+
img, x_start = self.p_sample(img, t, self_cond, model_forward_kwargs = model_forward_kwargs)
|
|
603
612
|
|
|
604
613
|
img = self.unnormalize(img)
|
|
605
|
-
|
|
614
|
+
|
|
615
|
+
if not return_noise:
|
|
616
|
+
return img
|
|
617
|
+
|
|
618
|
+
return img, noise
|
|
606
619
|
|
|
607
620
|
@torch.no_grad()
|
|
608
|
-
def ddim_sample(self, shape, clip_denoised = True):
|
|
621
|
+
def ddim_sample(self, shape, clip_denoised = True, model_forward_kwargs: dict = dict(), return_noise = False):
|
|
609
622
|
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
|
|
610
623
|
|
|
611
624
|
times = torch.linspace(-1, total_timesteps - 1, steps=sampling_timesteps + 1) # [-1, 0, 1, 2, ..., T-1] when sampling_timesteps == total_timesteps
|
|
612
625
|
times = list(reversed(times.int().tolist()))
|
|
613
626
|
time_pairs = list(zip(times[:-1], times[1:])) # [(T-1, T-2), (T-2, T-3), ..., (1, 0), (0, -1)]
|
|
614
627
|
|
|
615
|
-
|
|
628
|
+
noise = torch.randn(shape, device = device)
|
|
629
|
+
img = noise
|
|
616
630
|
|
|
617
631
|
x_start = None
|
|
618
632
|
|
|
619
633
|
for time, time_next in tqdm(time_pairs, desc = 'sampling loop time step'):
|
|
620
634
|
time_cond = torch.full((batch,), time, device=device, dtype=torch.long)
|
|
621
635
|
self_cond = x_start if self.self_condition else None
|
|
622
|
-
pred_noise, x_start, *_ = self.model_predictions(img, time_cond, self_cond, clip_x_start = clip_denoised)
|
|
636
|
+
pred_noise, x_start, *_ = self.model_predictions(img, time_cond, self_cond, clip_x_start = clip_denoised, model_forward_kwargs = model_forward_kwargs)
|
|
623
637
|
|
|
624
638
|
if time_next < 0:
|
|
625
639
|
img = x_start
|
|
@@ -638,15 +652,19 @@ class GaussianDiffusion1D(Module):
|
|
|
638
652
|
sigma * noise
|
|
639
653
|
|
|
640
654
|
img = self.unnormalize(img)
|
|
641
|
-
|
|
655
|
+
|
|
656
|
+
if not return_noise:
|
|
657
|
+
return img
|
|
658
|
+
|
|
659
|
+
return img, noise
|
|
642
660
|
|
|
643
661
|
@torch.no_grad()
|
|
644
|
-
def sample(self, batch_size = 16):
|
|
662
|
+
def sample(self, batch_size = 16, return_noise = False, model_forward_kwargs: dict = dict()):
|
|
645
663
|
seq_length, channels = self.seq_length, self.channels
|
|
646
664
|
sample_fn = self.p_sample_loop if not self.is_ddim_sampling else self.ddim_sample
|
|
647
665
|
|
|
648
666
|
shape = (batch_size, channels, seq_length) if self.channel_first else (batch_size, seq_length, channels)
|
|
649
|
-
return sample_fn(shape)
|
|
667
|
+
return sample_fn(shape, return_noise = return_noise, model_forward_kwargs = model_forward_kwargs)
|
|
650
668
|
|
|
651
669
|
@torch.no_grad()
|
|
652
670
|
def interpolate(self, x1, x2, t = None, lam = 0.5):
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = '2.2.2'
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = '2.2.0'
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|