First commit.
Add RES_Momentumized Add DPMPP_DualSDE_Momentumized add egotistical Clyb_4M_SDE_Momentumized add TTM (Doesn't seem to work well with current imp.?) add LCM Custom Noise
This commit is contained in:
+17
@@ -0,0 +1,17 @@
|
||||
from . import extra_samplers
|
||||
from . import nodes
|
||||
|
||||
extra_samplers.add_samplers()
|
||||
#extra_samplers.add_schedulers()
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"SamplerCustomNoise": nodes.SamplerCustomNoise,
|
||||
"SamplerCustomNoiseDuo": nodes.SamplerCustomNoiseDuo,
|
||||
"SamplerCustomModelMixtureDuo": nodes.SamplerCustomModelMixtureDuo,
|
||||
"SamplerRES_Momentumized": nodes.SamplerRES_MOMENTUMIZED,
|
||||
"SamplerDPMPP_DualSDE_Momentumized": nodes.SamplerDPMPP_DUALSDE_MOMENTUMIZED,
|
||||
"SamplerCLYB_4M_SDE_Momentumized": nodes.SamplerCLYB_4M_SDE_MOMENTUMIZED,
|
||||
"SamplerTTM": nodes.SamplerTTM,
|
||||
"SamplerLCMCustom": nodes.SamplerLCMCustom,
|
||||
}
|
||||
__all__ = ['NODE_CLASS_MAPPINGS']
|
||||
@@ -0,0 +1,689 @@
|
||||
import math
|
||||
|
||||
from scipy import integrate
|
||||
import torch
|
||||
from torch import nn
|
||||
import torchsde
|
||||
from tqdm.auto import trange, tqdm
|
||||
|
||||
import comfy.sample
|
||||
|
||||
import k_diffusion.sampling
|
||||
from k_diffusion.sampling import BrownianTreeNoiseSampler, PIDStepSizeController, get_ancestral_step, to_d, default_noise_sampler
|
||||
|
||||
# The following function adds the samplers during initialization, in __init__.py
|
||||
def add_samplers():
|
||||
from comfy.samplers import KSampler, k_diffusion_sampling
|
||||
for sampler in extra_samplers: #getattr(self, "sample_{}".format(extra_samplers))
|
||||
if sampler not in KSampler.SAMPLERS:
|
||||
try:
|
||||
idx = KSampler.SAMPLERS.index("uni_pc_bh2") # Last item in the samplers list
|
||||
KSampler.SAMPLERS.insert(idx+1, sampler) # Add our custom samplers
|
||||
setattr(k_diffusion_sampling, "sample_{}".format(sampler), extra_samplers[sampler])
|
||||
import importlib
|
||||
importlib.reload(k_diffusion_sampling)
|
||||
except ValueError as err:
|
||||
pass
|
||||
|
||||
# The following function adds the samplers during initialization, in __init__.py
|
||||
def add_schedulers():
|
||||
from comfy.samplers import KSampler, k_diffusion_sampling
|
||||
for scheduler in extra_schedulers: #getattr(self, "sample_{}".format(extra_samplers))
|
||||
if scheduler not in KSampler.SCHEDULERS:
|
||||
try:
|
||||
idx = KSampler.SCHEDULERS.index("ddim_uniform") # Last item in the samplers list
|
||||
KSampler.SCHEDULERS.insert(idx+1, scheduler) # Add our custom samplers
|
||||
setattr(k_diffusion_sampling, "get_sigmas_{}".format(scheduler), extra_schedulers[scheduler])
|
||||
import importlib
|
||||
importlib.reload(k_diffusion_sampling)
|
||||
except ValueError as err:
|
||||
pass
|
||||
|
||||
# Noise samplers
|
||||
from torch import Generator, Tensor, lerp
|
||||
from torch.nn.functional import unfold
|
||||
from typing import Callable, Tuple
|
||||
from math import pi
|
||||
|
||||
def get_positions(block_shape: Tuple[int, int]) -> Tensor:
|
||||
"""
|
||||
Generate position tensor.
|
||||
|
||||
Arguments:
|
||||
block_shape -- (height, width) of position tensor
|
||||
|
||||
Returns:
|
||||
position vector shaped (1, height, width, 1, 1, 2)
|
||||
"""
|
||||
bh, bw = block_shape
|
||||
positions = torch.stack(
|
||||
torch.meshgrid(
|
||||
[(torch.arange(b) + 0.5) / b for b in (bw, bh)],
|
||||
indexing="xy",
|
||||
),
|
||||
-1,
|
||||
).view(1, bh, bw, 1, 1, 2)
|
||||
return positions
|
||||
|
||||
|
||||
def unfold_grid(vectors: Tensor) -> Tensor:
|
||||
"""
|
||||
Unfold vector grid to batched vectors.
|
||||
|
||||
Arguments:
|
||||
vectors -- grid vectors
|
||||
|
||||
Returns:
|
||||
batched grid vectors
|
||||
"""
|
||||
batch_size, _, gpy, gpx = vectors.shape
|
||||
return (
|
||||
unfold(vectors, (2, 2))
|
||||
.view(batch_size, 2, 4, -1)
|
||||
.permute(0, 2, 3, 1)
|
||||
.view(batch_size, 4, gpy - 1, gpx - 1, 2)
|
||||
)
|
||||
|
||||
|
||||
def smooth_step(t: Tensor) -> Tensor:
|
||||
"""
|
||||
Smooth step function [0, 1] -> [0, 1].
|
||||
|
||||
Arguments:
|
||||
t -- input values (any shape)
|
||||
|
||||
Returns:
|
||||
output values (same shape as input values)
|
||||
"""
|
||||
return t * t * (3.0 - 2.0 * t)
|
||||
|
||||
|
||||
def perlin_noise_tensor(
|
||||
vectors: Tensor, positions: Tensor, step: Callable = None
|
||||
) -> Tensor:
|
||||
"""
|
||||
Generate perlin noise from batched vectors and positions.
|
||||
|
||||
Arguments:
|
||||
vectors -- batched grid vectors shaped (batch_size, 4, grid_height, grid_width, 2)
|
||||
positions -- batched grid positions shaped (batch_size or 1, block_height, block_width, grid_height or 1, grid_width or 1, 2)
|
||||
|
||||
Keyword Arguments:
|
||||
step -- smooth step function [0, 1] -> [0, 1] (default: `smooth_step`)
|
||||
|
||||
Raises:
|
||||
Exception: if position and vector shapes do not match
|
||||
|
||||
Returns:
|
||||
(batch_size, block_height * grid_height, block_width * grid_width)
|
||||
"""
|
||||
if step is None:
|
||||
step = smooth_step
|
||||
|
||||
batch_size = vectors.shape[0]
|
||||
# grid height, grid width
|
||||
gh, gw = vectors.shape[2:4]
|
||||
# block height, block width
|
||||
bh, bw = positions.shape[1:3]
|
||||
|
||||
for i in range(2):
|
||||
if positions.shape[i + 3] not in (1, vectors.shape[i + 2]):
|
||||
raise Exception(
|
||||
f"Blocks shapes do not match: vectors ({vectors.shape[1]}, {vectors.shape[2]}), positions {gh}, {gw})"
|
||||
)
|
||||
|
||||
if positions.shape[0] not in (1, batch_size):
|
||||
raise Exception(
|
||||
f"Batch sizes do not match: vectors ({vectors.shape[0]}), positions ({positions.shape[0]})"
|
||||
)
|
||||
|
||||
vectors = vectors.view(batch_size, 4, 1, gh * gw, 2)
|
||||
positions = positions.view(positions.shape[0], bh * bw, -1, 2)
|
||||
|
||||
step_x = step(positions[..., 0])
|
||||
step_y = step(positions[..., 1])
|
||||
|
||||
row0 = lerp(
|
||||
(vectors[:, 0] * positions).sum(dim=-1),
|
||||
(vectors[:, 1] * (positions - positions.new_tensor((1, 0)))).sum(dim=-1),
|
||||
step_x,
|
||||
)
|
||||
row1 = lerp(
|
||||
(vectors[:, 2] * (positions - positions.new_tensor((0, 1)))).sum(dim=-1),
|
||||
(vectors[:, 3] * (positions - positions.new_tensor((1, 1)))).sum(dim=-1),
|
||||
step_x,
|
||||
)
|
||||
noise = lerp(row0, row1, step_y)
|
||||
return (
|
||||
noise.view(
|
||||
batch_size,
|
||||
bh,
|
||||
bw,
|
||||
gh,
|
||||
gw,
|
||||
)
|
||||
.permute(0, 3, 1, 4, 2)
|
||||
.reshape(batch_size, gh * bh, gw * bw)
|
||||
)
|
||||
|
||||
|
||||
def perlin_noise(
|
||||
grid_shape: Tuple[int, int],
|
||||
out_shape: Tuple[int, int],
|
||||
batch_size: int = 1,
|
||||
generator: Generator = None,
|
||||
*args,
|
||||
**kwargs,
|
||||
) -> Tensor:
|
||||
"""
|
||||
Generate perlin noise with given shape. `*args` and `**kwargs` are forwarded to `Tensor` creation.
|
||||
|
||||
Arguments:
|
||||
grid_shape -- Shape of grid (height, width).
|
||||
out_shape -- Shape of output noise image (height, width).
|
||||
|
||||
Keyword Arguments:
|
||||
batch_size -- (default: {1})
|
||||
generator -- random generator used for grid vectors (default: {None})
|
||||
|
||||
Raises:
|
||||
Exception: if grid and out shapes do not match
|
||||
|
||||
Returns:
|
||||
Noise image shaped (batch_size, height, width)
|
||||
"""
|
||||
# grid height and width
|
||||
gh, gw = grid_shape
|
||||
# output height and width
|
||||
oh, ow = out_shape
|
||||
# block height and width
|
||||
bh, bw = oh // gh, ow // gw
|
||||
|
||||
if oh != bh * gh:
|
||||
raise Exception(f"Output height {oh} must be divisible by grid height {gh}")
|
||||
if ow != bw * gw != 0:
|
||||
raise Exception(f"Output width {ow} must be divisible by grid width {gw}")
|
||||
|
||||
angle = torch.empty(
|
||||
[batch_size] + [s + 1 for s in grid_shape], *args, **kwargs
|
||||
).uniform_(to=2.0 * pi, generator=generator)
|
||||
# random vectors on grid points
|
||||
vectors = unfold_grid(torch.stack((torch.cos(angle), torch.sin(angle)), dim=1))
|
||||
# positions inside grid cells [0, 1)
|
||||
positions = get_positions((bh, bw)).to(vectors)
|
||||
return perlin_noise_tensor(vectors, positions).squeeze(0)
|
||||
|
||||
def rand_perlin_like(x):
|
||||
noise = torch.randn_like(x) / 2.0
|
||||
noise_size_H = noise.size(dim=2)
|
||||
noise_size_W = noise.size(dim=3)
|
||||
perlin = None
|
||||
for i in range(2):
|
||||
noise += perlin_noise((noise_size_H, noise_size_W), (noise_size_H, noise_size_W), batch_size=4).to(x.device)
|
||||
#noise += perlin
|
||||
#print(noise)
|
||||
return noise / noise.std()
|
||||
|
||||
def uniform_noise_sampler(x): # Even distribution, seemingly produces more information in non-subject areas than the normal (gaussian) noise sampler
|
||||
return lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
|
||||
|
||||
from torch.distributions import StudentT
|
||||
def studentt_noise_sampler(x): # Produces more subject-focused outputs due to distribution, unsure if this works
|
||||
noise = StudentT(loc=0, scale=0.2, df=1).rsample(x.size())
|
||||
#noise *= 2 / (torch.max(torch.abs(noise)) + 1e-8)
|
||||
s: FloatTensor = torch.quantile(
|
||||
noise.flatten(start_dim=1).abs(),
|
||||
0.75,
|
||||
dim = -1
|
||||
)
|
||||
#s.clamp_(min = 1.)
|
||||
s = s.reshape(*s.shape, 1, 1, 1)
|
||||
noise = noise.clamp(-s, s)
|
||||
noise = torch.copysign(torch.pow(torch.abs(noise), 0.5), noise)
|
||||
print(s)
|
||||
return lambda sigma, sigma_next: noise.to(x.device) / (7/3)
|
||||
|
||||
def highres_pyramid_noise_like(x, discount=0.7):
|
||||
b, c, h, w = x.shape # EDIT: w and h get over-written, rename for a different variant!
|
||||
orig_h = h
|
||||
orig_w = w
|
||||
u = torch.nn.Upsample(size=(orig_h, orig_w), mode='bilinear')
|
||||
noise = (torch.rand_like(x) - 0.5) * 2 * 1.73 # Start with scaled uniform noise
|
||||
for i in range(4):
|
||||
r = random.random()*2+2 # Rather than always going 2x,
|
||||
h, w = min(orig_h*15, int(h*(r**i))), min(orig_w*15, int(w*(r**i)))
|
||||
noise += u(torch.randn(b, c, h, w).to(x)) * discount**i
|
||||
if h>=orig_h*15 or w>=orig_w*15: break # Lowest resolution is 1x1
|
||||
return noise/noise.std() # Scaled back to roughly unit variance
|
||||
|
||||
def green_noise_like(x):
|
||||
noise = torch.randn_like(x)
|
||||
width = noise.size(dim=2)
|
||||
height = noise.size(dim=3)
|
||||
scale = 1.0 / (width * height)
|
||||
fy = torch.fft.fftfreq(width, device=x.device)[:, None] ** 2
|
||||
fx = torch.fft.fftfreq(height, device=x.device) ** 2
|
||||
f = fy + fx
|
||||
power = torch.sqrt(f)
|
||||
power[0, 0] = 1
|
||||
noise = torch.fft.ifft2(torch.fft.fft2(noise) / torch.sqrt(power))
|
||||
noise *= scale / noise.std()
|
||||
noise = torch.real(noise).to(x.device)
|
||||
return noise / noise.std()
|
||||
|
||||
def green_noise_sampler(x): # This doesn't work properly right now
|
||||
width = x.size(dim=2)
|
||||
height = x.size(dim=3)
|
||||
noise = torch.randn(width, height)
|
||||
#scale = 1.0 / (width * height)
|
||||
fy = torch.fft.fftfreq(width)[:, None] ** 2
|
||||
fx = torch.fft.fftfreq(height) ** 2
|
||||
f = fy + fx
|
||||
power = torch.sqrt(f)
|
||||
power[0, 0] = 1
|
||||
noise = torch.fft.ifft2(torch.fft.fft2(noise) / torch.sqrt(power))
|
||||
#noise *= scale / noise.std()
|
||||
noise = torch.real(noise).to(x.device)
|
||||
mean = torch.mean(noise)
|
||||
std = torch.std(noise)
|
||||
|
||||
noise.sub_(mean).div_(std)
|
||||
print(noise)
|
||||
return lambda sigma, sigma_next: noise
|
||||
|
||||
def power_noise_sampler(tensor, alpha=2, k=1): # This doesn't work properly right now
|
||||
"""Generate 1/f noise for a given tensor.
|
||||
|
||||
Args:
|
||||
tensor: The tensor to add noise to.
|
||||
alpha: The parameter that determines the slope of the spectrum.
|
||||
k: A constant.
|
||||
|
||||
Returns:
|
||||
A tensor with the same shape as `tensor` containing 1/f noise.
|
||||
"""
|
||||
tensor = torch.randn_like(tensor)
|
||||
fft = torch.fft.fft2(tensor)
|
||||
freq = torch.arange(1, len(fft) + 1, dtype=torch.float)
|
||||
spectral_density = k / freq**alpha
|
||||
noise = torch.rand(tensor.shape) * spectral_density
|
||||
mean = torch.mean(noise, dim=(-2, -1), keepdim=True).to(tensor.device)
|
||||
std = torch.std(noise, dim=(-2, -1), keepdim=True).to(tensor.device)
|
||||
noise = noise.to(tensor.device).sub_(mean).div_(std)
|
||||
variance = torch.var(noise, dim=(-2, -1), keepdim=True)
|
||||
print(variance)
|
||||
return lambda sigma, sigma_next: noise / 3
|
||||
|
||||
def mixed_noise_sampler(x): # Meant for RES
|
||||
gaussian = torch.randn_like(x)
|
||||
uniform = ((torch.rand_like(x) - 0.5) * 2 * 1.73)
|
||||
|
||||
# Calculate variances
|
||||
#gaussian_variance = torch.var(gaussian, dim=(-2, -1), keepdim=True)
|
||||
#uniform_variance = torch.var(uniform, dim=(-2, -1), keepdim=True)
|
||||
|
||||
# Determine weights based on variances
|
||||
#total_variance = gaussian_variance + uniform_variance
|
||||
#gaussian_weight = gaussian_variance / total_variance
|
||||
#uniform_weight = uniform_variance / total_variance
|
||||
|
||||
#mixed_noise = gaussian * gaussian_weight + uniform * uniform_weight
|
||||
#print(gaussian_weight, uniform_weight)
|
||||
mixed_noise = (gaussian * uniform) / 4
|
||||
|
||||
# Return the final mixed noise sample
|
||||
return lambda sigma, sigma_next: mixed_noise
|
||||
|
||||
import random
|
||||
def pyramid_noise_like(x, discount=0.75):
|
||||
b, c, w, h = x.shape # EDIT: w and h get over-written, rename for a different variant!
|
||||
#noise_vector_magnitude = (torch.linalg.vector_norm(torch.randn_like(x), dim=(1)) + 0.0000000001)[:,None]
|
||||
gauss_noise = torch.randn_like(x)
|
||||
gn_mean = gauss_noise.mean() * 0
|
||||
gn_std = gauss_noise.std() / 2
|
||||
|
||||
noise = torch.nn.functional.interpolate((torch.normal(mean=0, std=0.25, size=(b, c, w // 2, h // 2)).to(x)), size=(w, h), mode='bicubic')
|
||||
noise_2 = torch.nn.functional.interpolate((torch.normal(mean=0, std=0.5, size=(b, c, w * 8, h * 8)).to(x)), size=(w, h), mode='bicubic')
|
||||
noise = noise + noise_2
|
||||
#gauss_mag = (torch.linalg.vector_norm(gauss_noise, dim=(1)) + 0.0000000001)[:,None]
|
||||
#noise_mag = (torch.linalg.vector_norm(noise, dim=(1)) + 0.0000000001)[:,None]
|
||||
#noise /= noise_mag
|
||||
#noise *= gauss_mag
|
||||
#noise_scaled = torch.copysign(torch.pow(torch.abs(noise / noise.max()), 0.95), noise) * 1.4
|
||||
#noise = torch.copysign(torch.pow(torch.abs(noise_scaled), 0.9), noise_scaled) / 1.15
|
||||
|
||||
return lambda sigma, sigma_next: noise# / 2# / 1.5#torch.copysign(torch.pow(torch.abs(noise), 1.25), noise) * 2 * 1.73#.sub_(noise.mean()).div_(noise.std()) # Scaled back to roughly unit variance
|
||||
|
||||
# Below this point are extra samplers
|
||||
@torch.no_grad()
|
||||
def sample_clyb_4m_sde_momentumized(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1.0, s_noise=1., noise_sampler=None, momentum=0.5):
|
||||
"""DPM-Solver++(3M) SDE, modified with an extra SDE, and momentumized in both the SDE and ODE(?). 'its a first' - Clybius 2023
|
||||
The expression for d1 is derived from the extrapolation formula given in the paper “Diffusion Monte Carlo with stochastic Hamiltonians” by M. Foulkes, L. Mitas, R. Needs, and G. Rajagopal. The formula is given as follows:
|
||||
d1 = d1_0 + (d1_0 - d1_1) * r2 / (r2 + r1) + ((d1_0 - d1_1) * r2 / (r2 + r1) - (d1_1 - d1_2) * r1 / (r0 + r1)) * r2 / ((r2 + r1) * (r0 + r1))
|
||||
(if this is an incorrect citing, we blame Google's Bard and OpenAI's ChatGPT for this and NOT me :^) )
|
||||
|
||||
where d1_0, d1_1, and d1_2 are defined as follows:
|
||||
d1_0 = (denoised - denoised_1) / r2
|
||||
d1_1 = (denoised_1 - denoised_2) / r1
|
||||
d1_2 = (denoised_2 - denoised_3) / r0
|
||||
|
||||
The variables r0, r1, and r2 are defined as follows:
|
||||
r0 = h_3 / h_2
|
||||
r1 = h_2 / h
|
||||
r2 = h / h_1
|
||||
"""
|
||||
|
||||
def momentum_func(diff, velocity, timescale=1.0, offset=-momentum / 2.0): # Diff is current diff, vel is previous diff
|
||||
if velocity is None:
|
||||
momentum_vel = diff
|
||||
else:
|
||||
momentum_vel = momentum * (timescale + offset) * velocity + (1 - momentum * (timescale + offset)) * diff
|
||||
return momentum_vel
|
||||
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
noise_sampler = rand_perlin_like(x) if noise_sampler is None else noise_sampler
|
||||
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
|
||||
denoised_1, denoised_2, denoised_3 = None, None, None
|
||||
h_1, h_2, h_3 = None, None, None
|
||||
vel, vel_sde = None, None
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
time = sigmas[i] / sigma_max
|
||||
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
if sigmas[i + 1] == 0:
|
||||
# Denoising step
|
||||
x = denoised
|
||||
else:
|
||||
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
|
||||
h = s - t
|
||||
h_eta = h * (eta + 1)
|
||||
x_diff = momentum_func((-h_eta).expm1().neg() * denoised, vel, time)
|
||||
vel = x_diff
|
||||
x = torch.exp(-h_eta) * x + vel
|
||||
|
||||
if h_3 is not None:
|
||||
r0 = h_3 / h_2
|
||||
r1 = h_2 / h
|
||||
r2 = h / h_1
|
||||
d1_0 = (denoised - denoised_1) / r2
|
||||
d1_1 = (denoised_1 - denoised_2) / r1
|
||||
d1_2 = (denoised_2 - denoised_3) / r0
|
||||
d1 = d1_0 + (d1_0 - d1_1) * r2 / (r2 + r1) + ((d1_0 - d1_1) * r2 / (r2 + r1) - (d1_1 - d1_2) * r1 / (r0 + r1)) * r2 / ((r2 + r1) * (r0 + r1))
|
||||
d2 = (d1_0 - d1_1) / (r2 + r1) + ((d1_0 - d1_1) * r2 / (r2 + r1) - (d1_1 - d1_2) * r1 / (r0 + r1)) / ((r2 + r1) * (r0 + r1))
|
||||
phi_3 = h_eta.neg().expm1() / h_eta + 1
|
||||
phi_4 = phi_3 / h_eta - 0.5
|
||||
sde_diff = momentum_func(phi_3 * d1 - phi_4 * d2, vel_sde, time)
|
||||
vel_sde = sde_diff
|
||||
x = x + vel_sde
|
||||
elif h_2 is not None:
|
||||
r0 = h_1 / h
|
||||
r1 = h_2 / h
|
||||
d1_0 = (denoised - denoised_1) / r0
|
||||
d1_1 = (denoised_1 - denoised_2) / r1
|
||||
d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1)
|
||||
d2 = (d1_0 - d1_1) / (r0 + r1)
|
||||
phi_2 = h_eta.neg().expm1() / h_eta + 1
|
||||
phi_3 = phi_2 / h_eta - 0.5
|
||||
sde_diff = momentum_func(phi_2 * d1 - phi_3 * d2, vel_sde, time)
|
||||
vel_sde = sde_diff
|
||||
x = x + vel_sde
|
||||
elif h_1 is not None:
|
||||
r = h_1 / h
|
||||
d = (denoised - denoised_1) / r
|
||||
phi_2 = h_eta.neg().expm1() / h_eta + 1
|
||||
sde_diff = momentum_func(phi_2 * d, vel_sde, time)
|
||||
vel_sde = sde_diff
|
||||
x = x + vel_sde
|
||||
|
||||
if eta:
|
||||
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * h * eta).expm1().neg().sqrt() * s_noise
|
||||
|
||||
denoised_1, denoised_2, denoised_3 = denoised, denoised_1, denoised_2
|
||||
h_1, h_2, h_3 = h, h_1, h_2
|
||||
return x
|
||||
|
||||
# Kat's Truncated Taylor Method sampler, by Katherine Crowson
|
||||
def sample_ttm_jvp(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
|
||||
"""Second order truncated Taylor method (torch.func.jvp() version)."""
|
||||
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
model_fn = lambda x, sigma: model(x, sigma * s_in, **extra_args)
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
denoised = model_fn(x, sigmas[i])
|
||||
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
|
||||
if sigmas[i + 1] == 0:
|
||||
# Denoising step
|
||||
x = denoised
|
||||
else:
|
||||
# 2nd order truncated Taylor method
|
||||
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
|
||||
h = s - t
|
||||
h_eta = h * (eta + 1)
|
||||
|
||||
eps = to_d(x, sigmas[i], denoised)
|
||||
_, denoised_prime = torch.func.jvp(model_fn, (x, sigmas[i]), (eps * -sigmas[i], -sigmas[i]))
|
||||
|
||||
phi_1 = -torch.expm1(-h_eta)
|
||||
#phi_2 = torch.expm1(-h_eta) + h_eta
|
||||
phi_2 = torch.expm1(-h) + h # seems to work better with eta > 0
|
||||
x = torch.exp(-h_eta) * x + phi_1 * denoised + phi_2 * denoised_prime
|
||||
|
||||
if eta:
|
||||
phi_1_noise = torch.sqrt(-torch.expm1(-2 * h * eta))
|
||||
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * phi_1_noise * s_noise
|
||||
|
||||
return x
|
||||
|
||||
|
||||
from .other_samplers.refined_exp_solver import sample_refined_exp_s
|
||||
def sample_res_solver(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler=None, denoise_to_zero=True, simple_phi_calc=False, c2=0.5, ita=torch.Tensor((0.25,)), momentum=0.5):
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
seed = extra_args.get("seed", None)
|
||||
match noise_sampler:
|
||||
case "brownian":
|
||||
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False)
|
||||
case "gaussian":
|
||||
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
|
||||
case "uniform":
|
||||
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
|
||||
case "highres-pyramid":
|
||||
noise_sampler = lambda sigma, sigma_next: highres_pyramid_noise_like(x)
|
||||
case "perlin":
|
||||
noise_sampler = lambda sigma, sigma_next: rand_perlin_like(x)
|
||||
case _:
|
||||
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
|
||||
return sample_refined_exp_s(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, noise_sampler=noise_sampler, denoise_to_zero=denoise_to_zero, simple_phi_calc=simple_phi_calc, c2=c2, ita=ita, momentum=momentum)
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_dpmpp_dualsde_momentum(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, r=1/2, momentum=0.0):
|
||||
"""DPM-Solver++ (Stochastic with Momentum). Personal modified sampler by Clybius"""
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
noise_sampler = rand_perlin_like(x) if noise_sampler is None else noise_sampler
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
sigma_fn = lambda t: t.neg().exp()
|
||||
t_fn = lambda sigma: sigma.log().neg()
|
||||
|
||||
denoisedsde_1, denoisedsde_2, denoisedsde_3 = None, None, None # new line
|
||||
h_1, h_2, h_3 = None, None, None # new line
|
||||
|
||||
#def momentum_func(diff, velocity, timescale=1.0, offset=momentum_offset): # Diff is current diff, vel is previous diff
|
||||
# if velocity is None:
|
||||
# momentum_vel = diff
|
||||
# #print("Setting up momentum")
|
||||
# else:
|
||||
# momentum_vel = momentum * (timescale - momentum_offset) * velocity + (1 - momentum * (timescale - momentum_offset)) * diff
|
||||
# #print("Calculating momentum at", momentum)
|
||||
# return momentum_vel
|
||||
|
||||
def momentum_func(diff, velocity, timescale=1.0, offset=-momentum / 2.0): # Diff is current diff, vel is previous diff
|
||||
if velocity is None:
|
||||
momentum_vel = diff
|
||||
else:
|
||||
momentum_vel = momentum * (timescale + offset) * velocity + (1 - momentum * (timescale + offset)) * diff
|
||||
return momentum_vel
|
||||
|
||||
vel = None
|
||||
vel_2 = None
|
||||
vel_sde = None
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
time = sigmas[i] / sigma_max
|
||||
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
if sigmas[i + 1] == 0:
|
||||
# Euler method
|
||||
d = to_d(x, sigmas[i], denoised)
|
||||
dt = sigmas[i + 1] - sigmas[i]
|
||||
x = x + d * dt
|
||||
else:
|
||||
# DPM-Solver++
|
||||
t, t_next = t_fn(sigmas[i]), t_fn(sigmas[i + 1])
|
||||
h = t_next - t
|
||||
h_eta = h * (eta + 1)
|
||||
s = t + h * r
|
||||
fac = 1 / (2 * r)
|
||||
|
||||
# Step 1
|
||||
sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(s), eta)
|
||||
s_ = t_fn(sd)
|
||||
diff_2 = momentum_func((t - s_).expm1() * denoised, vel_2, time)
|
||||
vel_2 = diff_2
|
||||
x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - diff_2
|
||||
x_2 = x_2 + noise_sampler(sigma_fn(t), sigma_fn(s)) * s_noise * su
|
||||
denoised_2 = model(x_2, sigma_fn(s) * s_in, **extra_args)
|
||||
|
||||
# Step 2
|
||||
sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(t_next), eta)
|
||||
t_next_ = t_fn(sd)
|
||||
denoised_d = (1 - fac) * denoised + fac * denoised_2
|
||||
diff = momentum_func((t - t_next_).expm1() * denoised_d, vel, time)
|
||||
vel = diff
|
||||
x = (sigma_fn(t_next_) / sigma_fn(t)) * x - diff
|
||||
|
||||
if h_3 is not None:
|
||||
r0 = h_3 / h_2
|
||||
r1 = h_2 / h
|
||||
r2 = h / h_1
|
||||
d1_0 = (denoised_d - denoisedsde_1) / r2
|
||||
d1_1 = (denoisedsde_1 - denoisedsde_2) / r1
|
||||
d1_2 = (denoisedsde_2 - denoisedsde_3) / r0
|
||||
d1 = d1_0 + (d1_0 - d1_1) * r2 / (r2 + r1) + ((d1_0 - d1_1) * r2 / (r2 + r1) - (d1_1 - d1_2) * r1 / (r0 + r1)) * r2 / ((r2 + r1) * (r0 + r1))
|
||||
d2 = (d1_0 - d1_1) / (r2 + r1) + ((d1_0 - d1_1) * r2 / (r2 + r1) - (d1_1 - d1_2) * r1 / (r0 + r1)) / ((r2 + r1) * (r0 + r1))
|
||||
phi_3 = h_eta.neg().expm1() / h_eta + 1
|
||||
phi_4 = phi_3 / h_eta - 0.5
|
||||
diff = momentum_func(phi_3 * d1 - phi_4 * d2, vel_sde, time)
|
||||
vel_sde = diff
|
||||
x = x + diff
|
||||
elif h_2 is not None:
|
||||
r0 = h_1 / h
|
||||
r1 = h_2 / h
|
||||
d1_0 = (denoised_d - denoisedsde_1) / r0
|
||||
d1_1 = (denoisedsde_1 - denoisedsde_2) / r1
|
||||
d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1)
|
||||
d2 = (d1_0 - d1_1) / (r0 + r1)
|
||||
phi_2 = h_eta.neg().expm1() / h_eta + 1
|
||||
phi_3 = phi_2 / h_eta - 0.5
|
||||
diff = momentum_func(phi_2 * d1 - phi_3 * d2, vel_sde, time)
|
||||
vel_sde = diff
|
||||
x = x + diff
|
||||
elif h_1 is not None:
|
||||
r = h_1 / h
|
||||
d = (denoised_d - denoisedsde_1) / r
|
||||
phi_2 = h_eta.neg().expm1() / h_eta + 1
|
||||
diff = momentum_func(phi_2 * d, vel_sde, time)
|
||||
vel_sde = diff
|
||||
x = x + diff
|
||||
|
||||
if eta:
|
||||
x = x + noise_sampler(sigma_fn(t), sigma_fn(t_next)) * s_noise * su
|
||||
#if 'denoised_d' in locals():
|
||||
denoisedsde_1, denoisedsde_2, denoisedsde_3 = denoised_d, denoisedsde_1, denoisedsde_2 # new line
|
||||
#if 'h' in locals():
|
||||
h_1, h_2, h_3 = h, h_1, h_2
|
||||
return x
|
||||
|
||||
def sample_dpmpp_dualsdemomentum(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, r=1/2, momentum=0.0):
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
seed = extra_args.get("seed", None)
|
||||
match noise_sampler:
|
||||
case "brownian":
|
||||
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False)
|
||||
case "gaussian":
|
||||
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
|
||||
case "uniform":
|
||||
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
|
||||
case "perlin":
|
||||
noise_sampler = lambda sigma, sigma_next: rand_perlin_like(x)
|
||||
case _:
|
||||
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
|
||||
return sample_dpmpp_dualsde_momentum(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler, r=r, momentum=momentum)
|
||||
|
||||
from .other_samplers.sample_ttm import sample_ttm_jvp
|
||||
def sample_ttmcustom(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
seed = extra_args.get("seed", None)
|
||||
match noise_sampler:
|
||||
case "gaussian":
|
||||
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
|
||||
case "uniform":
|
||||
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
|
||||
case "brownian":
|
||||
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False)
|
||||
case _:
|
||||
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
|
||||
return sample_ttm_jvp(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler)
|
||||
|
||||
from k_diffusion.sampling import sample_lcm
|
||||
def sample_lcmcustom(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler=None):
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
seed = extra_args.get("seed", None)
|
||||
match noise_sampler:
|
||||
case "gaussian":
|
||||
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
|
||||
case "uniform":
|
||||
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
|
||||
case "brownian":
|
||||
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False)
|
||||
case _:
|
||||
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
|
||||
return sample_lcm(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, noise_sampler=noise_sampler)
|
||||
|
||||
def sample_clyb_4m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler="brownian", momentum=0.5):
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
seed = extra_args.get("seed", None)
|
||||
match noise_sampler:
|
||||
case "brownian":
|
||||
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False)
|
||||
case "gaussian":
|
||||
noise_sampler = lambda sigma, sigma_next: torch.randn_like(x)
|
||||
case "uniform":
|
||||
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
|
||||
case "highres-pyramid":
|
||||
noise_sampler = lambda sigma, sigma_next: highres_pyramid_noise_like(x)
|
||||
case "perlin":
|
||||
noise_sampler = lambda sigma, sigma_next: rand_perlin_like(x)
|
||||
case _:
|
||||
noise_sampler = lambda sigma, sigma_next: (torch.rand_like(x) - 0.5) * 2 * 1.73
|
||||
return sample_clyb_4m_sde_momentumized(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler, momentum=momentum)
|
||||
|
||||
# Add your personal samplers below here, just for formatting purposes ;3
|
||||
|
||||
# Add any extra samplers to the following dictionary
|
||||
extra_samplers = {
|
||||
"res_momentumized": sample_res_solver,
|
||||
"dpmpp_dualsde_momentumized": sample_dpmpp_dualsdemomentum,
|
||||
"clyb_4m_sde_momentumized": sample_clyb_4m_sde,
|
||||
"ttm": sample_ttmcustom,
|
||||
"lcm_custom_noise": sample_lcmcustom,
|
||||
}
|
||||
@@ -0,0 +1,474 @@
|
||||
from .other_samplers.refined_exp_solver import sample_refined_exp_s
|
||||
import comfy.samplers
|
||||
import comfy.sample
|
||||
from comfy.k_diffusion import sampling as k_diffusion_sampling
|
||||
import latent_preview
|
||||
import torch
|
||||
import numpy as np
|
||||
from tqdm.auto import trange
|
||||
|
||||
import random
|
||||
def pyramid_noise_like(size, dtype, layout, generator, device="cpu", discount=0.8):
|
||||
b, c, h, w = size
|
||||
orig_h = h
|
||||
orig_w = w
|
||||
noise = torch.zeros(size=size, dtype=dtype, layout=layout, device=device)
|
||||
r = 1
|
||||
for i in range(5):
|
||||
r *= 2 # Rather than always going 2x,
|
||||
#w, h = max(1, int(w/(r**i))), max(1, int(h/(r**i)))
|
||||
noise += torch.nn.functional.interpolate((torch.normal(mean=0, std=0.5 ** i, size=(b, c, h * r, w * r), dtype=dtype, layout=layout, generator=generator, device=device)), size=(orig_h, orig_w), mode='nearest-exact') * discount**i
|
||||
#if w>=orig_w*16 or h>=orig_h*16: break
|
||||
return noise
|
||||
|
||||
def power_noise_sampler(size, dtype, layout, generator, device="cpu", alpha=2, k=1): # This doesn't work properly right now
|
||||
"""Generate 1/f noise for a given tensor.
|
||||
|
||||
Args:
|
||||
tensor: The tensor to add noise to.
|
||||
alpha: The parameter that determines the slope of the spectrum.
|
||||
k: A constant.
|
||||
|
||||
Returns:
|
||||
A tensor with the same shape as `tensor` containing 1/f noise.
|
||||
"""
|
||||
tensor = torch.randn(size=size, dtype=dtype, layout=layout, generator=generator, device=device)
|
||||
fft = torch.fft.fft2(tensor)
|
||||
freq = torch.arange(1, len(fft) + 1, dtype=torch.float)
|
||||
spectral_density = k / freq**alpha
|
||||
noise = torch.rand(size=size, dtype=dtype, layout=layout, generator=generator, device=device) * spectral_density
|
||||
mean = torch.mean(noise, dim=(-2, -1), keepdim=True).to(tensor.device)
|
||||
std = torch.std(noise, dim=(-2, -1), keepdim=True).to(tensor.device)
|
||||
noise = noise.to(tensor.device).sub_(mean).div_(std)
|
||||
return noise
|
||||
|
||||
def prepare_noise(latent_image, seed, noise_type, noise_inds=None): # From `sample.py`
|
||||
"""
|
||||
creates random noise given a latent image and a seed.
|
||||
optional arg skip can be used to skip and discard x number of noise generations for a given seed
|
||||
"""
|
||||
generator = torch.manual_seed(seed)
|
||||
match noise_type:
|
||||
case "gaussian":
|
||||
noise_func = torch.randn
|
||||
case "uniform":
|
||||
def uniform_rand(*size, **kwargs):
|
||||
return (torch.rand(*size, **kwargs) - 0.5) * 2 * 1.73
|
||||
noise_func = uniform_rand
|
||||
case "pyramid":
|
||||
noise_func = pyramid_noise_like
|
||||
case "power":
|
||||
noise_func = power_noise_sampler
|
||||
case _:
|
||||
noise_func = torch.randn
|
||||
if noise_inds is None:
|
||||
return noise_func(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, generator=generator, device="cpu")
|
||||
|
||||
unique_inds, inverse = np.unique(noise_inds, return_inverse=True)
|
||||
noises = []
|
||||
for i in range(unique_inds[-1]+1):
|
||||
noise = noise_func([1] + list(latent_image.size())[1:], dtype=latent_image.dtype, layout=latent_image.layout, generator=generator, device="cpu")
|
||||
if i in unique_inds:
|
||||
noises.append(noise)
|
||||
noises = [noises[i] for i in inverse]
|
||||
noises = torch.cat(noises, axis=0)
|
||||
return noises
|
||||
|
||||
class SamplerRES_MOMENTUMIZED:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"noise_sampler_type": (["gaussian", "uniform", "brownian", "highres-pyramid", "perlin"], ),
|
||||
"momentum": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step":0.01}),
|
||||
"denoise_to_zero": ("BOOLEAN", {"default": True}),
|
||||
"simple_phi_calc": ("BOOLEAN", {"default": False}),
|
||||
"ita": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 100.0, "step":0.01, "round": False}),
|
||||
"c2": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step":0.01, "round": False}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(self, noise_sampler_type, momentum, denoise_to_zero, simple_phi_calc, ita, c2):
|
||||
sampler = comfy.samplers.ksampler("res_momentumized", {"noise_sampler": noise_sampler_type, "denoise_to_zero": denoise_to_zero, "simple_phi_calc": simple_phi_calc, "c2": c2, "ita": torch.Tensor((ita,)), "momentum": momentum})
|
||||
return (sampler, )
|
||||
|
||||
class SamplerDPMPP_DUALSDE_MOMENTUMIZED:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"noise_sampler_type": (["gaussian", "uniform", "brownian", "perlin"], ),
|
||||
"momentum": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step":0.01}),
|
||||
"eta": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
"s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
"r": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(self, noise_sampler_type, momentum, eta, s_noise, r,):
|
||||
sampler = comfy.samplers.ksampler("dpmpp_dualsde_momentumized", {"noise_sampler": noise_sampler_type, "eta": eta, "s_noise": s_noise, "r": r, "momentum": momentum})
|
||||
return (sampler, )
|
||||
|
||||
class SamplerTTM:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"noise_sampler_type": (["gaussian", "uniform", "brownian"], ),
|
||||
"eta": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
"s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(self, noise_sampler_type, eta, s_noise):
|
||||
sampler = comfy.samplers.ksampler("ttm", {"noise_sampler": noise_sampler_type, "eta": eta, "s_noise": s_noise})
|
||||
return (sampler, )
|
||||
|
||||
|
||||
class SamplerLCMCustom:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"noise_sampler_type": (["gaussian", "uniform", "brownian"], ),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(self, noise_sampler_type):
|
||||
sampler = comfy.samplers.ksampler("lcm_custom_noise", {"noise_sampler": noise_sampler_type})
|
||||
return (sampler, )
|
||||
|
||||
class SamplerCLYB_4M_SDE_MOMENTUMIZED:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"noise_sampler_type": (["gaussian", "uniform", "brownian", "highres-pyramid", "perlin"], ),
|
||||
"momentum": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step":0.01}),
|
||||
"eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
"s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(self, noise_sampler_type, eta, s_noise, momentum):
|
||||
sampler = comfy.samplers.ksampler("clyb_4m_sde_momentumized", {"noise_sampler": noise_sampler_type, "eta": eta, "s_noise": s_noise, "momentum": momentum})
|
||||
return (sampler, )
|
||||
|
||||
from comfy import model_management
|
||||
import comfy.utils
|
||||
import comfy.conds
|
||||
from comfy.sample import prepare_sampling, cleanup_additional_models, get_models_from_cond
|
||||
|
||||
def mixture_sample(model, model2, noise, positive, positive2, negative, negative2, cfg, cfg2, device, device2, sampler, sampler2, sigmas, sigmas2, model_options={}, model_options2={}, latent_image=None, denoise_mask=None, denoise_mask2=None, callback=None, callback2=None, disable_pbar=False, seed=None):
|
||||
positive = positive[:]
|
||||
negative = negative[:]
|
||||
positive2 = positive2[:]
|
||||
negative2 = negative2[:]
|
||||
|
||||
comfy.samplers.resolve_areas_and_cond_masks(positive, noise.shape[2], noise.shape[3], device)
|
||||
comfy.samplers.resolve_areas_and_cond_masks(negative, noise.shape[2], noise.shape[3], device)
|
||||
comfy.samplers.resolve_areas_and_cond_masks(positive2, noise.shape[2], noise.shape[3], device2)
|
||||
comfy.samplers.resolve_areas_and_cond_masks(negative2, noise.shape[2], noise.shape[3], device2)
|
||||
|
||||
model_wrap = comfy.samplers.wrap_model(model)
|
||||
model_wrap2 = comfy.samplers.wrap_model(model2)
|
||||
|
||||
comfy.samplers.calculate_start_end_timesteps(model, negative)
|
||||
comfy.samplers.calculate_start_end_timesteps(model, positive)
|
||||
comfy.samplers.calculate_start_end_timesteps(model2, negative2)
|
||||
comfy.samplers.calculate_start_end_timesteps(model2, positive2)
|
||||
|
||||
if latent_image is not None:
|
||||
latent_image = model.process_latent_in(latent_image)
|
||||
|
||||
if hasattr(model, 'extra_conds'):
|
||||
positive = comfy.samplers.encode_model_conds(model.extra_conds, positive, noise, device, "positive", latent_image=latent_image, denoise_mask=denoise_mask, seed=seed)
|
||||
negative = comfy.samplers.encode_model_conds(model.extra_conds, negative, noise, device, "negative", latent_image=latent_image, denoise_mask=denoise_mask, seed=seed)
|
||||
if hasattr(model2, 'extra_conds'):
|
||||
positive = comfy.samplers.encode_model_conds(model2.extra_conds, positive2, noise, device2, "positive", latent_image=latent_image, denoise_mask=denoise_mask2, seed=seed)
|
||||
negative = comfy.samplers.encode_model_conds(model2.extra_conds, negative2, noise, device2, "negative", latent_image=latent_image, denoise_mask=denoise_mask2, seed=seed)
|
||||
|
||||
#make sure each cond area has an opposite one with the same area
|
||||
for c in positive:
|
||||
comfy.samplers.create_cond_with_same_area_if_none(negative, c)
|
||||
for c in negative:
|
||||
comfy.samplers.create_cond_with_same_area_if_none(positive, c)
|
||||
for c in positive2:
|
||||
comfy.samplers.create_cond_with_same_area_if_none(negative2, c)
|
||||
for c in negative2:
|
||||
comfy.samplers.create_cond_with_same_area_if_none(positive2, c)
|
||||
|
||||
comfy.samplers.pre_run_control(model, negative + positive)
|
||||
comfy.samplers.pre_run_control(model2, negative2 + positive2)
|
||||
|
||||
comfy.samplers.apply_empty_x_to_equal_area(list(filter(lambda c: c.get('control_apply_to_uncond', False) == True, positive)), negative, 'control', lambda cond_cnets, x: cond_cnets[x])
|
||||
comfy.samplers.apply_empty_x_to_equal_area(positive, negative, 'gligen', lambda cond_cnets, x: cond_cnets[x])
|
||||
|
||||
comfy.samplers.apply_empty_x_to_equal_area(list(filter(lambda c: c.get('control_apply_to_uncond', False) == True, positive2)), negative2, 'control', lambda cond_cnets, x: cond_cnets[x])
|
||||
comfy.samplers.apply_empty_x_to_equal_area(positive2, negative2, 'gligen', lambda cond_cnets, x: cond_cnets[x])
|
||||
|
||||
extra_args = {"cond":positive, "uncond":negative, "cond_scale": cfg, "model_options": model_options, "seed":seed}
|
||||
extra_args2 = {"cond":positive2, "uncond":negative2, "cond_scale": cfg, "model_options": model_options2, "seed":seed}
|
||||
samples = None
|
||||
temp_sigmas = sigmas
|
||||
temp_sigmas2 = sigmas2
|
||||
#samples = sampler.sample(model_wrap, sigmas, extra_args, callback, noise, latent_image, denoise_mask, True)
|
||||
for i in trange(len(sigmas) - 1, disable=disable_pbar):
|
||||
last_step = i + 1
|
||||
start_step = i
|
||||
if last_step is not None and last_step < (len(sigmas) - 1):
|
||||
temp_sigmas = sigmas[:last_step + 1]
|
||||
temp_sigmas2 = sigmas2[:last_step + 1]
|
||||
if start_step is not None:
|
||||
if start_step < (len(sigmas) - 1):
|
||||
temp_sigmas = temp_sigmas[start_step:]
|
||||
temp_sigmas2 = temp_sigmas2[start_step:]
|
||||
else:
|
||||
if latent_image is not None:
|
||||
return latent_image
|
||||
else:
|
||||
return torch.zeros_like(noise)
|
||||
if len(temp_sigmas) != 2:
|
||||
temp_sigmas = sigmas[-2:]
|
||||
temp_sigmas2 = sigmas2[-2:]
|
||||
if (i % 2) == 0:
|
||||
#print(temp_sigmas)
|
||||
samples = sampler.sample(model_wrap, temp_sigmas, extra_args, callback, noise.to(device) if i is 0 else torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device=device), samples if samples is not None else latent_image, denoise_mask, True)
|
||||
else:
|
||||
#print(temp_sigmas)
|
||||
samples = sampler2.sample(model_wrap2, temp_sigmas2, extra_args2, callback2, noise.to(device2) if i is 0 else torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device=device2), samples if samples is not None else latent_image, denoise_mask2, True)
|
||||
return model.process_latent_out(samples.to(torch.float32))
|
||||
|
||||
def sample_mixture(model, model2, noise, cfg, cfg2, sampler, sampler2, sigmas, sigmas2, positive, negative, latent_image, noise_mask=None, callback=None, callback2=None, disable_pbar=False, seed=None):
|
||||
real_model, positive_copy, negative_copy, noise_mask, models = prepare_sampling(model, noise.shape, positive, negative, noise_mask)
|
||||
real_model2, positive_copy2, negative_copy2, noise_mask2, models2 = prepare_sampling(model2, noise.shape, positive, negative, noise_mask)
|
||||
noise = noise.to(model.load_device)
|
||||
latent_image = latent_image.to(model.load_device)
|
||||
sigmas = sigmas.to(model.load_device)
|
||||
sigmas2 = sigmas2.to(model.load_device)
|
||||
|
||||
samples = mixture_sample(real_model, real_model2, noise, positive_copy, positive_copy2, negative_copy, negative_copy2, cfg, cfg2, model.load_device, model2.load_device, sampler, sampler2, sigmas, sigmas2, model_options=model.model_options, model_options2=model2.model_options, latent_image=latent_image, denoise_mask=noise_mask, denoise_mask2=noise_mask2, callback=callback, callback2=callback2, disable_pbar=disable_pbar, seed=seed)
|
||||
|
||||
samples = samples.to(comfy.model_management.intermediate_device())
|
||||
cleanup_additional_models(models)
|
||||
cleanup_additional_models(models2)
|
||||
cleanup_additional_models(set(get_models_from_cond(positive_copy, "control") + get_models_from_cond(negative_copy, "control")))
|
||||
cleanup_additional_models(set(get_models_from_cond(positive_copy2, "control") + get_models_from_cond(negative_copy2, "control")))
|
||||
return samples
|
||||
|
||||
class SamplerCustomNoise:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"model": ("MODEL",),
|
||||
"add_noise": ("BOOLEAN", {"default": True}),
|
||||
"noise_is_latent": ("BOOLEAN", {"default": False}),
|
||||
"noise_type": (["gaussian", "uniform", "pyramid", "power"], ),
|
||||
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.5, "round": 0.01}),
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"sampler": ("SAMPLER", ),
|
||||
"sigmas": ("SIGMAS", ),
|
||||
"latent_image": ("LATENT", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT","LATENT")
|
||||
RETURN_NAMES = ("output", "denoised_output")
|
||||
|
||||
FUNCTION = "sample"
|
||||
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
def sample(self, model, add_noise, noise_is_latent, noise_type, noise_seed, cfg, positive, negative, sampler, sigmas, latent_image):
|
||||
latent = latent_image
|
||||
latent_image = latent["samples"]
|
||||
if not add_noise:
|
||||
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
|
||||
else:
|
||||
batch_inds = latent["batch_index"] if "batch_index" in latent else None
|
||||
noise = prepare_noise(latent_image, noise_seed, noise_type, batch_inds)
|
||||
|
||||
if noise_is_latent:
|
||||
noise += latent_image.cpu()# * noise.std()
|
||||
noise.sub_(noise.mean()).div_(noise.std())
|
||||
|
||||
noise_mask = None
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent["noise_mask"]
|
||||
|
||||
x0_output = {}
|
||||
|
||||
callback = latent_preview.prepare_callback(model, sigmas.shape[-1] - 1, x0_output)
|
||||
|
||||
disable_pbar = False
|
||||
samples = comfy.sample.sample_custom(model, noise, cfg, sampler, sigmas, positive, negative, latent_image, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=noise_seed)
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = samples
|
||||
if "x0" in x0_output:
|
||||
out_denoised = latent.copy()
|
||||
out_denoised["samples"] = model.model.process_latent_out(x0_output["x0"].cpu())
|
||||
else:
|
||||
out_denoised = out
|
||||
return (out, out_denoised)
|
||||
|
||||
class SamplerCustomNoiseDuo:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"model": ("MODEL",),
|
||||
"add_noise": ("BOOLEAN", {"default": True}),
|
||||
"add_noise_pass2": ("BOOLEAN", {"default": True}),
|
||||
"return_noisy_pass1": ("BOOLEAN", {"default": False}),
|
||||
"noise_type": (["gaussian", "uniform", "pyramid", "power"], ),
|
||||
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
|
||||
"cfg2": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"sampler": ("SAMPLER", ),
|
||||
"sampler2": ("SAMPLER", ),
|
||||
"sigmas": ("SIGMAS", ),
|
||||
"sigmas2": ("SIGMAS", ),
|
||||
"hr_upscale": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 9.0, "step":0.1, "round": 0.01}),
|
||||
"latent_image": ("LATENT", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT","LATENT")
|
||||
RETURN_NAMES = ("output", "denoised_output")
|
||||
|
||||
FUNCTION = "sample"
|
||||
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
def sample(self, model, add_noise, add_noise_pass2, return_noisy_pass1, noise_type, noise_seed, cfg, cfg2, positive, negative, sampler, sampler2, sigmas, sigmas2, hr_upscale, latent_image):
|
||||
latent = latent_image
|
||||
latent_image = latent["samples"]
|
||||
if not add_noise:
|
||||
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
|
||||
else:
|
||||
batch_inds = latent["batch_index"] if "batch_index" in latent else None
|
||||
noise = prepare_noise(latent_image, noise_seed, noise_type, batch_inds)
|
||||
|
||||
noise_mask = None
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent["noise_mask"]
|
||||
|
||||
x0_output = {}
|
||||
|
||||
callback = latent_preview.prepare_callback(model, sigmas.shape[-1] - 1, x0_output)
|
||||
|
||||
disable_pbar = False
|
||||
|
||||
samples = comfy.sample.sample_custom(model, noise, cfg, sampler, sigmas, positive, negative, latent_image, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=noise_seed)
|
||||
|
||||
if not return_noisy_pass1:
|
||||
out_denoised = latent.copy()
|
||||
samples = model.model.process_latent_out(x0_output["x0"].cpu())
|
||||
if hr_upscale > 1.0:
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = comfy.utils.common_upscale(noise_mask, (int)(noise_mask.shape[-1] * hr_upscale), (int)(noise_mask.shape[-2] * hr_upscale), "bislerp", "disabled")
|
||||
samples = comfy.utils.common_upscale(samples, (int)(samples.shape[-1] * hr_upscale), (int)(samples.shape[-2] * hr_upscale), "bislerp", "disabled")
|
||||
noise = prepare_noise(samples, noise_seed, noise_type, batch_inds)
|
||||
|
||||
samples = comfy.sample.sample_custom(model, noise if add_noise_pass2 else torch.zeros(samples.size(), dtype=samples.dtype, layout=samples.layout, device="cpu"), cfg2, sampler2, sigmas2, positive, negative, samples, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=noise_seed)
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = samples
|
||||
if "x0" in x0_output:
|
||||
out_denoised = latent.copy()
|
||||
out_denoised["samples"] = model.model.process_latent_out(x0_output["x0"].cpu())
|
||||
else:
|
||||
out_denoised = out
|
||||
return (out, out_denoised)
|
||||
|
||||
class SamplerCustomModelMixtureDuo:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"model": ("MODEL",),
|
||||
"model2": ("MODEL",),
|
||||
"add_noise": ("BOOLEAN", {"default": True}),
|
||||
"add_noise_pass2": ("BOOLEAN", {"default": True}),
|
||||
"return_noisy_pass1": ("BOOLEAN", {"default": False}),
|
||||
"noise_type": (["gaussian", "uniform", "pyramid", "power"], ),
|
||||
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
|
||||
"cfg2": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"sampler": ("SAMPLER", ),
|
||||
"sampler2": ("SAMPLER", ),
|
||||
"sigmas": ("SIGMAS", ),
|
||||
"sigmas2": ("SIGMAS", ),
|
||||
"hr_upscale": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 9.0, "step":0.1, "round": 0.01}),
|
||||
"latent_image": ("LATENT", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT","LATENT")
|
||||
RETURN_NAMES = ("output", "denoised_output")
|
||||
|
||||
FUNCTION = "sample"
|
||||
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
def sample(self, model, model2, add_noise, add_noise_pass2, return_noisy_pass1, noise_type, noise_seed, cfg, cfg2, positive, negative, sampler, sampler2, sigmas, sigmas2, hr_upscale, latent_image):
|
||||
latent = latent_image
|
||||
latent_image = latent["samples"]
|
||||
if not add_noise:
|
||||
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
|
||||
else:
|
||||
batch_inds = latent["batch_index"] if "batch_index" in latent else None
|
||||
noise = prepare_noise(latent_image, noise_seed, noise_type, batch_inds)
|
||||
|
||||
noise_mask = None
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent["noise_mask"]
|
||||
|
||||
x0_output = {}
|
||||
|
||||
callback = latent_preview.prepare_callback(model, sigmas.shape[-1] - 1, x0_output)
|
||||
callback2 = latent_preview.prepare_callback(model2, sigmas.shape[-1] - 1, x0_output)
|
||||
|
||||
disable_pbar = False
|
||||
|
||||
samples = sample_mixture(model, model2, noise, cfg, cfg2, sampler, sampler2, sigmas, sigmas2, positive, negative, latent_image, noise_mask=noise_mask, callback=callback, callback2=callback2, disable_pbar=disable_pbar, seed=noise_seed)
|
||||
|
||||
#if not return_noisy_pass1:
|
||||
# out_denoised = latent.copy()
|
||||
# samples = model.model.process_latent_out(x0_output["x0"].cpu())
|
||||
#if hr_upscale > 1.0:
|
||||
# if "noise_mask" in latent:
|
||||
# noise_mask = comfy.utils.common_upscale(noise_mask, (int)(noise_mask.shape[-1] * hr_upscale), (int)(noise_mask.shape[-2] * hr_upscale), "bislerp", "disabled")
|
||||
# samples = comfy.utils.common_upscale(samples, (int)(samples.shape[-1] * hr_upscale), (int)(samples.shape[-2] * hr_upscale), "bislerp", "disabled")
|
||||
# noise = prepare_noise(samples, noise_seed, noise_type, batch_inds)
|
||||
|
||||
#samples = sample_mixture(model, model2, noise if add_noise_pass2 else torch.zeros(samples.size(), dtype=samples.dtype, layout=samples.layout, device="cpu"), cfg2, sampler2, sigmas2, positive, negative, samples, noise_mask=noise_mask, callback=callback, callback2=callback2, disable_pbar=disable_pbar, seed=noise_seed)
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = samples
|
||||
if "x0" in x0_output:
|
||||
out_denoised = latent.copy()
|
||||
out_denoised["samples"] = model.model.process_latent_out(x0_output["x0"].cpu())
|
||||
else:
|
||||
out_denoised = out
|
||||
return (out, out_denoised)
|
||||
@@ -0,0 +1,296 @@
|
||||
import torch
|
||||
from torch import no_grad, FloatTensor
|
||||
from tqdm import tqdm
|
||||
from itertools import pairwise
|
||||
from typing import Protocol, Optional, Dict, Any, TypedDict, NamedTuple, Union, List
|
||||
import math
|
||||
|
||||
class DenoiserModel(Protocol):
|
||||
def __call__(self, x: FloatTensor, t: FloatTensor, *args, **kwargs) -> FloatTensor: ...
|
||||
|
||||
class RefinedExpCallbackPayload(TypedDict):
|
||||
x: FloatTensor
|
||||
i: int
|
||||
sigma: FloatTensor
|
||||
sigma_hat: FloatTensor
|
||||
|
||||
class RefinedExpCallback(Protocol):
|
||||
def __call__(self, payload: RefinedExpCallbackPayload) -> None: ...
|
||||
|
||||
class NoiseSampler(Protocol):
|
||||
def __call__(self, x: FloatTensor) -> FloatTensor: ...
|
||||
|
||||
class StepOutput(NamedTuple):
|
||||
x_next: FloatTensor
|
||||
denoised: FloatTensor
|
||||
denoised2: FloatTensor
|
||||
vel: FloatTensor
|
||||
vel_2: FloatTensor
|
||||
|
||||
def _gamma(
|
||||
n: int,
|
||||
) -> int:
|
||||
"""
|
||||
https://en.wikipedia.org/wiki/Gamma_function
|
||||
for every positive integer n,
|
||||
Γ(n) = (n-1)!
|
||||
"""
|
||||
return math.factorial(n-1)
|
||||
|
||||
def _incomplete_gamma(
|
||||
s: int,
|
||||
x: float,
|
||||
gamma_s: Optional[int] = None
|
||||
) -> float:
|
||||
"""
|
||||
https://en.wikipedia.org/wiki/Incomplete_gamma_function#Special_values
|
||||
if s is a positive integer,
|
||||
Γ(s, x) = (s-1)!*∑{k=0..s-1}(x^k/k!)
|
||||
"""
|
||||
if gamma_s is None:
|
||||
gamma_s = _gamma(s)
|
||||
|
||||
sum_: float = 0
|
||||
# {k=0..s-1} inclusive
|
||||
for k in range(s):
|
||||
numerator: float = x**k
|
||||
denom: int = math.factorial(k)
|
||||
quotient: float = numerator/denom
|
||||
sum_ += quotient
|
||||
incomplete_gamma_: float = sum_ * math.exp(-x) * gamma_s
|
||||
return incomplete_gamma_
|
||||
|
||||
# by Katherine Crowson
|
||||
def _phi_1(neg_h: FloatTensor):
|
||||
return torch.nan_to_num(torch.expm1(neg_h) / neg_h, nan=1.0)
|
||||
|
||||
# by Katherine Crowson
|
||||
def _phi_2(neg_h: FloatTensor):
|
||||
return torch.nan_to_num((torch.expm1(neg_h) - neg_h) / neg_h**2, nan=0.5)
|
||||
|
||||
# by Katherine Crowson
|
||||
def _phi_3(neg_h: FloatTensor):
|
||||
return torch.nan_to_num((torch.expm1(neg_h) - neg_h - neg_h**2 / 2) / neg_h**3, nan=1 / 6)
|
||||
|
||||
def _phi(
|
||||
neg_h: float,
|
||||
j: int,
|
||||
):
|
||||
"""
|
||||
For j={1,2,3}: you could alternatively use Kat's phi_1, phi_2, phi_3 which perform fewer steps
|
||||
|
||||
Lemma 1
|
||||
https://arxiv.org/abs/2308.02157
|
||||
ϕj(-h) = 1/h^j*∫{0..h}(e^(τ-h)*(τ^(j-1))/((j-1)!)dτ)
|
||||
|
||||
https://www.wolframalpha.com/input?i=integrate+e%5E%28%CF%84-h%29*%28%CF%84%5E%28j-1%29%2F%28j-1%29%21%29d%CF%84
|
||||
= 1/h^j*[(e^(-h)*(-τ)^(-j)*τ(j))/((j-1)!)]{0..h}
|
||||
https://www.wolframalpha.com/input?i=integrate+e%5E%28%CF%84-h%29*%28%CF%84%5E%28j-1%29%2F%28j-1%29%21%29d%CF%84+between+0+and+h
|
||||
= 1/h^j*((e^(-h)*(-h)^(-j)*h^j*(Γ(j)-Γ(j,-h)))/(j-1)!)
|
||||
= (e^(-h)*(-h)^(-j)*h^j*(Γ(j)-Γ(j,-h))/((j-1)!*h^j)
|
||||
= (e^(-h)*(-h)^(-j)*(Γ(j)-Γ(j,-h))/(j-1)!
|
||||
= (e^(-h)*(-h)^(-j)*(Γ(j)-Γ(j,-h))/Γ(j)
|
||||
= (e^(-h)*(-h)^(-j)*(1-Γ(j,-h)/Γ(j))
|
||||
|
||||
requires j>0
|
||||
"""
|
||||
assert j > 0
|
||||
gamma_: float = _gamma(j)
|
||||
incomp_gamma_: float = _incomplete_gamma(j, neg_h, gamma_s=gamma_)
|
||||
|
||||
phi_: float = math.exp(neg_h) * neg_h**-j * (1-incomp_gamma_/gamma_)
|
||||
|
||||
return phi_
|
||||
|
||||
class RESDECoeffsSecondOrder(NamedTuple):
|
||||
a2_1: float
|
||||
b1: float
|
||||
b2: float
|
||||
|
||||
def _de_second_order(
|
||||
h: float,
|
||||
c2: float,
|
||||
simple_phi_calc = False,
|
||||
) -> RESDECoeffsSecondOrder:
|
||||
"""
|
||||
Table 3
|
||||
https://arxiv.org/abs/2308.02157
|
||||
ϕi,j := ϕi,j(-h) = ϕi(-cj*h)
|
||||
a2_1 = c2ϕ1,2
|
||||
= c2ϕ1(-c2*h)
|
||||
b1 = ϕ1 - ϕ2/c2
|
||||
"""
|
||||
if simple_phi_calc:
|
||||
# Kat computed simpler expressions for phi for cases j={1,2,3}
|
||||
a2_1: float = c2 * _phi_1(-c2*h)
|
||||
phi1: float = _phi_1(-h)
|
||||
phi2: float = _phi_2(-h)
|
||||
else:
|
||||
# I computed general solution instead.
|
||||
# they're close, but there are slight differences. not sure which would be more prone to numerical error.
|
||||
a2_1: float = c2 * _phi(j=1, neg_h=-c2*h)
|
||||
phi1: float = _phi(j=1, neg_h=-h)
|
||||
phi2: float = _phi(j=2, neg_h=-h)
|
||||
phi2_c2: float = phi2/c2
|
||||
b1: float = phi1 - phi2_c2
|
||||
b2: float = phi2_c2
|
||||
return RESDECoeffsSecondOrder(
|
||||
a2_1=a2_1,
|
||||
b1=b1,
|
||||
b2=b2,
|
||||
)
|
||||
|
||||
def _refined_exp_sosu_step(
|
||||
model: DenoiserModel,
|
||||
x: FloatTensor,
|
||||
sigma: FloatTensor,
|
||||
sigma_next: FloatTensor,
|
||||
c2 = 0.5,
|
||||
extra_args: Dict[str, Any] = {},
|
||||
pbar: Optional[tqdm] = None,
|
||||
simple_phi_calc = False,
|
||||
momentum = 0.0,
|
||||
vel = None,
|
||||
vel_2 = None,
|
||||
time = None
|
||||
) -> StepOutput:
|
||||
"""
|
||||
Algorithm 1 "RES Second order Single Update Step with c2"
|
||||
https://arxiv.org/abs/2308.02157
|
||||
|
||||
Parameters:
|
||||
model (`DenoiserModel`): a k-diffusion wrapped denoiser model (e.g. a subclass of DiscreteEpsDDPMDenoiser)
|
||||
x (`FloatTensor`): noised latents (or RGB I suppose), e.g. torch.randn((B, C, H, W)) * sigma[0]
|
||||
sigma (`FloatTensor`): timestep to denoise
|
||||
sigma_next (`FloatTensor`): timestep+1 to denoise
|
||||
c2 (`float`, *optional*, defaults to .5): partial step size for solving ODE. .5 = midpoint method
|
||||
extra_args (`Dict[str, Any]`, *optional*, defaults to `{}`): kwargs to pass to `model#__call__()`
|
||||
pbar (`tqdm`, *optional*, defaults to `None`): progress bar to update after each model call
|
||||
simple_phi_calc (`bool`, *optional*, defaults to `True`): True = calculate phi_i,j(-h) via simplified formulae specific to j={1,2}. False = Use general solution that works for any j. Mathematically equivalent, but could be numeric differences.
|
||||
"""
|
||||
|
||||
def momentum_func(diff, velocity, timescale=1.0, offset=-momentum / 2.0): # Diff is current diff, vel is previous diff
|
||||
if velocity is None:
|
||||
momentum_vel = diff
|
||||
else:
|
||||
momentum_vel = momentum * (timescale + offset) * velocity + (1 - momentum * (timescale + offset)) * diff
|
||||
return momentum_vel
|
||||
|
||||
lam_next, lam = (s.log().neg() for s in (sigma_next, sigma))
|
||||
|
||||
# type hints aren't strictly true regarding float vs FloatTensor.
|
||||
# everything gets promoted to `FloatTensor` after interacting with `sigma: FloatTensor`.
|
||||
# I will use float to indicate any variables which are scalars.
|
||||
h: float = lam_next - lam
|
||||
a2_1, b1, b2 = _de_second_order(h=h, c2=c2, simple_phi_calc=simple_phi_calc)
|
||||
|
||||
denoised: FloatTensor = model(x, sigma, **extra_args)
|
||||
if pbar is not None:
|
||||
pbar.update(0.5)
|
||||
|
||||
c2_h: float = c2*h
|
||||
|
||||
diff_2 = momentum_func(a2_1*h*denoised, vel_2, time)
|
||||
vel_2 = diff_2
|
||||
x_2: FloatTensor = math.exp(-c2_h)*x + diff_2
|
||||
lam_2: float = lam + c2_h
|
||||
sigma_2: float = lam_2.neg().exp()
|
||||
|
||||
denoised2: FloatTensor = model(x_2, sigma_2, **extra_args)
|
||||
if pbar is not None:
|
||||
pbar.update(0.5)
|
||||
|
||||
diff = momentum_func(h*(b1*denoised + b2*denoised2), vel, time)
|
||||
vel = diff
|
||||
|
||||
x_next: FloatTensor = math.exp(-h)*x + diff
|
||||
|
||||
return StepOutput(
|
||||
x_next=x_next,
|
||||
denoised=denoised,
|
||||
denoised2=denoised2,
|
||||
vel=vel,
|
||||
vel_2=vel_2,
|
||||
)
|
||||
|
||||
|
||||
@no_grad()
|
||||
def sample_refined_exp_s(
|
||||
model: FloatTensor,
|
||||
x: FloatTensor,
|
||||
sigmas: FloatTensor,
|
||||
denoise_to_zero: bool = True,
|
||||
extra_args: Dict[str, Any] = {},
|
||||
callback: Optional[RefinedExpCallback] = None,
|
||||
disable: Optional[bool] = None,
|
||||
ita: FloatTensor = torch.zeros((1,)),
|
||||
c2 = .5,
|
||||
noise_sampler: NoiseSampler = torch.randn_like,
|
||||
simple_phi_calc = False,
|
||||
momentum = 0.0,
|
||||
):
|
||||
"""
|
||||
Refined Exponential Solver (S).
|
||||
Algorithm 2 "RES Single-Step Sampler" with Algorithm 1 second-order step
|
||||
https://arxiv.org/abs/2308.02157
|
||||
|
||||
Parameters:
|
||||
model (`DenoiserModel`): a k-diffusion wrapped denoiser model (e.g. a subclass of DiscreteEpsDDPMDenoiser)
|
||||
x (`FloatTensor`): noised latents (or RGB I suppose), e.g. torch.randn((B, C, H, W)) * sigma[0]
|
||||
sigmas (`FloatTensor`): sigmas (ideally an exponential schedule!) e.g. get_sigmas_exponential(n=25, sigma_min=model.sigma_min, sigma_max=model.sigma_max)
|
||||
denoise_to_zero (`bool`, *optional*, defaults to `True`): whether to finish with a first-order step down to 0 (rather than stopping at sigma_min). True = fully denoise image. False = match Algorithm 2 in paper
|
||||
extra_args (`Dict[str, Any]`, *optional*, defaults to `{}`): kwargs to pass to `model#__call__()`
|
||||
callback (`RefinedExpCallback`, *optional*, defaults to `None`): you can supply this callback to see the intermediate denoising results, e.g. to preview each step of the denoising process
|
||||
disable (`bool`, *optional*, defaults to `False`): whether to hide `tqdm`'s progress bar animation from being printed
|
||||
ita (`FloatTensor`, *optional*, defaults to 0.): degree of stochasticity, η, for each timestep. tensor shape must be broadcastable to 1-dimensional tensor with length `len(sigmas) if denoise_to_zero else len(sigmas)-1`. each element should be from 0 to 1.
|
||||
c2 (`float`, *optional*, defaults to .5): partial step size for solving ODE. .5 = midpoint method
|
||||
noise_sampler (`NoiseSampler`, *optional*, defaults to `torch.randn_like`): method used for adding noise
|
||||
simple_phi_calc (`bool`, *optional*, defaults to `True`): True = calculate phi_i,j(-h) via simplified formulae specific to j={1,2}. False = Use general solution that works for any j. Mathematically equivalent, but could be numeric differences.
|
||||
"""
|
||||
#assert sigmas[-1] == 0
|
||||
ita = ita.to(x.device)
|
||||
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
|
||||
vel, vel_2 = None, None
|
||||
with tqdm(disable=disable, total=len(sigmas)-(1 if denoise_to_zero else 2)) as pbar:
|
||||
for i, (sigma, sigma_next) in enumerate(pairwise(sigmas[:-1].split(1))):
|
||||
time = sigmas[i] / sigma_max
|
||||
if 'sigma' not in locals():
|
||||
sigma = sigmas[i]
|
||||
eps = noise_sampler(sigma, sigma_next).float()
|
||||
sigma_hat = sigma * (1 + ita)
|
||||
x_hat = x + (sigma_hat ** 2 - sigma ** 2) ** .5 * eps
|
||||
x_next, denoised, denoised2, vel, vel_2 = _refined_exp_sosu_step(
|
||||
model,
|
||||
x_hat,
|
||||
sigma_hat,
|
||||
sigma_next,
|
||||
c2=c2,
|
||||
extra_args=extra_args,
|
||||
pbar=pbar,
|
||||
simple_phi_calc=simple_phi_calc,
|
||||
momentum = momentum,
|
||||
vel = vel,
|
||||
vel_2 = vel_2,
|
||||
time = time
|
||||
)
|
||||
if callback is not None:
|
||||
payload = RefinedExpCallbackPayload(
|
||||
x=x,
|
||||
i=i,
|
||||
sigma=sigma,
|
||||
sigma_hat=sigma_hat,
|
||||
denoised=denoised,
|
||||
denoised2=denoised2,
|
||||
)
|
||||
callback(payload)
|
||||
x = x_next
|
||||
if denoise_to_zero:
|
||||
eps = noise_sampler(sigma, sigma_next).float()
|
||||
sigma_hat = sigma * (1 + ita)
|
||||
x_hat = x + (sigma_hat ** 2 - sigma ** 2) ** .5 * eps
|
||||
x_next: FloatTensor = model(x_hat, sigma.to(x_hat.device), **extra_args)
|
||||
pbar.update()
|
||||
x = x_next
|
||||
return x
|
||||
@@ -0,0 +1,118 @@
|
||||
from k_diffusion.sampling import default_noise_sampler, to_d
|
||||
from tqdm import trange
|
||||
import torch
|
||||
from torch import enable_grad
|
||||
|
||||
# by Katherine Crowson
|
||||
@enable_grad()
|
||||
def sample_ttm(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
|
||||
import torch.autograd.forward_ad as fwAD
|
||||
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
if sigmas[i + 1] == 0:
|
||||
# Denoising step
|
||||
x = denoised
|
||||
else:
|
||||
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
|
||||
h = s - t
|
||||
h_eta = h * (eta + 1)
|
||||
|
||||
with fwAD.dual_level():
|
||||
eps = to_d(x, sigmas[i], denoised)
|
||||
dual_x = fwAD.make_dual(x, eps * -sigmas[i])
|
||||
dual_sigma = fwAD.make_dual(sigmas[i] * s_in, -sigmas[i] * s_in)
|
||||
dual_denoised = model(dual_x, dual_sigma, **extra_args)
|
||||
denoised_prime = fwAD.unpack_dual(dual_denoised).tangent
|
||||
|
||||
phi_1 = -torch.expm1(-h_eta)
|
||||
phi_2 = torch.expm1(-h_eta) + h_eta
|
||||
x = torch.exp(-h_eta) * x + phi_1 * denoised + phi_2 * denoised_prime
|
||||
|
||||
if eta:
|
||||
phi_1_noise = torch.sqrt(-torch.expm1(-2 * h * eta))
|
||||
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * phi_1_noise * s_noise
|
||||
|
||||
return x
|
||||
|
||||
# by Katherine Crowson
|
||||
def sample_ttm_jvp(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
|
||||
"""Second order truncated Taylor method (torch.func.jvp() version)."""
|
||||
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
model_fn = lambda x, sigma: model(x, sigma * s_in, **extra_args)
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
denoised = model_fn(x, sigmas[i])
|
||||
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
|
||||
if sigmas[i + 1] == 0:
|
||||
# Denoising step
|
||||
x = denoised
|
||||
else:
|
||||
# 2nd order truncated Taylor method
|
||||
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
|
||||
h = s - t
|
||||
h_eta = h * (eta + 1)
|
||||
|
||||
eps = to_d(x, sigmas[i], denoised)
|
||||
_, denoised_prime = torch.func.jvp(model_fn, (x, sigmas[i]), (eps * -sigmas[i], -sigmas[i]))
|
||||
|
||||
phi_1 = -torch.expm1(-h_eta)
|
||||
phi_2 = torch.expm1(-h_eta) + h_eta
|
||||
# phi_2 = torch.expm1(-h) + h # seems to work better with eta > 0
|
||||
x = torch.exp(-h_eta) * x + phi_1 * denoised + phi_2 * denoised_prime
|
||||
|
||||
if eta:
|
||||
phi_1_noise = torch.sqrt(-torch.expm1(-2 * h * eta))
|
||||
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * phi_1_noise * s_noise
|
||||
|
||||
return x
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_lcm_ttm_jvp(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler=None, eta=0.0):
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
model_fn = lambda x, sigma: model(x, sigma * s_in, **extra_args)
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
denoised = model_fn(x, sigmas[i])
|
||||
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
|
||||
if sigmas[i + 1] == 0:
|
||||
# Denoising step
|
||||
x = denoised
|
||||
else:
|
||||
# 2nd order truncated Taylor method
|
||||
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
|
||||
h = s - t
|
||||
h_eta = h * (eta + 1)
|
||||
|
||||
x = denoised
|
||||
x_2 = x + sigmas[i + 1] * noise_sampler(sigmas[i], sigmas[i + 1])
|
||||
|
||||
eps = to_d(x_2, sigmas[i + 1], denoised)
|
||||
_, denoised_prime = torch.func.jvp(model_fn, (x_2, sigmas[i + 1]), (eps * -sigmas[i + 1], -sigmas[i + 1]))
|
||||
|
||||
phi_1 = -torch.expm1(-h_eta)
|
||||
phi_2 = torch.expm1(-h_eta) + h_eta
|
||||
# phi_2 = torch.expm1(-h) + h # seems to work better with eta > 0
|
||||
x = torch.exp(-h_eta) * x + phi_1 * denoised + phi_2 * denoised_prime
|
||||
|
||||
if sigmas[i + 1] > 0:
|
||||
x = x + sigmas[i + 1] * noise_sampler(sigmas[i], sigmas[i + 1])
|
||||
|
||||
return x
|
||||
Reference in New Issue
Block a user