denoising-diffusion-pytorch 2.3.0__tar.gz → 2.3.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 (25) hide show
  1. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/PKG-INFO +35 -2
  2. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/README.md +33 -0
  3. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/classifier_free_guidance.py +26 -8
  4. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py +3 -3
  5. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +7 -3
  6. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/denoising_diffusion_pytorch_1d.py +11 -3
  7. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/elucidated_diffusion.py +2 -2
  8. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/v_param_continuous_time_gaussian_diffusion.py +2 -2
  9. denoising_diffusion_pytorch-2.3.2/denoising_diffusion_pytorch/version.py +1 -0
  10. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/xm.py +21 -1
  11. denoising_diffusion_pytorch-2.3.0/denoising_diffusion_pytorch/version.py +0 -1
  12. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/.gitignore +0 -0
  13. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/LICENSE +0 -0
  14. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/__init__.py +0 -0
  15. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/attend.py +0 -0
  16. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/fid_evaluation.py +0 -0
  17. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/guided_diffusion.py +0 -0
  18. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/karras_unet.py +0 -0
  19. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/karras_unet_1d.py +0 -0
  20. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/karras_unet_3d.py +0 -0
  21. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/learned_gaussian_diffusion.py +0 -0
  22. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/repaint.py +0 -0
  23. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/simple_diffusion.py +0 -0
  24. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/denoising_diffusion_pytorch/weighted_objective_gaussian_diffusion.py +0 -0
  25. {denoising_diffusion_pytorch-2.3.0 → denoising_diffusion_pytorch-2.3.2}/pyproject.toml +0 -0
@@ -1,6 +1,6 @@
1
- Metadata-Version: 2.4
1
+ Metadata-Version: 2.5
2
2
  Name: denoising-diffusion-pytorch
3
- Version: 2.3.0
3
+ Version: 2.3.2
4
4
  Summary: Denoising Diffusion Probabilistic Models - Pytorch
5
5
  Project-URL: Homepage, https://github.com/lucidrains/denoising-diffusion-pytorch
6
6
  Project-URL: Repository, https://github.com/lucidrains/denoising-diffusion-pytorch
@@ -115,6 +115,39 @@ trainer.train()
115
115
 
116
116
  Samples and model checkpoints will be logged to `./results` periodically
117
117
 
118
+ ## Explorative Modeling (Forward XM)
119
+
120
+ To use <a href="https://arxiv.org/abs/2607.27372">Explorative Modeling (Forward XM)</a> for multi-candidate loss calculation during training, wrap any diffusion model in `XMWrapper`.
121
+
122
+ ```python
123
+ import torch
124
+ from denoising_diffusion_pytorch import Unet, GaussianDiffusion, XMWrapper
125
+
126
+ model = Unet(
127
+ dim = 64,
128
+ dim_mults = (1, 2, 4, 8)
129
+ )
130
+
131
+ diffusion = GaussianDiffusion(
132
+ model,
133
+ image_size = 128,
134
+ timesteps = 1000
135
+ )
136
+
137
+ xm = XMWrapper(
138
+ diffusion,
139
+ candidates = 4 # generate 4 candidates per sample and pick minimum loss
140
+ )
141
+
142
+ training_images = torch.rand(8, 3, 128, 128)
143
+ loss = xm(training_images)
144
+ loss.backward()
145
+
146
+ # sampling works as usual
147
+
148
+ sampled_images = xm.sample(batch_size = 4)
149
+ ```
150
+
118
151
  ## Multi-GPU Training
119
152
 
120
153
  The `Trainer` class is now equipped with <a href="https://huggingface.co/docs/accelerate/accelerator">🤗 Accelerator</a>. You can easily do multi-gpu training in two steps using their `accelerate` CLI
@@ -87,6 +87,39 @@ trainer.train()
87
87
 
88
88
  Samples and model checkpoints will be logged to `./results` periodically
89
89
 
