(x, t_continuous, cond=None)
| 288 | return t_continuous |
| 289 | |
| 290 | def noise_pred_fn(x, t_continuous, cond=None): |
| 291 | t_input = get_model_input_time(t_continuous) |
| 292 | if cond is None: |
| 293 | output = model(x, t_input, **model_kwargs) |
| 294 | else: |
| 295 | output = model(x, t_input, cond, **model_kwargs) |
| 296 | if model_type == "noise": |
| 297 | return output |
| 298 | elif model_type == "x_start": |
| 299 | alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous) |
| 300 | return (x - alpha_t * output) / sigma_t |
| 301 | elif model_type == "v": |
| 302 | alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous) |
| 303 | return alpha_t * output + sigma_t * x |
| 304 | elif model_type == "score": |
| 305 | sigma_t = noise_schedule.marginal_std(t_continuous) |
| 306 | return -sigma_t * output |
| 307 | |
| 308 | def cond_grad_fn(x, t_input): |
| 309 | """ |
no test coverage detected