MCPcopy Create free account
hub / github.com/CompVis/diff2flow / make_beta_schedule

Function make_beta_schedule

diff2flow/utils/diffusion_utils.py:5–27  ·  view source on GitHub ↗
(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3)

Source from the content-addressed store, hash-verified

3
4
5def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3):
6 if schedule == "linear":
7 betas = (
8 torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=torch.float64) ** 2
9 )
10
11 elif schedule == "cosine":
12 timesteps = (
13 torch.arange(n_timestep + 1, dtype=torch.float64) / n_timestep + cosine_s
14 )
15 alphas = timesteps / (1 + cosine_s) * np.pi / 2
16 alphas = torch.cos(alphas).pow(2)
17 alphas = alphas / alphas[0]
18 betas = 1 - alphas[1:] / alphas[:-1]
19 betas = np.clip(betas, a_min=0, a_max=0.999)
20
21 elif schedule == "sqrt_linear":
22 betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64)
23 elif schedule == "sqrt":
24 betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64) ** 0.5
25 else:
26 raise ValueError(f"schedule '{schedule}' unknown.")
27 return betas.numpy()
28
29def enforce_zero_terminal_snr(betas):
30 # Copied from https://openaccess.thecvf.com/content/WACV2024/papers/Lin_Common_Diffusion_Noise_Schedules_and_Sample_Steps_Are_Flawed_WACV_2024_paper.pdf

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected