denoising-diffusion-pytorch 2.2.4__tar.gz → 2.2.6__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.4 → denoising_diffusion_pytorch-2.2.6}/PKG-INFO +1 -1
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/attend.py +6 -5
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/continuous_time_gaussian_diffusion.py +1 -1
- denoising_diffusion_pytorch-2.2.6/denoising_diffusion_pytorch/version.py +1 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch.egg-info/PKG-INFO +1 -1
- denoising_diffusion_pytorch-2.2.4/denoising_diffusion_pytorch/version.py +0 -1
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/LICENSE +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/README.md +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/__init__.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/classifier_free_guidance.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/denoising_diffusion_pytorch_1d.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/elucidated_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/fid_evaluation.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/guided_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/karras_unet.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/karras_unet_1d.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/karras_unet_3d.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/learned_gaussian_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/repaint.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/simple_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/v_param_continuous_time_gaussian_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch/weighted_objective_gaussian_diffusion.py +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch.egg-info/SOURCES.txt +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch.egg-info/dependency_links.txt +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch.egg-info/requires.txt +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/denoising_diffusion_pytorch.egg-info/top_level.txt +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/setup.cfg +0 -0
- {denoising_diffusion_pytorch-2.2.4 → denoising_diffusion_pytorch-2.2.6}/setup.py +0 -0
|
@@ -7,10 +7,11 @@ from torch import nn, einsum
|
|
|
7
7
|
import torch.nn.functional as F
|
|
8
8
|
|
|
9
9
|
from einops import rearrange
|
|
10
|
+
from torch.nn.attention import SDPBackend
|
|
10
11
|
|
|
11
12
|
# constants
|
|
12
13
|
|
|
13
|
-
AttentionConfig = namedtuple('AttentionConfig', ['
|
|
14
|
+
AttentionConfig = namedtuple('AttentionConfig', ['backends'])
|
|
14
15
|
|
|
15
16
|
# helpers
|
|
16
17
|
|
|
@@ -52,7 +53,7 @@ class Attend(nn.Module):
|
|
|
52
53
|
|
|
53
54
|
# determine efficient attention configs for cuda and cpu
|
|
54
55
|
|
|
55
|
-
self.cpu_config = AttentionConfig(
|
|
56
|
+
self.cpu_config = AttentionConfig([SDPBackend.FLASH_ATTENTION, SDPBackend.MATH, SDPBackend.EFFICIENT_ATTENTION])
|
|
56
57
|
self.cuda_config = None
|
|
57
58
|
|
|
58
59
|
if not torch.cuda.is_available() or not flash:
|
|
@@ -64,10 +65,10 @@ class Attend(nn.Module):
|
|
|
64
65
|
|
|
65
66
|
if device_version > version.parse('8.0'):
|
|
66
67
|
print_once('A100 GPU detected, using flash attention if input tensor is on cuda')
|
|
67
|
-
self.cuda_config = AttentionConfig(
|
|
68
|
+
self.cuda_config = AttentionConfig([SDPBackend.FLASH_ATTENTION])
|
|
68
69
|
else:
|
|
69
70
|
print_once('Non-A100 GPU detected, using math or mem efficient attention if input tensor is on cuda')
|
|
70
|
-
self.cuda_config = AttentionConfig(
|
|
71
|
+
self.cuda_config = AttentionConfig([SDPBackend.MATH, SDPBackend.EFFICIENT_ATTENTION])
|
|
71
72
|
|
|
72
73
|
def flash_attn(self, q, k, v):
|
|
73
74
|
_, heads, q_len, _, k_len, is_cuda, device = *q.shape, k.shape[-2], q.is_cuda, q.device
|
|
@@ -84,7 +85,7 @@ class Attend(nn.Module):
|
|
|
84
85
|
|
|
85
86
|
# pytorch 2.0 flash attn: q, k, v, mask, dropout, causal, softmax_scale
|
|
86
87
|
|
|
87
|
-
with torch.
|
|
88
|
+
with torch.nn.attention.sdpa_kernel(**config._asdict()):
|
|
88
89
|
out = F.scaled_dot_product_attention(
|
|
89
90
|
q, k, v,
|
|
90
91
|
dropout_p = self.dropout if self.training else 0.
|
|
@@ -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(
|
|
264
|
+
loss_weight = snr.clamp(max = self.min_snr_gamma) / snr
|
|
265
265
|
losses = losses * loss_weight
|
|
266
266
|
|
|
267
267
|
return losses.mean()
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = '2.2.6'
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = '2.2.4'
|
|
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
|