Implement SENS (DPM++2M/3M SDE Hybrid)
Implement IPNDM_VAPP (IPNDM_V with ancestral sampling and CFGPP) Implement STRIKE (A heavily modified Euler with denoised history and 'full ancestral sampling')
This commit is contained in:
@@ -15,6 +15,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ScaledCFGGuider": nodes.ScaledCFGGuider,
|
||||
"WarmupDecayCFGGuider": nodes.WarmupDecayCFGGuider,
|
||||
"MegaCFGGuider": nodes.MegaCFGGuider,
|
||||
### Noise
|
||||
"ImmiscibleNoise": nodes.ImmiscibleNoise,
|
||||
## Samplers
|
||||
"SamplerRES_Momentumized": nodes.SamplerRES_MOMENTUMIZED,
|
||||
"SamplerDPMPP_DualSDE_Momentumized": nodes.SamplerDPMPP_DUALSDE_MOMENTUMIZED,
|
||||
@@ -24,7 +26,12 @@ NODE_CLASS_MAPPINGS = {
|
||||
"SamplerEulerAncestralDancing_Experimental": nodes.SamplerEULER_ANCESTRAL_DANCING,
|
||||
"SamplerDPMPP_3M_SDE_DynETA": nodes.SamplerDPMPP_3M_SDE_DYN_ETA,
|
||||
"SamplerSupreme": nodes.SamplerSUPREME,
|
||||
"SamplerSENS": nodes.SamplerSENS,
|
||||
"SamplerIPNDM_VAPP": nodes.SamplerIPNDM_VAPP,
|
||||
"SamplerSTRIKE": nodes.SamplerSTRIKE,
|
||||
### Schedulers
|
||||
"SimpleExponentialScheduler": nodes.SimpleExponentialScheduler,
|
||||
"KLOptimalScheduler": nodes.KLOptimalScheduler,
|
||||
"SimpleKLOptimalScheduler": nodes.SimpleKLOptimalScheduler,
|
||||
}
|
||||
__all__ = ['NODE_CLASS_MAPPINGS']
|
||||
|
||||
+453
-215
@@ -8,6 +8,7 @@ from tqdm.auto import trange, tqdm
|
||||
import numpy as np
|
||||
|
||||
import comfy.sample
|
||||
import comfy.model_patcher
|
||||
|
||||
from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler, PIDStepSizeController, get_ancestral_step, to_d, default_noise_sampler, DPMSolver
|
||||
|
||||
@@ -49,13 +50,19 @@ def add_schedulers():
|
||||
|
||||
|
||||
# Noise samplers
|
||||
NOISE_SAMPLER_NAMES=("gaussian", "uniform", "brownian", "highres-pyramid", "pyramid", "perlin", "laplacian")
|
||||
IMMISCIBLE_NOISE_NAMES=("gaussian_1024", "perlin")
|
||||
NOISE_SAMPLER_NAMES=("gaussian", "uniform", "brownian", "highres-pyramid", "pyramid", "perlin", "laplacian", "immiscible_gaussian", "immiscible_gaussian_maximize", "immiscible_perlin")
|
||||
|
||||
def get_noise_sampler_names(default=None):
|
||||
if not default:
|
||||
return NOISE_SAMPLER_NAMES
|
||||
return (default,) + tuple(n for n in NOISE_SAMPLER_NAMES if n != default)
|
||||
|
||||
def get_immiscible_noise_sampler_names(default=None):
|
||||
if not default:
|
||||
return IMMISCIBLE_NOISE_NAMES
|
||||
return (default,) + tuple(n for n in IMMISCIBLE_NOISE_NAMES if n != default)
|
||||
|
||||
def mk_noise_sampler(x, fun):
|
||||
return lambda _sigma, _sigma_next: fun(x)
|
||||
|
||||
@@ -74,184 +81,145 @@ from math import pi
|
||||
def uniform_noise_like(x):
|
||||
return (torch.rand_like(x) - 0.5) * 2 * 1.73
|
||||
|
||||
def get_positions(block_shape: Tuple[int, int]) -> Tensor:
|
||||
"""
|
||||
Generate position tensor.
|
||||
from scipy.optimize import linear_sum_assignment
|
||||
def check_set_immiscible(x, noise_sampler_type, extra_args):
|
||||
if noise_sampler_type.startswith("immiscible"):
|
||||
match noise_sampler_type:
|
||||
case "immiscible_gaussian":
|
||||
immiscibility = make_immiscible("gaussian_1024", batching="channel") # FINISH THE REST
|
||||
extra_args = immiscibility.set_immiscible_extra_args(extra_args)
|
||||
noise_sampler = lambda _sigma, _sigma_next: immiscibility(x)
|
||||
return noise_sampler, extra_args
|
||||
case "immiscible_gaussian_maximize":
|
||||
immiscibility = make_immiscible("gaussian_1024", maximize=True, batching="channel") # FINISH THE REST
|
||||
extra_args = immiscibility.set_immiscible_extra_args(extra_args)
|
||||
noise_sampler = lambda _sigma, _sigma_next: immiscibility(x)
|
||||
return noise_sampler, extra_args
|
||||
case "immiscible_perlin":
|
||||
immiscibility = make_immiscible("perlin", immiscible_latents=8) # FINISH THE REST
|
||||
extra_args = immiscibility.set_immiscible_extra_args(extra_args)
|
||||
noise_sampler = lambda _sigma, _sigma_next: immiscibility(x)
|
||||
return noise_sampler, extra_args
|
||||
return None, extra_args
|
||||
|
||||
Arguments:
|
||||
block_shape -- (height, width) of position tensor
|
||||
class make_immiscible:
|
||||
def __init__(self, noise_func="gaussian_1024", immiscible_latents=1024, maximize=False, batching="batch"):
|
||||
self.noise_func = noise_func
|
||||
self.n_latents = immiscible_latents
|
||||
self.maximize = maximize
|
||||
self.updated_latent = None
|
||||
self.batching = batching
|
||||
|
||||
def __call__(self, latents):
|
||||
# "Immiscible Diffusion: Accelerating Diffusion Training with Noise Assignment" (2024) Li et al. arxiv.org/abs/2406.12303
|
||||
# Minimize latent-noise pairs over a batch
|
||||
# Code from https://github.com/kohya-ss/sd-scripts/pull/1395
|
||||
reference_latent = latents
|
||||
if self.updated_latent != None:
|
||||
reference_latent = self.updated_latent
|
||||
reference_latent = self.batch(reference_latent)
|
||||
n = self.n_latents # arg is an integer for how many noise tensors to generate
|
||||
noise = None
|
||||
match self.noise_func:
|
||||
case "gaussian_1024":
|
||||
#n = 1024
|
||||
size = [n] + list(reference_latent.shape[1:])
|
||||
noise = torch.randn(size, dtype=reference_latent.dtype, layout=reference_latent.layout, device=reference_latent.device)
|
||||
case "perlin":
|
||||
#n = n//32
|
||||
size = [n] + list(reference_latent.shape[1:])
|
||||
noise = torch.randn(size, dtype=reference_latent.dtype, layout=reference_latent.layout, device=reference_latent.device)
|
||||
for i in range(n):
|
||||
for j in range(reference_latent.size(dim=1)):
|
||||
noise_values = rand_perlin_2d_octaves((reference_latent.size(dim=-2), reference_latent.size(dim=-1)), (1,1), 1, 1).to(reference_latent.device)
|
||||
result = (1+0/10)*torch.erfinv(2 * noise_values - 1) * (2 ** 0.5)
|
||||
result = torch.where(torch.abs(result) > 5, noise[i, j, :, :], result)
|
||||
noise[i, j, :, :] = result
|
||||
latents_expanded = reference_latent.half().unsqueeze(1).expand(-1, n, *reference_latent.shape[1:])
|
||||
noise_expanded = noise.half().unsqueeze(0).expand(reference_latent.shape[0], *noise.shape)
|
||||
dist = (latents_expanded - noise_expanded)**2
|
||||
dist = dist.mean(list(range(2, dist.dim()))).cpu()
|
||||
assign_mat = linear_sum_assignment(dist, maximize=self.maximize)
|
||||
noise = noise[assign_mat[1]]
|
||||
return self.unbatch(noise, latents)
|
||||
|
||||
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 batch(self, ref):
|
||||
if self.batching == "batch":
|
||||
return ref
|
||||
rsz = ref.shape
|
||||
if len(rsz) != 4:
|
||||
raise ValueError("Reference must be four-dimensional")
|
||||
if self.batching == "channel":
|
||||
ref = ref.view(rsz[0] * rsz[1], *rsz[2:])
|
||||
return ref
|
||||
if self.batching == "row":
|
||||
ref = ref.view(rsz[0] * rsz[1] * rsz[2], rsz[3])
|
||||
return ref
|
||||
if self.batching == "column":
|
||||
ref = ref.permute(0, 1, 3, 2).reshape(rsz[0] * rsz[1] * rsz[3], rsz[2])
|
||||
return ref
|
||||
raise ValueError("Bad Immmiscible noise batching type")
|
||||
|
||||
def unbatch(self, noise, x_ref):
|
||||
xsz = x_ref.shape
|
||||
if self.batching == "column":
|
||||
return noise.view(*xsz[:2], xsz[3], xsz[2]).permute(0, 1, 3, 2)
|
||||
return noise.view(*xsz)
|
||||
|
||||
def set_immiscible_extra_args(self, extra_args):
|
||||
def immiscible_post_cfg_function(args):
|
||||
self.updated_latent = args["cond_denoised"]
|
||||
return args["denoised"]
|
||||
model_options = extra_args.get("model_options", {}).copy()
|
||||
extra_args["model_options"] = comfy.model_patcher.set_model_options_post_cfg_function(model_options, immiscible_post_cfg_function, disable_cfg1_optimization=True)
|
||||
return extra_args
|
||||
|
||||
def unfold_grid(vectors: Tensor) -> Tensor:
|
||||
"""
|
||||
Unfold vector grid to batched vectors.
|
||||
# From https://github.com/Extraltodeus/noise_latent_perlinpinpin/blob/main/latent_noisy_perlin.py
|
||||
# which was found at https://gist.github.com/vadimkantorov/ac1b097753f217c5c11bc2ff396e0a57
|
||||
# which was ported from https://github.com/pvigier/perlin-numpy/blob/master/perlin2d.py
|
||||
def rand_perlin_2d(shape, res, fade = lambda t: 6*t**5 - 15*t**4 + 10*t**3):
|
||||
delta = (res[0] / shape[0], res[1] / shape[1])
|
||||
d = (shape[0] // res[0], shape[1] // res[1])
|
||||
|
||||
grid = torch.stack(torch.meshgrid(torch.arange(0, res[0], delta[0]), torch.arange(0, res[1], delta[1])), dim = -1) % 1
|
||||
angles = 2*math.pi*torch.rand(res[0]+1, res[1]+1)
|
||||
gradients = torch.stack((torch.cos(angles), torch.sin(angles)), dim = -1)
|
||||
|
||||
tile_grads = lambda slice1, slice2: gradients[slice1[0]:slice1[1], slice2[0]:slice2[1]].repeat_interleave(d[0], 0).repeat_interleave(d[1], 1)
|
||||
dot = lambda grad, shift: (torch.stack((grid[:shape[0],:shape[1],0] + shift[0], grid[:shape[0],:shape[1], 1] + shift[1] ), dim = -1) * grad[:shape[0], :shape[1]]).sum(dim = -1)
|
||||
|
||||
n00 = dot(tile_grads([0, -1], [0, -1]), [0, 0])
|
||||
n10 = dot(tile_grads([1, None], [0, -1]), [-1, 0])
|
||||
n01 = dot(tile_grads([0, -1],[1, None]), [0, -1])
|
||||
n11 = dot(tile_grads([1, None], [1, None]), [-1,-1])
|
||||
t = fade(grid[:shape[0], :shape[1]])
|
||||
return math.sqrt(2) * torch.lerp(torch.lerp(n00, n10, t[..., 0]), torch.lerp(n01, n11, t[..., 0]), t[..., 1])
|
||||
|
||||
Arguments:
|
||||
vectors -- grid vectors
|
||||
def rand_perlin_2d_octaves(shape, res, octaves=1, persistence=0.5):
|
||||
noise = torch.zeros(shape)
|
||||
frequency = 1
|
||||
amplitude = 1
|
||||
for _ in range(octaves):
|
||||
noise += amplitude * rand_perlin_2d(shape, (frequency*res[0], frequency*res[1]))
|
||||
frequency *= 2
|
||||
amplitude *= persistence
|
||||
noise = torch.remainder(torch.abs(noise)*1000000,11)/11
|
||||
# noise = (torch.sin(torch.remainder(noise*1000000,83))+1)/2
|
||||
return noise
|
||||
|
||||
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 create_noisy_latents_perlin(x, detail_level=0):
|
||||
batch_size = x.size(dim=0)
|
||||
noise = torch.randn((batch_size, x.size(dim=1), x.size(dim=2), x.size(dim=3)), dtype=x.dtype, layout=x.layout, device=x.device)
|
||||
for i in range(batch_size):
|
||||
for j in range(x.size(dim=1)):
|
||||
noise_values = rand_perlin_2d_octaves((x.size(dim=2), x.size(dim=3)), (1,1), 1, 1).to(x.device)
|
||||
result = (1+detail_level/10)*torch.erfinv(2 * noise_values - 1) * (2 ** 0.5)
|
||||
result = torch.where(torch.abs(result) > 3, noise[i, j, :, :], result)
|
||||
noise[i, j, :, :] = result
|
||||
return noise
|
||||
|
||||
|
||||
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=x.shape[1]).to(x.device)
|
||||
#noise += perlin
|
||||
#print(noise)
|
||||
return noise / noise.std()
|
||||
def rand_perlin_like(x): # Even distribution, seemingly produces more information in non-subject areas than the normal (gaussian) noise sampler
|
||||
return create_noisy_latents_perlin(x)
|
||||
|
||||
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
|
||||
@@ -274,12 +242,8 @@ def studentt_noise_sampler(x): # Produces more subject-focused outputs due to di
|
||||
|
||||
from torch.distributions import Laplace
|
||||
def rand_laplacian_like(x):
|
||||
noise = torch.randn_like(x) / 4.0
|
||||
noise_size_H = noise.size(dim=2)
|
||||
noise_size_W = noise.size(dim=3)
|
||||
noise += Laplace(loc=0, scale=1.0).rsample(x.size()).to(noise.device)
|
||||
#noise += perlin
|
||||
#print(noise)
|
||||
noise = torch.zeros_like(x)#.div_(4.0)
|
||||
noise += Laplace(loc=0, scale=2 ** 0.5).rsample(x.size()).to(noise.device)
|
||||
return noise / noise.std()
|
||||
|
||||
def highres_pyramid_noise_like(x, discount=0.7):
|
||||
@@ -564,7 +528,10 @@ def sample_ttm_jvp(model, x, sigmas, extra_args=None, callback=None, disable=Non
|
||||
# Many thanks to Kat + Birch-San for this wonderful sampler implementation! https://github.com/Birch-san/sdxl-play/commits/res/
|
||||
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_type="gaussian", noise_sampler=None, denoise_to_zero=True, simple_phi_calc=False, c2=0.5, ita=torch.Tensor((0.25,)), momentum=0.0):
|
||||
return sample_refined_exp_s(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, noise_sampler=noise_sampler or get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), denoise_to_zero=denoise_to_zero, simple_phi_calc=simple_phi_calc, c2=c2, ita=ita, momentum=momentum)
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
|
||||
return sample_refined_exp_s(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), 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):
|
||||
@@ -667,18 +634,30 @@ def sample_dpmpp_dualsde_momentum(model, x, sigmas, extra_args=None, callback=No
|
||||
return x
|
||||
|
||||
def sample_dpmpp_dualsdemomentum(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, r=1/2, momentum=0.0):
|
||||
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 or get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), r=r, momentum=momentum)
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
|
||||
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 if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), 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_type="gaussian",noise_sampler=None):
|
||||
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 or get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args))
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
|
||||
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 if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args))
|
||||
|
||||
from comfy.k_diffusion.sampling import sample_lcm
|
||||
def sample_lcmcustom(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler_type="gaussian", noise_sampler=None):
|
||||
return sample_lcm(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, noise_sampler=noise_sampler or get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args))
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
|
||||
return sample_lcm(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args))
|
||||
|
||||
def sample_clyb_4m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="brownian", noise_sampler=None, momentum=0.0):
|
||||
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 or get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), momentum=momentum)
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
|
||||
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 if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), momentum=momentum)
|
||||
|
||||
|
||||
# This code works, but I'm currently experimenting with different methods
|
||||
@@ -729,7 +708,7 @@ def sampler_euler_ancestral_dancing(model, x, sigmas, extra_args=None, callback=
|
||||
return x
|
||||
|
||||
def sample_euler_ancestral_dancing(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, leap=2, eta_dance=1.0):
|
||||
return sampler_euler_ancestral_dancing(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler or get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), leap=leap, eta_dance=eta_dance)
|
||||
return sampler_euler_ancestral_dancing(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), leap=leap, eta_dance=eta_dance)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -792,7 +771,10 @@ def sampler_dpmpp_3m_sde_dynamic_eta(model, x, sigmas, extra_args=None, callback
|
||||
return x
|
||||
|
||||
def sample_dpmpp_3m_sde_dynamic_eta(model, x, sigmas, extra_args=None, callback=None, disable=None, eta_max=1.0, eta_min=0.0, s_noise=1., noise_sampler_type="brownian", noise_sampler=None):
|
||||
return sampler_dpmpp_3m_sde_dynamic_eta(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta_max=eta_max, eta_min=eta_min, s_noise=s_noise, noise_sampler=noise_sampler or get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args))
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
|
||||
return sampler_dpmpp_3m_sde_dynamic_eta(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta_max=eta_max, eta_min=eta_min, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args))
|
||||
|
||||
|
||||
from .other_samplers.refined_exp_solver import _de_second_order
|
||||
@@ -801,7 +783,7 @@ from .other_samplers.refined_exp_solver import _de_second_order
|
||||
SUPREME_ORDER = { "euler": 1, "dpm_1s": 1, "dpm_3s": 3, "rk4": 4, "reversible_heun_1s": 1, "rkf45": 6, "bogacki_shampine": 3, }
|
||||
|
||||
@torch.no_grad()
|
||||
def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=None, s_noise=1., noise_sampler=None, eta=1.0, step_method="euler", substep_method="euler", centralization=0.05, normalization=0.05, edge_enhancement=0.25, perphist=0.5, substeps=2, noise_modulation="intensity", modulation_strength=2.0, modulation_dims=3, reversible_eta=1.0):
|
||||
def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=None, s_noise=1., noise_sampler=None, eta=1.0, step_method="euler", substep_method="euler", warmup_method="euler", centralization=0.00, normalization=0.00, edge_enhancement=0.00, perphist=0.25, substeps=2, noise_modulation="none", modulation_strength=2., modulation_dims=3, reversible_eta=1.0, dyneta=True, reversible_dyneta=True, enable_free_reverse=True, free_reverse_eta=0.0, free_reverse_dyneta=True):
|
||||
"""
|
||||
Supreme Sampler, Euler steps. Based on no paper, purely interesting thoughts.
|
||||
|
||||
@@ -823,7 +805,9 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No
|
||||
noise_modulation: Method of changing the noise based on situations within the sampler
|
||||
modulation_strength: Strength of the modulation using a weighted sum between the modulation and noise sampler's noise.
|
||||
modulation_dims: Choose between (channel) modulation, (height, width) modulation, or (channels, height, width) modulation
|
||||
reversible_eta: Power scalar for increasing the strength of the reversible correction dynamically, along with eta and cond modification.
|
||||
reversible_eta: Ancestralness in the reversible component of reversible samplers.
|
||||
dyneta: Enable a dynamic eta based on sigma. Higher sigmas have a lower eta, while lower sigmas have a higher eta, max clamped to user-chosen eta.
|
||||
reversible_dyneta: Enable a dynamic reversible eta based on sigma. Higher sigmas have a lower eta, while lower sigmas have a higher eta, max clamped to user-chosen eta. Good for stability.
|
||||
"""
|
||||
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
@@ -857,8 +841,10 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No
|
||||
|
||||
# DynETA
|
||||
orig_eta = eta
|
||||
def dyneta_fn(original_eta, error):
|
||||
return original_eta * (1 / (1 + error))
|
||||
orig_reversible_eta = reversible_eta
|
||||
orig_free_reverse_eta = free_reverse_eta
|
||||
def dyneta_fn(original_eta, sigma, sigma_next):
|
||||
return torch.clamp(1 / (sigma**2 - sigma_next**2)**0.5, min=0.0, max=original_eta)
|
||||
|
||||
order, sub_order = SUPREME_ORDER.get(step_method, 2), SUPREME_ORDER.get(substep_method, 2)
|
||||
steps_per_sigma = order + sub_order * (substeps - 1)
|
||||
@@ -869,17 +855,17 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No
|
||||
|
||||
if edge_enhancement != 0:
|
||||
blur = (kornia.filters.joint_bilateral_blur(x, denoised, (3, 3), 0.1, (1.5, 1.5)) - x) # Blurs non-edges
|
||||
denoised += (kornia.filters.unsharp_mask(denoised, (3, 3), (1.5, 1.5)) - denoised) * (sigmas[i] - sigmas[i + 1]) * edge_enhancement / steps_per_sigma # Sharpens everything
|
||||
denoised += blur * (sigmas[i] - sigmas[i + 1]) * edge_enhancement / steps_per_sigma # Apply blur to non-edges, thus leaving edges sharpened
|
||||
denoised += (kornia.filters.unsharp_mask(denoised, (3, 3), (1.5, 1.5)) - denoised) * (sigmas[i] - sigmas[i + 1]) * edge_enhancement # Sharpens everything
|
||||
denoised += blur * (sigmas[i] - sigmas[i + 1]) * edge_enhancement # Apply blur to non-edges, thus leaving edges sharpened
|
||||
|
||||
if centralization != 0:
|
||||
denoised = centralize(denoised, centralization / steps_per_sigma, i)
|
||||
denoised = centralize(denoised, centralization, i)
|
||||
|
||||
if normalization != 0:
|
||||
denoised = normalize(denoised, normalization / steps_per_sigma, i)
|
||||
denoised = normalize(denoised, normalization, i)
|
||||
|
||||
if old_denoised != None and perphist != 0:
|
||||
denoised = perpadd(denoised, old_denoised, x, perphist / steps_per_sigma)
|
||||
denoised = perpadd(denoised, old_denoised, x, perphist)
|
||||
|
||||
return denoised
|
||||
|
||||
@@ -906,11 +892,13 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No
|
||||
sampler = step_method
|
||||
order = 1
|
||||
error = 0
|
||||
if iteration == 0 or prev_denoised == None: # Warmup with a RKF45 step, else use substep method for substeps
|
||||
if iteration == 0 or prev_denoised == None: # Warmup with the chosen warmup step, else use substep method for substeps
|
||||
if warmup_method == "none":
|
||||
return step_method
|
||||
if substep_iter > 0:
|
||||
return substep_method, 1, error
|
||||
order = 6
|
||||
return dynamic_order_samplers[order], order, error
|
||||
order = 2 # Chosen for simplicity
|
||||
return warmup_method, order, error
|
||||
|
||||
d = to_d(prev_x, sigmas[iteration - 1], prev_denoised)
|
||||
x_pred = prev_x + d * (sigmas[iteration] - sigmas[iteration - 1])
|
||||
@@ -1077,12 +1065,17 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No
|
||||
|
||||
dpm_solver = DPMSolver(model, extra_args)
|
||||
|
||||
# DynETA
|
||||
if dyneta: eta = dyneta_fn(orig_eta, sigmas[i], sigmas[i + 1])
|
||||
if reversible_dyneta: reversible_eta = dyneta_fn(orig_reversible_eta, sigmas[i], sigmas[i + 1])
|
||||
|
||||
# Renoising iterations
|
||||
z_avg = torch.zeros_like(x)
|
||||
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
|
||||
sigma_down_reversible, _ = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=reversible_eta)
|
||||
for k in range(substeps):
|
||||
z_k = x
|
||||
orig_zk = z_k
|
||||
eps_cache = {}
|
||||
|
||||
denoised = model(z_k, sigmas[i] * s_in, **extra_args)
|
||||
@@ -1093,9 +1086,6 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No
|
||||
|
||||
step_method_dyn, order, error = dynamic_step_method(step_method, model, prev_x, denoised, prev_denoised, i, k) #step_method, model, prev_x, denoised, prev_denoised, i, k
|
||||
|
||||
# DynETA
|
||||
#eta = dyneta_fn(orig_eta, error)
|
||||
|
||||
match step_method_dyn if sigmas[i + 1] != 0 else "euler":
|
||||
case "euler": # 1 model call
|
||||
d = to_d(z_k, sigmas[i], denoised)
|
||||
@@ -1153,10 +1143,10 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No
|
||||
dt_reversible = sigma_down_reversible - sigma_i
|
||||
|
||||
# Calculate the derivative using the model
|
||||
d_i_old = to_d(prev_x, sigma_i, prev_denoised) if prev_denoised is not None else to_d(prev_x, sigma_i, model(prev_x, sigma_i * s_in, **extra_args))
|
||||
d_i_old = to_d(z_k, sigma_i, prev_denoised) if prev_denoised is not None else to_d(z_k, sigma_i, model(z_k, sigma_i * s_in, **extra_args))
|
||||
|
||||
# Predict the sample at the next sigma using Euler step
|
||||
x_pred = prev_x + d_i_old * dt
|
||||
x_pred = z_k + d_i_old * dt
|
||||
|
||||
# Calculate the derivative at the next sigma
|
||||
d_i_plus_1 = to_d(x_pred, sigma_i_plus_1, denoised)
|
||||
@@ -1247,21 +1237,37 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No
|
||||
else:
|
||||
z_k = denoised
|
||||
case "RES":
|
||||
lam_next = sigma_down.log().neg() if eta != 0 else sigmas[i + 1].log().neg()
|
||||
lam = sigmas[i].log().neg()
|
||||
if sigmas[i + 1] > 0:
|
||||
lam_next = sigma_down.log().neg() if eta != 0 else sigmas[i + 1].log().neg()
|
||||
lam = sigmas[i].log().neg()
|
||||
|
||||
h = lam_next - lam
|
||||
a2_1, b1, b2 = _de_second_order(h=h, c2=0.5, simple_phi_calc=False)
|
||||
h = lam_next - lam
|
||||
a2_1, b1, b2 = _de_second_order(h=h, c2=0.5, simple_phi_calc=False)
|
||||
|
||||
c2_h = 0.5*h
|
||||
c2_h = 0.5*h
|
||||
|
||||
x_2 = math.exp(-c2_h)*z_k + a2_1*h*denoised
|
||||
lam_2 = lam + c2_h
|
||||
sigma_2 = lam_2.neg().exp()
|
||||
x_2 = math.exp(-c2_h)*z_k + a2_1*h*denoised
|
||||
lam_2 = lam + c2_h
|
||||
sigma_2 = lam_2.neg().exp()
|
||||
|
||||
denoised2 = model(x_2, sigma_2 * s_in, **extra_args)
|
||||
denoised2 = model(x_2, sigma_2 * s_in, **extra_args)
|
||||
|
||||
z_k = math.exp(-h)*z_k + h*(b1*denoised + b2*denoised2)
|
||||
z_k = math.exp(-h)*z_k + h*(b1*denoised + b2*denoised2)
|
||||
else:
|
||||
z_k = denoised
|
||||
|
||||
# Free Reverse
|
||||
if enable_free_reverse:
|
||||
if free_reverse_dyneta: free_reverse_eta = dyneta_fn(orig_free_reverse_eta, sigmas[i], sigmas[i + 1])
|
||||
sigma_down_freereversible, _ = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=free_reverse_eta)
|
||||
|
||||
d_i = to_d(orig_zk, sigmas[i], denoised)
|
||||
|
||||
dt_reversible = sigma_down_freereversible - sigmas[i]
|
||||
|
||||
d_i_old = to_d(prev_x, sigmas[i], prev_denoised) if prev_denoised is not None else to_d(prev_x, sigmas[i], model(prev_x, sigmas[i] * s_in, **extra_args))
|
||||
|
||||
z_k = z_k + (d_i - d_i_old) / 2 * dt - dt_reversible**2 * (d_i_old - d_i) / 2
|
||||
|
||||
z_avg += renoise_weights[k] * z_k
|
||||
if sigmas[i + 1] > 0: # Random noise for variance on ancestral samplers
|
||||
@@ -1303,8 +1309,212 @@ def sampler_supreme(model, x, sigmas, extra_args=None, callback=None, disable=No
|
||||
|
||||
return x
|
||||
|
||||
def sample_supreme(model, x, sigmas, extra_args=None, callback=None, disable=None, s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, eta=1.0, step_method="euler", substep_method="euler", centralization=0.05, normalization=0.05, edge_enhancement=0.25, perphist=0.5, substeps=2, noise_modulation="intensity", modulation_strength=2.0, modulation_dims=3, reversible_eta=1.0):
|
||||
return sampler_supreme(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, s_noise=s_noise, noise_sampler=noise_sampler or get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), eta=eta, step_method=step_method, substep_method=substep_method, centralization=centralization, normalization=normalization, edge_enhancement=edge_enhancement, perphist=perphist, substeps=substeps, noise_modulation=noise_modulation, modulation_strength=modulation_strength, modulation_dims=modulation_dims, reversible_eta=reversible_eta)
|
||||
def sample_supreme(model, x, sigmas, extra_args=None, callback=None, disable=None, s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, eta=1.0, step_method="RES", substep_method="euler", warmup_method="euler", centralization=0.00, normalization=0.00, edge_enhancement=0.00, perphist=0.25, substeps=2, noise_modulation="none", modulation_strength=2., modulation_dims=3, reversible_eta=1.0, dyneta=True, reversible_dyneta=True, enable_free_reverse=True, free_reverse_eta=0.0, free_reverse_dyneta=True):
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
|
||||
return sampler_supreme(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), eta=eta, step_method=step_method, substep_method=substep_method, warmup_method=warmup_method, centralization=centralization, normalization=normalization, edge_enhancement=edge_enhancement, perphist=perphist, substeps=substeps, noise_modulation=noise_modulation, modulation_strength=modulation_strength, modulation_dims=modulation_dims, reversible_eta=reversible_eta, dyneta=dyneta, reversible_dyneta=reversible_dyneta, enable_free_reverse=enable_free_reverse, free_reverse_eta=free_reverse_eta, free_reverse_dyneta=free_reverse_dyneta)
|
||||
|
||||
@torch.no_grad()
|
||||
def sampler_sens(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., rsde_eta=1., tsde_eta=1., s_noise=1., noise_sampler=None):
|
||||
"""SDE-Endowed Nimble Sampler. Based off of DPM-Solver++(2M) SDE and DPM-Solver++(3M) SDE. R-SDE for reversible SDE, T-SDE for tertiary SDE."""
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
|
||||
seed = extra_args.get("seed", None)
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=True) 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]])
|
||||
|
||||
old_denoised, old_denoised_2 = None, None
|
||||
h_last, h_last_2 = None, None
|
||||
h = None
|
||||
|
||||
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:
|
||||
# DPM-Solver++(2M) SDE
|
||||
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
|
||||
h = s - t
|
||||
eta_h = eta * h
|
||||
rsde_eta_h = rsde_eta * h
|
||||
tsde_eta_h = tsde_eta * h
|
||||
|
||||
x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + (-h - eta_h).expm1().neg() * denoised
|
||||
|
||||
if old_denoised is not None:
|
||||
r = h_last / h
|
||||
x = x + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * (1 / r) * (denoised - old_denoised) / 2 - ((-h - rsde_eta_h).expm1().neg() / (-h - rsde_eta_h) + 1)**2 * (1 / r) * (old_denoised - denoised) / 2
|
||||
|
||||
# DPM-Solver++(3M) SDE
|
||||
if h_last_2 is not None and tsde_eta:
|
||||
r = h_last_2 / h
|
||||
d = (old_denoised - old_denoised_2) / r
|
||||
d_2 = (old_denoised - denoised) / r
|
||||
|
||||
d_rev = (denoised - old_denoised) / r
|
||||
d_2_rev = (old_denoised_2 - old_denoised) / r
|
||||
|
||||
#phi = eta_h.neg().expm1() / eta_h + 1
|
||||
rphi = tsde_eta_h.neg().expm1() / tsde_eta_h + 1
|
||||
x = x + rphi * (d + d_2) / 2 - rphi**2 * (d_rev + d_2_rev) / 2
|
||||
|
||||
if eta:
|
||||
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise
|
||||
|
||||
old_denoised, old_denoised_2 = denoised, old_denoised
|
||||
h_last, h_last_2 = h, h_last
|
||||
return x
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_sens(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., rsde_eta=1., tsde_eta=1., s_noise=1., noise_sampler_type="brownian", noise_sampler=None):
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
|
||||
return sampler_sens(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, rsde_eta=rsde_eta, tsde_eta=tsde_eta, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args))
|
||||
|
||||
#From https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py
|
||||
#under Apache 2 license
|
||||
def sampler_ipndm_vapp(model, x, sigmas, extra_args=None, callback=None, disable=None, max_order=4, eta=1., s_noise=1., noise_sampler=None, pp_guidance=1.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
|
||||
|
||||
temp_uncond = [0]
|
||||
temp_cond = [0]
|
||||
def post_cfg_function(args):
|
||||
temp_uncond[0] = args["uncond_denoised"]
|
||||
temp_cond[0] = args["cond_denoised"]
|
||||
return args["denoised"]
|
||||
|
||||
model_options = extra_args.get("model_options", {}).copy()
|
||||
extra_args["model_options"] = comfy.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True)
|
||||
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
|
||||
x_next = x
|
||||
t_steps = sigmas
|
||||
|
||||
buffer_model = []
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
t_cur = sigmas[i]
|
||||
t_next = sigmas[i + 1]
|
||||
sigma_down, sigma_up = get_ancestral_step(t_cur, t_next, eta=eta)
|
||||
|
||||
x_cur = x_next
|
||||
|
||||
denoised = model(x_cur, t_cur * s_in, **extra_args)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
|
||||
faux_d_cur = (x_cur - temp_uncond[0]) / t_cur # CFG++
|
||||
#d_cur = ((x_cur - temp_cond[0]) - (denoised - temp_uncond[0])) / t_cur # 2x CFG
|
||||
d_cur = -temp_cond[0] / t_cur * pp_guidance + (x_cur - denoised) / t_cur + temp_uncond[0] / t_cur * pp_guidance
|
||||
# I've found that chhanging x_cur to `denoised` results in over-denoised samples, so we're sticking with this alt method
|
||||
|
||||
order = min(max_order, i+1)
|
||||
if order == 1: # First Euler step.
|
||||
x_next = x_cur + (sigma_down - t_cur) * d_cur # Modified t_next to sigma_down for ancestral capability.
|
||||
elif order == 2: # Use one history point.
|
||||
h_n = (t_next - t_cur)
|
||||
h_n_1 = (t_cur - t_steps[i-1])
|
||||
coeff1 = (2 + (h_n / h_n_1)) / 2
|
||||
coeff2 = -(h_n / h_n_1) / 2
|
||||
x_next = x_cur + (sigma_down - t_cur) * (coeff1 * d_cur + coeff2 * buffer_model[-1])
|
||||
elif order == 3: # Use two history points.
|
||||
h_n = (t_next - t_cur)
|
||||
h_n_1 = (t_cur - t_steps[i-1])
|
||||
h_n_2 = (t_steps[i-1] - t_steps[i-2])
|
||||
temp = (1 - h_n / (3 * (h_n + h_n_1)) * (h_n * (h_n + h_n_1)) / (h_n_1 * (h_n_1 + h_n_2))) / 2
|
||||
coeff1 = (2 + (h_n / h_n_1)) / 2 + temp
|
||||
coeff2 = -(h_n / h_n_1) / 2 - (1 + h_n_1 / h_n_2) * temp
|
||||
coeff3 = temp * h_n_1 / h_n_2
|
||||
x_next = x_cur + (sigma_down - t_cur) * (coeff1 * d_cur + coeff2 * buffer_model[-1] + coeff3 * buffer_model[-2])
|
||||
elif order == 4: # Use three history points.
|
||||
h_n = (t_next - t_cur)
|
||||
h_n_1 = (t_cur - t_steps[i-1])
|
||||
h_n_2 = (t_steps[i-1] - t_steps[i-2])
|
||||
h_n_3 = (t_steps[i-2] - t_steps[i-3])
|
||||
temp1 = (1 - h_n / (3 * (h_n + h_n_1)) * (h_n * (h_n + h_n_1)) / (h_n_1 * (h_n_1 + h_n_2))) / 2
|
||||
temp2 = ((1 - h_n / (3 * (h_n + h_n_1))) / 2 + (1 - h_n / (2 * (h_n + h_n_1))) * h_n / (6 * (h_n + h_n_1 + h_n_2))) \
|
||||
* (h_n * (h_n + h_n_1) * (h_n + h_n_1 + h_n_2)) / (h_n_1 * (h_n_1 + h_n_2) * (h_n_1 + h_n_2 + h_n_3))
|
||||
coeff1 = (2 + (h_n / h_n_1)) / 2 + temp1 + temp2
|
||||
coeff2 = -(h_n / h_n_1) / 2 - (1 + h_n_1 / h_n_2) * temp1 - (1 + (h_n_1 / h_n_2) + (h_n_1 * (h_n_1 + h_n_2) / (h_n_2 * (h_n_2 + h_n_3)))) * temp2
|
||||
coeff3 = temp1 * h_n_1 / h_n_2 + ((h_n_1 / h_n_2) + (h_n_1 * (h_n_1 + h_n_2) / (h_n_2 * (h_n_2 + h_n_3))) * (1 + h_n_2 / h_n_3)) * temp2
|
||||
coeff4 = -temp2 * (h_n_1 * (h_n_1 + h_n_2) / (h_n_2 * (h_n_2 + h_n_3))) * h_n_1 / h_n_2
|
||||
x_next = x_cur + (sigma_down - t_cur) * (coeff1 * d_cur + coeff2 * buffer_model[-1] + coeff3 * buffer_model[-2] + coeff4 * buffer_model[-3])
|
||||
|
||||
if eta and sigmas[i + 1] > 0:
|
||||
x_next = x_next + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
|
||||
|
||||
if len(buffer_model) == max_order - 1:
|
||||
for k in range(max_order - 2):
|
||||
buffer_model[k] = buffer_model[k+1]
|
||||
buffer_model[-1] = faux_d_cur.detach() # Utilize CFG++ as history points
|
||||
else:
|
||||
buffer_model.append(faux_d_cur.detach())
|
||||
|
||||
return x_next
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_ipndm_vapp(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., max_order=4, noise_sampler_type="brownian", noise_sampler=None, pp_guidance=1.0):
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
|
||||
return sampler_ipndm_vapp(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, max_order=max_order, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), pp_guidance=pp_guidance)
|
||||
|
||||
@torch.no_grad()
|
||||
def sampler_STRIKE(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, order=3):
|
||||
"""Full ancestral sampling with STRIKE (Stochastic/Temporal, Reversible, and Improvised K-Diffusion Experiment) steps."""
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
||||
|
||||
temp = [0]
|
||||
temp_cond = [0]
|
||||
def post_cfg_function(args):
|
||||
temp[0] = args["uncond_denoised"]
|
||||
temp_cond[0] = args["cond_denoised"]
|
||||
return args["denoised"]
|
||||
|
||||
model_options = extra_args.get("model_options", {}).copy()
|
||||
extra_args["model_options"] = comfy.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True)
|
||||
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
old_uncond, old_uncond_2 = None, None
|
||||
old_cond, old_cond_2 = None, None
|
||||
old_dt, old_dt_2 = None, None
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
d = to_d(x, sigmas[i], temp[0])
|
||||
d_2 = to_d(x, sigmas[i], temp_cond[0])
|
||||
# Euler method
|
||||
dt = sigma_down - sigmas[i]
|
||||
x = denoised + d * dt - d_2 * dt
|
||||
if old_uncond is not None and old_cond is not None and order >= 2:
|
||||
x = x + (old_cond - old_uncond) / (old_dt / dt)
|
||||
if old_uncond_2 is not None and old_cond_2 is not None and order >= 3:
|
||||
x = x + (old_cond_2 - old_uncond_2) / (old_dt_2 / old_dt) / (old_dt / dt)
|
||||
if sigmas[i + 1] > 0:
|
||||
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigmas[i + 1]
|
||||
old_uncond, old_uncond_2 = temp[0], old_uncond
|
||||
old_cond, old_cond_2 = temp_cond[0], old_cond
|
||||
old_dt, old_dt_2 = dt, old_dt
|
||||
return x
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_STRIKE(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="brownian", noise_sampler=None, order=3):
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
|
||||
return sampler_STRIKE(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler if noise_sampler is not None else get_noise_sampler(x, sigmas, noise_sampler_type, noise_sampler, extra_args), order=order)
|
||||
|
||||
# Add your personal samplers below here, just for formatting purposes ;3
|
||||
|
||||
@@ -1318,6 +1528,9 @@ extra_samplers = {
|
||||
"euler_ancestral_dancing": sample_euler_ancestral_dancing,
|
||||
"dpmpp_3m_sde_dynamic_eta": sample_dpmpp_3m_sde_dynamic_eta,
|
||||
"supreme": sample_supreme,
|
||||
"sens": sample_sens,
|
||||
"ipndm_vapp": sample_ipndm_vapp,
|
||||
"euler_clybtune": sample_euler_clybtune,
|
||||
}
|
||||
|
||||
discard_penultimate_sigma_samplers = set((
|
||||
@@ -1336,6 +1549,31 @@ def get_sigmas_simple_exponential(model, steps):
|
||||
exp = torch.exp(torch.log(torch.linspace(1, 0, steps + 1)))
|
||||
return sigs * exp
|
||||
|
||||
def get_sigmas_kl_optimal(model, steps):
|
||||
s = model.model_sampling
|
||||
sigs = []
|
||||
alpha_min = torch.arctan(s.sigma_min).item()
|
||||
alpha_max = torch.arctan(s.sigma_max).item()
|
||||
for x in range(steps+1):
|
||||
sigs += [torch.tan(torch.tensor(((x/steps) * alpha_min + (1.0-x/steps) * alpha_max)))]
|
||||
return torch.FloatTensor(sigs)
|
||||
|
||||
def get_sigmas_simple_kl_optimal(model, steps):
|
||||
s = model.model_sampling
|
||||
sigs = []
|
||||
idx_list = []
|
||||
ss = len(s.sigmas) / steps
|
||||
for x in range(steps):
|
||||
step = (x/steps) * math.atan(len(s.sigmas) / steps) + (x/steps) * math.atan(1 / steps)
|
||||
idx = int(len(s.sigmas) * (1.0 - math.atan(step))) - 1
|
||||
idx_list += [idx]
|
||||
sigs += [float(s.sigmas[idx])]
|
||||
#print(idx_list)
|
||||
sigs += [0.0]
|
||||
return torch.FloatTensor(sigs)
|
||||
|
||||
extra_schedulers = {
|
||||
"simple_exponential": get_sigmas_simple_exponential
|
||||
"simple_exponential": get_sigmas_simple_exponential,
|
||||
"kl_optimal": get_sigmas_kl_optimal,
|
||||
"simple_kl_optimal": get_sigmas_simple_kl_optimal,
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from .other_samplers.refined_exp_solver import sample_refined_exp_s
|
||||
from .extra_samplers import get_noise_sampler_names, prepare_noise
|
||||
from .extra_samplers import get_noise_sampler_names, get_immiscible_noise_sampler_names, prepare_noise, make_immiscible
|
||||
|
||||
import comfy.samplers
|
||||
import comfy.sample
|
||||
@@ -11,6 +11,7 @@ import latent_preview
|
||||
import torch
|
||||
import math
|
||||
from tqdm.auto import trange
|
||||
import numpy as np
|
||||
|
||||
import kornia
|
||||
|
||||
@@ -156,19 +157,25 @@ class SamplerSUPREME:
|
||||
NOISE_MODULATION_TYPES=["none", "intensity", "frequency", "spectral_signum"]
|
||||
return {"required":
|
||||
{"noise_sampler_type": (get_noise_sampler_names(),),
|
||||
"step_method": (STEP_METHODS, ),
|
||||
"substep_method": (SUBSTEP_METHODS, ),
|
||||
"step_method": (STEP_METHODS, {"default": "euler"}),
|
||||
"substep_method": (SUBSTEP_METHODS, {"default": "euler"}),
|
||||
"warmup_method": (SUBSTEP_METHODS, {"default": "euler"}),
|
||||
"eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
"centralization": ("FLOAT", {"default": 0.02, "min": -1.0, "max": 1.0, "step":0.01}),
|
||||
"normalization": ("FLOAT", {"default": 0.01, "min": -1.0, "max": 1.0, "step":0.01}),
|
||||
"edge_enhancement": ("FLOAT", {"default": 0.05, "min": -100.0, "max": 100.0, "step":0.01}),
|
||||
"perphist": ("FLOAT", {"default": 0, "min": -5.0, "max": 5.0, "step":0.01}),
|
||||
"centralization": ("FLOAT", {"default": 0.00, "min": -1.0, "max": 1.0, "step":0.01}),
|
||||
"normalization": ("FLOAT", {"default": 0.00, "min": -1.0, "max": 1.0, "step":0.01}),
|
||||
"edge_enhancement": ("FLOAT", {"default": 0.00, "min": -100.0, "max": 100.0, "step":0.01}),
|
||||
"perphist": ("FLOAT", {"default": 0.25, "min": -1.0, "max": 1.0, "step":0.01}),
|
||||
"substeps": ("INT", {"default": 2, "min": 1, "max": 100, "step":1}),
|
||||
"s_noise": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
"noise_modulation": (NOISE_MODULATION_TYPES, {"default": "intensity"}),
|
||||
"modulation_strength": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100.0, "step":0.01}),
|
||||
"noise_modulation": (NOISE_MODULATION_TYPES, {"default": "none"}),
|
||||
"modulation_strength": ("FLOAT", {"default": 2., "min": -100.0, "max": 100.0, "step":0.01}),
|
||||
"modulation_dims": ("INT", {"default": 3, "min": 1, "max": 3, "step":1}),
|
||||
"reversible_eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
"dyneta": ("BOOLEAN", {"default": True}),
|
||||
"reversible_dyneta": ("BOOLEAN", {"default": True}),
|
||||
"enable_free_reverse": ("BOOLEAN", {"default": True}),
|
||||
"free_reverse_eta": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
"free_reverse_dyneta": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
@@ -176,10 +183,123 @@ class SamplerSUPREME:
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(self, noise_sampler_type, step_method, substep_method, eta, centralization, normalization, edge_enhancement, perphist, substeps, noise_modulation, modulation_strength, modulation_dims, reversible_eta, s_noise):
|
||||
sampler = comfy.samplers.ksampler("supreme", {"noise_sampler_type": noise_sampler_type, "step_method": step_method, "eta": eta, "centralization": centralization, "normalization": normalization, "edge_enhancement": edge_enhancement, "perphist": perphist, "substeps": substeps, "substep_method": substep_method, "noise_modulation": noise_modulation, "modulation_strength": modulation_strength, "modulation_dims": modulation_dims, "reversible_eta": reversible_eta, "s_noise": s_noise})
|
||||
def get_sampler(self, noise_sampler_type, step_method, substep_method, warmup_method, eta, centralization, normalization, edge_enhancement, perphist, substeps, noise_modulation, modulation_strength, modulation_dims, reversible_eta, dyneta, reversible_dyneta, enable_free_reverse, free_reverse_eta, free_reverse_dyneta, s_noise):
|
||||
sampler = comfy.samplers.ksampler("supreme", {"noise_sampler_type": noise_sampler_type, "step_method": step_method, "eta": eta, "centralization": centralization, "normalization": normalization, "edge_enhancement": edge_enhancement, "perphist": perphist, "substeps": substeps, "substep_method": substep_method, "warmup_method": warmup_method, "noise_modulation": noise_modulation, "modulation_strength": modulation_strength, "modulation_dims": modulation_dims, "reversible_eta": reversible_eta, "dyneta": dyneta, "reversible_dyneta": reversible_dyneta, "enable_free_reverse": enable_free_reverse, "free_reverse_eta": free_reverse_eta, "free_reverse_dyneta": free_reverse_dyneta, "s_noise": s_noise})
|
||||
return (sampler, )
|
||||
|
||||
# SENS
|
||||
class SamplerSENS:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"noise_sampler_type": (get_noise_sampler_names(default="brownian"), ),
|
||||
"eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
"rsde_eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
"tsde_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/samplers"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(self, noise_sampler_type, eta, rsde_eta, tsde_eta, s_noise):
|
||||
sampler = comfy.samplers.ksampler("sens", {"noise_sampler_type": noise_sampler_type, "eta": eta, "rsde_eta": rsde_eta, "tsde_eta": tsde_eta, "s_noise": s_noise})
|
||||
return (sampler, )
|
||||
|
||||
# IPNDM_V Ancestral CFG++
|
||||
class SamplerIPNDM_VAPP:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"noise_sampler_type": (get_noise_sampler_names(default="brownian"), ),
|
||||
"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}),
|
||||
"max_order": ("INT", {"default": 4, "min": 1, "max": 4, "step":1}),
|
||||
"pp_guidance": ("FLOAT", {"default": 1, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(self, noise_sampler_type, eta, s_noise, max_order, pp_guidance):
|
||||
sampler = comfy.samplers.ksampler("ipndm_vapp", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "max_order": max_order, "pp_guidance": pp_guidance})
|
||||
return (sampler, )
|
||||
|
||||
# STRIKE (Stochastic Temporal Reversible Improvised K-Diffusion Experiment)
|
||||
class SamplerSTRIKE:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"noise_sampler_type": (get_noise_sampler_names(default="gaussian"), ),
|
||||
"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}),
|
||||
"order": ("INT", {"default": 3, "min": 1, "max": 3, "step":1}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(self, noise_sampler_type, eta, s_noise, order):
|
||||
sampler = comfy.samplers.ksampler("STRIKE", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "order": order})
|
||||
return (sampler, )
|
||||
|
||||
### Noise
|
||||
|
||||
class Noise_ImmiscibleNoise:
|
||||
def __init__(self, noise_type, seed, image_scaling, latent_image):
|
||||
self.noise_type = noise_type
|
||||
self.seed = seed
|
||||
self.image_scaling = image_scaling
|
||||
self.latent_image = latent_image
|
||||
|
||||
def generate_noise(self, input_latent):
|
||||
latent_image = input_latent["samples"]
|
||||
batch_inds = input_latent["batch_index"] if "batch_index" in input_latent else None
|
||||
generator = torch.manual_seed(self.seed)
|
||||
if batch_inds is None:
|
||||
gauss = torch.randn_like(latent_image)
|
||||
noise = make_immiscible(noise_func=self.noise_type)(latent_image if self.latent_image is None else self.latent_image["samples"])
|
||||
noise = gauss * (1.0 - self.image_scaling) + noise * self.image_scaling
|
||||
return noise
|
||||
#return torch.randn(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, generator=generator, device="cpu")
|
||||
|
||||
unique_inds, inverse = np.unique(batch_inds, return_inverse=True)
|
||||
noises = []
|
||||
for i in range(unique_inds[-1]+1):
|
||||
gauss = torch.randn_like(latent_image)
|
||||
noise = make_immiscible(noise_func=self.noise_type)(latent_image if self.latent_image is None else self.latent_image["samples"])
|
||||
noise = gauss * (1.0 - self.image_scaling) + noise * self.image_scaling
|
||||
if i in unique_inds:
|
||||
noises.append(noise)
|
||||
noises = [noises[i] for i in inverse]
|
||||
noises = torch.cat(noises, axis=0)
|
||||
return noises
|
||||
|
||||
from comfy_extras.nodes_custom_sampler import DisableNoise
|
||||
class ImmiscibleNoise(DisableNoise):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{
|
||||
"noise_type": (get_immiscible_noise_sampler_names(), ),
|
||||
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
},
|
||||
"optional":
|
||||
{
|
||||
"image_scaling": ("FLOAT", {"default": 1.0, "min": -1000.0, "max": 1000.0, "step":0.01, "round": 0.001}),
|
||||
"latent_image": ("LATENT", ),
|
||||
}
|
||||
}
|
||||
|
||||
def get_noise(self, noise_type, noise_seed, image_scaling, latent_image):
|
||||
return (Noise_ImmiscibleNoise(noise_type, noise_seed, image_scaling, latent_image),)
|
||||
|
||||
### Schedulers
|
||||
from .extra_samplers import get_sigmas_simple_exponential
|
||||
class SimpleExponentialScheduler:
|
||||
@@ -192,7 +312,7 @@ class SimpleExponentialScheduler:
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("SIGMAS",)
|
||||
CATEGORY = "clybNodes/schedulers"
|
||||
CATEGORY = "sampling/custom_sampling/schedulers"
|
||||
|
||||
FUNCTION = "get_sigmas"
|
||||
|
||||
@@ -205,6 +325,54 @@ class SimpleExponentialScheduler:
|
||||
sigmas = sigmas[-(steps + 1):]
|
||||
return (sigmas, )
|
||||
|
||||
from .extra_samplers import get_sigmas_simple_kl_optimal
|
||||
class SimpleKLOptimalScheduler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"model": ("MODEL",),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("SIGMAS",)
|
||||
CATEGORY = "sampling/custom_sampling/schedulers"
|
||||
|
||||
FUNCTION = "get_sigmas"
|
||||
|
||||
def get_sigmas(self, model, steps, denoise):
|
||||
total_steps = steps
|
||||
if denoise < 1.0:
|
||||
total_steps = int(steps/denoise)
|
||||
|
||||
sigmas = get_sigmas_simple_kl_optimal(model.model, total_steps).cpu()
|
||||
sigmas = sigmas[-(steps + 1):]
|
||||
return (sigmas, )
|
||||
|
||||
from .extra_samplers import get_sigmas_kl_optimal
|
||||
class KLOptimalScheduler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"model": ("MODEL",),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("SIGMAS",)
|
||||
CATEGORY = "sampling/custom_sampling/schedulers"
|
||||
|
||||
FUNCTION = "get_sigmas"
|
||||
|
||||
def get_sigmas(self, model, steps, denoise):
|
||||
total_steps = steps
|
||||
if denoise < 1.0:
|
||||
total_steps = int(steps/denoise)
|
||||
|
||||
sigmas = get_sigmas_kl_optimal(model.model, total_steps).cpu()
|
||||
sigmas = sigmas[-(steps + 1):]
|
||||
return (sigmas, )
|
||||
|
||||
### KSampler Nodes
|
||||
|
||||
from comfy import model_management
|
||||
@@ -730,15 +898,20 @@ class WarmupDecayCFGGuider:
|
||||
return (guider,)
|
||||
|
||||
class Guider_MegaCFG(comfy.samplers.CFGGuider):
|
||||
def set_cfg(self, model, cfg_max, cfg_min, warmup_percent, mean_cfg):
|
||||
def set_cfg(self, model, cfg_max, cfg_min, warmup_percent, mean_cfg, vector_rejection_scale):
|
||||
self.model = model
|
||||
self.cfg_max = cfg_max
|
||||
self.cfg_min = cfg_min
|
||||
self.warmup_percent = warmup_percent
|
||||
self.mean_cfg = mean_cfg
|
||||
|
||||
self.vector_rejection_scale = vector_rejection_scale
|
||||
|
||||
self.prev_cond = None
|
||||
self.prev_cfg = None
|
||||
|
||||
self.cond_result = None
|
||||
self.cfg_result = None
|
||||
|
||||
def set_conds(self, positive, negative):
|
||||
self.inner_set_conds({"positive": positive, "negative": negative})
|
||||
@@ -748,6 +921,17 @@ class Guider_MegaCFG(comfy.samplers.CFGGuider):
|
||||
self.image_weighting = image_weighting
|
||||
self.weight_scaling = weight_scaling
|
||||
self.latent_image = latent_image
|
||||
|
||||
def set_perphist_params(self, perphist):
|
||||
self.perphist = perphist
|
||||
|
||||
def perpadd(self, denoised_tensor, old_denoised_tensor, x, alpha):
|
||||
a_diff = x - (denoised_tensor - x)
|
||||
b_diff = x - (old_denoised_tensor - x)
|
||||
a_ortho = a_diff * (a_diff / torch.linalg.norm(a_diff) * (b_diff / torch.linalg.norm(a_diff))).sum()
|
||||
b_perp = b_diff - a_ortho
|
||||
res = denoised_tensor + alpha * b_perp
|
||||
return res
|
||||
|
||||
def post_cfg_reference_img(self, args):
|
||||
model = args["model"]
|
||||
@@ -776,6 +960,24 @@ class Guider_MegaCFG(comfy.samplers.CFGGuider):
|
||||
weight = ((-torch.cos((sigma / sigma_max) * math.pi) / 2) + 0.5)[:, None, None, None].clone()
|
||||
|
||||
return cfg_result + (cond_pred - ref) * self.image_guidance * (weight**self.weight_scaling)
|
||||
|
||||
def post_cfg_perphist(self, args):
|
||||
noise_pred = args["denoised"]
|
||||
if self.prev_cfg != None:
|
||||
noise_pred = self.perpadd(noise_pred, self.prev_cfg, args["input"], self.perphist)
|
||||
self.prev_cfg = args["denoised"]
|
||||
return noise_pred
|
||||
|
||||
def vect_rej(self, conditioning, unconditioning, x_input):
|
||||
def rej(a, b):
|
||||
"""
|
||||
Implements vector rejection for alternative diffusion.
|
||||
"""
|
||||
return a - b * torch.tensordot(a, b, dims=4) / torch.tensordot(b, b, dims=4)
|
||||
cond_ind_pred = conditioning - x_input
|
||||
neg_ind_pred = unconditioning - x_input
|
||||
noise_pred = rej(cond_ind_pred, neg_ind_pred) + (rej(x_input, cond_ind_pred) - rej(x_input, neg_ind_pred))
|
||||
return noise_pred
|
||||
|
||||
def predict_noise(self, x, timestep, model_options={}, seed=None):
|
||||
negative_cond = self.conds.get("negative", None)
|
||||
@@ -802,12 +1004,14 @@ class Guider_MegaCFG(comfy.samplers.CFGGuider):
|
||||
cfg_cos = (1 + -torch.cos((timestep / percent_sigma) * math.pi))
|
||||
mod_cfg = cfg_scale * cfg_cos + self.cfg_min
|
||||
|
||||
if self.vector_rejection_scale != 0:
|
||||
out[1] += self.vect_rej(out[1], out[0], x) * self.vector_rejection_scale
|
||||
|
||||
cfg = comfy.samplers.cfg_function(self.inner_model, out[1], out[0], mod_cfg, x, timestep, model_options=model_options, cond=positive_cond, uncond=negative_cond)
|
||||
|
||||
if self.mean_cfg != 0:
|
||||
cfg += out0_mean + (out1_mean - out0_mean) * self.mean_cfg
|
||||
|
||||
self.prev_cfg = cfg
|
||||
self.prev_cond = out[1]
|
||||
|
||||
return cfg
|
||||
@@ -823,6 +1027,8 @@ class MegaCFGGuider:
|
||||
"cfg_min": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
|
||||
"warmup_percent": ("FLOAT", {"default": 0.5, "min": 0.01, "max": 1.0, "step":0.01, "round": 0.001}),
|
||||
"mean_cfg": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
|
||||
"perphist": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step":0.01, "round": 0.001}),
|
||||
"vector_rejection_scale": ("FLOAT", {"default": 0.0, "min": -0.5, "max": 5.0, "step":0.001, "round": 0.0001}),
|
||||
},
|
||||
"optional":
|
||||
{
|
||||
@@ -838,13 +1044,17 @@ class MegaCFGGuider:
|
||||
FUNCTION = "get_guider"
|
||||
CATEGORY = "sampling/custom_sampling/guiders"
|
||||
|
||||
def get_guider(self, model, positive, negative, cfg_max, cfg_min, warmup_percent, mean_cfg,
|
||||
def get_guider(self, model, positive, negative, cfg_max, cfg_min, warmup_percent, mean_cfg, perphist, vector_rejection_scale,
|
||||
image_guidance, image_weighting, weight_scaling, latent_image = None):
|
||||
m = model.clone()
|
||||
copy_model = perphist != 0 or latent_image != 0
|
||||
m = model.clone() if copy_model else model
|
||||
guider = Guider_MegaCFG(m)
|
||||
guider.set_conds(positive, negative) # Conds
|
||||
guider.set_cfg(m, cfg_max, cfg_min, warmup_percent, mean_cfg) # Strengths
|
||||
guider.set_cfg(m, cfg_max, cfg_min, warmup_percent, mean_cfg, vector_rejection_scale) # Strengths
|
||||
if latent_image != None:
|
||||
guider.set_img_cfg(image_guidance, image_weighting, weight_scaling, latent_image)
|
||||
m.set_model_sampler_post_cfg_function(guider.post_cfg_reference_img)
|
||||
return (guider,)
|
||||
if perphist != 0:
|
||||
guider.set_perphist_params(perphist)
|
||||
m.set_model_sampler_post_cfg_function(guider.post_cfg_perphist)
|
||||
return (guider,)
|
||||
Reference in New Issue
Block a user