denoising-diffusion-pytorch 2.2.1__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.1 → denoising_diffusion_pytorch-2.2.2}/PKG-INFO +1 -1
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/denoising_diffusion_pytorch_1d.py +27 -13
- denoising_diffusion_pytorch-2.2.2/denoising_diffusion_pytorch/version.py +1 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/PKG-INFO +1 -1
- denoising_diffusion_pytorch-2.2.1/denoising_diffusion_pytorch/version.py +0 -1
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/LICENSE +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/README.md +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/__init__.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/attend.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/classifier_free_guidance.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/elucidated_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/fid_evaluation.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/guided_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/karras_unet.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/karras_unet_1d.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/karras_unet_3d.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/learned_gaussian_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/repaint.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/simple_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/v_param_continuous_time_gaussian_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/weighted_objective_gaussian_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/SOURCES.txt +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/dependency_links.txt +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/requires.txt +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/top_level.txt +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/setup.cfg +0 -0
- {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/setup.py +0 -0
|
@@ -574,8 +574,12 @@ class GaussianDiffusion1D(Module):
|
|
|
574
574
|
|
|
575
575
|
return ModelPrediction(pred_noise, x_start)
|
|
576
576
|
|
|
577
|
-
def p_mean_variance(self, x, t, x_self_cond = None, clip_denoised = True):
|
|
578
|
-
|
|
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)
|
|
579
583
|
x_start = preds.pred_x_start
|
|
580
584
|
|
|
581
585
|
if clip_denoised:
|
|
@@ -585,38 +589,44 @@ class GaussianDiffusion1D(Module):
|
|
|
585
589
|
return model_mean, posterior_variance, posterior_log_variance, x_start
|
|
586
590
|
|
|
587
591
|
@torch.no_grad()
|
|
588
|
-
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()):
|
|
589
593
|
b, *_, device = *x.shape, x.device
|
|
590
594
|
batched_times = torch.full((b,), t, device = x.device, dtype = torch.long)
|
|
591
|
-
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)
|
|
592
596
|
noise = torch.randn_like(x) if t > 0 else 0. # no noise if t == 0
|
|
593
597
|
pred_img = model_mean + (0.5 * model_log_variance).exp() * noise
|
|
594
598
|
return pred_img, x_start
|
|
595
599
|
|
|
596
600
|
@torch.no_grad()
|
|
597
|
-
def p_sample_loop(self, shape):
|
|
601
|
+
def p_sample_loop(self, shape, return_noise = False, model_forward_kwargs: dict = dict()):
|
|
598
602
|
batch, device = shape[0], self.betas.device
|
|
599
603
|
|
|
600
|
-
|
|
604
|
+
noise = torch.randn(shape, device=device)
|
|
605
|
+
img = noise
|
|
601
606
|
|
|
602
607
|
x_start = None
|
|
603
608
|
|
|
604
609
|
for t in tqdm(reversed(range(0, self.num_timesteps)), desc = 'sampling loop time step', total = self.num_timesteps):
|
|
605
610
|
self_cond = x_start if self.self_condition else None
|
|
606
|
-
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)
|
|
607
612
|
|
|
608
613
|
img = self.unnormalize(img)
|
|
609
|
-
|
|
614
|
+
|
|
615
|
+
if not return_noise:
|
|
616
|
+
return img
|
|
617
|
+
|
|
618
|
+
return img, noise
|
|
610
619
|
|
|
611
620
|
@torch.no_grad()
|
|
612
|
-
def ddim_sample(self, shape, clip_denoised = True, model_forward_kwargs: dict = dict()):
|
|
621
|
+
def ddim_sample(self, shape, clip_denoised = True, model_forward_kwargs: dict = dict(), return_noise = False):
|
|
613
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
|
|
614
623
|
|
|
615
624
|
times = torch.linspace(-1, total_timesteps - 1, steps=sampling_timesteps + 1) # [-1, 0, 1, 2, ..., T-1] when sampling_timesteps == total_timesteps
|
|
616
625
|
times = list(reversed(times.int().tolist()))
|
|
617
626
|
time_pairs = list(zip(times[:-1], times[1:])) # [(T-1, T-2), (T-2, T-3), ..., (1, 0), (0, -1)]
|
|
618
627
|
|
|
619
|
-
|
|
628
|
+
noise = torch.randn(shape, device = device)
|
|
629
|
+
img = noise
|
|
620
630
|
|
|
621
631
|
x_start = None
|
|
622
632
|
|
|
@@ -642,15 +652,19 @@ class GaussianDiffusion1D(Module):
|
|
|
642
652
|
sigma * noise
|
|
643
653
|
|
|
644
654
|
img = self.unnormalize(img)
|
|
645
|
-
|
|
655
|
+
|
|
656
|
+
if not return_noise:
|
|
657
|
+
return img
|
|
658
|
+
|
|
659
|
+
return img, noise
|
|
646
660
|
|
|
647
661
|
@torch.no_grad()
|
|
648
|
-
def sample(self, batch_size = 16, model_forward_kwargs: dict = dict()):
|
|
662
|
+
def sample(self, batch_size = 16, return_noise = False, model_forward_kwargs: dict = dict()):
|
|
649
663
|
seq_length, channels = self.seq_length, self.channels
|
|
650
664
|
sample_fn = self.p_sample_loop if not self.is_ddim_sampling else self.ddim_sample
|
|
651
665
|
|
|
652
666
|
shape = (batch_size, channels, seq_length) if self.channel_first else (batch_size, seq_length, channels)
|
|
653
|
-
return sample_fn(shape, model_forward_kwargs = model_forward_kwargs)
|
|
667
|
+
return sample_fn(shape, return_noise = return_noise, model_forward_kwargs = model_forward_kwargs)
|
|
654
668
|
|
|
655
669
|
@torch.no_grad()
|
|
656
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.1'
|
|
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
|