new project from v0.0.1
This commit is contained in:
@@ -0,0 +1,611 @@
|
||||
# -*- 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):
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
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):
|
||||
"""
|
||||
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
|
||||
Reference in New Issue
Block a user