MCPcopy Create free account
hub / github.com/VisionXLab/OF-Diff / noise_pred_fn

Function noise_pred_fn

ldm/models/diffusion/dpm_solver/dpm_solver.py:257–278  ·  view source on GitHub ↗
(x, t_continuous, cond=None)

Source from the content-addressed store, hash-verified

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 """

Callers 1

model_fnFunction · 0.85

Calls 4

get_model_input_timeFunction · 0.85
expand_dimsFunction · 0.85
marginal_alphaMethod · 0.80
marginal_stdMethod · 0.80

Tested by

no test coverage detected