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.
Files changed (29) hide show
  1. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/PKG-INFO +1 -1
  2. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/denoising_diffusion_pytorch_1d.py +27 -13
  3. denoising_diffusion_pytorch-2.2.2/denoising_diffusion_pytorch/version.py +1 -0
  4. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/PKG-INFO +1 -1
  5. denoising_diffusion_pytorch-2.2.1/denoising_diffusion_pytorch/version.py +0 -1
  6. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/LICENSE +0 -0
  7. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/README.md +0 -0
  8. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/__init__.py +0 -0
  9. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/attend.py +0 -0
  10. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/classifier_free_guidance.py +0 -0
  11. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py +0 -0
  12. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +0 -0
  13. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/elucidated_diffusion.py +0 -0
  14. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/fid_evaluation.py +0 -0
  15. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/guided_diffusion.py +0 -0
  16. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/karras_unet.py +0 -0
  17. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/karras_unet_1d.py +0 -0
  18. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/karras_unet_3d.py +0 -0
  19. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/learned_gaussian_diffusion.py +0 -0
  20. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/repaint.py +0 -0
  21. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/simple_diffusion.py +0 -0
  22. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/v_param_continuous_time_gaussian_diffusion.py +0 -0
  23. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch/weighted_objective_gaussian_diffusion.py +0 -0
  24. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/SOURCES.txt +0 -0
  25. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/dependency_links.txt +0 -0
  26. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/requires.txt +0 -0
  27. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/denoising_diffusion_pytorch.egg-info/top_level.txt +0 -0
  28. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/setup.cfg +0 -0
  29. {denoising_diffusion_pytorch-2.2.1 → denoising_diffusion_pytorch-2.2.2}/setup.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: denoising-diffusion-pytorch
3
- Version: 2.2.1
3
+ Version: 2.2.2
4
4
  Summary: Denoising Diffusion Probabilistic Models - Pytorch
5
5
  Home-page: https://github.com/lucidrains/denoising-diffusion-pytorch
6
6
  Author: Phil Wang
@@ -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
- preds = self.model_predictions(x, t, x_self_cond)
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
- img = torch.randn(shape, device=device)
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
- return img
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
- img = torch.randn(shape, device = device)
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
- return img
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,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: denoising-diffusion-pytorch
3
- Version: 2.2.1
3
+ Version: 2.2.2
4
4
  Summary: Denoising Diffusion Probabilistic Models - Pytorch
5
5
  Home-page: https://github.com/lucidrains/denoising-diffusion-pytorch
6
6
  Author: Phil Wang
@@ -1 +0,0 @@
1
- __version__ = '2.2.1'