90
+ ## Explorative Modeling (Forward XM)
91
+
92
+ To use <a href="https://arxiv.org/abs/2607.27372">Explorative Modeling (Forward XM)</a> for multi-candidate loss calculation during training, wrap any diffusion model in `XMWrapper`.
93
+
94
+ ```python
95
+ import torch
96
+ from denoising_diffusion_pytorch import Unet, GaussianDiffusion, XMWrapper
97
+
98
+ model = Unet(
99
+ dim = 64,
100
+ dim_mults = (1, 2, 4, 8)
101
+ )
102
+
103
+ diffusion = GaussianDiffusion(
104
+ model,
105
+ image_size = 128,
106
+ timesteps = 1000
107
+ )
108
+
109
+ xm = XMWrapper(
110
+ diffusion,
111
+ candidates = 4 # generate 4 candidates per sample and pick minimum loss
112
+ )
113
+
114
+ training_images = torch.rand(8, 3, 128, 128)
115
+ loss = xm(training_images)
116
+ loss.backward()
117
+
118
+ # sampling works as usual
119
+
120
+ sampled_images = xm.sample(batch_size = 4)
121
+ ```
122
+
90
123
  ## Multi-GPU Training
91
124
 
92
125
  The `Trainer` class is now equipped with <a href="https://huggingface.co/docs/accelerate/accelerator">🤗 Accelerator</a>. You can easily do multi-gpu training in two steps using their `accelerate` CLI
@@ -411,7 +411,8 @@ class Unet(nn.Module):
411
411
  x,
412
412
  time,
413
413
  classes,
414
- cond_drop_prob = None
414
+ cond_drop_prob = None,
415
+ cond_keep_mask = None
415
416
  ):
416
417
  batch, device = x.shape[0], x.device
417
418
 
@@ -421,8 +422,10 @@ class Unet(nn.Module):
421
422
 
422
423
  classes_emb = self.classes_emb(classes)
423
424
 
424
- if cond_drop_prob > 0:
425
- keep_mask = prob_mask_like((batch,), 1 - cond_drop_prob, device = device)
425
+ if cond_drop_prob > 0 or exists(cond_keep_mask):
426
+ keep_mask = default(cond_keep_mask, lambda: prob_mask_like((batch,), 1 - cond_drop_prob, device = device))
427
+ assert keep_mask.shape == (batch,)
428
+ keep_mask = keep_mask.to(device = device, dtype = torch.bool)
426
429
  null_classes_emb = repeat(self.null_classes_emb, 'd -> b d', b = batch)
427
430
 
