(beta_schedule, *, beta_start, beta_end, num_diffusion_timesteps)
| 3 | import numpy as np |
| 4 | |
| 5 | def get_beta_schedule(beta_schedule, *, beta_start, beta_end, num_diffusion_timesteps): |
| 6 | def sigmoid(x): |
| 7 | return 1 / (np.exp(-x) + 1) |
| 8 | |
| 9 | if beta_schedule == "quad": |
| 10 | betas = ( |
| 11 | np.linspace( |
| 12 | beta_start ** 0.5, |
| 13 | beta_end ** 0.5, |
| 14 | num_diffusion_timesteps, |
| 15 | dtype=np.float64, |
| 16 | ) |
| 17 | ** 2 |
| 18 | ) |
| 19 | elif beta_schedule == "linear": |
| 20 | betas = np.linspace( |
| 21 | beta_start, beta_end, num_diffusion_timesteps, dtype=np.float64 |
| 22 | ) |
| 23 | elif beta_schedule == "const": |
| 24 | betas = beta_end * np.ones(num_diffusion_timesteps, dtype=np.float64) |
| 25 | elif beta_schedule == "jsd": # 1/T, 1/(T-1), 1/(T-2), ..., 1 |
| 26 | betas = 1.0 / np.linspace( |
| 27 | num_diffusion_timesteps, 1, num_diffusion_timesteps, dtype=np.float64 |
| 28 | ) |
| 29 | elif beta_schedule == "sigmoid": |
| 30 | betas = np.linspace(-6, 6, num_diffusion_timesteps) |
| 31 | betas = sigmoid(betas) * (beta_end - beta_start) + beta_start |
| 32 | elif beta_schedule == "cosine": |
| 33 | return betas_for_alpha_bar( |
| 34 | num_diffusion_timesteps, |
| 35 | lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2, |
| 36 | ) |
| 37 | else: |
| 38 | raise NotImplementedError(beta_schedule) |
| 39 | assert betas.shape == (num_diffusion_timesteps,) |
| 40 | return betas |
| 41 | |
| 42 | def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999): |
| 43 | """ |
no test coverage detected