This is the deprecated API for creating beta schedules. See get_named_beta_schedule() for the new library of schedules.
(beta_schedule, *, beta_start, beta_end, num_diffusion_timesteps)
| 63 | |
| 64 | |
| 65 | def get_beta_schedule(beta_schedule, *, beta_start, beta_end, num_diffusion_timesteps): |
| 66 | """ |
| 67 | This is the deprecated API for creating beta schedules. |
| 68 | See get_named_beta_schedule() for the new library of schedules. |
| 69 | """ |
| 70 | if beta_schedule == "quad": |
| 71 | betas = ( |
| 72 | np.linspace( |
| 73 | beta_start ** 0.5, |
| 74 | beta_end ** 0.5, |
| 75 | num_diffusion_timesteps, |
| 76 | dtype=np.float64, |
| 77 | ) |
| 78 | ** 2 |
| 79 | ) |
| 80 | elif beta_schedule == "linear": |
| 81 | betas = np.linspace(beta_start, beta_end, num_diffusion_timesteps, dtype=np.float64) |
| 82 | elif beta_schedule == "warmup10": |
| 83 | betas = _warmup_beta(beta_start, beta_end, num_diffusion_timesteps, 0.1) |
| 84 | elif beta_schedule == "warmup50": |
| 85 | betas = _warmup_beta(beta_start, beta_end, num_diffusion_timesteps, 0.5) |
| 86 | elif beta_schedule == "const": |
| 87 | betas = beta_end * np.ones(num_diffusion_timesteps, dtype=np.float64) |
| 88 | elif beta_schedule == "jsd": # 1/T, 1/(T-1), 1/(T-2), ..., 1 |
| 89 | betas = 1.0 / np.linspace( |
| 90 | num_diffusion_timesteps, 1, num_diffusion_timesteps, dtype=np.float64 |
| 91 | ) |
| 92 | else: |
| 93 | raise NotImplementedError(beta_schedule) |
| 94 | assert betas.shape == (num_diffusion_timesteps,) |
| 95 | return betas |
| 96 | |
| 97 | |
| 98 | def get_named_beta_schedule(schedule_name, num_diffusion_timesteps): |
no test coverage detected