diff --git a/control/control.py b/control/control.py index ccf6283..ab260bd 100644 --- a/control/control.py +++ b/control/control.py @@ -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 diff --git a/control/latent_keyframe_nodes.py b/control/latent_keyframe_nodes.py index 57cd280..8152921 100644 --- a/control/latent_keyframe_nodes.py +++ b/control/latent_keyframe_nodes.py @@ -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: diff --git a/control/nodes.py b/control/nodes.py index 98c433c..1371f7a 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -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] 🛂🅐🅒🅝" } diff --git a/control/weight_nodes.py b/control/weight_nodes.py index 2015c8a..1bd2c37 100644 --- a/control/weight_nodes.py +++ b/control/weight_nodes.py @@ -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):