428
431
  classes_emb = torch.where(
@@ -606,6 +609,20 @@ class GaussianDiffusion(nn.Module):
606
609
  def device(self):
607
610
  return self.betas.device
608
611
 
612
+ def random_times(self, batch_size):
613
+ return torch.randint(0, self.num_timesteps, (batch_size,), device = self.device).long()
614
+
615
+ def random_cond_keep_mask(self, batch_size):
616
+ keep_prob = 1. - self.model.cond_drop_prob
617
+ return prob_mask_like((batch_size,), keep_prob, device = self.device)
618
+
619
+ def xm_shared_random_kwargs(self, batch_size):
620
+ # condition dropping is shared across Forward XM candidates
621
+
622
+ return dict(
623
+ cond_keep_mask = self.random_cond_keep_mask(batch_size)
624
+ )
625
+
609
626
  def predict_start_from_noise(self, x_t, t, noise):
610
627
  return (
611
628
  extract(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t -
@@ -774,7 +791,7 @@ class GaussianDiffusion(nn.Module):
774
791
  extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
775
792
  )
776
793
 
777
- def p_losses(self, x_start, t, *, classes, noise = None, loss_reduction = 'mean'):
794
+ def p_losses(self, x_start, t, *, classes, noise = None, cond_keep_mask = None, loss_reduction = 'mean'):
778
795
  b, c, h, w = x_start.shape
779
796
  noise = default(noise, lambda: torch.randn_like(x_start))
780
797
 
@@ -784,7 +801,8 @@ class GaussianDiffusion(nn.Module):
784
801
 
785
802
  # predict and take gradient step
786
803
 
787
- model_out = self.model(x, t, classes)
804
+ model_kwargs = dict(cond_keep_mask = cond_keep_mask) if exists(cond_keep_mask) else dict()
805
+ model_out = self.model(x, t, classes, **model_kwargs)
788
806
 
789
807
  if self.objective == 'pred_noise':
790
808
  target = noise
@@ -806,13 +824,13 @@ class GaussianDiffusion(nn.Module):
806
824
 
807
825
  return loss.mean()
808
826
 
809
- def forward(self, img, *args, loss_reduction = 'mean', **kwargs):
827
+ def forward(self, img, *args, times = None, cond_keep_mask = None, loss_reduction = 'mean', **kwargs):
810
828
  b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
811
829
  assert h == img_size and w == img_size, f'height and width of image must be {img_size}'
812
- t = torch.randint(0, self.num_timesteps, (b,), device=device).long()
830
+ times = default(times, lambda: self.random_times(b))
813
831
 
814
832
  img = normalize_to_neg_one_to_one(img)
815
- return self.p_losses(img, t, *args, loss_reduction = loss_reduction, **kwargs)
833
+ return self.p_losses(img, times, *args, cond_keep_mask = cond_keep_mask, loss_reduction = loss_reduction, **kwargs)
816
834
 
817
835
  # example
818
836
 
@@ -261,7 +261,7 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
261
261
 
262
262
  if self.min_snr_loss_weight:
263
263
  snr = log_snr.exp()
264
- loss_weight = snr.clamp(min = self.min_snr_gamma) / snr
264
+ loss_weight = snr.clamp(max = self.min_snr_gamma) / snr
265
265
  losses = losses * loss_weight
266
266
 
267
267
  if loss_reduction == 'none':
@@ -269,10 +269,10 @@ class ContinuousTimeGaussianDiffusion(nn.Module):
269
269
 
270
270
  return losses.mean()
271
271
 
272
- def forward(self, img, *args, loss_reduction = 'mean', **kwargs):
272
+ def forward(self, img, *args, times = None, loss_reduction = 'mean', **kwargs):
273
273
  b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
274
274
  assert h == img_size and w == img_size, f'height and width of image must be {img_size}'
275
275
 
276
- times = self.random_times(b)
276
+ times = default(times, self.random_times(b))
277
277
  img = normalize_to_neg_one_to_one(img)
278
278
  return self.p_losses(img, times, *args, loss_reduction = loss_reduction, **kwargs)
@@ -836,13 +836,17 @@ class GaussianDiffusion(Module):
836
836
 
837
837
  return loss.mean()
838
838
 
839
- def forward(self, img, *args, loss_reduction = 'mean', **kwargs):
839
+ def random_times(self, batch_size):
840
+ return torch.randint(0, self.num_timesteps, (batch_size,), device = self.device).long()
841
+
842
+ def forward(self, img, *args, times = None, loss_reduction = 'mean', **kwargs):
840
843
  b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
841
844
  assert h == img_size[0] and w == img_size[1], f'height and width of image must be {img_size}'
842
- t = torch.randint(0, self.num_timesteps, (b,), device=device).long()
845
+
846
+ times = default(times, lambda: self.random_times(b))
843
847
 
844
848
  img = self.normalize(img)
845
- return self.p_losses(img, t, *args, loss_reduction = loss_reduction, **kwargs)
849
+ return self.p_losses(img, times, *args, loss_reduction = loss_reduction, **kwargs)
846
850
 
847
851
  # dataset classes
848
852
 
@@ -545,6 +545,10 @@ 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
+ @property
549
+ def device(self):
550
+ return self.betas.device
551
+
548
552
  def model_predictions(self, x, t, x_self_cond = None, clip_x_start = False, rederive_pred_noise = False, model_forward_kwargs: dict = dict()):
549
553
 
550
554
  if exists(x_self_cond):
@@ -747,14 +751,18 @@ class GaussianDiffusion1D(Module):
747
751
 
748
752
  return loss.mean()
749
753
 
750
- def forward(self, img, *args, loss_reduction = 'mean', **kwargs):
754
+ def random_times(self, batch_size):
755
+ return torch.randint(0, self.num_timesteps, (batch_size,), device = self.device).long()
756
+
757
+ def forward(self, img, *args, times = None, loss_reduction = 'mean', **kwargs):
751
758
  b, n, device, seq_length, = img.shape[0], img.shape[self.seq_index], img.device, self.seq_length
752
759
 
753
760
  assert n == seq_length, f'seq length must be {seq_length}'
754
- t = torch.randint(0, self.num_timesteps, (b,), device=device).long()
761
+
762
+ times = default(times, lambda: self.random_times(b))
755
763
 
756
764
  img = self.normalize(img)
757
- return self.p_losses(img, t, *args, loss_reduction = loss_reduction, **kwargs)
765
+ return self.p_losses(img, times, *args, loss_reduction = loss_reduction, **kwargs)
758
766
 
759
767
  # trainer class
760
768
 
@@ -244,7 +244,7 @@ class ElucidatedDiffusion(nn.Module):
244
244
  def noise_distribution(self, batch_size):
245
245
  return (self.P_mean + self.P_std * torch.randn((batch_size,), device = self.device)).exp()
246
246
 
247
- def forward(self, images, loss_reduction = 'mean'):
247
+ def forward(self, images, *args, sigma = None, loss_reduction = 'mean', **kwargs):
248
248
  batch_size, c, h, w, device, image_size, channels = *images.shape, images.device, self.image_size, self.channels
249
249
 
250
250
  assert h == image_size and w == image_size, f'height and width of image must be {image_size}'
@@ -252,7 +252,7 @@ class ElucidatedDiffusion(nn.Module):
252
252
 
253
253
  images = normalize_to_neg_one_to_one(images)
254
254
 
255
- sigmas = self.noise_distribution(batch_size)
255
+ sigmas = default(sigma, lambda: self.noise_distribution(batch_size))
256
256
  padded_sigmas = rearrange(sigmas, 'b -> b 1 1 1')
257
257
 
258
258
  noise = torch.randn_like(images)
@@ -183,10 +183,10 @@ class VParamContinuousTimeGaussianDiffusion(nn.Module):
183
183
 
184
184
  return losses.mean()
185
185
 
186
- def forward(self, img, *args, loss_reduction = 'mean', **kwargs):
186
+ def forward(self, img, *args, times = None, loss_reduction = 'mean', **kwargs):
187
187
  b, c, h, w, device, img_size, = *img.shape, img.device, self.image_size
188
188
  assert h == img_size and w == img_size, f'height and width of image must be {img_size}'
189
189
 
190
- times = self.random_times(b)
190
+ times = default(times, self.random_times(b))
191
191
  img = normalize_to_neg_one_to_one(img)
192
192
  return self.p_losses(img, times, *args, loss_reduction = loss_reduction, **kwargs)
@@ -0,0 +1 @@
1
+ __version__ = '2.3.2'
@@ -1,5 +1,6 @@
1
1
  from inspect import signature
2
2
 
3
+ import torch
3
4
  from torch.nn import Module
4
5
  from torch import cat, is_tensor
5
6
  from torch.utils._pytree import tree_map, tree_flatten
@@ -30,7 +31,9 @@ class XMWrapper(Module):
30
31
  self,
31
32
  flow_model: Module,
32
33
  candidates = 1,
33
- max_batch_size = None
34
+ max_batch_size = None,
35
+ random_time_method = 'random_times',
36
+ random_time_kwarg = 'times'
34
37
  ):
35
38
  super().__init__()
36
39
  self.flow_model = flow_model
@@ -40,6 +43,9 @@ class XMWrapper(Module):
40
43
  self.max_batch_size = max_batch_size
41
44
  self.has_loss_reduction = 'loss_reduction' in signature(flow_model.forward).parameters
42
45
 
46
+ self.random_time_method = random_time_method
47
+ self.random_time_kwarg = random_time_kwarg
48
+
43
49
  @property
44
50
  def data_shape(self):
45
51
  return getattr(self.flow_model, 'data_shape', None)
@@ -71,6 +77,20 @@ class XMWrapper(Module):
71
77
  first_tensor = next(t for t in leaves if is_tensor(t))
72
78
  batch = first_tensor.shape[0]
73
79
 
80
+ # sample anything that must be shared before expanding candidates
81
+
82
+ if hasattr(self.flow_model, 'xm_shared_random_kwargs'):
83
+ random_kwargs = self.flow_model.xm_shared_random_kwargs(batch)
84
+
85
+ for key, value in random_kwargs.items():
86
+ if key not in kwargs:
87
+ kwargs[key] = value
88
+
89
+ if self.random_time_kwarg not in kwargs:
90
+ assert hasattr(self.flow_model, self.random_time_method), f'flow_model must have a {self.random_time_method} method'
91
+ fn = getattr(self.flow_model, self.random_time_method)
92
+ kwargs[self.random_time_kwarg] = fn(batch)
93
+
74
94
  # repeat inputs K candidates times
75
95
 
76
96
  args_K, kwargs_K = tree_map(
@@ -1 +0,0 @@
1
- __version__ = '2.3.0'