Update clybGuidance, clybSchedulers, and clybSamplers
clybGuidance adds new feature(s) clybSchedulers fixes a bug clybSamplers adds new functionality to BDF sampler
This commit is contained in:
@@ -4,7 +4,6 @@ from . import clyb_Samplers
|
||||
from . import clyb_Schedulers
|
||||
|
||||
clyb_Samplers.add_samplers()
|
||||
clyb_Schedulers.add_schedulers()
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"ChromaNAG": chroma_NAG.ChromaNAG,
|
||||
|
||||
+215
-156
@@ -1,4 +1,5 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import math
|
||||
import comfy
|
||||
import re
|
||||
@@ -26,16 +27,20 @@ class ClybGuidance:
|
||||
"var_rescale": ("BOOLEAN", {"default": False, "tooltip":"Whether we use torch.var (true) or torch.std (false) for rescaling."}),
|
||||
"scale_up_ratio": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip":"Weight of: Initiating CFG at guidance scale 1, increasing to your guidance scale in the middle of diffusion, and lower back to 1."}),
|
||||
"scale_up_shift": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 10.0, "step": 0.01, "tooltip":"Whether to shift to your CFG scale later (lower than 1.0) or earlier (higher than 1.0) in the schedule."}),
|
||||
"atan2sin_ratio": ("FLOAT", {"default": 0.0, "min": -100.0, "max": 100.0, "step": 0.01, "tooltip":"Applies standard deviation renormalization of CFG to cond at this rate."}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "patch"
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
def patch(self, model, eta, norm_threshold, momentum, momentum_beta, momentum_renorm, scalar_projection, scalar_logsumexp, rescale_phi, var_rescale, scale_up_ratio, scale_up_shift):
|
||||
def patch(self, model, eta, norm_threshold, momentum, momentum_beta, momentum_renorm, scalar_projection, scalar_logsumexp, rescale_phi, var_rescale, scale_up_ratio, scale_up_shift, atan2sin_ratio):
|
||||
running_avg = 0
|
||||
prev_sigma = None
|
||||
|
||||
m = model.clone()
|
||||
model_sampling = m.get_model_object("model_sampling")
|
||||
|
||||
def pre_cfg_function(args):
|
||||
nonlocal running_avg, prev_sigma
|
||||
|
||||
@@ -45,7 +50,6 @@ class ClybGuidance:
|
||||
uncond = args["conds_out"][1]
|
||||
sigma = args["sigma"][0]
|
||||
cond_scale = args["cond_scale"]
|
||||
model_sampling = model.get_model_object("model_sampling")
|
||||
flow = False
|
||||
if isinstance(model_sampling, comfy.model_sampling.CONST):
|
||||
flow = True
|
||||
@@ -65,20 +69,22 @@ class ClybGuidance:
|
||||
else:
|
||||
timestep_ratio = model_sampling.timestep(args["timestep"]).float() / float(num_timesteps - 1) # Ratio scaling from 0 to 1 as diffusion goes on.
|
||||
|
||||
guidance_multiplier = torch.lerp(torch.ones_like(args["timestep"]), torch.sin((1. - timestep_ratio**scale_up_shift) * math.pi), weight=scale_up_ratio) # Lerp from static 1.0 scale to bell-curve (sine wave) scale
|
||||
cfg_scalar = 1 / cond_scale + guidance_multiplier * ((cond_scale - 1) / cond_scale) # The guidance scale ought to be at least 1 (guidance is multiplied by cond scale, so ensure a min of 1/cond_scale)
|
||||
if atan2sin_ratio != 0:
|
||||
uncond = uncond.lerp(cond.atan().sin_().div_(uncond.atan().cos_()), weight=atan2sin_ratio)
|
||||
|
||||
if scalar_projection:
|
||||
cond_flat, uncond_flat = cond.view(cond.shape[0], -1), uncond.view(uncond.shape[0], -1)
|
||||
|
||||
dot_product = torch.logsumexp(cond * uncond, dim=1, keepdim=True) if scalar_logsumexp else torch.sum(cond_flat * uncond_flat, dim=1, keepdim=True)
|
||||
dot_product = torch.logsumexp(cond_flat * uncond_flat, dim=1, keepdim=True) if scalar_logsumexp else torch.sum(cond_flat * uncond_flat, dim=1, keepdim=True)
|
||||
|
||||
squared_norm = torch.logsumexp(uncond**2, dim=1, keepdim=True) if scalar_logsumexp else torch.sum(uncond_flat**2, dim=1, keepdim=True)
|
||||
squared_norm = torch.logsumexp(uncond_flat**2, dim=1, keepdim=True) if scalar_logsumexp else torch.sum(uncond_flat**2, dim=1, keepdim=True)
|
||||
|
||||
alpha = dot_product / squared_norm.clamp_min(1e-7)
|
||||
|
||||
uncond = uncond * alpha
|
||||
|
||||
guidance_multiplier = torch.lerp(torch.ones_like(args["timestep"]), torch.sin((1. - timestep_ratio**scale_up_shift) * math.pi), weight=scale_up_ratio) # Lerp from static 1.0 scale to bell-curve (sine wave) scale
|
||||
cfg_scalar = 1 / cond_scale + guidance_multiplier * ((cond_scale - 1) / cond_scale) # The guidance scale ought to be at least 1 (guidance is multiplied by cond scale, so ensure a min of 1/cond_scale)
|
||||
guidance = ((cond - uncond) * cfg_scalar) if scale_up_ratio != 0 else (cond - uncond) # Guidance is equivalent to (uncond -> cond)
|
||||
|
||||
if momentum != 0:
|
||||
@@ -86,22 +92,19 @@ class ClybGuidance:
|
||||
running_avg = guidance
|
||||
else:
|
||||
running_avg = running_avg.lerp(guidance, weight=1. - momentum_beta)#running_avg.lerp(guidance, weight=1. - abs(momentum))# Update running average
|
||||
#running_avg = running_avg * (guidance.pow(2).mean().sqrt_() / running_avg.pow(2).mean().sqrt_().clamp_min_(1e-8)) # Normalize running average to guidance
|
||||
momentumized_guidance = guidance.add(running_avg, alpha=momentum)
|
||||
momentumized_guidance_flat, guidance_flat = momentumized_guidance.view(momentumized_guidance.shape[0], -1), guidance.view(guidance.shape[0], -1)
|
||||
guidance = momentumized_guidance.lerp(momentumized_guidance * (guidance_flat.norm(1, dim=1, keepdim=True) / momentumized_guidance_flat.norm(1, dim=1, keepdim=True).clamp_min(1e-7)), weight=momentum_renorm)
|
||||
|
||||
if norm_threshold > 0:
|
||||
guidance_norm = guidance.view(guidance.shape[0], -1).norm(p=2, dim=1, keepdim=True)
|
||||
scale = torch.minimum(
|
||||
torch.ones_like(guidance_norm),
|
||||
norm_threshold / guidance_norm
|
||||
)
|
||||
guidance = guidance * scale
|
||||
guidance = momentumized_guidance.lerp(momentumized_guidance * (guidance_flat.norm(2, dim=1, keepdim=True) / momentumized_guidance_flat.norm(2, dim=1, keepdim=True).clamp_min(1e-7)), weight=momentum_renorm)
|
||||
|
||||
guidance_parallel, guidance_orthogonal = project(guidance, cond)
|
||||
modified_guidance = guidance_orthogonal + eta * guidance_parallel
|
||||
|
||||
if norm_threshold > 0:
|
||||
cond_norm = cond.norm(p=2, dim=tuple(range(1, len(cond.shape))), keepdim=True) * norm_threshold
|
||||
guidance_norm = (uncond + modified_guidance * cond_scale).norm(p=2, dim=tuple(range(1, len(cond.shape))), keepdim=True)
|
||||
if guidance_norm >= cond_norm:
|
||||
modified_guidance = modified_guidance * (cond_norm / guidance_norm)
|
||||
|
||||
modified_cond = (uncond + modified_guidance)
|
||||
if rescale_phi != 0:
|
||||
# Formulate CFG
|
||||
@@ -122,146 +125,202 @@ class ClybGuidance:
|
||||
|
||||
return [modified_cond, uncond] + args["conds_out"][2:]
|
||||
|
||||
m = model.clone()
|
||||
"""
|
||||
TODO: Rework these and add a selector
|
||||
def magnitude_guidance(args):
|
||||
cond = args["cond_denoised"]
|
||||
uncond = args["uncond_denoised"]
|
||||
cond_scale = args["cond_scale"]
|
||||
x = args['input']
|
||||
out = args["denoised"]
|
||||
|
||||
# Flatten cond and uncond, ensure cond scale is positive, utilize double precision
|
||||
device = cond.device
|
||||
b, c, h, w = cond.shape
|
||||
|
||||
# 1. Create the 2D Hann window kernel for convolution
|
||||
hann_1d = torch.signal.windows.hann(63, device=device)
|
||||
#hann_2d = torch.outer(hann_1d, hann_1d)
|
||||
# Normalize the kernel so that the sum of its elements is 1
|
||||
hann_1d /= hann_1d.sum()
|
||||
|
||||
# Reshape kernel for depthwise convolution: (out_channels, in_channels/groups, kH, kW)
|
||||
# We use groups=c to apply the same 2D filter to each channel independently.
|
||||
kernel = hann_1d.unsqueeze(0).unsqueeze(0)#.repeat(cond.shape[1], 1, 1, 1)
|
||||
|
||||
# 2. Calculate the local average magnitude of the `cond` tensor
|
||||
# We use the absolute value to measure magnitude, not the raw value.
|
||||
# 'same' padding ensures the output has the same HxW dimensions as the input.
|
||||
view_shape = (cond.shape[0], -1)
|
||||
cond_flat = cond.view(view_shape)
|
||||
uncond_flat = uncond.view(view_shape)
|
||||
local_avg_magnitude = F.conv1d((cond_flat - uncond_flat), kernel, padding='same')
|
||||
|
||||
# 3. Normalize the magnitude map for each image in the batch to the [0, 1] range
|
||||
# This makes the `strength` parameter behave consistently across different images.
|
||||
batch_mins = torch.min(local_avg_magnitude.view(view_shape), dim=-1)[0]#.view(cond.shape[0], 1, 1, 1)
|
||||
batch_maxs = torch.max(local_avg_magnitude.view(view_shape), dim=-1)[0]#.view(cond.shape[0], 1, 1, 1)
|
||||
|
||||
normalized_magnitude = (local_avg_magnitude - batch_mins) / (batch_maxs - batch_mins).clamp_min_(1e-16)
|
||||
|
||||
# 4. Create the local scale multiplier
|
||||
# The dampening is proportional to the normalized local magnitude and the `strength` param.
|
||||
# We clamp to ensure the multiplier stays within a reasonable [0, 1] range.
|
||||
dampening = torch.clamp(normalized_magnitude, 0, 1.0)
|
||||
|
||||
# Invert the dampening: high magnitude -> low multiplier, low magnitude -> high multiplier
|
||||
scale_multiplier = 1.0 - dampening
|
||||
|
||||
cond_normed = cond_flat / torch.linalg.norm(cond_flat,dim=1,keepdim=True)
|
||||
uncond_normed = uncond_flat / torch.linalg.norm(uncond_flat,dim=1,keepdim=True)
|
||||
dot_product = torch.sum(cond_normed*uncond_normed,dim=1,keepdim=True)
|
||||
# The final local scale is the base scale modulated by our multiplier
|
||||
local_scale = cond_scale * scale_multiplier * dot_product
|
||||
|
||||
#print(local_scale)
|
||||
|
||||
# 5. Apply the standard CFG formula, but with our dynamic local scale
|
||||
# uncond + local_scale * (cond - uncond)
|
||||
guided_tensor = cond + (local_scale * (cond_flat - uncond_flat)).view(cond.shape)
|
||||
|
||||
output = out.lerp(guided_tensor.view(cond.shape).to(cond.dtype), weight=atan2sin_ratio)
|
||||
|
||||
if norm_threshold > 0:
|
||||
guidance_norm = output.norm(p=2, dim=1, keepdim=True)
|
||||
output = torch.where(
|
||||
guidance_norm > norm_threshold,
|
||||
output * (norm_threshold / guidance_norm),
|
||||
output
|
||||
)
|
||||
return output
|
||||
|
||||
def frequency_guidance(args):
|
||||
cond = args["cond_denoised"]
|
||||
uncond = args["uncond_denoised"]
|
||||
cond_scale = args["cond_scale"]
|
||||
x = args['input']
|
||||
out = args["denoised"]
|
||||
|
||||
# 1. Move to Frequency domain using 2D Fast Fourier Transform
|
||||
# We use norm='ortho' to ensure the transform is unitary and preserves energy.
|
||||
fft_cond = torch.fft.fftshift(torch.fft.fftn(cond.to(torch.float64), norm='ortho'))
|
||||
fft_uncond = torch.fft.fftshift(torch.fft.fftn(uncond.to(torch.float64), norm='ortho'))
|
||||
fft_out = torch.fft.fftshift(torch.fft.fftn(out.to(torch.float64), norm='ortho'))
|
||||
|
||||
# 1. Create the 2D Hann window kernel for convolution
|
||||
hann_1d = torch.signal.windows.hann(5, device=cond.device)
|
||||
#hann_2d = torch.outer(hann_1d, hann_1d)
|
||||
# Normalize the kernel so that the sum of its elements is 1
|
||||
hann_1d /= hann_1d.sum()
|
||||
|
||||
kernel = hann_1d.unsqueeze(0).unsqueeze(0)#.repeat(cond.shape[1], 1, 1, 1)
|
||||
|
||||
# 2. Calculate the local average magnitude of the `cond` tensor
|
||||
# We use the absolute value to measure magnitude, not the raw value.
|
||||
# 'same' padding ensures the output has the same HxW dimensions as the input.
|
||||
view_shape = (fft_cond.shape[0], -1)
|
||||
|
||||
fft_cond_flat = fft_cond.view(view_shape)
|
||||
fft_uncond_flat = fft_uncond.view(view_shape)
|
||||
|
||||
fft_cond_real = fft_cond_flat.real
|
||||
fft_uncond_real = fft_uncond_flat.real
|
||||
#guidance_direction = (cond - uncond)
|
||||
local_avg_magnitude = F.conv1d(fft_cond_real, kernel.to(torch.float64), padding='same')
|
||||
|
||||
# 3. Normalize the magnitude map for each image in the batch to the [0, 1] range
|
||||
# This makes the `strength` parameter behave consistently across different images.
|
||||
#view_shape = (cond.shape[0], -1)
|
||||
batch_mins = torch.min(local_avg_magnitude.view(view_shape), dim=-1)[0]#.view(cond.shape[0], 1, 1, 1)
|
||||
batch_maxs = torch.max(local_avg_magnitude.view(view_shape), dim=-1)[0]#.view(cond.shape[0], 1, 1, 1)
|
||||
|
||||
normalized_magnitude = (local_avg_magnitude - batch_mins) / (batch_maxs - batch_mins).clamp_min_(1e-16)
|
||||
|
||||
dampening = torch.clamp(normalized_magnitude, 0, 1.0)
|
||||
|
||||
# Invert the dampening: high magnitude -> low multiplier, low magnitude -> high multiplier
|
||||
scale_multiplier = 1.0 - dampening
|
||||
|
||||
# The final local scale is the base scale modulated by our multiplier and dot product
|
||||
cond_normed = fft_cond_real / torch.linalg.norm(fft_cond_real,dim=1,keepdim=True).clamp_min_(1e-16)
|
||||
uncond_normed = fft_uncond_real / torch.linalg.norm(fft_uncond_real,dim=1,keepdim=True).clamp_min_(1e-16)
|
||||
dot_product = torch.sum(cond_normed*uncond_normed,dim=1,keepdim=True)
|
||||
|
||||
local_scale = cond_scale * scale_multiplier * dot_product
|
||||
|
||||
print(local_scale)
|
||||
|
||||
guided_tensor = fft_cond + (local_scale.to(torch.cdouble) * (fft_cond_flat - fft_uncond_flat)).view(fft_cond.shape)
|
||||
|
||||
guided_tensor = torch.fft.ifftshift(guided_tensor)
|
||||
guided_tensor = torch.fft.ifftn(guided_tensor, norm='ortho').real
|
||||
|
||||
output = out.lerp(guided_tensor.view(cond.shape).to(cond.dtype), weight=atan2sin_ratio)
|
||||
|
||||
if norm_threshold > 0:
|
||||
guidance_norm = output.norm(p=2, dim=1, keepdim=True)
|
||||
output = torch.where(
|
||||
guidance_norm > norm_threshold,
|
||||
output * (norm_threshold / guidance_norm),
|
||||
output
|
||||
)
|
||||
return output
|
||||
|
||||
def kron_guidance(args):
|
||||
cond = args["cond_denoised"]
|
||||
uncond = args["uncond_denoised"]
|
||||
cond_scale = args["cond_scale"]
|
||||
x = args['input']
|
||||
out = args["denoised"]
|
||||
|
||||
guidance = cond - uncond
|
||||
|
||||
device = guidance.device
|
||||
|
||||
# 2. To Frequency Domain
|
||||
# Apply 2D FFT to the spatial dimensions (H, W)
|
||||
guidance_fft = torch.fft.fftn(guidance)
|
||||
|
||||
# Shift the zero-frequency component to the center for easier mask creation
|
||||
guidance_fft_shifted = torch.fft.fftshift(guidance_fft).reshape(guidance_fft.shape[0], -1)
|
||||
|
||||
U, S, Vh = torch.linalg.svd(guidance_fft_shifted, full_matrices=False)
|
||||
|
||||
# 3. Truncate the SVD components to the desired rank
|
||||
U_k = U[:, :1]
|
||||
S_k = S[:1]
|
||||
Vh_k = Vh[:1, :]
|
||||
#S_k_sqrt = S_k.sqrt()
|
||||
|
||||
# 4. Apply the Mask
|
||||
# Multiply the shifted FFT of the guidance by the scale mask
|
||||
scaled_guidance_fft_shifted = (U_k @ (torch.diag(S_k).to(torch.cfloat) @ Vh_k)).reshape(guidance_fft.shape)
|
||||
|
||||
# 5. Back to Latent Domain
|
||||
# Inverse shift to move the zero-frequency back to the corner
|
||||
scaled_guidance_fft = torch.fft.ifftshift(scaled_guidance_fft_shifted)
|
||||
|
||||
# Inverse 2D FFT to get back to the spatial (latent) domain
|
||||
# The result of ifft2 will be complex; we take the real part. The imaginary
|
||||
# part should be negligible for real inputs.
|
||||
modified_guidance = torch.fft.ifftn(scaled_guidance_fft).real
|
||||
|
||||
# 6. Apply Guidance
|
||||
# Add the modified guidance to the unconditional prediction
|
||||
output = out + (guidance - modified_guidance) * atan2sin_ratio
|
||||
|
||||
#output = out.lerp(guided_tensor.view(cond.shape).to(cond.dtype), weight=atan2sin_ratio)
|
||||
|
||||
if norm_threshold > 0:
|
||||
guidance_norm = output.norm(p=2, dim=1, keepdim=True)
|
||||
output = torch.where(
|
||||
guidance_norm > norm_threshold,
|
||||
output * (norm_threshold / guidance_norm),
|
||||
output
|
||||
)
|
||||
return output
|
||||
"""
|
||||
|
||||
m.set_model_sampler_pre_cfg_function(pre_cfg_function)
|
||||
return (m,)
|
||||
|
||||
#args = {"conds":conds, "conds_out": out, "cond_scale": self.cfg, "timestep": timestep,
|
||||
# "input": x, "sigma": timestep, "model": self.inner_model, "model_options": model_options}
|
||||
#out = fn(args)
|
||||
|
||||
def create_number_range(range_str: str) -> list[int] | None:
|
||||
"""
|
||||
Creates a list of numbers from a string in "start-end" format.
|
||||
|
||||
Args:
|
||||
range_str: The input string (e.g., "1-5", "10-20", " 5 - 10 ").
|
||||
|
||||
Returns:
|
||||
A list of integers representing the range (inclusive),
|
||||
or None if the string format is invalid or start > end.
|
||||
"""
|
||||
# The core 're' module line to define and apply the pattern:
|
||||
# 1. r"..." denotes a raw string to avoid issues with backslashes.
|
||||
# 2. (\d+) is a capturing group for one or more digits (the start number).
|
||||
# 3. \s* matches zero or more whitespace characters (optional spaces around the hyphen).
|
||||
# 4. - matches the literal hyphen.
|
||||
# 5. \s* matches zero or more whitespace characters again.
|
||||
# 6. (\d+) is another capturing group for the end number.
|
||||
# re.match() attempts to match the pattern from the beginning of the string.
|
||||
match = re.match(r"(\d+)\s*-\s*(\d+)", range_str.strip())
|
||||
|
||||
if match:
|
||||
# Extract the captured groups and convert them to integers
|
||||
start_str, end_str = match.groups()
|
||||
start = int(start_str)
|
||||
end = int(end_str)
|
||||
|
||||
# Ensure the start is not greater than the end for a valid range
|
||||
if start <= end:
|
||||
return list(range(start, end + 1))
|
||||
else:
|
||||
# Handle cases like "5-1" if they should not produce a range
|
||||
print(f"Warning: Start ({start}) is greater than end ({end}) for '{range_str}'")
|
||||
return None # Or [] if an empty list is preferred for invalid ranges
|
||||
else:
|
||||
return None # String does not match the expected format
|
||||
|
||||
class ClybLayerGuidanceDiT:
|
||||
'''
|
||||
Enhance guidance towards detailed dtructure by having another set of CFG negative with skipped layers.
|
||||
Inspired by Perturbed Attention Guidance (https://arxiv.org/abs/2403.17377)
|
||||
Original experimental implementation for SD3 by Dango233@StabilityAI.
|
||||
'''
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"model": ("MODEL", ),
|
||||
"double_layers": ("STRING", {"default": "7, 8, 9", "multiline": False}),
|
||||
"single_layers": ("STRING", {"default": "7, 8, 9", "multiline": False}),
|
||||
"scale": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 10.0, "step": 0.1}),
|
||||
"start_percent": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"rescaling_scale": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"attn_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "skip_guidance"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
DESCRIPTION = "Generic version of ClybLayerGuidance node that can be used on every DiT model."
|
||||
|
||||
CATEGORY = "advanced/guidance"
|
||||
|
||||
def skip_guidance(self, model, scale, start_percent, end_percent, double_layers="", single_layers="", rescaling_scale=0, attn_scale=1.0):
|
||||
# check if layer is comma separated integers
|
||||
def skip(args, extra_args):
|
||||
print(f"ARGS: {args.items()}", "\n\n\n", f"EXTRA_ARGS: {extra_args.items()}", "\n\n\n")
|
||||
for x in args:
|
||||
if 'vec' in x:
|
||||
for y in x:
|
||||
if 'scale' in y:
|
||||
args[x][y] = args[x][y] * attn_scale
|
||||
#args[x] = y * attn_scale
|
||||
#print(x, y)
|
||||
return args
|
||||
#for x, y in args.items():
|
||||
# if 'img' in x:
|
||||
# return x
|
||||
|
||||
model_sampling = model.get_model_object("model_sampling")
|
||||
sigma_start = model_sampling.percent_to_sigma(start_percent)
|
||||
sigma_end = model_sampling.percent_to_sigma(end_percent)
|
||||
|
||||
#double_layers = re.findall(r'\d+', double_layers)
|
||||
#double_layers = [int(i) for i in double_layers]
|
||||
double_layers = create_number_range(double_layers)
|
||||
|
||||
#single_layers = re.findall(r'\d+', single_layers)
|
||||
#single_layers = [int(i) for i in single_layers]
|
||||
single_layers = create_number_range(single_layers)
|
||||
|
||||
if len(double_layers) == 0 and len(single_layers) == 0:
|
||||
return (model, )
|
||||
|
||||
def post_cfg_function(args):
|
||||
model = args["model"]
|
||||
cond_pred = args["cond_denoised"]
|
||||
cond = args["cond"]
|
||||
cfg_result = args["denoised"]
|
||||
sigma = args["sigma"]
|
||||
x = args["input"]
|
||||
model_options = args["model_options"].copy()
|
||||
#print(model_options)
|
||||
for layer in double_layers:
|
||||
model_options = comfy.model_patcher.set_model_options_patch_replace(model_options, skip, "dit", "double_block", layer)
|
||||
|
||||
for layer in single_layers:
|
||||
model_options = comfy.model_patcher.set_model_options_patch_replace(model_options, skip, "dit", "single_block", layer)
|
||||
|
||||
model_sampling.percent_to_sigma(start_percent)
|
||||
|
||||
sigma_ = sigma[0].item()
|
||||
if scale > 0 and sigma_ >= sigma_end and sigma_ <= sigma_start:
|
||||
(slg,) = comfy.samplers.calc_cond_batch(model, [cond], x, sigma, model_options)
|
||||
cfg_result = cfg_result + (cond_pred - slg) * scale
|
||||
if rescaling_scale != 0:
|
||||
factor = cond_pred.std() / cfg_result.std()
|
||||
factor = rescaling_scale * factor + (1 - rescaling_scale)
|
||||
cfg_result *= factor
|
||||
|
||||
return cfg_result
|
||||
|
||||
m = model.clone()
|
||||
m.set_model_sampler_post_cfg_function(post_cfg_function)
|
||||
|
||||
return (m,)
|
||||
|
||||
#NODE_CLASS_MAPPINGS = {
|
||||
# "ClybGuidance": ClybGuidance,
|
||||
# "ClybLayerGuidanceDiT": ClybLayerGuidanceDiT,
|
||||
#}
|
||||
|
||||
#NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
# "ClybGuidance": "ClybGuidance",
|
||||
# "ClybLayerGuidanceDiT": "ClybLayerGuidanceDiT",
|
||||
#}
|
||||
|
||||
+23
-10
@@ -7,7 +7,7 @@ from comfy.k_diffusion.sampling import default_noise_sampler
|
||||
import comfy.samplers
|
||||
|
||||
@torch.no_grad()
|
||||
def sampler_clyb_bdf(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, flow=False):
|
||||
def sampler_clyb_bdf(model, x, sigmas, extra_args=None, callback=None, disable=None, scalar="atan2sin+projection", eta=1., s_noise=1., noise_sampler=None, flow=False):
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
seed = extra_args.get("seed", None)
|
||||
noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler
|
||||
@@ -34,14 +34,25 @@ def sampler_clyb_bdf(model, x, sigmas, extra_args=None, callback=None, disable=N
|
||||
alpha_down = 1 - sigma_down
|
||||
renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5
|
||||
|
||||
first_denoised = prev_denoised if prev_denoised is not None else model(x, sigmas[i] * s_in, **extra_args)
|
||||
first_denoised = prev_denoised if (prev_denoised is not None and sigma_down > 0) else model(x, sigmas[i] * s_in, **extra_args)
|
||||
|
||||
if sigma_down > 0:
|
||||
if sigma_down > 0 and i > 0:
|
||||
x_faux = first_denoised.lerp(x, weight=sigma_down/sigmas[i])
|
||||
denoised2 = model(x_faux, sigma_down * s_in, **extra_args)
|
||||
second_denoised = (first_denoised + denoised2) / 2
|
||||
scalar = (denoised2 * second_denoised) / (second_denoised.pow(2).clamp_min(1e-6))
|
||||
denoised_prime = second_denoised * scalar
|
||||
match scalar:
|
||||
case "projection":
|
||||
scaling = (denoised2 * second_denoised) / (second_denoised.pow(2).clamp_min(1e-7))
|
||||
denoised_prime = second_denoised * scaling
|
||||
case "atan2sin":
|
||||
denoised_prime = denoised2.atan().sin_().div_(second_denoised.atan().cos_())
|
||||
case "atan2sin+projection":
|
||||
denoised_prime = denoised2.atan().sin_().div_(second_denoised.atan().cos_())
|
||||
scaling = (denoised2 * denoised_prime) / (denoised_prime.pow(2).clamp_min(1e-7))
|
||||
denoised_prime = denoised_prime * scaling
|
||||
case _:
|
||||
scaling = (denoised2 * second_denoised) / (second_denoised.pow(2).clamp_min(1e-7))
|
||||
denoised_prime = second_denoised * scaling
|
||||
else:
|
||||
denoised_prime = first_denoised
|
||||
|
||||
@@ -61,11 +72,11 @@ def sampler_clyb_bdf(model, x, sigmas, extra_args=None, callback=None, disable=N
|
||||
return x
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_clyb_bdf(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
|
||||
def sample_clyb_bdf(model, x, sigmas, extra_args=None, callback=None, disable=None, scalar="atan2sin+projection", eta=1., s_noise=1., noise_sampler=None):
|
||||
flow = False
|
||||
if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):
|
||||
flow = True
|
||||
return sampler_clyb_bdf(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler, flow=flow)
|
||||
return sampler_clyb_bdf(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, scalar=scalar, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler, flow=flow)
|
||||
|
||||
# The following function adds the samplers during initialization, in __init__.py
|
||||
def add_samplers():
|
||||
@@ -95,8 +106,10 @@ discard_penultimate_sigma_samplers = set(())
|
||||
class SamplerClyb_BDF:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
NOISE_SAMPLER_NAMES=("projection", "atan2sin", "atan2sin+projection")
|
||||
return {"required":
|
||||
{"eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
{"scalar": (NOISE_SAMPLER_NAMES, {"default": NOISE_SAMPLER_NAMES[2]}),
|
||||
"eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
"s_noise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step":0.01}),
|
||||
}
|
||||
}
|
||||
@@ -105,6 +118,6 @@ class SamplerClyb_BDF:
|
||||
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(self, eta, s_noise):
|
||||
sampler = comfy.samplers.ksampler("clyb_bdf", {"eta": eta, "s_noise": s_noise})
|
||||
def get_sampler(self, scalar, eta, s_noise):
|
||||
sampler = comfy.samplers.ksampler("clyb_bdf", {"scalar": scalar, "eta": eta, "s_noise": s_noise})
|
||||
return (sampler, )
|
||||
+7
-15
@@ -1,5 +1,6 @@
|
||||
import numpy
|
||||
import torch
|
||||
from comfy.samplers import SchedulerHandler, SCHEDULER_HANDLERS, SCHEDULER_NAMES
|
||||
|
||||
def inverse_squared_scheduler(model_sampling, steps):
|
||||
total_timesteps = (len(model_sampling.sigmas) - 1)
|
||||
@@ -55,21 +56,12 @@ class PrintSigmas:
|
||||
print(sigmas)
|
||||
return (sigmas, )
|
||||
|
||||
def add_schedulers():
|
||||
from comfy.samplers import KSampler, k_diffusion_sampling
|
||||
added = 0
|
||||
for scheduler in extra_schedulers: #getattr(self, "sample_{}".format(extra_samplers))
|
||||
if scheduler not in KSampler.SCHEDULERS:
|
||||
try:
|
||||
idx = KSampler.SCHEDULERS.index("ddim_uniform") # Last item in the samplers list
|
||||
KSampler.SCHEDULERS.insert(idx+1, scheduler) # Add our custom samplers
|
||||
setattr(k_diffusion_sampling, "get_sigmas_{}".format(scheduler), extra_schedulers[scheduler])
|
||||
added += 1
|
||||
except ValueError as err:
|
||||
pass
|
||||
if added > 0:
|
||||
import importlib
|
||||
importlib.reload(k_diffusion_sampling)
|
||||
scheduler_name = "inverse_squared"
|
||||
if scheduler_name not in SCHEDULER_HANDLERS:
|
||||
scheduler_handler = SchedulerHandler(handler=inverse_squared_scheduler, use_ms=True)
|
||||
SCHEDULER_HANDLERS[scheduler_name] = scheduler_handler
|
||||
if scheduler_name not in SCHEDULER_NAMES:
|
||||
SCHEDULER_NAMES.append(scheduler_name)
|
||||
|
||||
extra_schedulers = {
|
||||
"inverse_squared": inverse_squared_scheduler,
|
||||
|
||||
Reference in New Issue
Block a user