| 55 | |
| 56 | |
| 57 | def next_step(model_output: Union[torch.FloatTensor, np.ndarray], timestep: int, |
| 58 | sample: Union[torch.FloatTensor, np.ndarray], ddim_scheduler): |
| 59 | timestep, next_timestep = min( |
| 60 | timestep - ddim_scheduler.config.num_train_timesteps // ddim_scheduler.num_inference_steps, 999), timestep |
| 61 | alpha_prod_t = ddim_scheduler.alphas_cumprod[timestep] if timestep >= 0 else ddim_scheduler.final_alpha_cumprod |
| 62 | alpha_prod_t_next = ddim_scheduler.alphas_cumprod[next_timestep] |
| 63 | beta_prod_t = 1 - alpha_prod_t |
| 64 | next_original_sample = (sample - beta_prod_t ** 0.5 * model_output) / alpha_prod_t ** 0.5 |
| 65 | next_sample_direction = (1 - alpha_prod_t_next) ** 0.5 * model_output |
| 66 | next_sample = alpha_prod_t_next ** 0.5 * next_original_sample + next_sample_direction |
| 67 | return next_sample |
| 68 | |
| 69 | |
| 70 | def get_noise_pred_single(latents, t, context, unet): |