Skip to content

Commit 6441694

Browse files
gapatronyiyixuxugithub-actions[bot]
authored
Laplace Scheduler for DDPM (#11320)
* Add Laplace scheduler that samples more around mid-range noise levels (around log SNR=0), increasing performance (lower FID) with faster convergence speed, and robust to resolution and objective. Reference: https://arxiv.org/pdf/2407.03297. * Fix copies. * Apply style fixes --------- Co-authored-by: YiYi Xu <yixu310@gmail.com> Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
1 parent 632765a commit 6441694

26 files changed

Lines changed: 186 additions & 0 deletions

src/diffusers/schedulers/scheduling_consistency_decoder.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,13 @@ def betas_for_alpha_bar(
4040
def alpha_bar_fn(t):
4141
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
4242

43+
elif alpha_transform_type == "laplace":
44+
45+
def alpha_bar_fn(t):
46+
lmb = -0.5 * math.copysign(1, 0.5 - t) * math.log(1 - 2 * math.fabs(0.5 - t) + 1e-6)
47+
snr = math.exp(lmb)
48+
return math.sqrt(snr / (1 + snr))
49+
4350
elif alpha_transform_type == "exp":
4451

4552
def alpha_bar_fn(t):

src/diffusers/schedulers/scheduling_ddim.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,13 @@ def betas_for_alpha_bar(
7777
def alpha_bar_fn(t):
7878
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
7979

80+
elif alpha_transform_type == "laplace":
81+
82+
def alpha_bar_fn(t):
83+
lmb = -0.5 * math.copysign(1, 0.5 - t) * math.log(1 - 2 * math.fabs(0.5 - t) + 1e-6)
84+
snr = math.exp(lmb)
85+
return math.sqrt(snr / (1 + snr))
86+
8087
elif alpha_transform_type == "exp":
8188

8289
def alpha_bar_fn(t):

src/diffusers/schedulers/scheduling_ddim_cogvideox.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,13 @@ def betas_for_alpha_bar(
7777
def alpha_bar_fn(t):
7878
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
7979

80+
elif alpha_transform_type == "laplace":
81+
82+
def alpha_bar_fn(t):
83+
lmb = -0.5 * math.copysign(1, 0.5 - t) * math.log(1 - 2 * math.fabs(0.5 - t) + 1e-6)
84+
snr = math.exp(lmb)
85+
return math.sqrt(snr / (1 + snr))
86+
8087
elif alpha_transform_type == "exp":
8188

8289
def alpha_bar_fn(t):

src/diffusers/schedulers/scheduling_ddim_inverse.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -75,6 +75,13 @@ def betas_for_alpha_bar(
7575
def alpha_bar_fn(t):
7676
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
7777

78+
elif alpha_transform_type == "laplace":
79+
80+
def alpha_bar_fn(t):
81+
lmb = -0.5 * math.copysign(1, 0.5 - t) * math.log(1 - 2 * math.fabs(0.5 - t) + 1e-6)
82+
snr = math.exp(lmb)
83+
return math.sqrt(snr / (1 + snr))
84+
7885
elif alpha_transform_type == "exp":
7986

8087
def alpha_bar_fn(t):

src/diffusers/schedulers/scheduling_ddim_parallel.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,13 @@ def betas_for_alpha_bar(
7777
def alpha_bar_fn(t):
7878
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
7979

80+
elif alpha_transform_type == "laplace":
81+
82+
def alpha_bar_fn(t):
83+
lmb = -0.5 * math.copysign(1, 0.5 - t) * math.log(1 - 2 * math.fabs(0.5 - t) + 1e-6)
84+
snr = math.exp(lmb)
85+
return math.sqrt(snr / (1 + snr))
86+
8087
elif alpha_transform_type == "exp":
8188

8289
def alpha_bar_fn(t):

src/diffusers/schedulers/scheduling_ddpm.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -74,6 +74,13 @@ def betas_for_alpha_bar(
7474
def alpha_bar_fn(t):
7575
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
7676

77+
elif alpha_transform_type == "laplace":
78+
79+
def alpha_bar_fn(t):
80+
lmb = -0.5 * math.copysign(1, 0.5 - t) * math.log(1 - 2 * math.fabs(0.5 - t) + 1e-6)
81+
snr = math.exp(lmb)
82+
return math.sqrt(snr / (1 + snr))
83+
7784
elif alpha_transform_type == "exp":
7885

7986
def alpha_bar_fn(t):
@@ -207,6 +214,8 @@ def __init__(
207214
elif beta_schedule == "squaredcos_cap_v2":
208215
# Glide cosine schedule
209216
self.betas = betas_for_alpha_bar(num_train_timesteps)
217+
elif beta_schedule == "laplace":
218+
self.betas = betas_for_alpha_bar(num_train_timesteps, alpha_transform_type="laplace")
210219
elif beta_schedule == "sigmoid":
211220
# GeoDiff sigmoid schedule
212221
betas = torch.linspace(-6, 6, num_train_timesteps)

src/diffusers/schedulers/scheduling_ddpm_parallel.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -76,6 +76,13 @@ def betas_for_alpha_bar(
7676
def alpha_bar_fn(t):
7777
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
7878

79+
elif alpha_transform_type == "laplace":
80+
81+
def alpha_bar_fn(t):
82+
lmb = -0.5 * math.copysign(1, 0.5 - t) * math.log(1 - 2 * math.fabs(0.5 - t) + 1e-6)
83+
snr = math.exp(lmb)
84+
return math.sqrt(snr / (1 + snr))
85+
7986
elif alpha_transform_type == "exp":
8087

8188
def alpha_bar_fn(t):
@@ -217,6 +224,8 @@ def __init__(
217224
elif beta_schedule == "squaredcos_cap_v2":
218225
# Glide cosine schedule
219226
self.betas = betas_for_alpha_bar(num_train_timesteps)
227+
elif beta_schedule == "laplace":
228+
self.betas = betas_for_alpha_bar(num_train_timesteps, alpha_transform_type="laplace")
220229
elif beta_schedule == "sigmoid":
221230
# GeoDiff sigmoid schedule
222231
betas = torch.linspace(-6, 6, num_train_timesteps)

src/diffusers/schedulers/scheduling_deis_multistep.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -60,6 +60,13 @@ def betas_for_alpha_bar(
6060
def alpha_bar_fn(t):
6161
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
6262

63+
elif alpha_transform_type == "laplace":
64+
65+
def alpha_bar_fn(t):
66+
lmb = -0.5 * math.copysign(1, 0.5 - t) * math.log(1 - 2 * math.fabs(0.5 - t) + 1e-6)
67+
snr = math.exp(lmb)
68+
return math.sqrt(snr / (1 + snr))
69+
6370
elif alpha_transform_type == "exp":
6471

6572
def alpha_bar_fn(t):

src/diffusers/schedulers/scheduling_dpm_cogvideox.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -78,6 +78,13 @@ def betas_for_alpha_bar(
7878
def alpha_bar_fn(t):
7979
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
8080

81+
elif alpha_transform_type == "laplace":
82+
83+
def alpha_bar_fn(t):
84+
lmb = -0.5 * math.copysign(1, 0.5 - t) * math.log(1 - 2 * math.fabs(0.5 - t) + 1e-6)
85+
snr = math.exp(lmb)
86+
return math.sqrt(snr / (1 + snr))
87+
8188
elif alpha_transform_type == "exp":
8289

8390
def alpha_bar_fn(t):

src/diffusers/schedulers/scheduling_dpmsolver_multistep.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -60,6 +60,13 @@ def betas_for_alpha_bar(
6060
def alpha_bar_fn(t):
6161
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
6262

63+
elif alpha_transform_type == "laplace":
64+
65+
def alpha_bar_fn(t):
66+
lmb = -0.5 * math.copysign(1, 0.5 - t) * math.log(1 - 2 * math.fabs(0.5 - t) + 1e-6)
67+
snr = math.exp(lmb)
68+
return math.sqrt(snr / (1 + snr))
69+
6370
elif alpha_transform_type == "exp":
6471

6572
def alpha_bar_fn(t):

0 commit comments

Comments
 (0)