From 3a430c9e5f1c4fa9f42817eb41f3354e469e5a73 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 25 Oct 2023 08:15:11 -0500 Subject: [PATCH 01/15] Started work on new features --- control/control.py | 17 ++++++++++++++++- control/latent_keyframe_nodes.py | 11 ++++++----- control/nodes.py | 12 ++++++++++-- control/reference_nodes.py | 12 ++++++++++++ 4 files changed, 44 insertions(+), 8 deletions(-) create mode 100644 control/reference_nodes.py diff --git a/control/control.py b/control/control.py index 0508ae4..10960a9 100644 --- a/control/control.py +++ b/control/control.py @@ -11,6 +11,14 @@ ControlNetWeightsType = list[float] T2IAdapterWeightsType = list[float] +class StrengthInterpolation: + LINEAR = "linear" + EASE_IN = "ease-in" + EASE_OUT = "ease-out" + EASE_IN_OUT = "ease-in-out" + NONE = "none" + + class LatentKeyframe: def __init__(self, batch_index: int, strength: float) -> None: self.batch_index = batch_index @@ -50,11 +58,15 @@ class LatentKeyframeGroup: class TimestepKeyframe: def __init__(self, start_percent: float = 0.0, + strength: float = 1.0, + interpolation: str = StrengthInterpolation.LINEAR, control_net_weights: ControlNetWeightsType = None, t2i_adapter_weights: T2IAdapterWeightsType = None, latent_keyframes: LatentKeyframeGroup = None, default_latent_strength: float = 0.0) -> None: self.start_percent = start_percent + self.strength = strength + self.interpolation = interpolation self.control_net_weights = control_net_weights self.t2i_adapter_weights = t2i_adapter_weights self.latent_keyframes = latent_keyframes @@ -229,7 +241,10 @@ class ControlNetAdvanced(ControlNet): self.mask_cond_hint = self.mask_cond_hint.to(self.control_model.dtype).to(self.device) context = cond['c_crossattn'] - y = cond.get('c_adm', None) + # 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) control = self.control_model(x=x_noisy.to(self.control_model.dtype), hint=self.cond_hint, timesteps=t, context=context.to(self.control_model.dtype), y=y) diff --git a/control/latent_keyframe_nodes.py b/control/latent_keyframe_nodes.py index d12935c..57cd280 100644 --- a/control/latent_keyframe_nodes.py +++ b/control/latent_keyframe_nodes.py @@ -3,6 +3,7 @@ import numpy as np from collections.abc import Iterable from .control import LatentKeyframe, LatentKeyframeGroup +from .control import StrengthInterpolation as SI from .logger import logger @@ -143,7 +144,7 @@ class LatentKeyframeInterpolationNode: "batch_index_to_excl": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), "strength_from": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), "strength_to": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), - "interpolation": (["linear", "ease-in", "ease-out", "ease-in-out"], ), + "interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT], ), }, "optional": { "prev_latent_keyframe": ("LATENT_KEYFRAME", ), @@ -174,15 +175,15 @@ class LatentKeyframeInterpolationNode: steps = batch_index_to_excl - batch_index_from diff = strength_to - strength_from - if interpolation == "linear": + if interpolation == SI.LINEAR: weights = np.linspace(strength_from, strength_to, steps) - elif interpolation == "ease-in": + elif interpolation == SI.EASE_IN: index = np.linspace(0, 1, steps) weights = diff * np.power(index, 2) + strength_from - elif interpolation == "ease-out": + elif interpolation == SI.EASE_OUT: index = np.linspace(0, 1, steps) weights = diff * (1 - np.power(1 - index, 2)) + strength_from - elif interpolation == "ease-in-out": + elif interpolation == SI.EASE_IN_OUT: index = np.linspace(0, 1, steps) weights = diff * ((1 - np.cos(index * np.pi)) / 2) + strength_from diff --git a/control/nodes.py b/control/nodes.py index b43f78f..74aeddb 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -2,8 +2,9 @@ import numpy as np import folder_paths -from .control import ControlNetAdvanced, T2IAdapterAdvanced, load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType,\ +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, \ SoftT2IAdapterWeights, CustomT2IAdapterWeights from .latent_keyframe_nodes import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode @@ -17,6 +18,9 @@ class TimestepKeyframeNode: return { "required": { "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT, SI.NONE], ), + "default_latent_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), }, "optional": { "control_net_weights": ("CONTROL_NET_WEIGHTS", ), @@ -33,13 +37,17 @@ class TimestepKeyframeNode: def load_keyframe(self, start_percent: float, + strength: float, + interpolation: str, + default_latent_strength: float, control_net_weights: ControlNetWeightsType=None, t2i_adapter_weights: T2IAdapterWeightsType=None, latent_keyframe: LatentKeyframeGroup=None, prev_timestep_keyframe: TimestepKeyframeGroup=None): if not prev_timestep_keyframe: prev_timestep_keyframe = TimestepKeyframeGroup() - keyframe = TimestepKeyframe(start_percent, control_net_weights, t2i_adapter_weights, latent_keyframe) + keyframe = TimestepKeyframe(start_percent=start_percent, strength=strength, interpolation=interpolation, default_latent_strength=default_latent_strength, + control_net_weights=control_net_weights, t2i_adapter_weights=t2i_adapter_weights, latent_keyframes=latent_keyframe) prev_timestep_keyframe.add(keyframe) return (prev_timestep_keyframe,) diff --git a/control/reference_nodes.py b/control/reference_nodes.py new file mode 100644 index 0000000..6879f97 --- /dev/null +++ b/control/reference_nodes.py @@ -0,0 +1,12 @@ +class AnimateDiffLoaderWithContext: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "image": ("IMAGE",), + }, + } + + RETURN_TYPES = ("MODEL",) + CATEGORY = "" \ No newline at end of file From dbe8136cdcba589be30a06df965fe0a82fc618eb Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 16 Nov 2023 15:25:41 -0600 Subject: [PATCH 02/15] Fixed Apply Advanced ControlNet node - new ComfyUI update does not subtract start/end percentages from 1.0 --- control/nodes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/control/nodes.py b/control/nodes.py index 74aeddb..98c433c 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -140,7 +140,7 @@ class AdvancedControlNetApply: if prev_cnet in cnets: c_net = cnets[prev_cnet] else: - c_net = control_net.copy().set_cond_hint(control_hint, strength, (1.0 - start_percent, 1.0 - end_percent)) + 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: if is_advanced_controlnet(c_net): From 1f4fcff9671ca728709c9e86ab3794cfd107934b Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 26 Nov 2023 08:38:12 -0600 Subject: [PATCH 03/15] 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 --- control/control.py | 252 +++++++++++++++++++------------ control/latent_keyframe_nodes.py | 51 +++++-- control/nodes.py | 5 +- control/weight_nodes.py | 22 +++ 4 files changed, 214 insertions(+), 116 deletions(-) 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): From badf1d33b7c7860b68eb8896d3b963c9caf74098 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 28 Nov 2023 14:55:22 -0600 Subject: [PATCH 04/15] TimestepKeyframes work now, unified Weights, reorganized some noes (massive refactor all around) --- control/control.py | 436 +++++++++++++++++++++++++++--------- control/control_lllite.py | 1 + control/deprecated_nodes.py | 33 +++ control/nodes.py | 86 ++++--- control/weight_nodes.py | 89 +++++--- 5 files changed, 470 insertions(+), 175 deletions(-) create mode 100644 control/control_lllite.py diff --git a/control/control.py b/control/control.py index ab260bd..f70900d 100644 --- a/control/control.py +++ b/control/control.py @@ -4,11 +4,79 @@ import torch import comfy.utils import comfy.controlnet as comfy_cn -from comfy.controlnet import ControlNet, ControlLora, T2IAdapter, broadcast_image_to +from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, broadcast_image_to -ControlNetWeightsType = list[float] -T2IAdapterWeightsType = list[float] +def get_properly_arranged_t2i_weights(initial_weights: list[float]): + new_weights = [] + new_weights.extend([initial_weights[0]]*3) + new_weights.extend([initial_weights[1]]*3) + new_weights.extend([initial_weights[2]]*3) + new_weights.extend([initial_weights[3]]*3) + return new_weights + + +class ControlWeightType: + DEFAULT = "default" + UNIVERSAL = "universal" + T2IADAPTER = "t2iadapter" + CONTROLNET = "controlnet" + CONTROLLORA = "controllora" + CONTROLLLLITE = "controllllite" + + +class ControlWeights: + def __init__(self, weight_type: str, base_multiplier: float=1.0, flip_weights: bool=False, weights: list[float]=None, weight_mask: Tensor=None): + self.weight_type = weight_type + self.base_multiplier = base_multiplier + self.flip_weights = flip_weights + self.weights = weights + if self.weights is not None and self.flip_weights: + self.weights.reverse() + self.weight_mask = weight_mask + + def get(self, idx: int, shape: Tensor) -> Union[float, Tensor]: + # if weight_mask present, return normalized mask + if self.base_multiplier != 1.0 and self.weight_mask is not None: + # TODO: fill out + pass + # if weights is not none, return index + if self.weights is not None: + return self.weights[idx] + return 1.0 + + @classmethod + def default(cls): + return cls(ControlWeightType.DEFAULT) + + @classmethod + def universal(cls, base_multiplier: float, flip_weights: bool=False): + return cls(ControlWeightType.UNIVERSAL, base_multiplier=base_multiplier, flip_weights=flip_weights) + + @classmethod + def t2iadapter(cls, weights: list[float]=None, flip_weights: bool=False): + if weights is None: + weights = [1.0]*12 + return cls(ControlWeightType.T2IADAPTER, weights=weights,flip_weights=flip_weights) + + @classmethod + def controlnet(cls, weights: list[float]=None, flip_weights: bool=False): + if weights is None: + weights = [1.0]*13 + return cls(ControlWeightType.CONTROLNET, weights=weights, flip_weights=flip_weights) + + @classmethod + def controllora(cls, weights: list[float]=None, flip_weights: bool=False): + if weights is None: + weights = [1.0]*10 + return cls(ControlWeightType.CONTROLLORA, weights=weights, flip_weights=flip_weights) + + @classmethod + def controllllite(cls, weights: list[float]=None, flip_weights: bool=False): + if weights is None: + # TODO: make this have a real value + weights = [1.0]*200 + return cls(ControlWeightType.CONTROLLLLITE, weights=weights, flip_weights=flip_weights) class StrengthInterpolation: @@ -59,18 +127,27 @@ class TimestepKeyframe: def __init__(self, start_percent: float = 0.0, strength: float = 1.0, - interpolation: str = StrengthInterpolation.LINEAR, - control_net_weights: ControlNetWeightsType = None, - t2i_adapter_weights: T2IAdapterWeightsType = None, + interpolation: str = StrengthInterpolation.NONE, + control_weights: ControlWeights = None, latent_keyframes: LatentKeyframeGroup = None, - default_latent_strength: float = 0.0) -> None: + null_latent_kf_strength: float = 0.0, + inherit_missing: bool = True, + guarantee_usage: bool = False) -> None: self.start_percent = start_percent + self.start_t = 999999999.9 self.strength = strength self.interpolation = interpolation - self.control_net_weights = control_net_weights - self.t2i_adapter_weights = t2i_adapter_weights + self.control_weights = control_weights self.latent_keyframes = latent_keyframes - self.default_latent_strength = default_latent_strength + self.null_latent_kf_strength = null_latent_kf_strength + self.inherit_missing = inherit_missing + self.guarantee_usage = guarantee_usage + + def has_control_weights(self): + return self.control_weights is not None + + def has_latent_keyframes(self): + return self.latent_keyframes is not None @classmethod @@ -102,9 +179,15 @@ class TimestepKeyframeGroup: except IndexError: return None + def has_index(self, index: int) -> int: + return index >=0 and index < len(self.keyframes) + def __getitem__(self, index) -> TimestepKeyframe: return self.keyframes[index] + def __len__(self) -> int: + return len(self.keyframes) + def is_empty(self) -> bool: return len(self.keyframes) == 0 @@ -116,61 +199,13 @@ class TimestepKeyframeGroup: # used to inject ControlNetAdvanced and T2IAdapterAdvanced control_merge function -def control_merge_inject(self, control_input, control_output, control_prev, output_dtype): - out = {'input':[], 'middle':[], 'output': []} - - if control_input is not None: - for i in range(len(control_input)): - key = 'input' - x = control_input[i] - if x is not None: - self.apply_advanced_strengths_and_masks(x, self.current_timestep_keyframe, self.batched_number) - - x *= self.strength * self.weights[i] - if x.dtype != output_dtype: - x = x.to(output_dtype) - out[key].insert(0, x) - - if control_output is not None: - for i in range(len(control_output)): - if i == (len(control_output) - 1): - key = 'middle' - index = 0 - else: - key = 'output' - index = i - x = control_output[i] - if x is not None: - self.apply_advanced_strengths_and_masks(x, self.current_timestep_keyframe, self.batched_number) - - if self.global_average_pooling: - x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3]) - - x *= self.strength * self.weights[i] - if x.dtype != output_dtype: - x = x.to(output_dtype) - - out[key].append(x) - if control_prev is not None: - for x in ['input', 'middle', 'output']: - o = out[x] - for i in range(len(control_prev[x])): - prev_val = control_prev[x][i] - if i >= len(o): - o.append(prev_val) - elif prev_val is not None: - if o[i] is None: - o[i] = prev_val - else: - o[i] += prev_val - return out 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] + def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights): + self.base = base + self.compatible_weights = [ControlWeightType.UNIVERSAL] + self.add_compatible_weight(weights_default.weight_type) # mask for which parts of controlnet output to keep self.mask_cond_hint_original = None self.mask_cond_hint = None @@ -178,26 +213,133 @@ class AdvancedControlBase: 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)) + # timesteps + self.t: Tensor = None + self.batched_number: int = None + # weights + override + self.weights: ControlWeights = None + self.weights_default: ControlWeights = weights_default + self.weights_override: ControlWeights = None + # latent keyframe + override + self.latent_keyframes: LatentKeyframeGroup = None + self.latent_keyframe_override: LatentKeyframeGroup = None + # initialize timestep_keyframes + self.set_timestep_keyframes(timestep_keyframes) + # override some functions + self.get_control = self.get_control_inject + self.control_merge = self.control_merge_inject#.__get__(self, type(self)) + self.pre_run = self.pre_run_inject + self.cleanup = self.cleanup_inject + + def add_compatible_weight(self, control_weight_type: str): + self.compatible_weights.append(control_weight_type) + + def set_timestep_keyframes(self, timestep_keyframes: TimestepKeyframeGroup): + self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup() + # prepare first timestep_keyframe related stuff + self.current_timestep_keyframe = None + self.current_timestep_index = -1 + self.next_timestep_keyframe = None + self.weights = None + self.latent_keyframes = None + + def prepare_current_timestep(self, t: Tensor, batched_number: int): + self.t = t + self.batched_number = batched_number + # get current step percent + curr_t: float = t[0] + print(f"$$$$ {curr_t}") + prev_index = self.current_timestep_index + # if has next index, loop through and see if need to switch + if self.timestep_keyframes.has_index(self.current_timestep_index+1): + for i in range(self.current_timestep_index+1, len(self.timestep_keyframes)): + eval_tk = self.timestep_keyframes[i] + # check if start percent is less or equal to curr_t + if eval_tk.start_t >= curr_t: + self.current_timestep_index = i + self.current_timestep_keyframe = eval_tk + # keep track of weights and control weights, accounting for inherit_missing + if self.current_timestep_keyframe.has_control_weights(): + self.weights = self.current_timestep_keyframe.control_weights + elif not self.current_timestep_keyframe.inherit_missing: + self.weights = self.weights_default + if self.current_timestep_keyframe.has_latent_keyframes(): + self.latent_keyframes = self.current_timestep_keyframe.latent_keyframes + elif not self.current_timestep_keyframe.inherit_missing: + self.latent_keyframes = None + # if guarantee_usage, stop searching for other TKs + if self.current_timestep_keyframe.guarantee_usage: + break + # if eval_tk is outside of percent range, stop looking further + else: + break + + # if index changed, apply overrides + if prev_index != self.current_timestep_index: + if self.weights_override is not None: + self.weights = self.weights_override + if self.latent_keyframe_override is not None: + self.latent_keyframes = self.latent_keyframe_override + + # make sure weights and latent_keyframes are in a workable state + # Note: each AdvancedControlBase should create their own get_universal_weights class + self.prepare_weights() + + def prepare_weights(self): + if self.weights is None or self.weights.weight_type == ControlWeightType.DEFAULT: + self.weights = self.weights_default + elif self.weights.weight_type == ControlWeightType.UNIVERSAL: + self.weights = self.get_universal_weights() + + def get_universal_weights(self) -> ControlWeights: + return self.weights def set_cond_hint_mask(self, mask_hint): self.mask_cond_hint_original = mask_hint return self - 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 + def pre_run_inject(self, model, percent_to_timestep_function): + self.base.pre_run(model, percent_to_timestep_function) + self.pre_run_advanced(model, percent_to_timestep_function) + + def pre_run_advanced(self, model, percent_to_timestep_function): + # for each timestep keyframe, calculate the start_t + for tk in self.timestep_keyframes.keyframes: + tk.start_t = percent_to_timestep_function(tk.start_percent) + # clear variables + self.cleanup_advanced() + + def get_control_inject(self, x_noisy, t, cond, batched_number): + # prepare timestep and everything related + self.prepare_current_timestep(t=t, batched_number=batched_number) + # if should not perform any actions for the controlnet, exit without doing any work + if self.strength == 0.0 or self.current_timestep_keyframe.strength == 0.0: + control_prev = None + if self.previous_controlnet is not None: + control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number) + if control_prev is not None: + return control_prev + else: + return None + # otherwise, perform normal function + return self.get_control_advanced(x_noisy, t, cond, batched_number) + + def get_control_advanced(self, x_noisy, t, cond, batched_number): + pass + + def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int): + # apply strengths, and get batch indeces to null out # AKA latents that should not be influenced by ControlNet - if current_timestep_keyframe.latent_keyframes is not None: + if self.latent_keyframes is not None: latent_count = x.size(0)//batched_number - indeces_to_default = set(range(latent_count)) + indeces_to_null = set(range(latent_count)) mapped_indeces = None # if expecting subdivision, will need to translate between subset and actual idx values if self.sub_idxs: mapped_indeces = {} for i, actual in enumerate(self.sub_idxs): mapped_indeces[actual] = i - for keyframe in current_timestep_keyframe.latent_keyframes: + for keyframe in self.latent_keyframes: real_index = keyframe.batch_index # if negative, count from end if real_index < 0: @@ -205,30 +347,82 @@ class AdvancedControlBase: # if not mapping indeces, what you see is what you get if mapped_indeces is None: - if real_index in indeces_to_default: - indeces_to_default.remove(real_index) + if real_index in indeces_to_null: + indeces_to_null.remove(real_index) # otherwise, see if batch_index is even included in this set of latents else: real_index = mapped_indeces.get(real_index, None) if real_index is None: continue - indeces_to_default.remove(real_index) + indeces_to_null.remove(real_index) # apply strength for each batched cond/uncond for b in range(batched_number): x[(latent_count*b)+real_index] = x[(latent_count*b)+real_index] * keyframe.strength - # default them out by multiplying by default_latent_strength - for batch_index in indeces_to_default: - # apply default for each batched cond/uncond + # null them out by multiplying by null_latent_kf_strength + for batch_index in indeces_to_null: + # apply null for each batched cond/uncond for b in range(batched_number): - x[(latent_count*b)+batch_index] = x[(latent_count*b)+batch_index] * current_timestep_keyframe.default_latent_strength + x[(latent_count*b)+batch_index] = x[(latent_count*b)+batch_index] * self.current_timestep_keyframe.null_latent_kf_strength # apply masks if self.mask_cond_hint is not None: # first, resize mask to required dims masks = prepare_mask_batch(self.mask_cond_hint, x.shape) x[:] = x[:] * masks + # apply timestep keyframe strengths + if self.current_timestep_keyframe.strength != 1.0: + x[:] *= self.current_timestep_keyframe.strength + def control_merge_inject(self: 'AdvancedControlBase', control_input, control_output, control_prev, output_dtype): + out = {'input':[], 'middle':[], 'output': []} + + if control_input is not None: + for i in range(len(control_input)): + key = 'input' + x = control_input[i] + if x is not None: + self.apply_advanced_strengths_and_masks(x, self.batched_number) + + x *= self.strength * self.weights.get(i, x.shape) + if x.dtype != output_dtype: + x = x.to(output_dtype) + out[key].insert(0, x) + + if control_output is not None: + for i in range(len(control_output)): + if i == (len(control_output) - 1): + key = 'middle' + index = 0 + else: + key = 'output' + index = i + x = control_output[i] + if x is not None: + self.apply_advanced_strengths_and_masks(x, self.batched_number) + + if self.global_average_pooling: + x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3]) + + x *= self.strength * self.weights.get(i, x.shape) + if x.dtype != output_dtype: + x = x.to(output_dtype) + + out[key].append(x) + if control_prev is not None: + for x in ['input', 'middle', 'output']: + o = out[x] + for i in range(len(control_prev[x])): + prev_val = control_prev[x][i] + if i >= len(o): + o.append(prev_val) + elif prev_val is not None: + if o[i] is None: + o[i] = prev_val + else: + o[i] += prev_val + return out + 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: @@ -251,27 +445,40 @@ class AdvancedControlBase: dtype = x_noisy.dtype self.mask_cond_hint = self.mask_cond_hint.to(dtype=dtype).to(self.device) + def cleanup_inject(self): + self.base.cleanup() + self.cleanup_advanced() + def cleanup_advanced(self): self.sub_idxs = None self.full_latent_length = 0 self.context_length = 0 + self.t = None + self.batched_number = None + self.weights = None + self.latent_keyframes = None + # timestep stuff + self.current_timestep_keyframe = None + self.next_timestep_keyframe = None + self.current_timestep_index = -1 def copy_to_advanced(self, copied: 'AdvancedControlBase'): copied.mask_cond_hint_original = self.mask_cond_hint_original + copied.weights_override = self.weights_override + copied.latent_keyframe_override = self.latent_keyframe_override 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 + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controlnet()) - 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 + def get_universal_weights(self) -> ControlWeights: + raw_weights = [(self.weights.base_multiplier ** float(12 - i)) for i in range(13)] + # TODO: account for masks? + return ControlWeights.controlnet(raw_weights, self.weights.flip_weights) + + def get_control_advanced(self, x_noisy, t, cond, batched_number): # perform special version of get_control that supports sliding context and masks return self.sliding_get_control(x_noisy, t, cond, batched_number) @@ -324,10 +531,6 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase): self.copy_to(c) self.copy_to_advanced(c) return c - - def cleanup(self): - super().cleanup() - self.cleanup_advanced() @staticmethod def from_vanilla(v: ControlNet, timestep_keyframe: TimestepKeyframeGroup=None) -> 'ControlNetAdvanced': @@ -338,15 +541,18 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase): 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) - 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 - - 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 + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.t2iadapter()) + + def get_universal_weights(self) -> ControlWeights: + raw_weights = [(self.weights.base_multiplier ** float(7 - i)) for i in range(8)] + raw_weights = [raw_weights[-8], raw_weights[-3], raw_weights[-2], raw_weights[-1]] + raw_weights = get_properly_arranged_t2i_weights(raw_weights) + # TODO: account for masks? + return ControlWeights.t2iadapter(raw_weights, self.weights.flip_weights) + + def get_control_advanced(self, x_noisy, t, cond, batched_number): + # prepare timestep and everything related + self.prepare_current_timestep(t=t, batched_number=batched_number) try: # if sub indexes present, replace original hint with subsection if self.sub_idxs is not None: @@ -382,13 +588,16 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): 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 + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllora()) # use some functions from ControlNetAdvanced - self.get_control = ControlNetAdvanced.get_control.__get__(self, type(self)) + self.get_control_advanced = ControlNetAdvanced.get_control_advanced.__get__(self, type(self)) self.sliding_get_control = ControlNetAdvanced.sliding_get_control.__get__(self, type(self)) + def get_universal_weights(self) -> ControlWeights: + raw_weights = [(self.weights.base_multiplier ** float(9 - i)) for i in range(10)] + # TODO: account for masks? + return ControlWeights.controllora(raw_weights, self.weights.flip_weights) + def copy(self): c = ControlLoraAdvanced(self.control_weights, self.timestep_keyframes, global_average_pooling=self.global_average_pooling) self.copy_to(c) @@ -405,19 +614,30 @@ class ControlLoraAdvanced(ControlLora, AdvancedControlBase): 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 +class ControlLLLiteAdvanced(ControlNet, AdvancedControlBase): + def __init__(self, control_weights, timestep_keyframes: TimestepKeyframeGroup, device=None): + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite()) def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None): - # TODO: support controlnet-lllite control = comfy_cn.load_controlnet(ckpt_path, model=model) + # TODO: support controlnet-lllite + # if is None, see if is a non-vanilla ControlNet + # if control is None: + # controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) + # # check if lllite + # if "lllite_unet" in controlnet_data: + # pass + return convert_to_advanced(control, timestep_keyframe=timestep_keyframe) + + +def convert_to_advanced(control, timestep_keyframe: TimestepKeyframeGroup=None): + # if already advanced, leave it be + if is_advanced_controlnet(control): + return control # 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) + return ControlNetAdvanced.from_vanilla(v=control, timestep_keyframe=timestep_keyframe) # if exactly ControlLora returned, transform it into ControlLoraAdvanced elif type(control) == ControlLora: return ControlLoraAdvanced.from_vanilla(v=control, timestep_keyframe=timestep_keyframe) diff --git a/control/control_lllite.py b/control/control_lllite.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/control/control_lllite.py @@ -0,0 +1 @@ + diff --git a/control/deprecated_nodes.py b/control/deprecated_nodes.py index 25b7169..a64ac9b 100644 --- a/control/deprecated_nodes.py +++ b/control/deprecated_nodes.py @@ -4,6 +4,7 @@ import torch import numpy as np from PIL import Image, ImageOps +from .control import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe from .logger import logger @@ -68,3 +69,35 @@ class LoadImagesFromDirectory: raise FileNotFoundError(f"No images could be loaded from directory '{directory}'.") return (torch.cat(images, dim=0), torch.stack(masks, dim=0), image_count) + + +class TimestepKeyframeNodeDeprecated: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), + }, + "optional": { + "control_net_weights": ("CONTROL_NET_WEIGHTS", ), + "t2i_adapter_weights": ("T2I_ADAPTER_WEIGHTS", ), + "latent_keyframe": ("LATENT_KEYFRAME", ), + "prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ), + } + } + + RETURN_TYPES = ("TIMESTEP_KEYFRAME", ) + FUNCTION = "load_keyframe" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" + + def load_keyframe(self, + start_percent: float, + control_net_weights: ControlWeights=None, + latent_keyframe: LatentKeyframeGroup=None, + prev_timestep_keyframe: TimestepKeyframeGroup=None): + if not prev_timestep_keyframe: + prev_timestep_keyframe = TimestepKeyframeGroup() + keyframe = TimestepKeyframe(start_percent, control_net_weights, latent_keyframe) + prev_timestep_keyframe.add(keyframe) + return (prev_timestep_keyframe,) diff --git a/control/nodes.py b/control/nodes.py index 1371f7a..e3803b3 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -2,10 +2,10 @@ import numpy as np import folder_paths -from .control import load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType,\ +from .control import load_controlnet, convert_to_advanced, ControlWeights, ControlWeightType,\ LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup, is_advanced_controlnet from .control import StrengthInterpolation as SI -from .weight_nodes import ScaledSoftControlLoraWeights, ScaledSoftControlNetWeights, SoftControlNetWeights, CustomControlNetWeights, \ +from .weight_nodes import DefaultWeights, ScaledSoftControlLoraWeights, ScaledSoftControlNetWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, \ SoftT2IAdapterWeights, CustomT2IAdapterWeights from .latent_keyframe_nodes import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode from .deprecated_nodes import LoadImagesFromDirectory @@ -19,17 +19,19 @@ class TimestepKeyframeNode: "required": { "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT, SI.NONE], ), - "default_latent_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), }, "optional": { "control_net_weights": ("CONTROL_NET_WEIGHTS", ), - "t2i_adapter_weights": ("T2I_ADAPTER_WEIGHTS", ), "latent_keyframe": ("LATENT_KEYFRAME", ), "prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ), + "null_latent_kf_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "inherit_missing": ("BOOLEAN", {"default": True}, ), + "guarantee_usage": ("BOOLEAN", {"default": True}, ), + #"interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT, SI.NONE], {"default": SI.NONE}, ), } } + RETURN_NAMES = ("TIMESTEP_KF", ) RETURN_TYPES = ("TIMESTEP_KEYFRAME", ) FUNCTION = "load_keyframe" @@ -38,16 +40,17 @@ class TimestepKeyframeNode: def load_keyframe(self, start_percent: float, strength: float, - interpolation: str, - default_latent_strength: float, - control_net_weights: ControlNetWeightsType=None, - t2i_adapter_weights: T2IAdapterWeightsType=None, + control_net_weights: ControlWeights=None, latent_keyframe: LatentKeyframeGroup=None, - prev_timestep_keyframe: TimestepKeyframeGroup=None): + prev_timestep_keyframe: TimestepKeyframeGroup=None, + null_latent_kf_strength: float=0.0, + inherit_missing=True, + guarantee_usage=True, + interpolation: str=SI.NONE,): if not prev_timestep_keyframe: prev_timestep_keyframe = TimestepKeyframeGroup() - keyframe = TimestepKeyframe(start_percent=start_percent, strength=strength, interpolation=interpolation, default_latent_strength=default_latent_strength, - control_net_weights=control_net_weights, t2i_adapter_weights=t2i_adapter_weights, latent_keyframes=latent_keyframe) + keyframe = TimestepKeyframe(start_percent=start_percent, strength=strength, interpolation=interpolation, null_latent_kf_strength=null_latent_kf_strength, + control_weights=control_net_weights, latent_keyframes=latent_keyframe, inherit_missing=inherit_missing, guarantee_usage=guarantee_usage) prev_timestep_keyframe.add(keyframe) return (prev_timestep_keyframe,) @@ -63,11 +66,11 @@ class ControlNetLoaderAdvanced: "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), } } - + RETURN_TYPES = ("CONTROL_NET", ) FUNCTION = "load_controlnet" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" def load_controlnet(self, control_net_name, timestep_keyframe: TimestepKeyframeGroup=None): controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) @@ -91,9 +94,9 @@ class DiffControlNetLoaderAdvanced: RETURN_TYPES = ("CONTROL_NET", ) FUNCTION = "load_controlnet" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" - def load_controlnet(self, control_net_name, timestep_keyframe: TimestepKeyframeGroup, model): + def load_controlnet(self, control_net_name, model, timestep_keyframe: TimestepKeyframeGroup=None): controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) controlnet = load_controlnet(controlnet_path, timestep_keyframe, model) return (controlnet,) @@ -114,6 +117,9 @@ class AdvancedControlNetApply: }, "optional": { "mask_optional": ("MASK", ), + "timestep_kf": ("TIMESTEP_KEYFRAME", ), + "latent_kf_override": ("LATENT_KEYFRAME", ), + "cn_weights_override": ("CONTROL_NET_WEIGHTS", ), } } @@ -121,9 +127,12 @@ class AdvancedControlNetApply: RETURN_NAMES = ("positive", "negative") FUNCTION = "apply_controlnet" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/conditioning" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" - def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, mask_optional=None): + def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, + mask_optional=None, + timestep_kf: TimestepKeyframeGroup=None, latent_kf_override: LatentKeyframeGroup=None, + weights_override: ControlWeights=None): if strength == 0: return (positive, negative) @@ -140,11 +149,18 @@ 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: - if is_advanced_controlnet(c_net): + # copy, convert to advanced if needed, and set cond + c_net = convert_to_advanced(control_net.copy()).set_cond_hint(control_hint, strength, (start_percent, end_percent)) + if is_advanced_controlnet(c_net): + # apply optional parameters and overrides, if provided + if timestep_kf is not None: + c_net.set_timestep_keyframes(timestep_kf) + if latent_kf_override is not None: + c_net.latent_keyframe_override = latent_kf_override + if weights_override is not None: + c_net.weights_override = weights_override + # set cond hint mask + if mask_optional is not None: # if not in the form of a batch, make it so if len(mask_optional.shape) < 3: mask_optional = mask_optional.unsqueeze(0) @@ -168,18 +184,20 @@ NODE_CLASS_MAPPINGS = { "LatentKeyframeGroup": LatentKeyframeGroupNode, "LatentKeyframeBatchedGroup": LatentKeyframeBatchedGroupNode, "LatentKeyframeTiming": LatentKeyframeInterpolationNode, + # Conditioning + "ACN_AdvancedControlNetApply": AdvancedControlNetApply, # Loaders "ControlNetLoaderAdvanced": ControlNetLoaderAdvanced, "DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced, - # Conditioning - "ACN_AdvancedControlNetApply": AdvancedControlNetApply, # Weights + "ScaledSoftUniversalWeights": ScaledSoftUniversalWeights, "ScaledSoftControlNetWeights": ScaledSoftControlNetWeights, "SoftControlNetWeights": SoftControlNetWeights, "CustomControlNetWeights": CustomControlNetWeights, "SoftT2IAdapterWeights": SoftT2IAdapterWeights, "CustomT2IAdapterWeights": CustomT2IAdapterWeights, "ScaledSoftControlLoraWeights": ScaledSoftControlLoraWeights, + "ACN_DefaultUniversalWeights": DefaultWeights, # Image "LoadImagesFromDirectory": LoadImagesFromDirectory } @@ -191,18 +209,20 @@ NODE_DISPLAY_NAME_MAPPINGS = { "LatentKeyframeGroup": "Latent Keyframe Group 🛂🅐🅒🅝", "LatentKeyframeBatchedGroup": "Latent Keyframe Batched Group 🛂🅐🅒🅝", "LatentKeyframeTiming": "Latent Keyframe Interpolation 🛂🅐🅒🅝", + # Conditioning + "ACN_AdvancedControlNetApply": "Apply Advanced ControlNet 🛂🅐🅒🅝", # Loaders "ControlNetLoaderAdvanced": "Load ControlNet Model (Advanced) 🛂🅐🅒🅝", "DiffControlNetLoaderAdvanced": "Load ControlNet Model (diff Advanced) 🛂🅐🅒🅝", - # Conditioning - "ACN_AdvancedControlNetApply": "Apply Advanced ControlNet 🛂🅐🅒🅝", # Weights - "ScaledSoftControlNetWeights": "Scaled Soft ControlNet Weights 🛂🅐🅒🅝", - "SoftControlNetWeights": "Soft ControlNet Weights 🛂🅐🅒🅝", - "CustomControlNetWeights": "Custom ControlNet Weights 🛂🅐🅒🅝", - "SoftT2IAdapterWeights": "Soft T2IAdapter Weights 🛂🅐🅒🅝", - "CustomT2IAdapterWeights": "Custom T2IAdapter Weights 🛂🅐🅒🅝", - "ScaledSoftControlLoraWeights": "Scaled Soft ControlLora Weights 🛂🅐🅒🅝", + "ScaledSoftUniversalWeights": "Scaled Soft Weights 🛂🅐🅒🅝", + "ScaledSoftControlNetWeights": "ControlNet Scaled Soft Weights 🛂🅐🅒🅝", + "SoftControlNetWeights": "ControlNet Soft Weights 🛂🅐🅒🅝", + "CustomControlNetWeights": "ControlNet Custom Weights 🛂🅐🅒🅝", + "SoftT2IAdapterWeights": "T2IAdapter Soft Weights 🛂🅐🅒🅝", + "CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝", + "ScaledSoftControlLoraWeights": "ControlLora Scaled Soft Weights 🛂🅐🅒🅝", + "ACN_DefaultUniversalWeights": "Force Default Weights 🛂🅐🅒🅝", # Image "LoadImagesFromDirectory": "Load Images [DEPRECATED] 🛂🅐🅒🅝" } diff --git a/control/weight_nodes.py b/control/weight_nodes.py index 1bd2c37..8d95593 100644 --- a/control/weight_nodes.py +++ b/control/weight_nodes.py @@ -1,14 +1,41 @@ -from .control import TimestepKeyframe, TimestepKeyframeGroup +from .control import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights from .logger import logger -def get_properly_arranged_t2i_weights(initial_weights: list[float]): - new_weights = [] - new_weights.extend([initial_weights[0]]*3) - new_weights.extend([initial_weights[1]]*3) - new_weights.extend([initial_weights[2]]*3) - new_weights.extend([initial_weights[3]]*3) - return new_weights +class DefaultWeights: + @classmethod + def INPUT_TYPES(s): + return { + } + + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + FUNCTION = "load_weights" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + + def load_weights(self): + weights = ControlWeights.default() + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) + + +class ScaledSoftUniversalWeights: + @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 = ControlWeights.universal(base_multiplier=base_multiplier, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) class ScaledSoftControlLoraWeights: @@ -24,13 +51,12 @@ class ScaledSoftControlLoraWeights: RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) FUNCTION = "load_weights" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlLoRA" 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))) + weights = ControlWeights.controllora(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) class ScaledSoftControlNetWeights: @@ -46,13 +72,12 @@ class ScaledSoftControlNetWeights: RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) FUNCTION = "load_weights" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" def load_weights(self, base_multiplier, flip_weights): weights = [(base_multiplier ** float(12 - i)) for i in range(13)] - if flip_weights: - weights.reverse() - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_net_weights=weights))) + weights = ControlWeights.controlnet(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) class SoftControlNetWeights: @@ -80,15 +105,14 @@ class SoftControlNetWeights: RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) FUNCTION = "load_weights" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights): weights = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, weight_07, weight_08, weight_09, weight_10, weight_11, weight_12] - if flip_weights: - weights.reverse() - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_net_weights=weights))) + weights = ControlWeights.controlnet(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) class CustomControlNetWeights: @@ -116,15 +140,14 @@ class CustomControlNetWeights: RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) FUNCTION = "load_weights" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights): weights = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, weight_07, weight_08, weight_09, weight_10, weight_11, weight_12] - if flip_weights: - weights.reverse() - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_net_weights=weights))) + weights = ControlWeights.controlnet(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) class SoftT2IAdapterWeights: @@ -140,17 +163,16 @@ class SoftT2IAdapterWeights: }, } - RETURN_TYPES = ("T2I_ADAPTER_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) FUNCTION = "load_weights" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter" def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights): weights = [weight_00, weight_01, weight_02, weight_03] - if flip_weights: - weights.reverse() weights = get_properly_arranged_t2i_weights(weights) - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(t2i_adapter_weights=weights))) + weights = ControlWeights.t2iadapter(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) class CustomT2IAdapterWeights: @@ -166,14 +188,13 @@ class CustomT2IAdapterWeights: }, } - RETURN_TYPES = ("T2I_ADAPTER_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) FUNCTION = "load_weights" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter" def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights): weights = [weight_00, weight_01, weight_02, weight_03] - if flip_weights: - weights.reverse() weights = get_properly_arranged_t2i_weights(weights) - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(t2i_adapter_weights=weights))) + weights = ControlWeights.t2iadapter(weights, flip_weights=flip_weights) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) From 1e2d9a54916d19020b3ac744ff29473e5da9f1e5 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 29 Nov 2023 03:36:36 -0600 Subject: [PATCH 05/15] Added masks to TimestepKeyframes, fixed duplication bug between runs with LatentKeyframes and TimestepKeyframes --- control/control.py | 81 ++++++++++++++++++++++++++++++-- control/latent_keyframe_nodes.py | 8 ++++ control/nodes.py | 7 ++- 3 files changed, 92 insertions(+), 4 deletions(-) diff --git a/control/control.py b/control/control.py index f70900d..180c9e9 100644 --- a/control/control.py +++ b/control/control.py @@ -122,6 +122,12 @@ class LatentKeyframeGroup: def is_empty(self) -> bool: return len(self.keyframes) == 0 + def clone(self) -> 'LatentKeyframeGroup': + cloned = LatentKeyframeGroup() + for tk in self.keyframes: + cloned.add(tk) + return cloned + class TimestepKeyframe: def __init__(self, @@ -132,7 +138,8 @@ class TimestepKeyframe: latent_keyframes: LatentKeyframeGroup = None, null_latent_kf_strength: float = 0.0, inherit_missing: bool = True, - guarantee_usage: bool = False) -> None: + guarantee_usage: bool = True, + mask_hint_orig: Tensor = None) -> None: self.start_percent = start_percent self.start_t = 999999999.9 self.strength = strength @@ -142,6 +149,7 @@ class TimestepKeyframe: self.null_latent_kf_strength = null_latent_kf_strength self.inherit_missing = inherit_missing self.guarantee_usage = guarantee_usage + self.mask_hint_orig = mask_hint_orig def has_control_weights(self): return self.control_weights is not None @@ -149,6 +157,9 @@ class TimestepKeyframe: def has_latent_keyframes(self): return self.latent_keyframes is not None + def has_mask_hint(self): + return self.mask_hint_orig is not None + @classmethod def default(cls) -> 'TimestepKeyframe': @@ -191,6 +202,12 @@ class TimestepKeyframeGroup: def is_empty(self) -> bool: return len(self.keyframes) == 0 + def clone(self) -> 'TimestepKeyframeGroup': + cloned = TimestepKeyframeGroup() + for tk in self.keyframes: + cloned.add(tk) + return cloned + @classmethod def default(cls, keyframe: TimestepKeyframe) -> 'TimestepKeyframeGroup': group = cls() @@ -209,6 +226,9 @@ class AdvancedControlBase: # mask for which parts of controlnet output to keep self.mask_cond_hint_original = None self.mask_cond_hint = None + self.tk_mask_cond_hint_original = None + self.tk_mask_cond_hint = None + self.weight_mask_cond_hint = None # actual index values self.sub_idxs = None self.full_latent_length = 0 @@ -267,6 +287,11 @@ class AdvancedControlBase: self.latent_keyframes = self.current_timestep_keyframe.latent_keyframes elif not self.current_timestep_keyframe.inherit_missing: self.latent_keyframes = None + if self.current_timestep_keyframe.has_mask_hint(): + self.tk_mask_cond_hint_original = self.current_timestep_keyframe.mask_hint_orig + elif not self.current_timestep_keyframe.inherit_missing: + del self.tk_mask_cond_hint_original + self.tk_mask_cond_hint_original = None # if guarantee_usage, stop searching for other TKs if self.current_timestep_keyframe.guarantee_usage: break @@ -365,11 +390,13 @@ class AdvancedControlBase: # apply null for each batched cond/uncond for b in range(batched_number): x[(latent_count*b)+batch_index] = x[(latent_count*b)+batch_index] * self.current_timestep_keyframe.null_latent_kf_strength - # apply masks + # apply masks, resizing mask to required dims if self.mask_cond_hint is not None: - # first, resize mask to required dims masks = prepare_mask_batch(self.mask_cond_hint, x.shape) x[:] = x[:] * masks + if self.tk_mask_cond_hint is not None: + masks = prepare_mask_batch(self.tk_mask_cond_hint, x.shape) + x[:] = x[:] * masks # apply timestep keyframe strengths if self.current_timestep_keyframe.strength != 1.0: x[:] *= self.current_timestep_keyframe.strength @@ -444,6 +471,41 @@ class AdvancedControlBase: if dtype is None: dtype = x_noisy.dtype self.mask_cond_hint = self.mask_cond_hint.to(dtype=dtype).to(self.device) + # prepare other masks + self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype) + + def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): + return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype) + + def prepare_weight_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): + return self._prepare_mask("weight_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype) + + def _prepare_mask(self, attr_name, orig_mask: Tensor, x_noisy: Tensor, t, cond, batched_number, dtype=None): + if orig_mask is not None: + out_mask = getattr(self, attr_name) + if self.sub_idxs is not None or out_mask is None or x_noisy.shape[2] * 8 != out_mask.shape[1] or x_noisy.shape[3] * 8 != out_mask.shape[2]: + self._reset_attr(attr_name) + del out_mask + # TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM + # resize mask and match batch count + out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=8) + actual_latent_length = x_noisy.shape[0] // batched_number + out_mask = comfy.utils.repeat_to_batch_size(out_mask, actual_latent_length if self.sub_idxs is None else self.full_latent_length) + if self.sub_idxs is not None: + out_mask = out_mask[self.sub_idxs] + # make cond_hint_mask length match x_noise + if x_noisy.shape[0] != out_mask.shape[0]: + out_mask = broadcast_image_to(out_mask, x_noisy.shape[0], batched_number) + # default dtype to be same as x_noisy + if dtype is None: + dtype = x_noisy.dtype + setattr(self, attr_name, out_mask.to(dtype=dtype).to(self.device)) + del out_mask + + def _reset_attr(self, attr_name, new_value=None): + if hasattr(self, attr_name): + delattr(self, attr_name) + setattr(self, attr_name, new_value) def cleanup_inject(self): self.base.cleanup() @@ -461,6 +523,19 @@ class AdvancedControlBase: self.current_timestep_keyframe = None self.next_timestep_keyframe = None self.current_timestep_index = -1 + # clear mask hints + if self.mask_cond_hint is not None: + del self.mask_cond_hint + self.mask_cond_hint = None + if self.tk_mask_cond_hint_original is not None: + del self.tk_mask_cond_hint_original + self.tk_mask_cond_hint_original = None + if self.tk_mask_cond_hint is not None: + del self.tk_mask_cond_hint + self.tk_mask_cond_hint = None + if self.weight_mask_cond_hint is not None: + del self.weight_mask_cond_hint + self.weight_mask_cond_hint = None def copy_to_advanced(self, copied: 'AdvancedControlBase'): copied.mask_cond_hint_original = self.mask_cond_hint_original diff --git a/control/latent_keyframe_nodes.py b/control/latent_keyframe_nodes.py index 8152921..d9b5856 100644 --- a/control/latent_keyframe_nodes.py +++ b/control/latent_keyframe_nodes.py @@ -31,6 +31,8 @@ class LatentKeyframeNode: prev_latent_keyframe: LatentKeyframeGroup=None): if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() keyframe = LatentKeyframe(batch_index, strength) prev_latent_keyframe.add(keyframe) return (prev_latent_keyframe,) @@ -126,6 +128,8 @@ class LatentKeyframeGroupNode: print_keyframes=False): if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() curr_latent_keyframe = LatentKeyframeGroup() latent_count = -1 @@ -185,6 +189,8 @@ class LatentKeyframeInterpolationNode: if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() curr_latent_keyframe = LatentKeyframeGroup() steps = batch_index_to_excl - batch_index_from @@ -236,6 +242,8 @@ class LatentKeyframeBatchedGroupNode: 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() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() curr_latent_keyframe = LatentKeyframeGroup() # if received a normal float input, do nothing diff --git a/control/nodes.py b/control/nodes.py index e3803b3..3d03faf 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -27,6 +27,7 @@ class TimestepKeyframeNode: "null_latent_kf_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "inherit_missing": ("BOOLEAN", {"default": True}, ), "guarantee_usage": ("BOOLEAN", {"default": True}, ), + "mask_optional": ("MASK", ), #"interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT, SI.NONE], {"default": SI.NONE}, ), } } @@ -46,11 +47,15 @@ class TimestepKeyframeNode: null_latent_kf_strength: float=0.0, inherit_missing=True, guarantee_usage=True, + mask_optional=None, interpolation: str=SI.NONE,): if not prev_timestep_keyframe: prev_timestep_keyframe = TimestepKeyframeGroup() + else: + prev_timestep_keyframe = prev_timestep_keyframe.clone() keyframe = TimestepKeyframe(start_percent=start_percent, strength=strength, interpolation=interpolation, null_latent_kf_strength=null_latent_kf_strength, - control_weights=control_net_weights, latent_keyframes=latent_keyframe, inherit_missing=inherit_missing, guarantee_usage=guarantee_usage) + control_weights=control_net_weights, latent_keyframes=latent_keyframe, inherit_missing=inherit_missing, guarantee_usage=guarantee_usage, + mask_hint_orig=mask_optional) prev_timestep_keyframe.add(keyframe) return (prev_timestep_keyframe,) From 583d2c59b77c2dc44baa50a5e0cccceea05b6430 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 29 Nov 2023 11:39:23 -0600 Subject: [PATCH 06/15] Scaled Soft Mask Weights added, fixed latent keyframe error when batch_index out of range, removed duplicate weight nodes, renamed TK outputs on weight nodes --- control/control.py | 75 +++++++++++++++++++----------------- control/nodes.py | 22 +++++------ control/weight_nodes.py | 85 ++++++++++++++++++++--------------------- 3 files changed, 93 insertions(+), 89 deletions(-) diff --git a/control/control.py b/control/control.py index 180c9e9..7607112 100644 --- a/control/control.py +++ b/control/control.py @@ -35,11 +35,7 @@ class ControlWeights: self.weights.reverse() self.weight_mask = weight_mask - def get(self, idx: int, shape: Tensor) -> Union[float, Tensor]: - # if weight_mask present, return normalized mask - if self.base_multiplier != 1.0 and self.weight_mask is not None: - # TODO: fill out - pass + def get(self, idx: int) -> Union[float, Tensor]: # if weights is not none, return index if self.weights is not None: return self.weights[idx] @@ -53,6 +49,10 @@ class ControlWeights: def universal(cls, base_multiplier: float, flip_weights: bool=False): return cls(ControlWeightType.UNIVERSAL, base_multiplier=base_multiplier, flip_weights=flip_weights) + @classmethod + def universal_mask(cls, weight_mask: Tensor): + return cls(ControlWeightType.UNIVERSAL, weight_mask=weight_mask) + @classmethod def t2iadapter(cls, weights: list[float]=None, flip_weights: bool=False): if weights is None: @@ -268,7 +268,6 @@ class AdvancedControlBase: self.batched_number = batched_number # get current step percent curr_t: float = t[0] - print(f"$$$$ {curr_t}") prev_index = self.current_timestep_index # if has next index, loop through and see if need to switch if self.timestep_keyframes.has_index(self.current_timestep_index+1): @@ -278,7 +277,8 @@ class AdvancedControlBase: if eval_tk.start_t >= curr_t: self.current_timestep_index = i self.current_timestep_keyframe = eval_tk - # keep track of weights and control weights, accounting for inherit_missing + # keep track of control weights, latent keyframes, and masks, + # accounting for inherit_missing if self.current_timestep_keyframe.has_control_weights(): self.weights = self.current_timestep_keyframe.control_weights elif not self.current_timestep_keyframe.inherit_missing: @@ -314,6 +314,9 @@ class AdvancedControlBase: if self.weights is None or self.weights.weight_type == ControlWeightType.DEFAULT: self.weights = self.weights_default elif self.weights.weight_type == ControlWeightType.UNIVERSAL: + # if universal and weight_mask present, no need to convert + if self.weights.weight_mask is not None: + return self.weights = self.get_universal_weights() def get_universal_weights(self) -> ControlWeights: @@ -352,6 +355,14 @@ class AdvancedControlBase: def get_control_advanced(self, x_noisy, t, cond, batched_number): pass + def calc_weight(self, idx: int, x: Tensor, layers: int) -> Union[float, Tensor]: + if self.weights.weight_mask is not None: + # prepare weight mask + self.prepare_weight_mask_cond_hint(x, self.batched_number) + # adjust mask for current layer and return + return torch.pow(self.weight_mask_cond_hint, (layers-1)-idx) + return self.weights.get(idx=idx) + def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int): # apply strengths, and get batch indeces to null out # AKA latents that should not be influenced by ControlNet @@ -381,6 +392,10 @@ class AdvancedControlBase: continue indeces_to_null.remove(real_index) + # if real_index is outside the bounds of latents, don't apply + if real_index >= latent_count or real_index < 0: + continue + # apply strength for each batched cond/uncond for b in range(batched_number): x[(latent_count*b)+real_index] = x[(latent_count*b)+real_index] * keyframe.strength @@ -411,7 +426,7 @@ class AdvancedControlBase: if x is not None: self.apply_advanced_strengths_and_masks(x, self.batched_number) - x *= self.strength * self.weights.get(i, x.shape) + x *= self.strength * self.calc_weight(i, x, len(control_input)) if x.dtype != output_dtype: x = x.to(output_dtype) out[key].insert(0, x) @@ -431,7 +446,7 @@ class AdvancedControlBase: if self.global_average_pooling: x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3]) - x *= self.strength * self.weights.get(i, x.shape) + x *= self.strength * self.calc_weight(i, x, len(control_output)) if x.dtype != output_dtype: x = x.to(output_dtype) @@ -451,36 +466,17 @@ class AdvancedControlBase: return out 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) - # prepare other masks + self._prepare_mask("mask_cond_hint", self.mask_cond_hint_original, x_noisy, t, cond, batched_number, dtype) self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype) def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype) - def prepare_weight_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): - return self._prepare_mask("weight_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype) + def prepare_weight_mask_cond_hint(self, x_noisy: Tensor, batched_number, dtype=None): + return self._prepare_mask("weight_mask_cond_hint", self.weights.weight_mask, x_noisy, t=None, cond=None, batched_number=batched_number, dtype=dtype, direct_attn=True) - def _prepare_mask(self, attr_name, orig_mask: Tensor, x_noisy: Tensor, t, cond, batched_number, dtype=None): + def _prepare_mask(self, attr_name, orig_mask: Tensor, x_noisy: Tensor, t, cond, batched_number, dtype=None, direct_attn=False): + # make mask appropriate dimensions, if present if orig_mask is not None: out_mask = getattr(self, attr_name) if self.sub_idxs is not None or out_mask is None or x_noisy.shape[2] * 8 != out_mask.shape[1] or x_noisy.shape[3] * 8 != out_mask.shape[2]: @@ -488,7 +484,8 @@ class AdvancedControlBase: del out_mask # TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM # resize mask and match batch count - out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=8) + multiplier = 1 if direct_attn else 8 + out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=multiplier) actual_latent_length = x_noisy.shape[0] // batched_number out_mask = comfy.utils.repeat_to_batch_size(out_mask, actual_latent_length if self.sub_idxs is None else self.full_latent_length) if self.sub_idxs is not None: @@ -734,3 +731,13 @@ def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim if match_dim1: mask = torch.cat([mask] * shape[1], dim=1) return mask + + +# applies min-max normalization, from: +# https://stackoverflow.com/questions/68791508/min-max-normalization-of-a-tensor-in-pytorch +def normalize_min_max(x: Tensor, new_min = 0.0, new_max = 1.0): + x_min, x_max = x.min(), x.max() + return (((x - x_min)/(x_max - x_min)) * (new_max - new_min)) + new_min + +def linear_conversion(x, x_min=0.0, x_max=1.0, new_min=0.0, new_max=1.0): + return (((x - x_min)/(x_max - x_min)) * (new_max - new_min)) + new_min diff --git a/control/nodes.py b/control/nodes.py index 3d03faf..f4c194c 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -1,11 +1,12 @@ import numpy as np +from torch import Tensor import folder_paths from .control import load_controlnet, convert_to_advanced, ControlWeights, ControlWeightType,\ LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup, is_advanced_controlnet from .control import StrengthInterpolation as SI -from .weight_nodes import DefaultWeights, ScaledSoftControlLoraWeights, ScaledSoftControlNetWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, \ +from .weight_nodes import DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, \ SoftT2IAdapterWeights, CustomT2IAdapterWeights from .latent_keyframe_nodes import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode from .deprecated_nodes import LoadImagesFromDirectory @@ -124,7 +125,7 @@ class AdvancedControlNetApply: "mask_optional": ("MASK", ), "timestep_kf": ("TIMESTEP_KEYFRAME", ), "latent_kf_override": ("LATENT_KEYFRAME", ), - "cn_weights_override": ("CONTROL_NET_WEIGHTS", ), + "weights_override": ("CONTROL_NET_WEIGHTS", ), } } @@ -135,7 +136,7 @@ class AdvancedControlNetApply: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, - mask_optional=None, + mask_optional: Tensor=None, timestep_kf: TimestepKeyframeGroup=None, latent_kf_override: LatentKeyframeGroup=None, weights_override: ControlWeights=None): if strength == 0: @@ -166,6 +167,7 @@ class AdvancedControlNetApply: c_net.weights_override = weights_override # set cond hint mask if mask_optional is not None: + mask_optional = mask_optional.clone() # if not in the form of a batch, make it so if len(mask_optional.shape) < 3: mask_optional = mask_optional.unsqueeze(0) @@ -195,13 +197,12 @@ NODE_CLASS_MAPPINGS = { "ControlNetLoaderAdvanced": ControlNetLoaderAdvanced, "DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced, # Weights - "ScaledSoftUniversalWeights": ScaledSoftUniversalWeights, - "ScaledSoftControlNetWeights": ScaledSoftControlNetWeights, + "ScaledSoftControlNetWeights": ScaledSoftUniversalWeights, + "ScaledSoftMaskedUniversalWeights": ScaledSoftMaskedUniversalWeights, "SoftControlNetWeights": SoftControlNetWeights, "CustomControlNetWeights": CustomControlNetWeights, "SoftT2IAdapterWeights": SoftT2IAdapterWeights, "CustomT2IAdapterWeights": CustomT2IAdapterWeights, - "ScaledSoftControlLoraWeights": ScaledSoftControlLoraWeights, "ACN_DefaultUniversalWeights": DefaultWeights, # Image "LoadImagesFromDirectory": LoadImagesFromDirectory @@ -217,16 +218,15 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Conditioning "ACN_AdvancedControlNetApply": "Apply Advanced ControlNet 🛂🅐🅒🅝", # Loaders - "ControlNetLoaderAdvanced": "Load ControlNet Model (Advanced) 🛂🅐🅒🅝", - "DiffControlNetLoaderAdvanced": "Load ControlNet Model (diff Advanced) 🛂🅐🅒🅝", + "ControlNetLoaderAdvanced": "Load Advanced ControlNet Model 🛂🅐🅒🅝", + "DiffControlNetLoaderAdvanced": "Load Advanced ControlNet Model (diff) 🛂🅐🅒🅝", # Weights - "ScaledSoftUniversalWeights": "Scaled Soft Weights 🛂🅐🅒🅝", - "ScaledSoftControlNetWeights": "ControlNet Scaled Soft Weights 🛂🅐🅒🅝", + "ScaledSoftControlNetWeights": "Scaled Soft Weights 🛂🅐🅒🅝", + "ScaledSoftMaskedUniversalWeights": "Scaled Soft Masked Weights 🛂🅐🅒🅝", "SoftControlNetWeights": "ControlNet Soft Weights 🛂🅐🅒🅝", "CustomControlNetWeights": "ControlNet Custom Weights 🛂🅐🅒🅝", "SoftT2IAdapterWeights": "T2IAdapter Soft Weights 🛂🅐🅒🅝", "CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝", - "ScaledSoftControlLoraWeights": "ControlLora Scaled Soft Weights 🛂🅐🅒🅝", "ACN_DefaultUniversalWeights": "Force Default Weights 🛂🅐🅒🅝", # Image "LoadImagesFromDirectory": "Load Images [DEPRECATED] 🛂🅐🅒🅝" diff --git a/control/weight_nodes.py b/control/weight_nodes.py index 8d95593..1245397 100644 --- a/control/weight_nodes.py +++ b/control/weight_nodes.py @@ -1,7 +1,11 @@ -from .control import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights +from torch import Tensor +from .control import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights, linear_conversion from .logger import logger +WEIGHTS_RETURN_NAMES = ("CONTROL_NET_WEIGHTS", "TK_SHORTCUT") + + class DefaultWeights: @classmethod def INPUT_TYPES(s): @@ -9,6 +13,7 @@ class DefaultWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" @@ -18,17 +23,47 @@ class DefaultWeights: return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) +class ScaledSoftMaskedUniversalWeights: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "mask": ("MASK", ), + "min_base_multiplier": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), + "max_base_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), + #"lock_min": ("BOOLEAN", {"default": False}, ), + #"lock_max": ("BOOLEAN", {"default": False}, ), + }, + } + + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES + FUNCTION = "load_weights" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + + def load_weights(self, mask: Tensor, min_base_multiplier: float, max_base_multiplier: float, lock_min=False, lock_max=False): + # normalize mask + mask = mask.clone() + x_min = 0.0 if lock_min else mask.min() + x_max = 1.0 if lock_max else mask.max() + mask = linear_conversion(mask, x_min, x_max, min_base_multiplier, max_base_multiplier) + weights = ControlWeights.universal_mask(weight_mask=mask) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) + + class ScaledSoftUniversalWeights: @classmethod def INPUT_TYPES(s): return { "required": { - "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 1.0, "step": 0.001}, ), "flip_weights": ("BOOLEAN", {"default": False}), }, } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" @@ -38,48 +73,6 @@ class ScaledSoftUniversalWeights: return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=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/ControlLoRA" - - def load_weights(self, base_multiplier, flip_weights): - weights = [(base_multiplier ** float(9 - i)) for i in range(10)] - weights = ControlWeights.controllora(weights, flip_weights=flip_weights) - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) - - -class ScaledSoftControlNetWeights: - @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/ControlNet" - - def load_weights(self, base_multiplier, flip_weights): - weights = [(base_multiplier ** float(12 - i)) for i in range(13)] - weights = ControlWeights.controlnet(weights, flip_weights=flip_weights) - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) - - class SoftControlNetWeights: @classmethod def INPUT_TYPES(s): @@ -103,6 +96,7 @@ class SoftControlNetWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" @@ -138,6 +132,7 @@ class CustomControlNetWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" @@ -164,6 +159,7 @@ class SoftT2IAdapterWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter" @@ -189,6 +185,7 @@ class CustomT2IAdapterWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter" From db7c535dc6d7678580a25b02e5403a6bb390b874 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 29 Nov 2023 11:53:05 -0600 Subject: [PATCH 07/15] Fixed T2IAdvanced trying to copy itself as a ControlLoraAdvanced --- control/control.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/control/control.py b/control/control.py index 7607112..9474185 100644 --- a/control/control.py +++ b/control/control.py @@ -643,7 +643,7 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): del full_cond_hint_original def copy(self): - c = ControlLoraAdvanced(self.t2i_model, self.timestep_keyframes, self.channels_in) + c = T2IAdapterAdvanced(self.t2i_model, self.timestep_keyframes, self.channels_in) self.copy_to(c) self.copy_to_advanced(c) return c From 6c8216a8891d555132a7d61cc476b51b0b5528e0 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 29 Nov 2023 13:57:25 -0600 Subject: [PATCH 08/15] Verify weights are compatible with loaded CN, fixed Scaled Soft Masked Weights in the case of a mask tensor with uniform max value, made Scaled Soft Masked Weights use same scaling for T2IAdapter as normal Scaled Soft Weights --- control/control.py | 35 +++++++++++++++++++++++++++++++---- control/nodes.py | 4 ++++ control/weight_nodes.py | 6 +++++- 3 files changed, 40 insertions(+), 5 deletions(-) diff --git a/control/control.py b/control/control.py index 9474185..18cbffd 100644 --- a/control/control.py +++ b/control/control.py @@ -254,6 +254,21 @@ class AdvancedControlBase: def add_compatible_weight(self, control_weight_type: str): self.compatible_weights.append(control_weight_type) + def verify_all_weights(self, throw_error=True): + # first, check if override exists - if so, only need to check the override + if self.weights_override is not None: + if self.weights_override.weight_type not in self.compatible_weights: + msg = f"Weight override is type {self.weights_override.weight_type}, but loaded {type(self).__name__}" + \ + f"only supports {self.compatible_weights} weights." + raise WeightTypeException(msg) + # otherwise, check all timestep keyframe weights + else: + for tk in self.timestep_keyframes.keyframes: + if tk.has_control_weights() and tk.control_weights.weight_type not in self.compatible_weights: + msg = f"Weight on Timestep Keyframe with start_percent={tk.start_percent} is type" + \ + f"{tk.control_weights.weight_type}, but loaded {type(self).__name__} only supports {self.compatible_weights} weights." + raise WeightTypeException(msg) + def set_timestep_keyframes(self, timestep_keyframes: TimestepKeyframeGroup): self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup() # prepare first timestep_keyframe related stuff @@ -360,8 +375,11 @@ class AdvancedControlBase: # prepare weight mask self.prepare_weight_mask_cond_hint(x, self.batched_number) # adjust mask for current layer and return - return torch.pow(self.weight_mask_cond_hint, (layers-1)-idx) + return torch.pow(self.weight_mask_cond_hint, self.get_calc_pow(idx=idx, layers=layers)) return self.weights.get(idx=idx) + + def get_calc_pow(self, idx: int, layers: int) -> int: + return (layers-1)-idx def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int): # apply strengths, and get batch indeces to null out @@ -547,7 +565,6 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase): def get_universal_weights(self) -> ControlWeights: raw_weights = [(self.weights.base_multiplier ** float(12 - i)) for i in range(13)] - # TODO: account for masks? return ControlWeights.controlnet(raw_weights, self.weights.flip_weights) def get_control_advanced(self, x_noisy, t, cond, batched_number): @@ -619,9 +636,15 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): raw_weights = [(self.weights.base_multiplier ** float(7 - i)) for i in range(8)] raw_weights = [raw_weights[-8], raw_weights[-3], raw_weights[-2], raw_weights[-1]] raw_weights = get_properly_arranged_t2i_weights(raw_weights) - # TODO: account for masks? return ControlWeights.t2iadapter(raw_weights, self.weights.flip_weights) + def get_calc_pow(self, idx: int, layers: int) -> int: + # match how T2IAdapterAdvanced deals with universal weights + indeces = [7 - i for i in range(8)] + indeces = [indeces[-8], indeces[-3], indeces[-2], indeces[-1]] + indeces = get_properly_arranged_t2i_weights(indeces) + return indeces[idx] + def get_control_advanced(self, x_noisy, t, cond, batched_number): # prepare timestep and everything related self.prepare_current_timestep(t=t, batched_number=batched_number) @@ -667,7 +690,6 @@ class ControlLoraAdvanced(ControlLora, AdvancedControlBase): def get_universal_weights(self) -> ControlWeights: raw_weights = [(self.weights.base_multiplier ** float(9 - i)) for i in range(10)] - # TODO: account for masks? return ControlWeights.controllora(raw_weights, self.weights.flip_weights) def copy(self): @@ -741,3 +763,8 @@ def normalize_min_max(x: Tensor, new_min = 0.0, new_max = 1.0): def linear_conversion(x, x_min=0.0, x_max=1.0, new_min=0.0, new_max=1.0): return (((x - x_min)/(x_max - x_min)) * (new_max - new_min)) + new_min + + +class WeightTypeException(TypeError): + "Raised when weight not compatible with AdvancedControlBase object" + pass diff --git a/control/nodes.py b/control/nodes.py index f4c194c..edaa220 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -105,6 +105,8 @@ class DiffControlNetLoaderAdvanced: def load_controlnet(self, control_net_name, model, timestep_keyframe: TimestepKeyframeGroup=None): controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) controlnet = load_controlnet(controlnet_path, timestep_keyframe, model) + if is_advanced_controlnet(controlnet): + controlnet.verify_all_weights() return (controlnet,) @@ -165,6 +167,8 @@ class AdvancedControlNetApply: c_net.latent_keyframe_override = latent_kf_override if weights_override is not None: c_net.weights_override = weights_override + # verify weights are compatible + c_net.verify_all_weights() # set cond hint mask if mask_optional is not None: mask_optional = mask_optional.clone() diff --git a/control/weight_nodes.py b/control/weight_nodes.py index 1245397..f4cc9f0 100644 --- a/control/weight_nodes.py +++ b/control/weight_nodes.py @@ -1,4 +1,5 @@ from torch import Tensor +import torch from .control import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights, linear_conversion from .logger import logger @@ -47,7 +48,10 @@ class ScaledSoftMaskedUniversalWeights: mask = mask.clone() x_min = 0.0 if lock_min else mask.min() x_max = 1.0 if lock_max else mask.max() - mask = linear_conversion(mask, x_min, x_max, min_base_multiplier, max_base_multiplier) + if x_min == x_max: + mask = torch.ones_like(mask) * max_base_multiplier + else: + mask = linear_conversion(mask, x_min, x_max, min_base_multiplier, max_base_multiplier) weights = ControlWeights.universal_mask(weight_mask=mask) return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) From fb38e126ef2304690d1d95394e102bcab9aef66c Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 30 Nov 2023 04:13:50 -0600 Subject: [PATCH 09/15] Made strength an optional param to keep old workflows from needing the Timestep Keyframe node to be manually replaced --- control/nodes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/control/nodes.py b/control/nodes.py index edaa220..c7a01ab 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -19,9 +19,9 @@ class TimestepKeyframeNode: return { "required": { "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), - "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), }, "optional": { + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "control_net_weights": ("CONTROL_NET_WEIGHTS", ), "latent_keyframe": ("LATENT_KEYFRAME", ), "prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ), From 7e2d8b08fcaa89be2a5994a2cf7df185b266af18 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 30 Nov 2023 04:36:21 -0600 Subject: [PATCH 10/15] Moved prev_timestep_keyframe param on TimestepKeyframeNode to the top to make it less annoying to connect adjacent keyframes --- control/nodes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/control/nodes.py b/control/nodes.py index c7a01ab..829f53d 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -21,10 +21,10 @@ class TimestepKeyframeNode: "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), }, "optional": { + "prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ), "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "control_net_weights": ("CONTROL_NET_WEIGHTS", ), "latent_keyframe": ("LATENT_KEYFRAME", ), - "prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ), "null_latent_kf_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "inherit_missing": ("BOOLEAN", {"default": True}, ), "guarantee_usage": ("BOOLEAN", {"default": True}, ), From 772bef72d22634031c4acfdf918b941e0accd112 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 30 Nov 2023 07:51:56 -0600 Subject: [PATCH 11/15] Update README.md - starting to add stuff --- README.md | 74 +++++++++++++++++++++++++++++++++++++++++++++++++++---- 1 file changed, 69 insertions(+), 5 deletions(-) diff --git a/README.md b/README.md index ed3607b..ca5a5b3 100644 --- a/README.md +++ b/README.md @@ -1,11 +1,75 @@ # ComfyUI-Advanced-ControlNet -These custom nodes allow for scheduling ControlNet strength across latents in the same batch (WORKING) and across timesteps (IN PROGRESS). +Nodes for scheduling ControlNet strength across timesteps and batched latents, as well as applying custom weights and attention masks. The ControlNet nodes here fully support sliding context sampling, like the one used in the [ComfyUI-AnimateDiff-Evolved](https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved) nodes. Currently supports ControlNets, T2IAdapters, and ControlLoRAs. Kohya Controllllite support coming soon. -Custom weights can also be applied to ControlNets and T2IAdapters to mimic the "My prompt is more important" functionality in AUTOMATIC1111's ControlNet extension. +Custom weights allow replication of the "My prompt is more important" feature of Auto1111's sd-webui ControlNet extension. -TODO: -- Other handy nodes -- Finish and update this README for other workflows +ControlNet preprocessors are available through [comfyui_controlnet_aux](https://github.com/Fannovel16/comfyui_controlnet_aux) nodes + +## Features +- Timestep and latent strength scheduling +- Attention masks +- Soft weights to replicate "My prompt is more important" feature from sd-webui ControlNet extension, and also change the scaling. +- ControlNet, T2IAdapter, and ControlLoRA support for sliding context windows. + +## Table of Contents: +- [Scheduling Explanation](#scheduling-explanation) +- [Nodes](#nodes) +- [Usage](#usage) + + +# Scheduling Explanation + +The two core concepts for scheduling are ***Timestep Keyframes*** and ***Latent Keyframes***. + +***Timestep Keyframes*** hold the values that guide the settings for a controlnet, and begin to take effect based on their start_percent, which corresponds to the percentage of the sampling process. They can contain masks for the strengths of each latent, control_net_weights, and latent_keyframes (specific strengths for each latent), all optional. + +***Latent Keyframes*** determine the strength of the controlnet for specific latents - all they contain is the batch_index of the latent, and the strength the controlnet should apply for that latent. As a concept, latent keyframes achieve the same affect as a uniform mask with the chosen strength value. + +![advcn_image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/e6275264-6c3f-4246-a319-111ee48f4cd9) + +# Nodes + +The ControlNet nodes provided here are the ***Apply Advanced ControlNet*** and ***Load Advanced ControlNet Model*** (or diff) nodes. The vanilla ControlNet nodes are also compatible, and can be used almost interchangeably - the only difference is that **at least one of these nodes must be used** for Advanced versions of ControlNets to be used (important for sliding context sampling, like with AnimateDiff-Evolved). + +Key: + +- 🟩 - required inputs +- 🟨 - optional inputs +- 🟥 - optional input/output, but not recommended to use unless needed +- 🟦 - start as widgets, can be converted to inputs + +## Apply Advanced ControlNet +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/dc541d41-70df-4a71-b832-efa65af98f06) + +Same functionality as the vanilla Apply Advanced ControlNet (Advanced) node, except with Advanced ControlNet features added to it. Automatically converts any ControlNet from ControlNet loaders into Advanced versions. + +### Inputs +- 🟩***positive***: conditioning (positive). +- 🟩***negative***: conditioning (negative). +- 🟩***control_net***: loaded controlnet; will be converted to Advanced version automatically by this node, if it's a supported type. +- 🟩***image***: images to guide controlnets - if the loaded controlnet requires it, they must preprocessed images. If one image provided, will be used for all latents. If more images provided, will use each image separately for each latent. If not enough images to meet latent count, will repeat the images from the beginning to match vanilla ControlNet functionality. +- 🟨***mask_optional***: attention masks to apply to controlnets; basically, decides what part of the image the controlnet to apply to (and the relative strength, if the mask is not binary). Same as image input, if you provide more than one mask, each can apply to a different latent. +- 🟨***timestep_kf***: timestep keyframes to guide controlnet effect throughout sampling steps. +- 🟨***latent_kf_override***: override for latent keyframes, useful if no other features from timestep keyframes is needed. *NOTE: this latent keyframe will be applied to ALL timesteps, regardless if there are other latent keyframes attached to connected timestep keyframes.* +- 🟨***weights_override***: override for weights, useful if no other features from timestep keyframes is needed. *NOTE: this weight will be applied to ALL timesteps, regardless if there are other weights attached to connected timestep keyframes.* +- 🟦***strength***: strength of controlnet; 1.0 is full strength, 0.0 is no effect at all. +- 🟦***start_percent***: sampling step percentage at which controlnet should start to be applied - no matter what start_percent is set on timestep keyframes, they won't take effect until this start_percent is reached. +- 🟦***stop_percent***: sampling step percentage at which controlnet should stop being applied - no matter what start_percent is set on timestep keyframes, they won't take effect once this end_percent is reached. + +### Outputs +- ***positive***: conditioning (positive) with applied controlnets +- ***negative***: conditioning (negative) with applied controlnets + +## Load Advanced ControlNet Model +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/4a7f58a9-783d-4da4-bf82-bc9c167e4722) + +Loads a ControlNet model and converts it into an Advanced version that supports all the features in this repo. When used with **Apply Advanced ControlNet** node, there is no reason to use the timestep_keyframe input on this node - use timestep_kf on the Apply node instead. + +### Inputs +- 🟥***timestep_keyframe***: optional and likely unnecessary input to have ControlNet use selected timestep_keyframes - should not be used unless you need to. Useful if this node is not attached to **Apply Advanced ControlNet** node, but still want to use Timestep Keyframe, or to use TK_SHORTCUT outputs from ControlWeights in the same scenario. Will be overriden by the timestep_kf input on **Apply Advanced ControlNet** node, if one is provided there. +- 🟨***model***: model to plug into the diff version of the node. Some controlnets are designed for receive the model; if you don't know what this does, you probably don't want tot use the diff version of the node. + +### Outputs ## Workflows From 941cd5c611e0e32b1932b02f4e3d20f9a04f0210 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 30 Nov 2023 10:01:04 -0600 Subject: [PATCH 12/15] Update README.md - more changes --- README.md | 42 ++++++++++++++++++++++++++++++++++++++---- 1 file changed, 38 insertions(+), 4 deletions(-) diff --git a/README.md b/README.md index ca5a5b3..ce38f68 100644 --- a/README.md +++ b/README.md @@ -32,11 +32,11 @@ The two core concepts for scheduling are ***Timestep Keyframes*** and ***Latent The ControlNet nodes provided here are the ***Apply Advanced ControlNet*** and ***Load Advanced ControlNet Model*** (or diff) nodes. The vanilla ControlNet nodes are also compatible, and can be used almost interchangeably - the only difference is that **at least one of these nodes must be used** for Advanced versions of ControlNets to be used (important for sliding context sampling, like with AnimateDiff-Evolved). Key: - - 🟩 - required inputs - 🟨 - optional inputs -- 🟥 - optional input/output, but not recommended to use unless needed - 🟦 - start as widgets, can be converted to inputs +- 🟥 - optional input/output, but not recommended to use unless needed +- 🟪 - output ## Apply Advanced ControlNet ![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/dc541d41-70df-4a71-b832-efa65af98f06) @@ -57,8 +57,8 @@ Same functionality as the vanilla Apply Advanced ControlNet (Advanced) node, exc - 🟦***stop_percent***: sampling step percentage at which controlnet should stop being applied - no matter what start_percent is set on timestep keyframes, they won't take effect once this end_percent is reached. ### Outputs -- ***positive***: conditioning (positive) with applied controlnets -- ***negative***: conditioning (negative) with applied controlnets +- 🟪***positive***: conditioning (positive) with applied controlnets +- 🟪***negative***: conditioning (negative) with applied controlnets ## Load Advanced ControlNet Model ![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/4a7f58a9-783d-4da4-bf82-bc9c167e4722) @@ -70,6 +70,40 @@ Loads a ControlNet model and converts it into an Advanced version that supports - 🟨***model***: model to plug into the diff version of the node. Some controlnets are designed for receive the model; if you don't know what this does, you probably don't want tot use the diff version of the node. ### Outputs +- 🟪***CONTROL_NET***: loaded Advanced ControlNet + +## Timestep Keyframe +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/3950d14d-6757-42f2-a615-9ba81a74e0ca) + +Scheduling node across timesteps (sampling steps) based on the set start_percent. Chaining Timestep Keyframes allows ControlNet scheduling across sampling steps (percentage-wise), through a timestep keyframe schedule. + +### Inputs +- 🟨***prev_timestep_keyframe***: used to chain Timestep Keyframes together to create a schedule. The order does not matter - the Timestep Keyframes sort themselves automatically by their start_percent. *Any Timestep Keyframe contained in the prev_timestep_keyframe that contains the same start_percent as the Timestep Keyframe will be overwritten.* +- 🟨***control_net_weights***: weights to apply to controlnet while this Timestep Keyframe is in effect. Must be compatible with the loaded controlnet, or will throw an error explaining what weight types are compatible. If inherit_missing is True, if no control_net_weight is passed in, will attempt to reuse the last-used weights in the timestep keyframe schedule. *If Apply Advanced ControlNet node has a weight_override, the weight_override will be used during sampling instead of control_net_weight.* +- 🟨***latent_keyframe***: latent keyframes to apply to controlnet while this Timestep Keyframe is in effect. If inherit_missing is True, if no latent_keyframe is passed in, will attempt to reuse the last-used weights in the timestep keyframe schedule. *If Apply Advanced ControlNet node has a latent_kf_override, the latent_lf_override will be used during sampling instead of latent_keyframe.* +- 🟨***mask_optional***: attention masks to apply to controlnets; basically, decides what part of the image the controlnet to apply to (and the relative strength, if the mask is not binary). Same as mask_optional on the Apply Advanced ControlNet node, can apply either one maks to all latents, or individual masks for each latent. If inherit_missing is True, if no mask_optional is passed in, will attempt to reuse the last-used mask_optional in the timestep keyframe schedule. It is NOT overriden by mask_optional on the Apply Advanced ControlNet node; will be used together. +- 🟦***start_percent***: sampling step percentage at which this Timestep Keyframe qualifies to be used. Acts as the 'key' for the Timestep Keyframe in the timestep keyframe schedule. +- 🟦***strength***: strength of the controlnet; multiplies the controlnet by this value, basically, applied alongside the strength on the Apply ControlNet node. If set to 0.0 will not have any effect during the duration of this Timestep Keyframe's effect, and will increase sampling speed by not doing any work. +- 🟦***null_latent_kf_strength***: strength to assign to latents that are unaccounted for in the passed in latent_keyframes. Has no effect if no latent_keyframes are passed in, or no batch_indeces are unaccounted in the latent_keyframes for during sampling. +- 🟦***inherit_missing***: determines if should reuse values from previous Timestep Keyframes for optional values (control_net_weights, latent_keyframe, and mask_option) that are not included on this TimestepKeyframe. To inherit only specific inputs, use default inputs. +- 🟦***guarantee_usage***: when true, even if a Timestep Keyframe's start_percent ahead of this one in the schedule is closer to current sampling percentage, this Timestep Keyframe will still be used for one step before moving on to the next selected Timestep Keyframe in the following step. Whether the Timestep Keyframe is used or not, its inputs will still be accounted for inherit_missing purposes. + +### Outputs +- 🟪***TIMESTEP_KF***: the created Timestep Keyframe, that can either be linked to another or into a Timestep Keyframe input. + +## Latent Keyframe +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/3de8b718-9f4c-46a9-957e-4ff120938b7e) + +A singular Latent Keyframe, selects the strength for a specific batch_index. If batch_index is not present during sampling, will simply have no effect. Can be chained with any other Latent Keyframe-type node to create a latent keyframe schedule. + +### Inputs +- 🟨***prev_latent_keyframe***: used to chain Latent Keyframes together to create a schedule. *If a Latent Keyframe contained in prev_latent_keyframes have the same batch_index as this Latent Keyframe, they will take priority over this node's value.* +- 🟦***batch_index***: index of latent in batch to apply controlnet strength to. Acts as the 'key' for the Latent Keyframe in the latent keyframe schedule. +- 🟦***strength***: strength of controlnet to apply to the corresponding latent. + +### Outputs + + ## Workflows From 764409115d874686f830d6195506e65d25c87488 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 30 Nov 2023 10:07:02 -0600 Subject: [PATCH 13/15] Shortened some input/output names while maintaining backward compatibility, set default strength so that old workflows will work fine --- control/latent_keyframe_nodes.py | 39 ++++++++++++++++++++++---------- control/nodes.py | 26 +++++++++++++-------- control/weight_nodes.py | 2 +- 3 files changed, 45 insertions(+), 22 deletions(-) diff --git a/control/latent_keyframe_nodes.py b/control/latent_keyframe_nodes.py index d9b5856..2fde61e 100644 --- a/control/latent_keyframe_nodes.py +++ b/control/latent_keyframe_nodes.py @@ -13,13 +13,14 @@ class LatentKeyframeNode: return { "required": { "batch_index": ("INT", {"default": 0, "min": -1000, "max": 1000, "step": 1}), - "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.00001}, ), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), }, "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + "prev_latent_kf": ("LATENT_KEYFRAME", ), } } + RETURN_NAMES = ("LATENT_KF", ) RETURN_TYPES = ("LATENT_KEYFRAME", ) FUNCTION = "load_keyframe" @@ -28,7 +29,10 @@ class LatentKeyframeNode: def load_keyframe(self, batch_index: int, strength: float, - prev_latent_keyframe: LatentKeyframeGroup=None): + prev_latent_kf: LatentKeyframeGroup=None, + prev_latent_keyframe: LatentKeyframeGroup=None, # old name + ): + prev_latent_keyframe = prev_latent_keyframe if prev_latent_keyframe else prev_latent_kf if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() else: @@ -46,12 +50,13 @@ class LatentKeyframeGroupNode: "index_strengths": ("STRING", {"multiline": True, "default": ""}), }, "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + "prev_latent_kf": ("LATENT_KEYFRAME", ), "latent_optional": ("LATENT", ), "print_keyframes": ("BOOLEAN", {"default": False}) } } + RETURN_NAMES = ("LATENT_KF", ) RETURN_TYPES = ("LATENT_KEYFRAME", ) FUNCTION = "load_keyframes" @@ -123,9 +128,11 @@ class LatentKeyframeGroupNode: def load_keyframes(self, index_strengths: str, - prev_latent_keyframe: LatentKeyframeGroup=None, + prev_latent_kf: LatentKeyframeGroup=None, + prev_latent_keyframe: LatentKeyframeGroup=None, # old name latent_image_opt=None, print_keyframes=False): + prev_latent_keyframe = prev_latent_keyframe if prev_latent_keyframe else prev_latent_kf if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() else: @@ -158,16 +165,17 @@ class LatentKeyframeInterpolationNode: "required": { "batch_index_from": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), "batch_index_to_excl": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), - "strength_from": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), - "strength_to": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001}, ), + "strength_from": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "strength_to": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT], ), }, "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + "prev_latent_kf": ("LATENT_KEYFRAME", ), "print_keyframes": ("BOOLEAN", {"default": False}) } } + RETURN_NAMES = ("LATENT_KF", ) RETURN_TYPES = ("LATENT_KEYFRAME", ) FUNCTION = "load_keyframe" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" @@ -178,7 +186,8 @@ class LatentKeyframeInterpolationNode: batch_index_to_excl: int, strength_to: float, interpolation: str, - prev_latent_keyframe: LatentKeyframeGroup=None, + prev_latent_kf: LatentKeyframeGroup=None, + prev_latent_keyframe: LatentKeyframeGroup=None, # old name print_keyframes=False): if (batch_index_from > batch_index_to_excl): @@ -187,6 +196,7 @@ class LatentKeyframeInterpolationNode: if (batch_index_from < 0 and batch_index_to_excl >= 0): raise ValueError("batch_index_from and batch_index_to must be either both positive or both negative.") + prev_latent_keyframe = prev_latent_keyframe if prev_latent_keyframe else prev_latent_kf if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() else: @@ -227,19 +237,24 @@ class LatentKeyframeBatchedGroupNode: def INPUT_TYPES(s): return { "required": { - "float_strengths": ("FLOAT", {"default": -1, "min": -1, "step": 0.0001, "forceInput": True}), + "float_strengths": ("FLOAT", {"default": -1, "min": -1, "step": 0.001, "forceInput": True}), }, "optional": { - "prev_latent_keyframe": ("LATENT_KEYFRAME", ), + "prev_latent_kf": ("LATENT_KEYFRAME", ), "print_keyframes": ("BOOLEAN", {"default": False}) } } + RETURN_NAMES = ("LATENT_KF", ) RETURN_TYPES = ("LATENT_KEYFRAME", ) FUNCTION = "load_keyframe" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes" - def load_keyframe(self, float_strengths: Union[float, list[float]], prev_latent_keyframe: LatentKeyframeGroup=None, print_keyframes=False): + def load_keyframe(self, float_strengths: Union[float, list[float]], + prev_latent_kf: LatentKeyframeGroup=None, + prev_latent_keyframe: LatentKeyframeGroup=None, # old name + print_keyframes=False): + prev_latent_keyframe = prev_latent_keyframe if prev_latent_keyframe else prev_latent_kf if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() else: diff --git a/control/nodes.py b/control/nodes.py index 829f53d..0ca0740 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -21,9 +21,9 @@ class TimestepKeyframeNode: "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), }, "optional": { - "prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ), + "prev_timestep_kf": ("TIMESTEP_KEYFRAME", ), "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "control_net_weights": ("CONTROL_NET_WEIGHTS", ), + "cn_weights": ("CONTROL_NET_WEIGHTS", ), "latent_keyframe": ("LATENT_KEYFRAME", ), "null_latent_kf_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "inherit_missing": ("BOOLEAN", {"default": True}, ), @@ -41,15 +41,17 @@ class TimestepKeyframeNode: def load_keyframe(self, start_percent: float, - strength: float, - control_net_weights: ControlWeights=None, + strength: float=1.0, + cn_weights: ControlWeights=None, control_net_weights: ControlWeights=None, # old name latent_keyframe: LatentKeyframeGroup=None, - prev_timestep_keyframe: TimestepKeyframeGroup=None, + prev_timestep_kf: TimestepKeyframeGroup=None, prev_timestep_keyframe: TimestepKeyframeGroup=None, # old name null_latent_kf_strength: float=0.0, inherit_missing=True, guarantee_usage=True, mask_optional=None, interpolation: str=SI.NONE,): + control_net_weights = control_net_weights if control_net_weights else cn_weights + prev_timestep_keyframe = prev_timestep_keyframe if prev_timestep_keyframe else prev_timestep_kf if not prev_timestep_keyframe: prev_timestep_keyframe = TimestepKeyframeGroup() else: @@ -69,7 +71,7 @@ class ControlNetLoaderAdvanced: "control_net_name": (folder_paths.get_filename_list("controlnet"), ), }, "optional": { - "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), + "timestep_kf": ("TIMESTEP_KEYFRAME", ), } } @@ -78,7 +80,10 @@ class ControlNetLoaderAdvanced: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" - def load_controlnet(self, control_net_name, timestep_keyframe: TimestepKeyframeGroup=None): + def load_controlnet(self, control_net_name, + timestep_kf: TimestepKeyframeGroup=None, timestep_keyframe: TimestepKeyframeGroup=None # old name + ): + timestep_keyframe = timestep_keyframe if timestep_keyframe else timestep_kf controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) controlnet = load_controlnet(controlnet_path, timestep_keyframe) return (controlnet,) @@ -93,7 +98,7 @@ class DiffControlNetLoaderAdvanced: "control_net_name": (folder_paths.get_filename_list("controlnet"), ) }, "optional": { - "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), + "timestep_kf": ("TIMESTEP_KEYFRAME", ), } } @@ -102,7 +107,10 @@ class DiffControlNetLoaderAdvanced: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" - def load_controlnet(self, control_net_name, model, timestep_keyframe: TimestepKeyframeGroup=None): + def load_controlnet(self, control_net_name, model, + timestep_kf: TimestepKeyframeGroup=None, timestep_keyframe: TimestepKeyframeGroup=None # old name + ): + timestep_keyframe = timestep_keyframe if timestep_keyframe else timestep_kf controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) controlnet = load_controlnet(controlnet_path, timestep_keyframe, model) if is_advanced_controlnet(controlnet): diff --git a/control/weight_nodes.py b/control/weight_nodes.py index f4cc9f0..f80d607 100644 --- a/control/weight_nodes.py +++ b/control/weight_nodes.py @@ -4,7 +4,7 @@ from .control import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, ge from .logger import logger -WEIGHTS_RETURN_NAMES = ("CONTROL_NET_WEIGHTS", "TK_SHORTCUT") +WEIGHTS_RETURN_NAMES = ("CN_WEIGHTS", "TK_SHORTCUT") class DefaultWeights: From ea3c846eae1aa399d138aba0dff77a235d87fcc5 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 30 Nov 2023 10:14:39 -0600 Subject: [PATCH 14/15] Brought back ControlNetLoaderAdvanced/Diff timestep_keyframe names to avoid node name cutoff --- control/nodes.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/control/nodes.py b/control/nodes.py index 0ca0740..3794ac7 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -71,7 +71,7 @@ class ControlNetLoaderAdvanced: "control_net_name": (folder_paths.get_filename_list("controlnet"), ), }, "optional": { - "timestep_kf": ("TIMESTEP_KEYFRAME", ), + "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), } } @@ -81,9 +81,8 @@ class ControlNetLoaderAdvanced: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" def load_controlnet(self, control_net_name, - timestep_kf: TimestepKeyframeGroup=None, timestep_keyframe: TimestepKeyframeGroup=None # old name + timestep_keyframe: TimestepKeyframeGroup=None ): - timestep_keyframe = timestep_keyframe if timestep_keyframe else timestep_kf controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) controlnet = load_controlnet(controlnet_path, timestep_keyframe) return (controlnet,) @@ -98,7 +97,7 @@ class DiffControlNetLoaderAdvanced: "control_net_name": (folder_paths.get_filename_list("controlnet"), ) }, "optional": { - "timestep_kf": ("TIMESTEP_KEYFRAME", ), + "timestep_keyframe": ("TIMESTEP_KEYFRAME", ), } } @@ -108,9 +107,8 @@ class DiffControlNetLoaderAdvanced: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" def load_controlnet(self, control_net_name, model, - timestep_kf: TimestepKeyframeGroup=None, timestep_keyframe: TimestepKeyframeGroup=None # old name + timestep_keyframe: TimestepKeyframeGroup=None ): - timestep_keyframe = timestep_keyframe if timestep_keyframe else timestep_kf controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) controlnet = load_controlnet(controlnet_path, timestep_keyframe, model) if is_advanced_controlnet(controlnet): From 419ae2983b1bf4685b556b497431a2c7df95f988 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 30 Nov 2023 10:44:32 -0600 Subject: [PATCH 15/15] Update README.md --- README.md | 57 ++++++++++++++++++++++++++++++++++++++++++++++--------- 1 file changed, 48 insertions(+), 9 deletions(-) diff --git a/README.md b/README.md index ce38f68..f0ef159 100644 --- a/README.md +++ b/README.md @@ -14,7 +14,7 @@ ControlNet preprocessors are available through [comfyui_controlnet_aux](https:// ## Table of Contents: - [Scheduling Explanation](#scheduling-explanation) - [Nodes](#nodes) -- [Usage](#usage) +- [Usage](#usage) (will fill this out soon) # Scheduling Explanation @@ -73,13 +73,13 @@ Loads a ControlNet model and converts it into an Advanced version that supports - 🟪***CONTROL_NET***: loaded Advanced ControlNet ## Timestep Keyframe -![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/3950d14d-6757-42f2-a615-9ba81a74e0ca) +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/c6f2a86e-fc96-4f8b-b976-7c2062a6eba2) Scheduling node across timesteps (sampling steps) based on the set start_percent. Chaining Timestep Keyframes allows ControlNet scheduling across sampling steps (percentage-wise), through a timestep keyframe schedule. ### Inputs -- 🟨***prev_timestep_keyframe***: used to chain Timestep Keyframes together to create a schedule. The order does not matter - the Timestep Keyframes sort themselves automatically by their start_percent. *Any Timestep Keyframe contained in the prev_timestep_keyframe that contains the same start_percent as the Timestep Keyframe will be overwritten.* -- 🟨***control_net_weights***: weights to apply to controlnet while this Timestep Keyframe is in effect. Must be compatible with the loaded controlnet, or will throw an error explaining what weight types are compatible. If inherit_missing is True, if no control_net_weight is passed in, will attempt to reuse the last-used weights in the timestep keyframe schedule. *If Apply Advanced ControlNet node has a weight_override, the weight_override will be used during sampling instead of control_net_weight.* +- 🟨***prev_timestep_kf***: used to chain Timestep Keyframes together to create a schedule. The order does not matter - the Timestep Keyframes sort themselves automatically by their start_percent. *Any Timestep Keyframe contained in the prev_timestep_keyframe that contains the same start_percent as the Timestep Keyframe will be overwritten.* +- 🟨***cn_weights***: weights to apply to controlnet while this Timestep Keyframe is in effect. Must be compatible with the loaded controlnet, or will throw an error explaining what weight types are compatible. If inherit_missing is True, if no control_net_weight is passed in, will attempt to reuse the last-used weights in the timestep keyframe schedule. *If Apply Advanced ControlNet node has a weight_override, the weight_override will be used during sampling instead of control_net_weight.* - 🟨***latent_keyframe***: latent keyframes to apply to controlnet while this Timestep Keyframe is in effect. If inherit_missing is True, if no latent_keyframe is passed in, will attempt to reuse the last-used weights in the timestep keyframe schedule. *If Apply Advanced ControlNet node has a latent_kf_override, the latent_lf_override will be used during sampling instead of latent_keyframe.* - 🟨***mask_optional***: attention masks to apply to controlnets; basically, decides what part of the image the controlnet to apply to (and the relative strength, if the mask is not binary). Same as mask_optional on the Apply Advanced ControlNet node, can apply either one maks to all latents, or individual masks for each latent. If inherit_missing is True, if no mask_optional is passed in, will attempt to reuse the last-used mask_optional in the timestep keyframe schedule. It is NOT overriden by mask_optional on the Apply Advanced ControlNet node; will be used together. - 🟦***start_percent***: sampling step percentage at which this Timestep Keyframe qualifies to be used. Acts as the 'key' for the Timestep Keyframe in the timestep keyframe schedule. @@ -92,21 +92,60 @@ Scheduling node across timesteps (sampling steps) based on the set start_percent - 🟪***TIMESTEP_KF***: the created Timestep Keyframe, that can either be linked to another or into a Timestep Keyframe input. ## Latent Keyframe -![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/3de8b718-9f4c-46a9-957e-4ff120938b7e) +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/7eb2cc4c-255c-4f32-b09b-699f713fada3) A singular Latent Keyframe, selects the strength for a specific batch_index. If batch_index is not present during sampling, will simply have no effect. Can be chained with any other Latent Keyframe-type node to create a latent keyframe schedule. ### Inputs -- 🟨***prev_latent_keyframe***: used to chain Latent Keyframes together to create a schedule. *If a Latent Keyframe contained in prev_latent_keyframes have the same batch_index as this Latent Keyframe, they will take priority over this node's value.* +- 🟨***prev_latent_kf***: used to chain Latent Keyframes together to create a schedule. *If a Latent Keyframe contained in prev_latent_keyframes have the same batch_index as this Latent Keyframe, they will take priority over this node's value.* - 🟦***batch_index***: index of latent in batch to apply controlnet strength to. Acts as the 'key' for the Latent Keyframe in the latent keyframe schedule. - 🟦***strength***: strength of controlnet to apply to the corresponding latent. ### Outputs +- 🟪***LATENT_KF***: the created Latent Keyframe, that can either be linked to another or into a Latent Keyframe input. +## Latent Keyframe Group +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/5ce3b795-f5fc-4dc3-ae30-a4c7f87e278c) +Allows to create Latent Keyframes via individual indeces or python-style ranges. -## Workflows +### Inputs +- 🟨***prev_latent_kf***: used to chain Latent Keyframes together to create a schedule. *If any Latent Keyframes contained in prev_latent_keyframes have the same batch_index as a this Latent Keyframe, they will take priority over this node's version.* +- 🟨***latent_optional***: the latents expected to be passed in for sampling; only required if you wish to use negative indeces (will be automatically converted to real values). +- 🟦***index_strengths***: string list of indeces or python-style ranges of indeces to assign strengths to. If latent_optional is passed in, can contain negative indeces or ranges that contain negative numbers, python-style. The different indeces must be comma separated. Individual latents can be specified by ```batch_index=strength```, like ```0=0.9```. Ranges can be specified by ```start_index_inclusive:end_index_exclusive=strength```, like ```0:8=strength```. Negative indeces are possible when latents_optional has an input, with a string such as ```0,-4=0.25```. +- 🟦***print_keyframes***: if True, will print the Latent Keyframes generated by this node for debugging purposes. -### AnimateDiff Workflows -***Latent Keyframes*** identify which latents in a batch the ControlNet should apply to, and at what strength. They connect to a ***Timestep Keyframe*** to identify at what point in the generation to kick in (for basic use, start_percent on the Timestep Keyframe should be 0.0). Latent Keyframe nodes can be chained to apply the ControlNet to multiple keyframes at various strengths. +### Outputs +- 🟪***LATENT_KF***: the created Latent Keyframe, that can either be linked to another or into a Latent Keyframe input. +## Latent Keyframe Interpolation +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/7986c737-83b9-46bc-aab0-ae4c368df446) + +Allows to create Latent Keyframes with interpolated values in a range. + +### Inputs +- 🟨***prev_latent_kf***: used to chain Latent Keyframes together to create a schedule. *If any Latent Keyframes contained in prev_latent_keyframes have the same batch_index as a this Latent Keyframe, they will take priority over this node's version.* +- 🟦***batch_index_from***: starting batch_index of range, included. +- 🟦***batch_index_to***: end batch_index of range, excluded (python-style range). +- 🟦***strength_from***: starting strength of interpolation. +- 🟦***strength_to***: end strength of interpolation. +- 🟦***interpolation***: the method of interpolation. +- 🟦***print_keyframes***: if True, will print the Latent Keyframes generated by this node for debugging purposes. + +### Outputs +- 🟪***LATENT_KF***: the created Latent Keyframe, that can either be linked to another or into a Latent Keyframe input. + +## Latent Keyframe Batched Group +![image](https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet/assets/7365912/6cec701f-6183-4aeb-af5c-cac76f5591b7) + +Allows to create Latent Keyframes via a list of floats, such as with Batch Value Schedule from [ComfyUI_FizzNodes](https://github.com/FizzleDorf/ComfyUI_FizzNodes) nodes. + +### Inputs +- 🟨***prev_latent_kf***: used to chain Latent Keyframes together to create a schedule. *If any Latent Keyframes contained in prev_latent_keyframes have the same batch_index as a this Latent Keyframe, they will take priority over this node's version.* +- 🟦***float_strengths***: a list of floats, that will correspond to the strength of each Latent Keyframe; the batch_index is the index of each float value in the list. +- 🟦***print_keyframes***: if True, will print the Latent Keyframes generated by this node for debugging purposes. + +### Outputs +- 🟪***LATENT_KF***: the created Latent Keyframe, that can either be linked to another or into a Latent Keyframe input. + +# There are more nodes to document and show usage - will add this soon! TODO