DPM-Solver++(2M).
(model, x, sigmas, extra_args=None, callback=None, disable=None)
| 583 | |
| 584 | @torch.no_grad() |
| 585 | def 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() |
nothing calls this directly
no outgoing calls
no test coverage detected