(x, t_continuous, cond=None)
| 255 | return t_continuous |
| 256 | |
| 257 | def noise_pred_fn(x, t_continuous, cond=None): |
| 258 | if t_continuous.reshape((-1,)).shape[0] == 1: |
| 259 | t_continuous = t_continuous.expand((x.shape[0])) |
| 260 | t_input = get_model_input_time(t_continuous) |
| 261 | if cond is None: |
| 262 | output = model(x, t_input, **model_kwargs) |
| 263 | else: |
| 264 | output = model(x, t_input, cond, **model_kwargs) |
| 265 | if model_type == "noise": |
| 266 | return output |
| 267 | elif model_type == "x_start": |
| 268 | alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous) |
| 269 | dims = x.dim() |
| 270 | return (x - expand_dims(alpha_t, dims) * output) / expand_dims(sigma_t, dims) |
| 271 | elif model_type == "v": |
| 272 | alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous) |
| 273 | dims = x.dim() |
| 274 | return expand_dims(alpha_t, dims) * output + expand_dims(sigma_t, dims) * x |
| 275 | elif model_type == "score": |
| 276 | sigma_t = noise_schedule.marginal_std(t_continuous) |
| 277 | dims = x.dim() |
| 278 | return -expand_dims(sigma_t, dims) * output |
| 279 | |
| 280 | def cond_grad_fn(x, t_input): |
| 281 | """ |
no test coverage detected