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