(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3)
| 3 | |
| 4 | |
| 5 | def 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 | |
| 29 | def 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 |
no outgoing calls
no test coverage detected