update v0.0.4
This commit is contained in:
@@ -13,8 +13,7 @@ from .schedules import karras_schedule
|
||||
from .solvers import (sample_ddim, sample_dpm_2, sample_dpm_2_ancestral,
|
||||
sample_dpmpp_2m, sample_dpmpp_2m_sde,
|
||||
sample_dpmpp_2s_ancestral, sample_dpmpp_sde,
|
||||
sample_euler, sample_euler_ancestral, sample_heun,
|
||||
sample_img2img_euler, sample_img2img_euler_ancestral)
|
||||
sample_euler, sample_euler_ancestral, sample_heun)
|
||||
|
||||
__all__ = ['GaussianDiffusion']
|
||||
|
||||
@@ -27,6 +26,148 @@ def _i(tensor, t, x):
|
||||
return tensor[t.to(tensor.device)].view(shape).to(x.device)
|
||||
|
||||
|
||||
def _unpack_2d_ks(kernel_size):
|
||||
if isinstance(kernel_size, int):
|
||||
ky = kx = kernel_size
|
||||
else:
|
||||
assert len(
|
||||
kernel_size) == 2, '2D Kernel size should have a length of 2.'
|
||||
ky, kx = kernel_size
|
||||
|
||||
ky = int(ky)
|
||||
kx = int(kx)
|
||||
return ky, kx
|
||||
|
||||
|
||||
def _compute_zero_padding(kernel_size):
|
||||
ky, kx = _unpack_2d_ks(kernel_size)
|
||||
return (ky - 1) // 2, (kx - 1) // 2
|
||||
|
||||
|
||||
def _bilateral_blur(
|
||||
input,
|
||||
guidance,
|
||||
kernel_size,
|
||||
sigma_color,
|
||||
sigma_space,
|
||||
border_type='reflect',
|
||||
color_distance_type='l1',
|
||||
):
|
||||
|
||||
if isinstance(sigma_color, torch.Tensor):
|
||||
sigma_color = sigma_color.to(device=input.device,
|
||||
dtype=input.dtype).view(-1, 1, 1, 1, 1)
|
||||
|
||||
ky, kx = _unpack_2d_ks(kernel_size)
|
||||
pad_y, pad_x = _compute_zero_padding(kernel_size)
|
||||
|
||||
padded_input = torch.nn.functional.pad(input, (pad_x, pad_x, pad_y, pad_y),
|
||||
mode=border_type)
|
||||
unfolded_input = padded_input.unfold(2, ky, 1).unfold(3, kx, 1).flatten(
|
||||
-2) # (B, C, H, W, Ky x Kx)
|
||||
|
||||
if guidance is None:
|
||||
guidance = input
|
||||
unfolded_guidance = unfolded_input
|
||||
else:
|
||||
padded_guidance = torch.nn.functional.pad(guidance,
|
||||
(pad_x, pad_x, pad_y, pad_y),
|
||||
mode=border_type)
|
||||
unfolded_guidance = padded_guidance.unfold(2, ky, 1).unfold(
|
||||
3, kx, 1).flatten(-2) # (B, C, H, W, Ky x Kx)
|
||||
|
||||
diff = unfolded_guidance - guidance.unsqueeze(-1)
|
||||
if color_distance_type == 'l1':
|
||||
color_distance_sq = diff.abs().sum(1, keepdim=True).square()
|
||||
elif color_distance_type == 'l2':
|
||||
color_distance_sq = diff.square().sum(1, keepdim=True)
|
||||
else:
|
||||
raise ValueError('color_distance_type only acceps l1 or l2')
|
||||
color_kernel = (-0.5 / sigma_color**2 *
|
||||
color_distance_sq).exp() # (B, 1, H, W, Ky x Kx)
|
||||
|
||||
space_kernel = get_gaussian_kernel2d(kernel_size,
|
||||
sigma_space,
|
||||
device=input.device,
|
||||
dtype=input.dtype)
|
||||
space_kernel = space_kernel.view(-1, 1, 1, 1, kx * ky)
|
||||
|
||||
kernel = space_kernel * color_kernel
|
||||
out = (unfolded_input * kernel).sum(-1) / kernel.sum(-1)
|
||||
return out
|
||||
|
||||
|
||||
def get_gaussian_kernel1d(
|
||||
kernel_size,
|
||||
sigma,
|
||||
force_even,
|
||||
*,
|
||||
device=None,
|
||||
dtype=None,
|
||||
):
|
||||
|
||||
return gaussian(kernel_size, sigma, device=device, dtype=dtype)
|
||||
|
||||
|
||||
def gaussian(window_size, sigma, *, device=None, dtype=None):
|
||||
|
||||
batch_size = sigma.shape[0]
|
||||
|
||||
x = (torch.arange(window_size, device=sigma.device, dtype=sigma.dtype) -
|
||||
window_size // 2).expand(batch_size, -1)
|
||||
|
||||
if window_size % 2 == 0:
|
||||
x = x + 0.5
|
||||
|
||||
gauss = torch.exp(-x.pow(2.0) / (2 * sigma.pow(2.0)))
|
||||
|
||||
return gauss / gauss.sum(-1, keepdim=True)
|
||||
|
||||
|
||||
def get_gaussian_kernel2d(
|
||||
kernel_size,
|
||||
sigma,
|
||||
force_even=False,
|
||||
*,
|
||||
device=None,
|
||||
dtype=None,
|
||||
):
|
||||
|
||||
sigma = torch.Tensor([[sigma, sigma]]).to(device=device, dtype=dtype)
|
||||
|
||||
ksize_y, ksize_x = _unpack_2d_ks(kernel_size)
|
||||
sigma_y, sigma_x = sigma[:, 0, None], sigma[:, 1, None]
|
||||
|
||||
kernel_y = get_gaussian_kernel1d(ksize_y,
|
||||
sigma_y,
|
||||
force_even,
|
||||
device=device,
|
||||
dtype=dtype)[..., None]
|
||||
kernel_x = get_gaussian_kernel1d(ksize_x,
|
||||
sigma_x,
|
||||
force_even,
|
||||
device=device,
|
||||
dtype=dtype)[..., None]
|
||||
|
||||
return kernel_y * kernel_x.view(-1, 1, ksize_x)
|
||||
|
||||
|
||||
def adaptive_anisotropic_filter(x, g=None):
|
||||
if g is None:
|
||||
g = x
|
||||
s, m = torch.std_mean(g, dim=(1, 2, 3), keepdim=True)
|
||||
s = s + 1e-5
|
||||
guidance = (g - m) / s
|
||||
y = _bilateral_blur(x,
|
||||
guidance,
|
||||
kernel_size=(13, 13),
|
||||
sigma_color=3.0,
|
||||
sigma_space=3.0,
|
||||
border_type='reflect',
|
||||
color_distance_type='l1')
|
||||
return y
|
||||
|
||||
|
||||
class GaussianDiffusion(object):
|
||||
def __init__(self, sigmas, prediction_type='eps'):
|
||||
assert prediction_type in {'x0', 'eps', 'v'}
|
||||
@@ -53,6 +194,7 @@ class GaussianDiffusion(object):
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
clamp=None,
|
||||
sharpness=0.0,
|
||||
percentile=None,
|
||||
cat_uc=False,
|
||||
**kwargs):
|
||||
@@ -79,54 +221,99 @@ class GaussianDiffusion(object):
|
||||
|
||||
# prediction
|
||||
if guide_scale is None:
|
||||
assert isinstance(model_kwargs, dict)
|
||||
out = model(xt, t=t, **model_kwargs, **kwargs)
|
||||
if isinstance(model_kwargs, dict):
|
||||
out = model(xt, t=t, **model_kwargs, **kwargs)
|
||||
elif isinstance(model_kwargs, list) and len(model_kwargs) > 0:
|
||||
out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
else:
|
||||
raise Exception('Error')
|
||||
else:
|
||||
# classifier-free guidance (arXiv:2207.12598)
|
||||
# model_kwargs[0]: conditional kwargs
|
||||
# model_kwargs[1]: non-conditional kwargs
|
||||
assert isinstance(model_kwargs, list) and len(model_kwargs) == 2
|
||||
|
||||
if guide_scale == 1.:
|
||||
out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
else:
|
||||
if cat_uc:
|
||||
|
||||
def parse_model_kwargs(prev_value, value):
|
||||
if isinstance(value, torch.Tensor):
|
||||
prev_value = torch.cat([prev_value, value], dim=0)
|
||||
elif isinstance(value, dict):
|
||||
for k, v in value.items():
|
||||
prev_value[k] = parse_model_kwargs(
|
||||
prev_value[k], v)
|
||||
elif isinstance(value, list):
|
||||
for idx, v in enumerate(value):
|
||||
prev_value[idx] = parse_model_kwargs(
|
||||
prev_value[idx], v)
|
||||
return prev_value
|
||||
|
||||
all_model_kwargs = copy.deepcopy(model_kwargs[0])
|
||||
for model_kwarg in model_kwargs[1:]:
|
||||
for key, value in model_kwarg.items():
|
||||
all_model_kwargs[key] = parse_model_kwargs(
|
||||
all_model_kwargs[key], value)
|
||||
all_out = model(xt.repeat(2, 1, 1, 1),
|
||||
t=t.repeat(2),
|
||||
**all_model_kwargs,
|
||||
**kwargs)
|
||||
y_out, u_out = all_out.chunk(2)
|
||||
assert isinstance(model_kwargs, list) and len(model_kwargs) >= 2
|
||||
if isinstance(guide_scale, float) or isinstance(guide_scale, int):
|
||||
assert len(model_kwargs) == 2
|
||||
if guide_scale == 1.:
|
||||
out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
else:
|
||||
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
u_out = model(xt, t=t, **model_kwargs[1], **kwargs)
|
||||
out = u_out + guide_scale * (y_out - u_out)
|
||||
if cat_uc:
|
||||
|
||||
# rescale the output according to arXiv:2305.08891
|
||||
if guide_rescale is not None:
|
||||
assert guide_rescale >= 0 and guide_rescale <= 1
|
||||
ratio = (y_out.flatten(1).std(dim=1) /
|
||||
(out.flatten(1).std(dim=1) +
|
||||
1e-12)).view((-1, ) + (1, ) * (y_out.ndim - 1))
|
||||
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
|
||||
def parse_model_kwargs(prev_value, value):
|
||||
if isinstance(value, torch.Tensor):
|
||||
prev_value = torch.cat([prev_value, value],
|
||||
dim=0)
|
||||
elif isinstance(value, dict):
|
||||
for k, v in value.items():
|
||||
prev_value[k] = parse_model_kwargs(
|
||||
prev_value[k], v)
|
||||
elif isinstance(value, list):
|
||||
for idx, v in enumerate(value):
|
||||
prev_value[idx] = parse_model_kwargs(
|
||||
prev_value[idx], v)
|
||||
return prev_value
|
||||
|
||||
all_model_kwargs = copy.deepcopy(model_kwargs[0])
|
||||
for model_kwarg in model_kwargs[1:]:
|
||||
for key, value in model_kwarg.items():
|
||||
all_model_kwargs[key] = parse_model_kwargs(
|
||||
all_model_kwargs[key], value)
|
||||
all_out = model(xt.repeat(2, 1, 1, 1),
|
||||
t=t.repeat(2),
|
||||
**all_model_kwargs,
|
||||
**kwargs)
|
||||
y_out, u_out = all_out.chunk(2)
|
||||
else:
|
||||
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
u_out = model(xt, t=t, **model_kwargs[1], **kwargs)
|
||||
# todo sharpness
|
||||
# sharpness sampling
|
||||
if sharpness is not None and sharpness > 0:
|
||||
positive_x0 = alphas * xt - sigmas * y_out
|
||||
negative_x0 = alphas * xt - sigmas * u_out
|
||||
|
||||
positive_eps = xt - positive_x0
|
||||
negative_eps = xt - negative_x0
|
||||
|
||||
global_diffusion_progress = (
|
||||
1 - t / 999.0).detach().cpu().numpy().tolist()[0]
|
||||
alpha = 0.001 * sharpness * global_diffusion_progress
|
||||
positive_eps_degraded = adaptive_anisotropic_filter(
|
||||
x=positive_eps, g=positive_x0)
|
||||
positive_eps_degraded_weighted = positive_eps_degraded * alpha + positive_eps * (
|
||||
1.0 - alpha)
|
||||
|
||||
final_eps = negative_eps + guide_scale * (
|
||||
positive_eps_degraded_weighted - negative_eps)
|
||||
final_x0 = xt - final_eps
|
||||
out = (alphas * xt - final_x0) / sigmas
|
||||
else:
|
||||
out = u_out + guide_scale * (y_out - u_out)
|
||||
elif isinstance(guide_scale, dict):
|
||||
assert len(model_kwargs) == 3
|
||||
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
m_out = model(xt, t=t, **model_kwargs[1], **kwargs)
|
||||
u_out = model(xt, t=t, **model_kwargs[2], **kwargs)
|
||||
out = u_out + guide_scale['image'] * (
|
||||
m_out - u_out) + guide_scale['text'] * (y_out - m_out)
|
||||
elif isinstance(guide_scale, list):
|
||||
assert len(guide_scale) == len(model_kwargs) - 1
|
||||
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
outs = [y_out]
|
||||
for i in range(1, len(model_kwargs)):
|
||||
outs.append(model(xt, t=t, **model_kwargs[i], **kwargs))
|
||||
out = outs[-1]
|
||||
for i in range(len(guide_scale)):
|
||||
out += guide_scale[i] * (outs[-i - 2] - outs[-i - 1])
|
||||
|
||||
# rescale the output according to arXiv:2305.08891
|
||||
if guide_rescale is not None and guide_rescale > 0.0:
|
||||
assert guide_rescale >= 0 and guide_rescale <= 1
|
||||
ratio = (
|
||||
y_out.flatten(1).std(dim=1) /
|
||||
(out.flatten(1).std(dim=1) + 1e-12)).view((-1, ) + (1, ) *
|
||||
(y_out.ndim - 1))
|
||||
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
|
||||
# compute x0
|
||||
if self.prediction_type == 'x0':
|
||||
x0 = out
|
||||
@@ -197,6 +384,7 @@ class GaussianDiffusion(object):
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
clamp=None,
|
||||
sharpness=0.0,
|
||||
percentile=None,
|
||||
solver='euler_a',
|
||||
steps=20,
|
||||
@@ -209,12 +397,16 @@ class GaussianDiffusion(object):
|
||||
seed=-1,
|
||||
intermediate_callback=None,
|
||||
cat_uc=False,
|
||||
add_noise=False,
|
||||
free_steps=None,
|
||||
step_offset=None,
|
||||
**kwargs):
|
||||
# sanity check
|
||||
assert isinstance(steps, (int, torch.LongTensor))
|
||||
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
|
||||
assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1)
|
||||
assert discretization in (None, 'leading', 'linspace', 'trailing')
|
||||
assert discretization in (None, 'leading', 'linspace', 'trailing',
|
||||
'free')
|
||||
assert discard_penultimate_step in (None, True, False)
|
||||
assert return_intermediate in (None, 'x0', 'xt')
|
||||
|
||||
@@ -255,17 +447,50 @@ class GaussianDiffusion(object):
|
||||
def model_fn(xt, sigma):
|
||||
# denoising
|
||||
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
|
||||
x0 = self.denoise(xt,
|
||||
t,
|
||||
None,
|
||||
model,
|
||||
model_kwargs,
|
||||
guide_scale,
|
||||
guide_rescale,
|
||||
clamp,
|
||||
percentile,
|
||||
cat_uc=cat_uc,
|
||||
**kwargs)[-2]
|
||||
|
||||
if isinstance(
|
||||
model_kwargs[0]['cond'], dict) and \
|
||||
'tar_x0' in model_kwargs[0]['cond'] and \
|
||||
'tar_mask_latent' in model_kwargs[0]['cond']:
|
||||
tar_x0 = model_kwargs[0]['cond']['tar_x0']
|
||||
tar_mask = model_kwargs[0]['cond']['tar_mask_latent']
|
||||
|
||||
tar_xt = self.diffuse(x0=tar_x0, t=t)
|
||||
xt = tar_xt * (1.0 - tar_mask) + xt * tar_mask
|
||||
|
||||
if isinstance(model_kwargs[0]['cond'],
|
||||
dict) and 'ref_x0' in model_kwargs[0]['cond']:
|
||||
model_kwargs[0]['cond']['ref_xt'] = self.diffuse(
|
||||
x0=model_kwargs[0]['cond']['ref_x0'], t=t)
|
||||
model_kwargs[1]['cond']['ref_xt'] = self.diffuse(
|
||||
x0=model_kwargs[1]['cond']['ref_x0'], t=t)
|
||||
|
||||
if solver in ('onestep', 'multistep', 'multistep2', 'multistep3'):
|
||||
x0 = self.denoise(xt,
|
||||
t,
|
||||
None,
|
||||
model,
|
||||
model_kwargs,
|
||||
guide_scale,
|
||||
guide_rescale,
|
||||
clamp,
|
||||
sharpness,
|
||||
percentile,
|
||||
cat_uc=cat_uc,
|
||||
**kwargs)[-3]
|
||||
else:
|
||||
x0 = self.denoise(xt,
|
||||
t,
|
||||
None,
|
||||
model,
|
||||
model_kwargs,
|
||||
guide_scale,
|
||||
guide_rescale,
|
||||
clamp,
|
||||
sharpness,
|
||||
percentile,
|
||||
cat_uc=cat_uc,
|
||||
**kwargs)[-2]
|
||||
|
||||
# collect intermediate outputs
|
||||
if return_intermediate == 'xt':
|
||||
@@ -291,10 +516,14 @@ class GaussianDiffusion(object):
|
||||
elif discretization == 'trailing':
|
||||
steps = torch.arange(t_max, t_min - 1,
|
||||
-((t_max - t_min + 1) / steps))
|
||||
elif discretization == 'free':
|
||||
steps = torch.tensor(free_steps)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'{discretization} discretization not implemented')
|
||||
steps = steps.clamp_(t_min, t_max)
|
||||
elif isinstance(steps, list):
|
||||
steps = torch.tensor(steps)
|
||||
steps = torch.as_tensor(steps,
|
||||
dtype=torch.float32,
|
||||
device=noise.device)
|
||||
@@ -335,6 +564,23 @@ class GaussianDiffusion(object):
|
||||
if discard_penultimate_step:
|
||||
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
|
||||
kwargs['seed'] = seed
|
||||
# add noise to x0
|
||||
if add_noise:
|
||||
if 'dm_steps' in kwargs:
|
||||
if step_offset:
|
||||
add_noise_step = -kwargs['dm_steps'] + step_offset
|
||||
if add_noise_step < 0:
|
||||
noise = self.diffuse(
|
||||
noise,
|
||||
torch.full((noise.shape[0], 1),
|
||||
steps[add_noise_step],
|
||||
dtype=torch.int))
|
||||
else:
|
||||
noise = self.diffuse(
|
||||
noise,
|
||||
torch.full((noise.shape[0], 1),
|
||||
steps[-kwargs['dm_steps'] - 1],
|
||||
dtype=torch.int))
|
||||
# sampling
|
||||
x0 = solver_fn(noise,
|
||||
model_fn,
|
||||
@@ -373,168 +619,6 @@ class GaussianDiffusion(object):
|
||||
| torch.isinf(log_sigma)] = float('inf')
|
||||
return log_sigma.exp()
|
||||
|
||||
@torch.no_grad()
|
||||
def stochastic_encode(self, x0, t, steps):
|
||||
# fast, but does not allow for exact reconstruction
|
||||
# t serves as an index to gather the correct alphas
|
||||
|
||||
t_max = None
|
||||
t_min = None
|
||||
|
||||
# discretization method
|
||||
discretization = 'trailing' if self.prediction_type == 'v' else 'leading'
|
||||
|
||||
# timesteps
|
||||
if isinstance(steps, int):
|
||||
t_max = self.num_timesteps - 1 if t_max is None else t_max
|
||||
t_min = 0 if t_min is None else t_min
|
||||
steps = discretize_timesteps(t_max, t_min, steps, discretization)
|
||||
steps = torch.as_tensor(steps).round().long().flip(0).to(x0.device)
|
||||
# steps = torch.as_tensor(steps).round().long().to(x0.device)
|
||||
|
||||
# self.alphas_bar = torch.cumprod(1 - self.sigmas ** 2, dim=0)
|
||||
# print('sigma: ', self.sigmas, len(self.sigmas))
|
||||
# print('alpha_bar: ', self.alphas_bar, len(self.alphas_bar))
|
||||
# print('steps: ', steps, len(steps))
|
||||
# sqrt_alphas_cumprod = torch.sqrt(self.alphas_bar).to(x0.device)[steps]
|
||||
# sqrt_one_minus_alphas_cumprod = torch.sqrt(1 - self.alphas_bar).to(x0.device)[steps]
|
||||
|
||||
sqrt_alphas_cumprod = self.alphas.to(x0.device)[steps]
|
||||
sqrt_one_minus_alphas_cumprod = self.sigmas.to(x0.device)[steps]
|
||||
# print('sigma: ', self.sigmas, len(self.sigmas))
|
||||
# print('alpha: ', self.alphas, len(self.alphas))
|
||||
# print('steps: ', steps, len(steps))
|
||||
|
||||
noise = torch.randn_like(x0)
|
||||
return (
|
||||
extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 +
|
||||
extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) *
|
||||
noise)
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_img2img(self,
|
||||
x,
|
||||
noise,
|
||||
model,
|
||||
denoising_strength=1,
|
||||
model_kwargs={},
|
||||
condition_fn=None,
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
solver='euler_a',
|
||||
steps=20,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discretization=None,
|
||||
discard_penultimate_step=None,
|
||||
return_intermediate=None,
|
||||
show_progress=False,
|
||||
seed=-1,
|
||||
**kwargs):
|
||||
# sanity check
|
||||
assert isinstance(steps, (int, torch.LongTensor))
|
||||
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
|
||||
assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1)
|
||||
assert discretization in (None, 'leading', 'linspace', 'trailing')
|
||||
assert discard_penultimate_step in (None, True, False)
|
||||
assert return_intermediate in (None, 'x0', 'xt')
|
||||
# function of diffusion solver
|
||||
solver_fn = {
|
||||
'euler_ancestral': sample_img2img_euler_ancestral,
|
||||
'euler': sample_img2img_euler,
|
||||
}[solver]
|
||||
# options
|
||||
schedule = 'karras' if 'karras' in solver else None
|
||||
discretization = discretization or 'linspace'
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**31)
|
||||
if isinstance(steps, torch.LongTensor):
|
||||
discard_penultimate_step = False
|
||||
if discard_penultimate_step is None:
|
||||
discard_penultimate_step = True if solver in (
|
||||
'dpm2', 'dpm2_ancestral', 'dpmpp_2m_sde', 'dpm2_karras',
|
||||
'dpm2_ancestral_karras', 'dpmpp_2m_sde_karras') else False
|
||||
|
||||
# function for denoising xt to get x0
|
||||
intermediates = []
|
||||
|
||||
def get_scalings(sigma):
|
||||
c_out = -sigma
|
||||
c_in = 1 / (sigma**2 + 1.**2)**0.5
|
||||
return c_out, c_in
|
||||
|
||||
def model_fn(xt, sigma):
|
||||
# denoising
|
||||
c_out, c_in = get_scalings(sigma)
|
||||
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
|
||||
|
||||
x0 = self.denoise(xt * c_in, t, None, model, model_kwargs,
|
||||
guide_scale, guide_rescale, clamp, percentile,
|
||||
**kwargs)[-2]
|
||||
# collect intermediate outputs
|
||||
if return_intermediate == 'xt':
|
||||
intermediates.append(xt)
|
||||
elif return_intermediate == 'x0':
|
||||
intermediates.append(x0)
|
||||
return xt + x0 * c_out
|
||||
|
||||
# get timesteps
|
||||
if isinstance(steps, int):
|
||||
steps += 1 if discard_penultimate_step else 0
|
||||
t_max = self.num_timesteps - 1 if t_max is None else t_max
|
||||
t_min = 0 if t_min is None else t_min
|
||||
# discretize timesteps
|
||||
if discretization == 'leading':
|
||||
steps = torch.arange(t_min, t_max + 1,
|
||||
(t_max - t_min + 1) / steps).flip(0)
|
||||
elif discretization == 'linspace':
|
||||
steps = torch.linspace(t_max, t_min, steps)
|
||||
elif discretization == 'trailing':
|
||||
steps = torch.arange(t_max, t_min - 1,
|
||||
-((t_max - t_min + 1) / steps))
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'{discretization} discretization not implemented')
|
||||
steps = steps.clamp_(t_min, t_max)
|
||||
steps = torch.as_tensor(steps, dtype=torch.float32, device=x.device)
|
||||
# get sigmas
|
||||
sigmas = self._t_to_sigma(steps)
|
||||
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
|
||||
t_enc = int(min(denoising_strength, 0.999) * len(steps))
|
||||
sigmas = sigmas[len(steps) - t_enc - 1:]
|
||||
noise = x + noise * sigmas[0]
|
||||
|
||||
if schedule == 'karras':
|
||||
if sigmas[0] == float('inf'):
|
||||
sigmas = karras_schedule(
|
||||
n=len(steps) - 1,
|
||||
sigma_min=sigmas[sigmas > 0].min().item(),
|
||||
sigma_max=sigmas[sigmas < float('inf')].max().item(),
|
||||
rho=7.).to(sigmas)
|
||||
sigmas = torch.cat([
|
||||
sigmas.new_tensor([float('inf')]), sigmas,
|
||||
sigmas.new_zeros([1])
|
||||
])
|
||||
else:
|
||||
sigmas = karras_schedule(
|
||||
n=len(steps),
|
||||
sigma_min=sigmas[sigmas > 0].min().item(),
|
||||
sigma_max=sigmas.max().item(),
|
||||
rho=7.).to(sigmas)
|
||||
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
|
||||
if discard_penultimate_step:
|
||||
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
|
||||
|
||||
# sampling
|
||||
x0 = solver_fn(noise,
|
||||
model_fn,
|
||||
sigmas,
|
||||
seed=seed,
|
||||
show_progress=show_progress,
|
||||
**kwargs)
|
||||
return (x0, intermediates) if return_intermediate is not None else x0
|
||||
|
||||
|
||||
def extract_into_tensor(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
|
||||
Reference in New Issue
Block a user