Added ControlLora full support, completed T2IAdapter full support, fixed Latent Keyframe Group node not working properly when optional latents are not passed in, made keyframe print statements optional+better

This commit is contained in:
Jedrzej Kosinski
2023-11-26 08:38:12 -06:00
parent 1bd0764359
commit 1f4fcff967
4 changed files with 214 additions and 116 deletions
+152 -100
View File
@@ -4,7 +4,7 @@ import torch
import comfy.utils
import comfy.controlnet as comfy_cn
from comfy.controlnet import ControlNet, T2IAdapter, broadcast_image_to
from comfy.controlnet import ControlNet, ControlLora, T2IAdapter, broadcast_image_to
ControlNetWeightsType = list[float]
@@ -166,14 +166,11 @@ def control_merge_inject(self, control_input, control_output, control_prev, outp
return out
class ControlNetAdvanced(ControlNet):
def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None):
super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, device=device)
class AdvancedControlBase:
def __init__(self, timestep_keyframes: TimestepKeyframeGroup):
# initialize timestep_keyframes
self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup()
self.current_timestep_keyframe = self.timestep_keyframes.keyframes[0]
# initialize weights
self.weights = self.timestep_keyframes.keyframes[0].control_net_weights if self.timestep_keyframes.keyframes[0].control_net_weights else [1.0]*13
# mask for which parts of controlnet output to keep
self.mask_cond_hint_original = None
self.mask_cond_hint = None
@@ -188,73 +185,6 @@ class ControlNetAdvanced(ControlNet):
self.mask_cond_hint_original = mask_hint
return self
def get_control(self, x_noisy, t, cond, batched_number):
# need to reference t and batched_number later
self.t = t
self.batched_number = batched_number
# TODO: choose TimestepKeyframe based on t
# perform special version of get_control that supports sliding context and masks
return self.sliding_get_control(x_noisy, t, cond, batched_number)
def sliding_get_control(self, x_noisy: Tensor, t, cond, batched_number):
control_prev = None
if self.previous_controlnet is not None:
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
if self.timestep_range is not None:
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
if control_prev is not None:
return control_prev
else:
return None
output_dtype = x_noisy.dtype
# make cond_hint appropriate dimensions
# TODO: change this to not require cond_hint upscaling every step when self.sub_idxs are present
if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] * 8 != self.cond_hint.shape[2] or x_noisy.shape[3] * 8 != self.cond_hint.shape[3]:
if self.cond_hint is not None:
del self.cond_hint
self.cond_hint = None
# if self.cond_hint_original length matches real latent count, need to subdivide it
if self.cond_hint_original.size(0) == self.full_latent_length:
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(self.control_model.dtype).to(self.device)
else:
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(self.control_model.dtype).to(self.device)
if x_noisy.shape[0] != self.cond_hint.shape[0]:
self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number)
# make mask appropriate dimensions, if present
if self.mask_cond_hint_original is not None:
if self.sub_idxs is not None or self.mask_cond_hint is None or x_noisy.shape[2] * 8 != self.mask_cond_hint.shape[1] or x_noisy.shape[3] * 8 != self.mask_cond_hint.shape[2]:
if self.mask_cond_hint is not None:
del self.mask_cond_hint
self.mask_cond_hint = None
# TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM
# resize mask and match batch count
self.mask_cond_hint = prepare_mask_batch(self.mask_cond_hint_original, x_noisy.shape, multiplier=8)
actual_latent_length = x_noisy.shape[0] // batched_number
self.mask_cond_hint = comfy.utils.repeat_to_batch_size(self.mask_cond_hint, actual_latent_length if self.sub_idxs is None else self.full_latent_length)
if self.sub_idxs is not None:
self.mask_cond_hint = self.mask_cond_hint[self.sub_idxs]
# make cond_hint_mask length match x_noise
if x_noisy.shape[0] != self.mask_cond_hint.shape[0]:
self.mask_cond_hint = broadcast_image_to(self.mask_cond_hint, x_noisy.shape[0], batched_number)
self.mask_cond_hint = self.mask_cond_hint.to(self.control_model.dtype).to(self.device)
context = cond['c_crossattn']
# uses 'y' in new ComfyUI update
y = cond.get('y', None)
if y is None: # TODO: remove this in the future since no longer used by newest ComfyUI
y = cond.get('c_adm', None)
if y is not None:
y = y.to(self.control_model.dtype)
timestep = self.model_sampling_current.timestep(t)
x_noisy = self.model_sampling_current.calculate_input(t, x_noisy)
control = self.control_model(x=x_noisy.to(self.control_model.dtype), hint=self.cond_hint, timesteps=timestep.float(), context=context.to(self.control_model.dtype), y=y)
return self.control_merge(None, control, control_prev, output_dtype)
def apply_advanced_strengths_and_masks(self, x: Tensor, current_timestep_keyframe: TimestepKeyframe, batched_number: int):
# apply strengths, and get batch indeces to default out
# AKA latents that should not be influenced by ControlNet
@@ -298,34 +228,119 @@ class ControlNetAdvanced(ControlNet):
# first, resize mask to required dims
masks = prepare_mask_batch(self.mask_cond_hint, x.shape)
x[:] = x[:] * masks
def prepare_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None):
# make mask appropriate dimensions, if present
if self.mask_cond_hint_original is not None:
if self.sub_idxs is not None or self.mask_cond_hint is None or x_noisy.shape[2] * 8 != self.mask_cond_hint.shape[1] or x_noisy.shape[3] * 8 != self.mask_cond_hint.shape[2]:
if self.mask_cond_hint is not None:
del self.mask_cond_hint
self.mask_cond_hint = None
# TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM
# resize mask and match batch count
self.mask_cond_hint = prepare_mask_batch(self.mask_cond_hint_original, x_noisy.shape, multiplier=8)
actual_latent_length = x_noisy.shape[0] // batched_number
self.mask_cond_hint = comfy.utils.repeat_to_batch_size(self.mask_cond_hint, actual_latent_length if self.sub_idxs is None else self.full_latent_length)
if self.sub_idxs is not None:
self.mask_cond_hint = self.mask_cond_hint[self.sub_idxs]
# make cond_hint_mask length match x_noise
if x_noisy.shape[0] != self.mask_cond_hint.shape[0]:
self.mask_cond_hint = broadcast_image_to(self.mask_cond_hint, x_noisy.shape[0], batched_number)
# default dtype to be same as x_noisy
if dtype is None:
dtype = x_noisy.dtype
self.mask_cond_hint = self.mask_cond_hint.to(dtype=dtype).to(self.device)
def cleanup_advanced(self):
self.sub_idxs = None
self.full_latent_length = 0
self.context_length = 0
def copy_to_advanced(self, copied: 'AdvancedControlBase'):
copied.mask_cond_hint_original = self.mask_cond_hint_original
class ControlNetAdvanced(ControlNet, AdvancedControlBase):
def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None):
super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, device=device)
AdvancedControlBase.__init__(self, timestep_keyframes=timestep_keyframes)
# initialize weights
self.weights = self.timestep_keyframes.keyframes[0].control_net_weights if self.timestep_keyframes.keyframes[0].control_net_weights else [1.0]*13
def get_control(self, x_noisy, t, cond, batched_number):
# need to reference t and batched_number later
self.t = t
self.batched_number = batched_number
# TODO: choose TimestepKeyframe based on t
# perform special version of get_control that supports sliding context and masks
return self.sliding_get_control(x_noisy, t, cond, batched_number)
def sliding_get_control(self, x_noisy: Tensor, t, cond, batched_number):
control_prev = None
if self.previous_controlnet is not None:
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
if self.timestep_range is not None:
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
if control_prev is not None:
return control_prev
else:
return None
output_dtype = x_noisy.dtype
# make cond_hint appropriate dimensions
# TODO: change this to not require cond_hint upscaling every step when self.sub_idxs are present
if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] * 8 != self.cond_hint.shape[2] or x_noisy.shape[3] * 8 != self.cond_hint.shape[3]:
if self.cond_hint is not None:
del self.cond_hint
self.cond_hint = None
# if self.cond_hint_original length greater or equal to real latent count, subdivide it before scaling
if self.sub_idxs is not None and self.cond_hint_original.size(0) >= self.full_latent_length:
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(self.control_model.dtype).to(self.device)
else:
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(self.control_model.dtype).to(self.device)
if x_noisy.shape[0] != self.cond_hint.shape[0]:
self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number)
# prepare mask_cond_hint
self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, dtype=self.control_model.dtype)
context = cond['c_crossattn']
# uses 'y' in new ComfyUI update
y = cond.get('y', None)
if y is None: # TODO: remove this in the future since no longer used by newest ComfyUI
y = cond.get('c_adm', None)
if y is not None:
y = y.to(self.control_model.dtype)
timestep = self.model_sampling_current.timestep(t)
x_noisy = self.model_sampling_current.calculate_input(t, x_noisy)
control = self.control_model(x=x_noisy.to(self.control_model.dtype), hint=self.cond_hint, timesteps=timestep.float(), context=context.to(self.control_model.dtype), y=y)
return self.control_merge(None, control, control_prev, output_dtype)
def copy(self):
c = ControlNetAdvanced(self.control_model, self.timestep_keyframes, global_average_pooling=self.global_average_pooling)
self.copy_to(c)
self.copy_to_advanced(c)
return c
def cleanup(self):
super().cleanup()
self.sub_idxs = None
self.full_latent_length = 0
self.context_length = 0
self.cleanup_advanced()
@staticmethod
def from_vanilla(v: ControlNet, timestep_keyframe: TimestepKeyframeGroup=None) -> 'ControlNetAdvanced':
return ControlNetAdvanced(control_model=v.control_model, timestep_keyframes=timestep_keyframe,
global_average_pooling=v.global_average_pooling, device=v.device)
class T2IAdapterAdvanced(T2IAdapter):
class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase):
def __init__(self, t2i_model, timestep_keyframes: TimestepKeyframeGroup, channels_in, device=None):
super().__init__(t2i_model=t2i_model, channels_in=channels_in, device=device)
self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup()
self.current_timestep_keyframe = self.timestep_keyframes.keyframes[0]
AdvancedControlBase.__init__(self, timestep_keyframes=timestep_keyframes)
first_weight = self.timestep_keyframes.keyframes[0].t2i_adapter_weights if self.timestep_keyframes.get_index(0) else None
self.weights = first_weight if first_weight else [1.0]*12
# mask for which parts of controlnet output to keep
self.cond_hint_mask = None
# actual index values
self.sub_idxs = None
self.full_latent_length = 0
self.context_length = 0
# override control_merge
self.control_merge = control_merge_inject.__get__(self, type(self))
def get_control(self, x_noisy, t, cond, batched_number):
# need to reference t and batched_number later
@@ -335,10 +350,13 @@ class T2IAdapterAdvanced(T2IAdapter):
try:
# if sub indexes present, replace original hint with subsection
if self.sub_idxs is not None:
# cond hints
full_cond_hint_original = self.cond_hint_original
del self.cond_hint
self.cond_hint = None
self.cond_hint_original = full_cond_hint_original[self.sub_idxs]
# mask hints
self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number)
return super().get_control(x_noisy, t, cond, batched_number)
finally:
if self.sub_idxs is not None:
@@ -346,38 +364,72 @@ class T2IAdapterAdvanced(T2IAdapter):
self.cond_hint_original = full_cond_hint_original
del full_cond_hint_original
def apply_advanced_strengths_and_masks(self, x, current_timestep_keyframe: TimestepKeyframe, batched_number: int):
# For now, do nothing; need to figure out LatentKeyframe control is even possible for T2I Adapters
# TODO: support masks
return
def copy(self):
c = T2IAdapterAdvanced(self.t2i_model, self.timestep_keyframes, self.channels_in)
c = ControlLoraAdvanced(self.t2i_model, self.timestep_keyframes, self.channels_in)
self.copy_to(c)
self.copy_to_advanced(c)
return c
def cleanup(self):
super().cleanup()
self.sub_idxs = None
self.full_latent_length = 0
self.context_length = 0
self.cleanup_advanced()
@staticmethod
def from_vanilla(v: T2IAdapter, timestep_keyframe: TimestepKeyframeGroup=None) -> 'T2IAdapterAdvanced':
return T2IAdapterAdvanced(t2i_model=v.t2i_model, timestep_keyframes=timestep_keyframe, channels_in=v.channels_in, device=v.device)
class ControlLoraAdvanced(ControlLora, AdvancedControlBase):
def __init__(self, control_weights, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None):
super().__init__(control_weights=control_weights, global_average_pooling=global_average_pooling, device=device)
AdvancedControlBase.__init__(self, timestep_keyframes=timestep_keyframes)
# initialize weights
self.weights = self.timestep_keyframes.keyframes[0].control_net_weights if self.timestep_keyframes.keyframes[0].control_net_weights else [1.0]*10
# use some functions from ControlNetAdvanced
self.get_control = ControlNetAdvanced.get_control.__get__(self, type(self))
self.sliding_get_control = ControlNetAdvanced.sliding_get_control.__get__(self, type(self))
def copy(self):
c = ControlLoraAdvanced(self.control_weights, self.timestep_keyframes, global_average_pooling=self.global_average_pooling)
self.copy_to(c)
self.copy_to_advanced(c)
return c
def cleanup(self):
super().cleanup()
self.cleanup_advanced()
@staticmethod
def from_vanilla(v: ControlLora, timestep_keyframe: TimestepKeyframeGroup=None) -> 'ControlLoraAdvanced':
return ControlLoraAdvanced(control_weights=v.control_weights, timestep_keyframes=timestep_keyframe,
global_average_pooling=v.global_average_pooling, device=v.device)
class ControlLLLiteAdvanced(AdvancedControlBase):
def __init__(self, timestep_keyframes: TimestepKeyframeGroup):
AdvancedControlBase.__init__(self, timestep_keyframes=timestep_keyframes)
# TODO: see if can use weights with ControlLLLite
self.weights = [1.0]*100
def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None):
# TODO: support controlnet-lllite
control = comfy_cn.load_controlnet(ckpt_path, model=model)
# if exactly ControlNet returned, transform it into ControlNetAdvanced
if type(control) == ControlNet:
return ControlNetAdvanced(control.control_model, timestep_keyframe, global_average_pooling=control.global_average_pooling)
# if exactly ControlLora returned, transform it into ControlLoraAdvanced
elif type(control) == ControlLora:
return ControlLoraAdvanced.from_vanilla(v=control, timestep_keyframe=timestep_keyframe)
# if T2IAdapter returned, transform it into T2IAdapterAdvanced
elif isinstance(control, T2IAdapter):
return T2IAdapterAdvanced(control.t2i_model, timestep_keyframe, control.channels_in)
# otherwise, leave it be - probably a ControlLora for SDXL (no support for advanced stuff yet from here)
# TODO add ControlLoraAdvanced
return T2IAdapterAdvanced.from_vanilla(v=control, timestep_keyframe=timestep_keyframe)
# otherwise, leave it be - might be something I am not supporting yet
return control
def is_advanced_controlnet(input_object):
return isinstance(input_object, ControlNetAdvanced) or isinstance(input_object, T2IAdapterAdvanced)
return hasattr(input_object, "sub_idxs")
# adapted from comfy/sample.py
+36 -15
View File
@@ -46,6 +46,7 @@ class LatentKeyframeGroupNode:
"optional": {
"prev_latent_keyframe": ("LATENT_KEYFRAME", ),
"latent_optional": ("LATENT", ),
"print_keyframes": ("BOOLEAN", {"default": False})
}
}
@@ -81,7 +82,7 @@ class LatentKeyframeGroupNode:
def convert_to_latent_keyframes(self, latent_indeces: str, latent_count: int) -> set[LatentKeyframe]:
if not latent_indeces:
return set()
all_indeces = [i for i in range(0, latent_count)]
int_latent_indeces = [i for i in range(0, latent_count)]
allow_negative = latent_count > 0
chosen_indeces = set()
# parse string - allow positive ints, negative ints, and ranges separated by ':'
@@ -105,8 +106,14 @@ class LatentKeyframeGroupNode:
index_range = [r.strip() for r in index_range]
start_index = self.convert_to_index_int(index_range[0], latent_count=latent_count, is_range=True, allow_negative=allow_negative)
end_index = self.convert_to_index_int(index_range[1], latent_count=latent_count, is_range=True, allow_negative=allow_negative)
for i in all_indeces[start_index:end_index]:
chosen_indeces.add(LatentKeyframe(i, strength))
# if latents were passed in, base indeces on known latent count
if len(int_latent_indeces) > 0:
for i in int_latent_indeces[start_index:end_index]:
chosen_indeces.add(LatentKeyframe(i, strength))
# otherwise, assume indeces are valid
else:
for i in range(start_index, end_index):
chosen_indeces.add(LatentKeyframe(i, strength))
# parse individual indeces
else:
chosen_indeces.add(LatentKeyframe(self.convert_to_index_int(g, latent_count=latent_count, allow_negative=allow_negative), strength))
@@ -115,7 +122,8 @@ class LatentKeyframeGroupNode:
def load_keyframes(self,
index_strengths: str,
prev_latent_keyframe: LatentKeyframeGroup=None,
latent_image_opt=None):
latent_image_opt=None,
print_keyframes=False):
if not prev_latent_keyframe:
prev_latent_keyframe = LatentKeyframeGroup()
curr_latent_keyframe = LatentKeyframeGroup()
@@ -126,9 +134,13 @@ class LatentKeyframeGroupNode:
latent_keyframes = self.convert_to_latent_keyframes(index_strengths, latent_count=latent_count)
for latent_keyframe in latent_keyframes:
logger.info(f"keyframe {latent_keyframe.batch_index}:{latent_keyframe.strength}")
curr_latent_keyframe.add(latent_keyframe)
if print_keyframes:
for keyframe in curr_latent_keyframe.keyframes:
logger.info(f"keyframe {keyframe.batch_index}:{keyframe.strength}")
# replace values with prev_latent_keyframes
for latent_keyframe in prev_latent_keyframe.keyframes:
curr_latent_keyframe.add(latent_keyframe)
@@ -148,6 +160,7 @@ class LatentKeyframeInterpolationNode:
},
"optional": {
"prev_latent_keyframe": ("LATENT_KEYFRAME", ),
"print_keyframes": ("BOOLEAN", {"default": False})
}
}
@@ -161,7 +174,8 @@ class LatentKeyframeInterpolationNode:
batch_index_to_excl: int,
strength_to: float,
interpolation: str,
prev_latent_keyframe: LatentKeyframeGroup=None):
prev_latent_keyframe: LatentKeyframeGroup=None,
print_keyframes=False):
if (batch_index_from > batch_index_to_excl):
raise ValueError("batch_index_from must be less than or equal to batch_index_to.")
@@ -189,8 +203,11 @@ class LatentKeyframeInterpolationNode:
for i in range(steps):
keyframe = LatentKeyframe(batch_index_from + i, float(weights[i]))
logger.info(f"keyframe {batch_index_from + i}:{weights[i]}")
curr_latent_keyframe.add(keyframe)
if print_keyframes:
for keyframe in curr_latent_keyframe.keyframes:
logger.info(f"keyframe {keyframe.batch_index}:{keyframe.strength}")
# replace values with prev_latent_keyframes
for latent_keyframe in prev_latent_keyframe.keyframes:
@@ -204,10 +221,11 @@ class LatentKeyframeBatchedGroupNode:
def INPUT_TYPES(s):
return {
"required": {
"strengths": ("FLOAT", {"default": -1, "min": -1, "step": 0.0001}),
"float_strengths": ("FLOAT", {"default": -1, "min": -1, "step": 0.0001, "forceInput": True}),
},
"optional": {
"prev_latent_keyframe": ("LATENT_KEYFRAME", ),
"print_keyframes": ("BOOLEAN", {"default": False})
}
}
@@ -215,22 +233,25 @@ class LatentKeyframeBatchedGroupNode:
FUNCTION = "load_keyframe"
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes"
def load_keyframe(self, strengths: Union[float, list[float]], prev_latent_keyframe: LatentKeyframeGroup=None):
def load_keyframe(self, float_strengths: Union[float, list[float]], prev_latent_keyframe: LatentKeyframeGroup=None, print_keyframes=False):
if not prev_latent_keyframe:
prev_latent_keyframe = LatentKeyframeGroup()
curr_latent_keyframe = LatentKeyframeGroup()
# if received a normal float input, do nothing
if type(strengths) in (float, int):
logger.info("No batched strengths passed into Latent Keyframe Batch Group node; will not create any new keyframes.")
if type(float_strengths) in (float, int):
logger.info("No batched float_strengths passed into Latent Keyframe Batch Group node; will not create any new keyframes.")
# if iterable, attempt to create LatentKeyframes with chosen strengths
elif isinstance(strengths, Iterable):
for idx, strength in enumerate(strengths):
elif isinstance(float_strengths, Iterable):
for idx, strength in enumerate(float_strengths):
keyframe = LatentKeyframe(idx, strength)
curr_latent_keyframe.add(keyframe)
logger.info(f"keyframe {keyframe.batch_index}:{keyframe.strength}")
else:
raise ValueError(f"Expected strengths to be an iterable input, but was {type(strengths).__repr__}.")
raise ValueError(f"Expected strengths to be an iterable input, but was {type(float_strengths).__repr__}.")
if print_keyframes:
for keyframe in curr_latent_keyframe.keyframes:
logger.info(f"keyframe {keyframe.batch_index}:{keyframe.strength}")
# replace values with prev_latent_keyframes
for latent_keyframe in prev_latent_keyframe.keyframes:
+4 -1
View File
@@ -5,7 +5,7 @@ import folder_paths
from .control import load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType,\
LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup, is_advanced_controlnet
from .control import StrengthInterpolation as SI
from .weight_nodes import ScaledSoftControlNetWeights, SoftControlNetWeights, CustomControlNetWeights, \
from .weight_nodes import ScaledSoftControlLoraWeights, ScaledSoftControlNetWeights, SoftControlNetWeights, CustomControlNetWeights, \
SoftT2IAdapterWeights, CustomT2IAdapterWeights
from .latent_keyframe_nodes import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode
from .deprecated_nodes import LoadImagesFromDirectory
@@ -140,6 +140,7 @@ class AdvancedControlNetApply:
if prev_cnet in cnets:
c_net = cnets[prev_cnet]
else:
# TODO: attempt to convert to Advanced versions
c_net = control_net.copy().set_cond_hint(control_hint, strength, (start_percent, end_percent))
# set cond hint mask
if mask_optional is not None:
@@ -178,6 +179,7 @@ NODE_CLASS_MAPPINGS = {
"CustomControlNetWeights": CustomControlNetWeights,
"SoftT2IAdapterWeights": SoftT2IAdapterWeights,
"CustomT2IAdapterWeights": CustomT2IAdapterWeights,
"ScaledSoftControlLoraWeights": ScaledSoftControlLoraWeights,
# Image
"LoadImagesFromDirectory": LoadImagesFromDirectory
}
@@ -200,6 +202,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"CustomControlNetWeights": "Custom ControlNet Weights 🛂🅐🅒🅝",
"SoftT2IAdapterWeights": "Soft T2IAdapter Weights 🛂🅐🅒🅝",
"CustomT2IAdapterWeights": "Custom T2IAdapter Weights 🛂🅐🅒🅝",
"ScaledSoftControlLoraWeights": "Scaled Soft ControlLora Weights 🛂🅐🅒🅝",
# Image
"LoadImagesFromDirectory": "Load Images [DEPRECATED] 🛂🅐🅒🅝"
}
+22
View File
@@ -11,6 +11,28 @@ def get_properly_arranged_t2i_weights(initial_weights: list[float]):
return new_weights
class ScaledSoftControlLoraWeights:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ),
"flip_weights": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",)
FUNCTION = "load_weights"
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights"
def load_weights(self, base_multiplier, flip_weights):
weights = [(base_multiplier ** float(9 - i)) for i in range(10)]
if flip_weights:
weights.reverse()
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_net_weights=weights)))
class ScaledSoftControlNetWeights:
@classmethod
def INPUT_TYPES(s):