MCPcopy Create free account
hub / github.com/Meshcapade/difflocks / sample_dpmpp_2m

Function sample_dpmpp_2m

k_diffusion/sampling.py:585–607  ·  view source on GitHub ↗

DPM-Solver++(2M).

(model, x, sigmas, extra_args=None, callback=None, disable=None)

Source from the content-addressed store, hash-verified

583
584@torch.no_grad()
585def sample_dpmpp_2m(model, x, sigmas, extra_args=None, callback=None, disable=None):
586 """DPM-Solver++(2M)."""
587 extra_args = {} if extra_args is None else extra_args
588 s_in = x.new_ones([x.shape[0]])
589 sigma_fn = lambda t: t.neg().exp()
590 t_fn = lambda sigma: sigma.log().neg()
591 old_denoised = None
592
593 for i in trange(len(sigmas) - 1, disable=disable):
594 denoised = model(x, sigmas[i] * s_in, **extra_args)
595 if callback is not None:
596 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
597 t, t_next = t_fn(sigmas[i]), t_fn(sigmas[i + 1])
598 h = t_next - t
599 if old_denoised is None or sigmas[i + 1] == 0:
600 x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised
601 else:
602 h_last = t - t_fn(sigmas[i - 1])
603 r = h_last / h
604 denoised_d = (1 + 1 / (2 * r)) * denoised - (1 / (2 * r)) * old_denoised
605 x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_d
606 old_denoised = denoised
607 return x
608
609
610@torch.no_grad()

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected