633 lines
22 KiB
Python
633 lines
22 KiB
Python
# -*- coding: utf-8 -*-
|
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
|
"""
|
|
ODE/SDE solver for denoising diffusion models under either variation preserving (VP) or
|
|
variation exploding (VE) settings. Under the VE setting, the diffusion process is:
|
|
|
|
q(x_t | x_0) = N(x_t | x_0, sigma_t^2 I),
|
|
|
|
where 0 <= sigma_t <= inf; while under the VP setting, the diffusion process is:
|
|
|
|
q(x_t | x_0) = N(x_t | alpha_t x_0, sigma_t^2 I),
|
|
|
|
where 0 <= sigma_t <= 1 and alpha_t^2 = 1 - sigma_t^2.
|
|
"""
|
|
import torch
|
|
from tqdm.auto import trange
|
|
|
|
__all__ = [
|
|
'sample_euler', 'sample_euler_ancestral', 'sample_heun', 'sample_dpm_2',
|
|
'sample_dpm_2_ancestral', 'sample_dpmpp_2s_ancestral', 'sample_dpmpp_sde',
|
|
'sample_dpmpp_2m', 'sample_dpmpp_2m_sde', 'sample_ddim'
|
|
]
|
|
|
|
# -------------------- variation exploding (VE) solver --------------------#
|
|
|
|
|
|
def get_ancestral_step(sigma_from, sigma_to, eta=1.):
|
|
"""
|
|
Calculates the noise level (sigma_down) to step down to and the amount
|
|
of noise to add (sigma_up) when doing an ancestral sampling step.
|
|
"""
|
|
if not eta:
|
|
return sigma_to, 0.
|
|
sigma_up = min(
|
|
sigma_to,
|
|
eta * (sigma_to**2 *
|
|
(sigma_from**2 - sigma_to**2) / sigma_from**2)**0.5)
|
|
sigma_down = (sigma_to**2 - sigma_up**2)**0.5
|
|
return sigma_down, sigma_up
|
|
|
|
|
|
def get_scalings(sigma):
|
|
c_out = -sigma
|
|
c_in = 1 / (sigma**2 + 1.**2)**0.5
|
|
return c_out, c_in
|
|
|
|
|
|
@torch.no_grad()
|
|
def sample_euler(noise,
|
|
model,
|
|
sigmas,
|
|
s_churn=0.,
|
|
s_tmin=0.,
|
|
s_tmax=float('inf'),
|
|
s_noise=1.,
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
Implements Algorithm 2 (Euler steps) from Karras et al. (2022).
|
|
"""
|
|
x = noise * sigmas[0]
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
gamma = 0.
|
|
if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'):
|
|
gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1)
|
|
eps = torch.randn_like(x) * s_noise
|
|
sigma_hat = sigmas[i] * (gamma + 1)
|
|
if gamma > 0:
|
|
x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5
|
|
# Euler method
|
|
if sigmas[i] == float('inf'):
|
|
denoised = model(noise, sigma_hat)
|
|
x = denoised + sigmas[i + 1] * (gamma + 1) * noise
|
|
else:
|
|
_, c_in = get_scalings(sigma_hat)
|
|
denoised = model(x * c_in, sigma_hat)
|
|
d = (x - denoised) / sigma_hat
|
|
dt = sigmas[i + 1] - sigma_hat
|
|
x = x + d * dt
|
|
return x
|
|
|
|
|
|
@torch.no_grad()
|
|
def sample_euler_ancestral(noise,
|
|
model,
|
|
sigmas,
|
|
eta=1.,
|
|
s_noise=1.,
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
Ancestral sampling with Euler method steps.
|
|
"""
|
|
x = noise * sigmas[0]
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
sigma_down, sigma_up = get_ancestral_step(sigmas[i],
|
|
sigmas[i + 1],
|
|
eta=eta)
|
|
# Euler method
|
|
if sigmas[i] == float('inf'):
|
|
denoised = model(noise, sigmas[i])
|
|
x = denoised + sigmas[i + 1] * noise
|
|
else:
|
|
_, c_in = get_scalings(sigmas[i])
|
|
denoised = model(x * c_in, sigmas[i])
|
|
d = (x - denoised) / sigmas[i]
|
|
dt = sigma_down - sigmas[i]
|
|
x = x + d * dt
|
|
if sigmas[i + 1] > 0:
|
|
x = x + torch.randn_like(x) * s_noise * sigma_up
|
|
return x
|
|
|
|
|
|
@torch.no_grad()
|
|
def sample_heun(noise,
|
|
model,
|
|
sigmas,
|
|
s_churn=0.,
|
|
s_tmin=0.,
|
|
s_tmax=float('inf'),
|
|
s_noise=1.,
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
Implements Algorithm 2 (Heun steps) from Karras et al. (2022).
|
|
"""
|
|
x = noise * sigmas[0]
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
gamma = 0.
|
|
if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'):
|
|
gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1)
|
|
eps = torch.randn_like(x) * s_noise
|
|
sigma_hat = sigmas[i] * (gamma + 1)
|
|
if gamma > 0:
|
|
x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5
|
|
if sigmas[i] == float('inf'):
|
|
# Euler method
|
|
denoised = model(noise, sigma_hat)
|
|
x = denoised + sigmas[i + 1] * (gamma + 1) * noise
|
|
else:
|
|
_, c_in = get_scalings(sigma_hat)
|
|
denoised = model(x * c_in, sigma_hat)
|
|
d = (x - denoised) / sigma_hat
|
|
dt = sigmas[i + 1] - sigma_hat
|
|
if sigmas[i + 1] == 0:
|
|
# Euler method
|
|
x = x + d * dt
|
|
else:
|
|
# Heun's method
|
|
x_2 = x + d * dt
|
|
_, c_in = get_scalings(sigmas[i + 1])
|
|
denoised_2 = model(x_2 * c_in, sigmas[i + 1])
|
|
d_2 = (x_2 - denoised_2) / sigmas[i + 1]
|
|
d_prime = (d + d_2) / 2
|
|
x = x + d_prime * dt
|
|
return x
|
|
|
|
|
|
@torch.no_grad()
|
|
def sample_dpm_2(noise,
|
|
model,
|
|
sigmas,
|
|
s_churn=0.,
|
|
s_tmin=0.,
|
|
s_tmax=float('inf'),
|
|
s_noise=1.,
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
A sampler inspired by DPM-Solver-2 and Algorithm 2 from Karras et al. (2022).
|
|
"""
|
|
x = noise * sigmas[0]
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
gamma = 0.
|
|
if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'):
|
|
gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1)
|
|
eps = torch.randn_like(x) * s_noise
|
|
sigma_hat = sigmas[i] * (gamma + 1)
|
|
if gamma > 0:
|
|
x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5
|
|
if sigmas[i] == float('inf'):
|
|
# Euler method
|
|
denoised = model(noise, sigma_hat)
|
|
x = denoised + sigmas[i + 1] * (gamma + 1) * noise
|
|
else:
|
|
_, c_in = get_scalings(sigma_hat)
|
|
denoised = model(x * c_in, sigma_hat)
|
|
d = (x - denoised) / sigma_hat
|
|
if sigmas[i + 1] == 0:
|
|
# Euler method
|
|
dt = sigmas[i + 1] - sigma_hat
|
|
x = x + d * dt
|
|
else:
|
|
# DPM-Solver-2
|
|
sigma_mid = sigma_hat.log().lerp(sigmas[i + 1].log(),
|
|
0.5).exp()
|
|
dt_1 = sigma_mid - sigma_hat
|
|
dt_2 = sigmas[i + 1] - sigma_hat
|
|
x_2 = x + d * dt_1
|
|
_, c_in = get_scalings(sigma_mid)
|
|
denoised_2 = model(x_2 * c_in, sigma_mid)
|
|
d_2 = (x_2 - denoised_2) / sigma_mid
|
|
x = x + d_2 * dt_2
|
|
return x
|
|
|
|
|
|
@torch.no_grad()
|
|
def sample_dpm_2_ancestral(noise,
|
|
model,
|
|
sigmas,
|
|
eta=1.,
|
|
s_noise=1.,
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
Ancestral sampling with DPM-Solver second-order steps.
|
|
"""
|
|
x = noise * sigmas[0]
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
sigma_down, sigma_up = get_ancestral_step(sigmas[i],
|
|
sigmas[i + 1],
|
|
eta=eta)
|
|
if sigmas[i] == float('inf'):
|
|
# Euler method
|
|
denoised = model(noise, sigmas[i])
|
|
x = denoised + sigmas[i + 1] * noise
|
|
else:
|
|
_, c_in = get_scalings(sigmas[i])
|
|
denoised = model(x * c_in, sigmas[i])
|
|
d = (x - denoised) / sigmas[i]
|
|
if sigma_down == 0:
|
|
# Euler method
|
|
dt = sigma_down - sigmas[i]
|
|
x = x + d * dt
|
|
else:
|
|
# DPM-Solver-2
|
|
sigma_mid = sigmas[i].log().lerp(sigma_down.log(), 0.5).exp()
|
|
dt_1 = sigma_mid - sigmas[i]
|
|
dt_2 = sigma_down - sigmas[i]
|
|
x_2 = x + d * dt_1
|
|
_, c_in = get_scalings(sigma_mid)
|
|
denoised_2 = model(x_2 * c_in, sigma_mid)
|
|
d_2 = (x_2 - denoised_2) / sigma_mid
|
|
x = x + d_2 * dt_2
|
|
x = x + torch.randn_like(x) * s_noise * sigma_up
|
|
return x
|
|
|
|
|
|
@torch.no_grad()
|
|
def sample_dpmpp_2s_ancestral(noise,
|
|
model,
|
|
sigmas,
|
|
eta=1.,
|
|
s_noise=1.,
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
Ancestral sampling with DPM-Solver++ (2S) second-order steps.
|
|
"""
|
|
def t_to_sigma(t):
|
|
return t.neg().exp()
|
|
|
|
def sigma_to_t(sigma):
|
|
return sigma.log().neg()
|
|
|
|
# x = noise * sigmas[0]
|
|
x = noise * torch.sqrt(1.0 + sigmas[0]**2.0)
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
sigma_down, sigma_up = get_ancestral_step(sigmas[i],
|
|
sigmas[i + 1],
|
|
eta=eta)
|
|
if sigmas[i] == float('inf'):
|
|
# Euler method
|
|
denoised = model(noise, sigmas[i])
|
|
x = denoised + sigma_down * noise
|
|
else:
|
|
_, c_in = get_scalings(sigmas[i])
|
|
denoised = model(x * c_in, sigmas[i])
|
|
if sigma_down == 0:
|
|
# Euler method
|
|
d = (x - denoised) / sigmas[i]
|
|
dt = sigma_down - sigmas[i]
|
|
x = x + d * dt
|
|
else:
|
|
# DPM-Solver++(2S)
|
|
t, t_next = sigma_to_t(sigmas[i]), sigma_to_t(sigma_down)
|
|
r = 1 / 2
|
|
h = t_next - t
|
|
s = t + r * h
|
|
x_2 = (t_to_sigma(s) /
|
|
t_to_sigma(t)) * x - (-h * r).expm1() * denoised
|
|
_, c_in = get_scalings(t_to_sigma(s))
|
|
denoised_2 = model(x_2 * c_in, t_to_sigma(s))
|
|
x = (t_to_sigma(t_next) /
|
|
t_to_sigma(t)) * x - (-h).expm1() * denoised_2
|
|
# Noise addition
|
|
if sigmas[i + 1] > 0:
|
|
x = x + torch.randn_like(x) * s_noise * sigma_up
|
|
return x
|
|
|
|
|
|
class BatchedBrownianTree:
|
|
"""
|
|
A wrapper around torchsde.BrownianTree that enables batches of entropy.
|
|
"""
|
|
def __init__(self, x, t0, t1, seed=None, **kwargs):
|
|
import torchsde
|
|
t0, t1, self.sign = self.sort(t0, t1)
|
|
w0 = kwargs.get('w0', torch.zeros_like(x))
|
|
if seed is None:
|
|
seed = torch.randint(0, 2**63 - 1, []).item()
|
|
self.batched = True
|
|
try:
|
|
assert len(seed) == x.shape[0]
|
|
w0 = w0[0]
|
|
except TypeError:
|
|
seed = [seed]
|
|
self.batched = False
|
|
self.trees = [
|
|
torchsde.BrownianTree(t0, w0, t1, entropy=s, **kwargs)
|
|
for s in seed
|
|
]
|
|
|
|
@staticmethod
|
|
def sort(a, b):
|
|
return (a, b, 1) if a < b else (b, a, -1)
|
|
|
|
def __call__(self, t0, t1):
|
|
t0, t1, sign = self.sort(t0, t1)
|
|
w = torch.stack([tree(t0, t1)
|
|
for tree in self.trees]) * (self.sign * sign)
|
|
return w if self.batched else w[0]
|
|
|
|
|
|
class BrownianTreeNoiseSampler:
|
|
"""
|
|
A noise sampler backed by a torchsde.BrownianTree.
|
|
|
|
Args:
|
|
x (Tensor): The tensor whose shape, device and dtype to use to generate
|
|
random samples.
|
|
sigma_min (float): The low end of the valid interval.
|
|
sigma_max (float): The high end of the valid interval.
|
|
seed (int or List[int]): The random seed. If a list of seeds is
|
|
supplied instead of a single integer, then the noise sampler will
|
|
use one BrownianTree per batch item, each with its own seed.
|
|
transform (callable): A function that maps sigma to the sampler's
|
|
internal timestep.
|
|
"""
|
|
def __init__(self,
|
|
x,
|
|
sigma_min,
|
|
sigma_max,
|
|
seed=None,
|
|
transform=lambda x: x):
|
|
self.transform = transform
|
|
t0 = self.transform(torch.as_tensor(sigma_min))
|
|
t1 = self.transform(torch.as_tensor(sigma_max))
|
|
self.tree = BatchedBrownianTree(x, t0, t1, seed)
|
|
|
|
def __call__(self, sigma, sigma_next):
|
|
t0 = self.transform(torch.as_tensor(sigma))
|
|
t1 = self.transform(torch.as_tensor(sigma_next))
|
|
return self.tree(t0, t1) / (t1 - t0).abs().sqrt()
|
|
|
|
|
|
@torch.no_grad()
|
|
def sample_dpmpp_sde(noise,
|
|
model,
|
|
sigmas,
|
|
eta=1.,
|
|
s_noise=1.,
|
|
r=1 / 2,
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
DPM-Solver++ (stochastic).
|
|
"""
|
|
def t_to_sigma(t):
|
|
return t.neg().exp()
|
|
|
|
def sigma_to_t(sigma):
|
|
return sigma.log().neg()
|
|
|
|
x = noise * sigmas[0]
|
|
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas[
|
|
sigmas < float('inf')].max()
|
|
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed)
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
if sigmas[i] == float('inf'):
|
|
# Euler method
|
|
denoised = model(noise, sigmas[i])
|
|
x = denoised + sigmas[i + 1] * noise
|
|
else:
|
|
_, c_in = get_scalings(sigmas[i])
|
|
denoised = model(x * c_in, sigmas[i])
|
|
if sigmas[i + 1] == 0:
|
|
# Euler method
|
|
d = (x - denoised) / sigmas[i]
|
|
dt = sigmas[i + 1] - sigmas[i]
|
|
x = x + d * dt
|
|
else:
|
|
# DPM-Solver++
|
|
t, t_next = sigma_to_t(sigmas[i]), sigma_to_t(sigmas[i + 1])
|
|
h = t_next - t
|
|
s = t + h * r
|
|
fac = 1 / (2 * r)
|
|
|
|
# Step 1
|
|
sd, su = get_ancestral_step(t_to_sigma(t), t_to_sigma(s), eta)
|
|
s_ = sigma_to_t(sd)
|
|
x_2 = (t_to_sigma(s_) /
|
|
t_to_sigma(t)) * x - (t - s_).expm1() * denoised
|
|
x_2 = x_2 + noise_sampler(t_to_sigma(t),
|
|
t_to_sigma(s)) * s_noise * su
|
|
_, c_in = get_scalings(t_to_sigma(s))
|
|
denoised_2 = model(x_2 * c_in, t_to_sigma(s))
|
|
|
|
# Step 2
|
|
sd, su = get_ancestral_step(t_to_sigma(t), t_to_sigma(t_next),
|
|
eta)
|
|
t_next_ = sigma_to_t(sd)
|
|
denoised_d = (1 - fac) * denoised + fac * denoised_2
|
|
x = (t_to_sigma(t_next_) / t_to_sigma(t)) * x - \
|
|
(t - t_next_).expm1() * denoised_d
|
|
x = x + noise_sampler(t_to_sigma(t),
|
|
t_to_sigma(t_next)) * s_noise * su
|
|
return x
|
|
|
|
|
|
@torch.no_grad()
|
|
def sample_dpmpp_2m(noise,
|
|
model,
|
|
sigmas,
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
DPM-Solver++ (2M).
|
|
"""
|
|
def t_to_sigma(t):
|
|
return t.neg().exp()
|
|
|
|
def sigma_to_t(sigma):
|
|
return sigma.log().neg()
|
|
|
|
x = noise * sigmas[0]
|
|
old_denoised = None
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
if sigmas[i] == float('inf'):
|
|
# Euler method
|
|
denoised = model(noise, sigmas[i])
|
|
x = denoised + sigmas[i + 1] * noise
|
|
else:
|
|
_, c_in = get_scalings(sigmas[i])
|
|
denoised = model(x * c_in, sigmas[i])
|
|
t, t_next = sigma_to_t(sigmas[i]), sigma_to_t(sigmas[i + 1])
|
|
h = t_next - t
|
|
if (old_denoised is None or sigmas[i - 1] == float('inf')
|
|
or sigmas[i + 1] == 0):
|
|
x = (t_to_sigma(t_next) /
|
|
t_to_sigma(t)) * x - (-h).expm1() * denoised
|
|
else:
|
|
h_last = t - sigma_to_t(sigmas[i - 1])
|
|
r = h_last / h
|
|
denoised_d = (1 + 1 /
|
|
(2 * r)) * denoised - (1 /
|
|
(2 * r)) * old_denoised
|
|
x = (t_to_sigma(t_next) /
|
|
t_to_sigma(t)) * x - (-h).expm1() * denoised_d
|
|
old_denoised = denoised
|
|
return x
|
|
|
|
|
|
@torch.no_grad()
|
|
def sample_dpmpp_2m_sde(noise,
|
|
model,
|
|
sigmas,
|
|
eta=1.,
|
|
s_noise=1.,
|
|
solver_type='midpoint',
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
DPM-Solver++ (2M) SDE.
|
|
"""
|
|
assert solver_type in {'heun', 'midpoint'}
|
|
|
|
x = noise * sigmas[0]
|
|
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas[
|
|
sigmas < float('inf')].max()
|
|
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed)
|
|
old_denoised = None
|
|
h_last = None
|
|
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
if sigmas[i] == float('inf'):
|
|
# Euler method
|
|
denoised = model(noise, sigmas[i])
|
|
x = denoised + sigmas[i + 1] * noise
|
|
else:
|
|
_, c_in = get_scalings(sigmas[i])
|
|
denoised = model(x * c_in, sigmas[i])
|
|
if sigmas[i + 1] == 0:
|
|
# Denoising step
|
|
x = denoised
|
|
else:
|
|
# DPM-Solver++(2M) SDE
|
|
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
|
|
h = s - t
|
|
eta_h = eta * h
|
|
|
|
x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + \
|
|
(-h - eta_h).expm1().neg() * denoised
|
|
|
|
if old_denoised is not None:
|
|
r = h_last / h
|
|
if solver_type == 'heun':
|
|
x = x + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * \
|
|
(1 / r) * (denoised - old_denoised)
|
|
elif solver_type == 'midpoint':
|
|
x = x + 0.5 * (-h - eta_h).expm1().neg() * \
|
|
(1 / r) * (denoised - old_denoised)
|
|
|
|
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[
|
|
i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise
|
|
|
|
old_denoised = denoised
|
|
h_last = h
|
|
return x
|
|
|
|
|
|
# -------------------- variation preserving (VP) solver --------------------#
|
|
@torch.no_grad()
|
|
def sample_ddim(noise,
|
|
model,
|
|
sigmas,
|
|
eta=0.,
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
DDIM solver steps.
|
|
"""
|
|
x = noise
|
|
sigmas_vp = (sigmas**2 / (1 + sigmas**2))**0.5
|
|
sigmas_vp[sigmas == float('inf')] = 1.
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
denoised = model(x, sigmas[i])
|
|
noise_factor = eta * (sigmas_vp[i + 1]**2 / sigmas_vp[i]**2 *
|
|
(1 - (1 - sigmas_vp[i]**2) /
|
|
(1 - sigmas_vp[i + 1]**2)))
|
|
d = (x - (1 - sigmas_vp[i]**2)**0.5 * denoised) / sigmas_vp[i]
|
|
x = (1 - sigmas_vp[i + 1] ** 2) ** 0.5 * denoised + \
|
|
(sigmas_vp[i + 1] ** 2 - noise_factor ** 2) ** 0.5 * d
|
|
if sigmas_vp[i + 1] > 0:
|
|
x += noise_factor * torch.randn_like(x)
|
|
return x
|
|
|
|
|
|
@torch.no_grad()
|
|
def sample_img2img_euler(noise,
|
|
model,
|
|
sigmas,
|
|
s_churn=0.,
|
|
s_tmin=0.,
|
|
s_tmax=float('inf'),
|
|
s_noise=1.,
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
Implements Algorithm 2 (Euler steps) from Karras et al. (2022).
|
|
"""
|
|
x = noise
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
gamma = 0.
|
|
if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'):
|
|
gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1)
|
|
eps = torch.randn_like(x) * s_noise
|
|
sigma_hat = sigmas[i] * (gamma + 1)
|
|
if gamma > 0:
|
|
x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5
|
|
# Euler method
|
|
if sigmas[i] == float('inf'):
|
|
denoised = model(noise, sigma_hat)
|
|
x = denoised + sigmas[i + 1] * (gamma + 1) * noise
|
|
else:
|
|
denoised = model(x, sigma_hat)
|
|
d = (x - denoised) / sigma_hat
|
|
dt = sigmas[i + 1] - sigma_hat
|
|
x = x + d * dt
|
|
return x
|
|
|
|
|
|
@torch.no_grad()
|
|
def sample_img2img_euler_ancestral(noise,
|
|
model,
|
|
sigmas,
|
|
eta=1.,
|
|
s_noise=1.,
|
|
seed=None,
|
|
show_progress=True,
|
|
**kwargs):
|
|
"""
|
|
Ancestral sampling with Euler method steps.
|
|
"""
|
|
x = noise
|
|
for i in trange(len(sigmas) - 1, disable=not show_progress):
|
|
sigma_down, sigma_up = get_ancestral_step(sigmas[i],
|
|
sigmas[i + 1],
|
|
eta=eta)
|
|
# Euler method
|
|
if sigmas[i] == float('inf'):
|
|
denoised = model(noise, sigmas[i])
|
|
x = denoised + sigmas[i + 1] * noise
|
|
else:
|
|
denoised = model(x, sigmas[i])
|
|
d = (x - denoised) / sigmas[i]
|
|
dt = sigma_down - sigmas[i]
|
|
x = x + d * dt
|
|
if sigmas[i + 1] > 0:
|
|
x = x + torch.randn_like(x) * s_noise * sigma_up
|
|
return x
|