idk whats even in this but i'll document it and push it later™️

This commit is contained in:
Clybius
2024-11-15 11:21:09 -06:00
parent bc2f7264bd
commit efffcbdc3d
3 changed files with 841 additions and 43 deletions
+5
View File
@@ -15,6 +15,7 @@ NODE_CLASS_MAPPINGS = {
"ScaledCFGGuider": nodes.ScaledCFGGuider,
"WarmupDecayCFGGuider": nodes.WarmupDecayCFGGuider,
"MegaCFGGuider": nodes.MegaCFGGuider,
"APGGuider": nodes.APGGuider,
### Noise
"ImmiscibleNoise": nodes.ImmiscibleNoise,
## Samplers
@@ -29,6 +30,10 @@ NODE_CLASS_MAPPINGS = {
"SamplerSENS": nodes.SamplerSENS,
"SamplerIPNDM_VAPP": nodes.SamplerIPNDM_VAPP,
"SamplerSHIDS": nodes.SamplerSHIDS,
"SamplerDPMPP_2M_SDE_EMA": nodes.SamplerDPMPP_2M_SDE_EMA,
"SamplerBiScope": nodes.SamplerBiScope,
"SamplerEuler_G": nodes.SamplerEuler_G,
"SamplerLeaping_Euler": nodes.SamplerLeaping_Euler,
### Schedulers
"SimpleExponentialScheduler": nodes.SimpleExponentialScheduler,
"KLOptimalScheduler": nodes.KLOptimalScheduler,
+602 -36
View File
@@ -50,8 +50,8 @@ def add_schedulers():
# Noise samplers
IMMISCIBLE_NOISE_NAMES=("gaussian_1024", "perlin")
NOISE_SAMPLER_NAMES=("gaussian", "uniform", "brownian", "highres-pyramid", "pyramid", "perlin", "laplacian", "immiscible_gaussian", "immiscible_gaussian_maximize", "immiscible_perlin")
IMMISCIBLE_NOISE_NAMES=("gaussian", "perlin")
NOISE_SAMPLER_NAMES=("gaussian", "uniform", "brownian", "highres-pyramid", "pyramid", "perlin", "laplacian", "immiscible_gaussian", "immiscible_gaussian_maximize", "immiscible_perlin", "immiscible_perlin_maximize")
def get_noise_sampler_names(default=None):
if not default:
@@ -86,30 +86,35 @@ 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
immiscibility = make_immiscible("gaussian") # 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
immiscibility = make_immiscible("gaussian", maximize=True) # 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
immiscibility = make_immiscible("perlin") # 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_maximize":
immiscibility = make_immiscible("perlin", maximize=True) # 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
class make_immiscible:
def __init__(self, noise_func="gaussian_1024", immiscible_latents=1024, maximize=False, batching="batch"):
def __init__(self, noise_func="gaussian", immiscible_latents=1024, maximize=False):
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
@@ -166,6 +171,36 @@ class make_immiscible:
return noise.view(*xsz[:2], xsz[3], xsz[2]).permute(0, 1, 3, 2)
return noise.view(*xsz)
"""
def __call__(self, latents):
reference_latent = latents
if self.updated_latent != None:
reference_latent = self.updated_latent
batch_size = latents.shape[0] if self.n_latents is None else self.n_latents
size = [batch_size] + list(latents.shape[1:])
#noise = torch.randn_like(latents) # [B, C, H, W]
match self.noise_func:
case "gaussian":
noise = torch.randn(size, dtype=latents.dtype, layout=latents.layout, device=latents.device)
case "perlin":
noise = create_noisy_latents_perlin(torch.randn(size, dtype=latents.dtype, layout=latents.layout, device=latents.device))
# Distance calculation (simplified for single process)
distance = torch.linalg.vector_norm(
0.10 * latents.to(torch.float16).flatten(start_dim=1).unsqueeze(1) -
0.10 * noise.to(torch.float16).flatten(start_dim=1).unsqueeze(0),
dim=2
) # [B, B]
# Noise Assignment (simplified for single process)
_, col_ind = linear_sum_assignment(distance.cpu().numpy(), maximize=self.maximize)
noise = noise[col_ind].to(latents.device) # Assign the permuted noise
return noise
def set_immiscible_extra_args(self, extra_args):
def immiscible_post_cfg_function(args):
self.updated_latent = args["cond_denoised"]
@@ -1316,7 +1351,7 @@ def sample_supreme(model, x, sigmas, extra_args=None, callback=None, disable=Non
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):
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, flow=False):
"""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
@@ -1346,6 +1381,13 @@ def sampler_sens(model, x, sigmas, extra_args=None, callback=None, disable=None,
rsde_eta_h = rsde_eta * h
tsde_eta_h = tsde_eta * h
# If/for flow model
downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta
sigma_down = sigmas[i+1] * downstep_ratio
alpha_ip1 = 1 - sigmas[i+1]
alpha_down = 1 - sigma_down
renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5
x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + (-h - eta_h).expm1().neg() * denoised
if old_denoised is not None:
@@ -1365,8 +1407,10 @@ def sampler_sens(model, x, sigmas, extra_args=None, callback=None, disable=None,
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:
if eta and not flow:
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise
elif eta and flow:
x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
old_denoised, old_denoised_2 = denoised, old_denoised
h_last, h_last_2 = h, h_last
@@ -1376,8 +1420,11 @@ def sampler_sens(model, x, sigmas, extra_args=None, callback=None, disable=None,
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
flow = False
if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):
flow = True
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))
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), flow=flow)
#From https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py
#under Apache 2 license
@@ -1468,8 +1515,11 @@ def sample_ipndm_vapp(model, x, sigmas, extra_args=None, callback=None, disable=
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)
import functools
import operator
@torch.no_grad()
def sampler_SHIDS(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, order=16):
def sampler_SHIDS(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, order=16, eta_order=1., solver_method="weighted_projection", flow=False):
"""Full ancestral sampling with SHIDS (Stochastic, Historical, Improvised Sampling) 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
@@ -1490,65 +1540,577 @@ def sampler_SHIDS(model, x, sigmas, extra_args=None, callback=None, disable=None
old_dt, old_dt_2 = None, None
buffer_model_cond = []
buffer_model_uncond = []
buffer_model_dt = []
#buffer_model_uncond = []
#buffer_model_dt = []
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)
_, sigma_up_order = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta_order)
# If/for flow model
downstep_ratio = None
sigma_down_rf = None
alpha_ip1 = None
alpha_down = None
renoise_coeff = None
if flow:
downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta
sigma_down_rf = sigmas[i+1] * downstep_ratio
alpha_ip1 = 1 - sigmas[i+1]
alpha_down = 1 - sigma_down_rf
renoise_coeff = (sigmas[i+1]**2 - sigma_down_rf**2*alpha_ip1**2/alpha_down**2)**0.5
if callback is not None:
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
d_step = to_d(x, sigmas[i], denoised)
d_full = to_d(x, sigmas[i], denoised)
d = to_d(x, sigmas[i], temp_uncond[0])
#d_2 = to_d(x, sigmas[i], temp_cond[0])
# Euler method
dt = sigma_down - sigmas[i] # Time Difference between now and next step (negative)
x_full = denoised + d_step * sigma_down
x_full = denoised + d_full * sigma_down
x_step = denoised + d * sigma_down
# Project denoised onto a line between (primarily) x_step (cfgpp), and x_full (normal cfg)
ba = x_step - denoised
ca = x_full - denoised
alpha = (ba * ca) / (ba ** 2 + 1e-8)
x = (1 - alpha)*denoised + alpha*x_step
match solver_method:
case "weighted_projection":
ba = x_step - denoised
ca = x_full - denoised
alpha = (ba * ca) / (ba ** 2 + 1e-8)
x = (1 - alpha)*denoised + alpha*x_step
case "qr_decomposition":
original_shape = x_step.shape
if not original_shape:
shape_2d = (1, 1)
elif len(original_shape) == 4:
shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:]))
else:
shape_2d = (-1, original_shape[-1])
A = x_step.reshape(shape_2d)
B = x_full.reshape(shape_2d)
C = denoised.reshape(shape_2d)
Q, _ = torch.qr(A - C)
# Compute the mapping matrix
mapping_matrix = torch.mm(Q.t(), B - C)
mapped_tensor = torch.mm(Q, mapping_matrix)
x = (C + mapped_tensor).reshape(original_shape)
case "svd_lowrank":
original_shape = x_step.shape
if not original_shape:
shape_2d = (1, 1)
elif len(original_shape) == 4:
shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:]))
else:
shape_2d = (-1, original_shape[-1])
A = x_step.reshape(shape_2d)
B = x_full.reshape(shape_2d)
C = denoised.reshape(shape_2d)
Ua, Sa, Va = torch.svd_lowrank(A - C, q=6, niter=2)
A_lowrank = torch.mm(Ua, torch.mm(torch.diag(Sa), Va.t()))
A_diff = (A - C) - A_lowrank
Qb, _ = torch.qr(B - C)
A_diff_projected = torch.mm(Qb, torch.mm(Qb.t(), A_diff))
x = (B + A_diff_projected).reshape(original_shape)
case "svd":
original_shape = x_step.shape
if not original_shape:
shape_2d = (1, 1)
elif len(original_shape) == 4:
shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:]))
else:
shape_2d = (-1, original_shape[-1])
A = x_step.reshape(shape_2d)
B = x_full.reshape(shape_2d)
C = denoised.reshape(shape_2d)
Ua, Sa, Va = torch.linalg.svd(A - C, full_matrices=False, driver="gesvd")
Ub, Sb, Vb = torch.linalg.svd(B - C, full_matrices=False, driver="gesvd")#Sb = torch.linalg.svdvals(B - C, driver="gesvd")#
A_lowrank = torch.mm(Ub, torch.mm(torch.diag_embed(Sa), Vb))
A_diff = (A - C) - A_lowrank
#Qb, _ = torch.qr(B - C)
#A_diff_projected = torch.mm(Ua, torch.mm(Ua.t(), A_diff))
x = (B + A_diff).reshape(original_shape)
# Create a list of order multipliers
multipliers = [i for i in range(1, len(buffer_model_cond))]
# Normalize so that they're summed up to a total of 1
total = sum(multipliers)
normalized_multipliers = [m / total for m in multipliers]
for iteration in range(len(buffer_model_cond) - 1):
#x = x - (buffer_model_uncond[iteration] - buffer_model_cond[iteration]) / (buffer_model_dt[iteration + 1] / buffer_model_dt[iteration])
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
#denoised = model(x, sigmas[i + 1] * s_in, **extra_args)
ba = x - buffer_model_cond[iteration]
ca = x_step - buffer_model_cond[iteration]
alpha = (ba * ca) / (ba ** 2 + 1e-8)
x = (1 - alpha)*buffer_model_cond[iteration] + alpha*x
if not flow:
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up_order * normalized_multipliers[iteration]
elif flow and eta_order:
downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta_order
sigma_down_rf = sigmas[i+1] * downstep_ratio
alpha_ip1 = 1 - sigmas[i+1]
alpha_down = 1 - sigma_down_rf
renoise_coeff = (sigmas[i+1]**2 - sigma_down_rf**2*alpha_ip1**2/alpha_down**2)**0.5
x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
match solver_method:
case "weighted_projection":
ba = x - buffer_model_cond[iteration]
ca = x_step - buffer_model_cond[iteration]
alpha = (ba * ca) / (ba ** 2 + 1e-8)
x = (1 - alpha)*buffer_model_cond[iteration] + alpha*x
case "qr_decomposition":
original_shape = x_step.shape
if not original_shape:
shape_2d = (1, 1)
elif len(original_shape) == 4:
shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:]))
else:
shape_2d = (-1, original_shape[-1])
A = x.reshape(shape_2d)
B = x_step.reshape(shape_2d)
C = buffer_model_cond[iteration].reshape(shape_2d)
Q, _ = torch.qr(A - C)
# Compute the mapping matrix
mapping_matrix = torch.mm(Q.t(), B - C)
mapped_tensor = torch.mm(Q, mapping_matrix)
x = (C + mapped_tensor).reshape(original_shape)
case "svd_lowrank":
original_shape = x_step.shape
if not original_shape:
shape_2d = (1, 1)
elif len(original_shape) == 4:
shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:]))
else:
shape_2d = (-1, original_shape[-1])
A = x.reshape(shape_2d)
B = x_step.reshape(shape_2d)
C = buffer_model_cond[iteration].reshape(shape_2d)
Ua, Sa, Va = torch.svd_lowrank(A - C, q=6, niter=2)
A_lowrank = torch.mm(Ua, torch.mm(torch.diag(Sa), Va.t()))
A_diff = (A - C) - A_lowrank
#Qb, _ = torch.qr(B - C)
#A_diff_projected = torch.mm(Qb, torch.mm(Qb.t(), A_diff))
x = (B + A_diff).reshape(original_shape)
case "svd":
original_shape = x.shape
if not original_shape:
shape_2d = (1, 1)
elif len(original_shape) == 4:
shape_2d = (-1, functools.reduce(operator.mul, original_shape[1:]))
else:
shape_2d = (-1, original_shape[-1])
A = x.reshape(shape_2d)
B = x_step.reshape(shape_2d)
C = buffer_model_cond[iteration].reshape(shape_2d)
Ua, Sa, Va = torch.linalg.svd(A - C, full_matrices=False, driver="gesvd")
Ub, Sb, Vb = torch.linalg.svd(B - C, full_matrices=False, driver="gesvd")#Sb = torch.linalg.svdvals(B - C, driver="gesvd")#
A_lowrank = torch.mm(Ub, torch.mm(torch.diag_embed(Sa), Vb))
A_diff = (A - C) - A_lowrank
#Qb, _ = torch.qr(B - C)
#A_diff_projected = torch.mm(Ua, torch.mm(Ua.t(), A_diff))
x = (B + A_diff).reshape(original_shape)
if len(buffer_model_cond) == max(order - 1, 1):
for k in range(order - 2):
buffer_model_cond[k] = buffer_model_cond[k+1]
buffer_model_uncond[k] = buffer_model_uncond[k+1]
buffer_model_dt[k] = buffer_model_dt[k+1]
#buffer_model_uncond[k] = buffer_model_uncond[k+1]
#buffer_model_dt[k] = buffer_model_dt[k+1]
buffer_model_cond[-1] = denoised.detach()
buffer_model_uncond[-1] = temp_uncond[0].detach()
buffer_model_dt[-1] = dt.detach()
#buffer_model_uncond[-1] = temp_uncond[0].detach()
#buffer_model_dt[-1] = dt.detach()
else:
buffer_model_cond.append(denoised.detach())
buffer_model_uncond.append(temp_uncond[0].detach())
buffer_model_dt.append(dt.detach())
#buffer_model_uncond.append(temp_uncond[0].detach())
#buffer_model_dt.append(dt.detach())
#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:
if sigmas[i + 1] > 0 and not flow:
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
elif sigmas[i + 1] > 0 and flow:
x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
#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_SHIDS(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, order=16):
def sample_SHIDS(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, order=16, eta_order=1., solver_method="weighted_projection"):
if len(sigmas) <= 1:
return x
flow = False
if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):
flow = True
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
return sampler_SHIDS(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)
return sampler_SHIDS(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, eta_order=eta_order, solver_method=solver_method, flow=flow)
@torch.no_grad()
def sampler_dpmpp_2m_sde_ema(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, amp_fac=2., beta1=0.8, beta2=0.95, weight_decay=0.1, centralization=1.0, normalization=1.0, flow=False):
"""DPM-Solver++(2M) SDE, with EMA uncond."""
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 = None
h_last = None
h = None
ema = torch.zeros_like(x)
ema_squared = torch.zeros_like(x)
grad = None
temp_cond = [0]
temp_uncond = [0]
#alpha = torch.linspace(1.0, 0.0, steps=len(sigmas)) ** amp_fac
alpha = [0]
def ema_retrieve_uncond_alpha(args):
temp_cond[0] = args["cond_denoised"]
temp_uncond[0] = args["uncond_denoised"]
alpha[0] = model.inner_model.inner_model.model_sampling.timestep(args["sigma"]) / 999.0
#alpha[0] = args["sigma"]
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, ema_retrieve_uncond_alpha, disable_cfg1_optimization=True)
ema = torch.zeros_like(x)
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
# If/for flow model
downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta
sigma_down = sigmas[i+1] * downstep_ratio
alpha_ip1 = 1 - sigmas[i+1]
alpha_down = 1 - sigma_down
renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5
grad = denoised
# Centralization
if centralization != 0:
grad.sub_(
grad.mean(dim=tuple(range(1, grad.dim())), keepdim=True).mul_(centralization)
)
# Lerp EMA
ema.lerp_(grad, 1. - beta1)
# Normalization
ema.lerp_(ema.div(ema.std(dim=tuple(range(1, grad.dim())), keepdim=True)), weight=normalization)
# Apply EMA onto grad (denoised)
grad.lerp_(ema, beta2)
if weight_decay != 0:
# Perform stepweight decay
wd_mult = 1 / (1 + weight_decay * (sigmas[i] - sigmas[i + 1]))
grad.mul_(wd_mult)
ema += (x - temp_uncond[0]) / sigmas[i] * amp_fac * (sigmas[i] - sigmas[i + 1])
ema -= (x - temp_cond[0]) / sigmas[i] * amp_fac * (sigmas[i] - sigmas[i + 1])
x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + (-h - eta_h).expm1().neg() * grad
if old_denoised is not None:
r = h_last / h
x = x + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * (1 / r) * (grad - old_denoised)
if eta and not flow:
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise
elif eta and flow:
x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
old_denoised = denoised
h_last = h
return x
@torch.no_grad()
def sample_dpmpp_2m_sde_ema(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="brownian", noise_sampler=None, amp_fac=2., beta1=0.8, beta2=0.95, weight_decay=0.1, centralization=1.0, normalization=1.0):
if len(sigmas) <= 1:
return x
flow = False
if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):
flow = True
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
return sampler_dpmpp_2m_sde_ema(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), amp_fac=amp_fac, beta1=beta1, beta2=beta2, weight_decay=weight_decay, centralization=centralization, normalization=normalization, flow=flow)
@torch.no_grad()
def sampler_biscope(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, amp_fac=2.0, local_smoothing_fac=4, smoothing_fac=0.75, ema_fac=0.9, flow=False):
"""Solving for the Compass model's noise problem using the Compass-like training procedure as an inference sampler."""
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]])
local_smoothing = []
smoothing = None
smoothing_diff = None
ema = None
prev_denoised = None
for i in trange(len(sigmas) - 1, disable=disable):
denoised = model(x, sigmas[i] * s_in, **extra_args)
grad = denoised
if len(local_smoothing) == max(local_smoothing_fac, 1):
for k in range(local_smoothing_fac - 1):
local_smoothing[k] = local_smoothing[k+1]
local_smoothing[-1] = grad.detach()
else:
local_smoothing.append(grad.detach())
#print(local_smoothing)
local_grad = torch.mean(torch.stack(local_smoothing), dim=0)# if len(local_smoothing) > 1 else grad
if smoothing is None:
smoothing = local_grad
smoothing.mul_(smoothing_fac).add_(local_grad, alpha=1 - smoothing_fac)
diff_grad = local_grad - smoothing
if smoothing_diff is None:
smoothing_diff = diff_grad
smoothing_diff.mul_(smoothing_fac).add_(diff_grad, alpha=1 - smoothing_fac)
local_grad.add_(smoothing_diff, alpha=amp_fac)
if ema is None:
ema = local_grad
ema.mul_(ema_fac).add_(local_grad, alpha=1 - ema_fac)
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
# Flow
downstep_ratio = None
alpha_ip1 = None
alpha_down = None
renoise_coeff = None
if flow:
# If/for flow model
downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta
sigma_down = sigmas[i+1] * downstep_ratio
alpha_ip1 = 1 - sigmas[i+1]
alpha_down = 1 - sigma_down
renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5
if callback is not None:
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': ema})
d = to_d(x, sigmas[i], ema)
# Euler method
dt = sigma_down - sigmas[i]
x = x + d * dt
if sigmas[i + 1] > 0 and not flow:
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
elif sigmas[i + 1] > 0 and flow:
x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
prev_denoised = denoised
return x
@torch.no_grad()
def sample_biscope(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, amp_fac=2.0, local_smoothing_fac=4, smoothing_fac=0.5, ema_fac=0.5):
if len(sigmas) <= 1:
return x
flow = False
if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):
flow = True
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
return sampler_biscope(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), amp_fac=amp_fac, local_smoothing_fac=local_smoothing_fac, smoothing_fac=smoothing_fac, ema_fac=ema_fac, flow=flow)
def gaussian_kernel_2d(kernel_size, sigma):
"""Generates a 2D Gaussian kernel."""
k = kernel_size // 2
x, y = torch.meshgrid(torch.arange(-k, k + 1, dtype=torch.float32), torch.arange(-k, k + 1, dtype=torch.float32))
gaussian = torch.exp(-(x**2 + y**2) / (2 * sigma**2))
return gaussian / gaussian.sum()
@torch.no_grad()
def sampler_euler_g(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, g_eta=1.0, sigma=5.0, order=2, flow=False):
"""Ancestral sampling with Euler method 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
s_in = x.new_ones([x.shape[0]])
neighborhood_size = min(x.shape[-2], x.shape[-1]) * 2 + 1
padding = neighborhood_size // 2
kernel = gaussian_kernel_2d(neighborhood_size, sigma).unsqueeze(0).unsqueeze(0).repeat(x.shape[1], 1, 1, 1).to(x.device)
x_buffer = []
denoised_buffer = []
for i in trange(len(sigmas) - 1, disable=disable):
denoised = model(x, sigmas[i] * s_in, **extra_args)
if sigmas[i + 1] == 0:
return denoised
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
# Flow
downstep_ratio = None
alpha_ip1 = None
alpha_down = None
renoise_coeff = None
if flow:
# If/for flow model
downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta
sigma_down = sigmas[i+1] * downstep_ratio
alpha_ip1 = 1 - sigmas[i+1]
alpha_down = 1 - sigma_down
renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5
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], denoised)
# Euler method
dt = sigma_down - sigmas[i]
x = x + d * dt
"""
for curr_order in range(1, order):
if sigmas[i + 1] > 0 and not flow:
faux_x = torch.nn.functional.conv2d(x, kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * (sigma_up)
elif sigmas[i + 1] > 0 and flow:
#x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
#faux_x = (alpha_ip1/(1 - sigmas[i])) * torch.nn.functional.conv2d(x, kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * (sigmas[i+1]**2 - sigmas[i]**2*alpha_ip1**2/(1 - sigmas[i])**2)**0.5
faux_x = (alpha_ip1/alpha_down) * torch.nn.functional.conv2d(x, kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
#gauss_x = F.conv2d(faux_x, kernel, padding=padding, groups=x.shape[1])
faux_denoised = model(faux_x, sigmas[i + 1] * s_in, **extra_args)
faux_d = to_d(faux_x, sigmas[i + 1], faux_denoised)
x = x - faux_d * (sigmas[i + 1] - sigmas[i]) * g_eta / (order - 1)
"""
if sigmas[i + 1] > 0:
# Create a list of order multipliers
multipliers = [i for i in range(1, len(x_buffer))]
# Normalize so that they're summed up to a total of 1
total = sum(multipliers)
normalized_multipliers = [m / total for m in multipliers]
for iteration in range(len(x_buffer) - 1):
if not flow:
faux_x = torch.nn.functional.conv2d(x_buffer[iteration], kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
else:
#x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
#faux_x = (alpha_ip1/(1 - sigmas[i])) * torch.nn.functional.conv2d(x, kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * (sigmas[i+1]**2 - sigmas[i]**2*alpha_ip1**2/(1 - sigmas[i])**2)**0.5
faux_x = (alpha_ip1/alpha_down) * torch.nn.functional.conv2d(x_buffer[iteration], kernel, padding=padding, groups=x.shape[1]) + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
#gauss_x = F.conv2d(faux_x, kernel, padding=padding, groups=x.shape[1])
faux_denoised = model(faux_x, sigmas[i + 1] * s_in, **extra_args)
faux_d = to_d(faux_x, sigmas[i + 1], faux_denoised)
x = x - faux_d * (sigmas[i + 1] - sigmas[i]) * g_eta * normalized_multipliers[iteration]
if len(x_buffer) == max(order - 1, 1):
for k in range(order - 2):
x_buffer[k] = x_buffer[k+1]
denoised_buffer[k] = denoised_buffer[k+1]
x_buffer[-1] = x.detach()
denoised_buffer[-1] = denoised.detach()
else:
x_buffer.append(x.detach())
denoised_buffer.append(denoised.detach())
if flow:
x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
else:
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
return x
@torch.no_grad()
def sample_euler_g(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None, g_eta=1.0, sigma=5.0, order=2):
if len(sigmas) <= 1:
return x
flow = False
if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):
flow = True
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
return sampler_euler_g(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), g_eta=g_eta, sigma=sigma, order=order, flow=flow)
@torch.no_grad()
def sampler_leaping_euler(model, x, sigmas, extra_args=None, callback=None, disable=None, leap=1, eta=1., s_noise=1., noise_sampler=None, flow=False):
#if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):
# return sample_euler_ancestral_RF(model, x, sigmas, extra_args, callback, disable, eta, s_noise, noise_sampler)
"""Ancestral sampling with Euler method 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
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)
do_dance = i < (len(sigmas) - (2 + leap))
if not do_dance:
leap -= 1
do_dance = True
sigma_next = sigmas[i + (1 + leap)] if do_dance else sigmas[i + 1]
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
# Flow
downstep_ratio = None
alpha_ip1 = None
alpha_down = None
renoise_coeff = None
if flow:
# If/for flow model
downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta
sigma_down = sigmas[i+1] * downstep_ratio
alpha_ip1 = 1 - sigmas[i+1]
alpha_down = 1 - sigma_down
renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5
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], denoised)
# Euler method
dt = sigma_next - sigmas[i]
x_2 = x + d * dt
if do_dance:
reverse_denoised = model(x_2, sigma_next * s_in, **extra_args)
_, r_sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
r_d = to_d(x_2, sigma_next, reverse_denoised)
r_dt = sigma_down - sigma_next
x_2 = x + d * dt + r_d * r_dt
if sigmas[i + 1] > 0 and not flow:
x_2 = x_2 + noise_sampler(sigma_next, sigmas[i+1]) * s_noise * sigma_up
elif flow:
x_2 = (alpha_ip1/alpha_down) * x_2 + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
x = x_2
return x
@torch.no_grad()
def sample_leaping_euler(model, x, sigmas, extra_args=None, callback=None, disable=None, leap=1, eta=1., s_noise=1., noise_sampler_type="gaussian", noise_sampler=None):
if len(sigmas) <= 1:
return x
flow = False
if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):
flow = True
noise_sampler, extra_args = check_set_immiscible(x, noise_sampler_type, extra_args)
return sampler_leaping_euler(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, leap=leap, 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), flow=flow)
# Add your personal samplers below here, just for formatting purposes ;3
@@ -1565,6 +2127,10 @@ extra_samplers = {
"sens": sample_sens,
"ipndm_vapp": sample_ipndm_vapp,
"SHIDS": sample_SHIDS,
"dpmpp_2m_sde_ema": sample_dpmpp_2m_sde_ema,
"biscope": sample_biscope,
"euler_g": sample_euler_g,
"leaping_euler": sample_leaping_euler,
}
discard_penultimate_sigma_samplers = set((
+234 -7
View File
@@ -238,6 +238,8 @@ class SamplerSHIDS:
"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": 16, "min": 1, "max": 100, "step":1}),
"eta_order": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
"solver_method": (["weighted_projection", "qr_decomposition", "svd_lowrank", "svd"], {"default": "weighted_projection"}),
}
}
RETURN_TYPES = ("SAMPLER",)
@@ -245,18 +247,149 @@ class SamplerSHIDS:
FUNCTION = "get_sampler"
def get_sampler(self, noise_sampler_type, eta, s_noise, order):
sampler = comfy.samplers.ksampler("SHIDS", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "order": order})
def get_sampler(self, noise_sampler_type, eta, s_noise, order, eta_order, solver_method):
sampler = comfy.samplers.ksampler("SHIDS", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "order": order, "eta_order": eta_order, "solver_method": solver_method})
return (sampler, )
# EMA DPM++ 2M SDE (Compass Optimizer-like implementation)
class SamplerDPMPP_2M_SDE_EMA:
@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}),
"amp_fac": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100.0, "step":0.01}),
"beta1": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 0.999, "step":0.01}),
"beta2": ("FLOAT", {"default": 0.95, "min": 0.0, "max": 0.999, "step":0.01}),
"weight_decay": ("FLOAT", {"default": 0.1, "min": 0.0, "max": 100.0, "step":0.01}),
"centralization": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step":0.01}),
"normalization": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.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, amp_fac, beta1, beta2, weight_decay, centralization, normalization):
sampler = comfy.samplers.ksampler("dpmpp_2m_sde_ema", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "amp_fac": amp_fac, "beta1": beta1, "beta2": beta2, "weight_decay": weight_decay, "centralization": centralization, "normalization": normalization})
return (sampler, )
class SamplerEuler_3EMA:
@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}),
"amp_fac": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100, "step":0.01}),
"smoothing_fac": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 0.999, "step":0.01}),
"ema_fac": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 0.999, "step":0.01}),
"beta": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 0.999, "step":0.01}),
}
}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling/samplers"
FUNCTION = "get_sampler"
def get_sampler(self, noise_sampler_type, eta, s_noise, amp_fac, smoothing_fac, ema_fac, beta):
sampler = comfy.samplers.ksampler("euler_3ema", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "amp_fac": amp_fac, "smoothing_fac": smoothing_fac, "ema_fac": ema_fac, "beta": beta})
return (sampler, )
class SamplerScope:
@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}),
"amp_fac": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100, "step":0.01}),
"smoothing_fac": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 0.999, "step":0.01}),
"ema_fac": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 0.999, "step":0.01}),
}
}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling/samplers"
FUNCTION = "get_sampler"
def get_sampler(self, noise_sampler_type, eta, s_noise, amp_fac, smoothing_fac, ema_fac):
sampler = comfy.samplers.ksampler("scope", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "amp_fac": amp_fac, "smoothing_fac": smoothing_fac, "ema_fac": ema_fac})
return (sampler, )
class SamplerBiScope:
@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}),
"amp_fac": ("FLOAT", {"default": 2.0, "min": -100.0, "max": 100, "step":0.01}),
"local_smoothing_fac": ("INT", {"default": 4, "min": 1, "max": 100, "step":1}),
"smoothing_fac": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 0.999, "step":0.01}),
"ema_fac": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 0.999, "step":0.01}),
}
}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling/samplers"
FUNCTION = "get_sampler"
def get_sampler(self, noise_sampler_type, eta, s_noise, amp_fac, local_smoothing_fac, smoothing_fac, ema_fac):
sampler = comfy.samplers.ksampler("biscope", {"noise_sampler_type": noise_sampler_type, "eta": eta, "s_noise": s_noise, "amp_fac": amp_fac, "local_smoothing_fac": local_smoothing_fac, "smoothing_fac": smoothing_fac, "ema_fac": ema_fac})
return (sampler, )
class SamplerEuler_G:
@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}),
"g_eta": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step":0.01}),
"sigma": ("FLOAT", {"default": 1, "min": 0.01, "max": 100.0, "step":0.01}),
"order": ("INT", {"default": 3, "min": 3, "max": 100, "step":1}),
"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, g_eta, sigma, order, s_noise):
sampler = comfy.samplers.ksampler("euler_g", {"noise_sampler_type": noise_sampler_type, "eta": eta, "g_eta": g_eta, "sigma": sigma, "order": order, "s_noise": s_noise})
return (sampler, )
class SamplerLeaping_Euler:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"noise_sampler_type": (get_noise_sampler_names(default="gaussian"), ),
"leap": ("INT", {"default": 1, "min": 1, "max": 8, "step":1}),
"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, leap, eta, s_noise):
sampler = comfy.samplers.ksampler("leaping_euler", {"noise_sampler_type": noise_sampler_type, "leap": leap, "eta": eta, "s_noise": s_noise})
return (sampler, )
### Noise
class Noise_ImmiscibleNoise:
def __init__(self, noise_type, seed, image_scaling, latent_image):
def __init__(self, noise_type, seed, image_scaling, latent_image, n_latents = 1024):
self.noise_type = noise_type
self.seed = seed
self.image_scaling = image_scaling
self.latent_image = latent_image
self.n_latents = n_latents
def generate_noise(self, input_latent):
latent_image = input_latent["samples"]
@@ -264,7 +397,7 @@ class Noise_ImmiscibleNoise:
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 = make_immiscible(noise_func=self.noise_type, immiscible_latents=self.n_latents)(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")
@@ -273,7 +406,7 @@ class Noise_ImmiscibleNoise:
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 = make_immiscible(noise_func=self.noise_type, immiscible_latents=self.n_latents)(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)
@@ -288,6 +421,7 @@ class ImmiscibleNoise(DisableNoise):
return {"required":
{
"noise_type": (get_immiscible_noise_sampler_names(), ),
"n_latents": ("INT", {"default": 1024, "min": 1, "max": 16384, "step":1}),
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
},
"optional":
@@ -297,8 +431,8 @@ class ImmiscibleNoise(DisableNoise):
}
}
def get_noise(self, noise_type, noise_seed, image_scaling, latent_image):
return (Noise_ImmiscibleNoise(noise_type, noise_seed, image_scaling, latent_image),)
def get_noise(self, noise_type, n_latents, noise_seed, image_scaling, latent_image):
return (Noise_ImmiscibleNoise(noise_type, noise_seed, image_scaling, latent_image, n_latents),)
### Schedulers
from .extra_samplers import get_sigmas_simple_exponential
@@ -1057,4 +1191,97 @@ class MegaCFGGuider:
if perphist != 0:
guider.set_perphist_params(perphist)
m.set_model_sampler_post_cfg_function(guider.post_cfg_perphist)
return (guider,)
class Guider_APG(comfy.samplers.CFGGuider):
class MomentumBuffer:
def __init__(self, momentum: float):
self.momentum = momentum
self.running_average = 0
def reset(self):
self.running_average = 0
def update(self, update_value: torch.Tensor):
new_average = self.momentum * self.running_average
self.running_average = update_value + new_average
def project(self, v0: torch.Tensor, v1: torch.Tensor):
dtype = v0.dtype
v0, v1 = v0.double(), v1.double()
v1 = torch.nn.functional.normalize(v1, dim=[-3, -2, -1])
v0_parallel = (v0 * v1).sum(dim=[-3, -2, -1], keepdim=True) * v1
v0_orthogonal = v0 - v0_parallel
return v0_parallel.to(dtype), v0_orthogonal.to(dtype)
def set_cfg(self, cfg_scale, apg_scale, eta, norm_threshold, momentum_buffer):
self.cfg_scale = cfg_scale
self.apg_scale = apg_scale
self.eta = eta
self.norm_threshold = norm_threshold
self.momentum_buffer = momentum_buffer
self.curr_timestep = 1.
def set_conds(self, positive, negative):
self.inner_set_conds({"positive": positive, "negative": negative})
def normalized_guidance(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor, guidance_scale: float, cfg: torch.Tensor = None, momentum_buffer: MomentumBuffer = None, eta: float = 1.0, norm_threshold: float = 0.0):
diff = pred_cond - pred_uncond
if momentum_buffer is not None:
momentum_buffer.update(diff)
diff = momentum_buffer.running_average
if norm_threshold > 0:
ones = torch.ones_like(diff)
diff_norm = diff.norm(p=2, dim=[-3, -2, -1], keepdim=True)
scale_factor = torch.minimum(ones, norm_threshold / diff_norm)
diff = diff * scale_factor
diff_parallel, diff_orthogonal = self.project(diff, pred_cond)
normalized_update = diff_orthogonal + eta * diff_parallel
if cfg is None:
cfg = pred_cond
pred_guided = cfg + (guidance_scale - 1) * normalized_update
return pred_guided
def predict_noise(self, x, timestep, model_options={}, seed=None):
negative = self.conds.get("negative", None)
positive_cond = self.conds.get("positive", None)
# Weird way to workaround momentum buffer sticking from run to run, this should automatically reset it in a majority of cases.
if timestep > (self.curr_timestep - 1e-6):
self.momentum_buffer.reset()
self.curr_timestep = timestep
out = comfy.samplers.calc_cond_batch(self.inner_model, [negative, positive_cond], x, timestep, model_options)
cfg = comfy.samplers.cfg_function(self.inner_model, out[1], out[0], self.cfg_scale, x, timestep, model_options=model_options, cond=positive_cond, uncond=negative)
apg = self.normalized_guidance(out[1], out[0], self.apg_scale, cfg, self.momentum_buffer, self.eta, self.norm_threshold)
return apg
class APGGuider:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"model": ("MODEL",),
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"apg_scale": ("FLOAT", {"default": 4.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"norm_threshold": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"momentum": ("FLOAT", {"default": -0.5, "min": -100.0, "max": 100.0, "step":0.1, "round": 0.01, "lazy": False}),
}
}
RETURN_TYPES = ("GUIDER",)
FUNCTION = "get_guider"
CATEGORY = "sampling/custom_sampling/guiders"
def get_guider(self, model, positive, negative, cfg_scale, apg_scale, eta, norm_threshold, momentum):
guider = Guider_APG(model)
guider.set_conds(positive, negative) # Conds
momentum_buffer = guider.MomentumBuffer(momentum)
guider.set_cfg(cfg_scale, apg_scale, eta, norm_threshold, momentum_buffer) # Strengths
return (guider,)