From fbf005f31d1f708cfe1c46d2df09c14c1eeb9c16 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 31 Jan 2024 23:19:16 -0600 Subject: [PATCH] Added full ControlLLLite support, refactored some code to eliminate code repetition and support controls that require model patching --- adv_control/control.py | 175 ++++++++++++++++++---------- adv_control/control_lllite.py | 208 +++++++++++++++------------------- adv_control/nodes.py | 21 +++- adv_control/utils.py | 64 ++++++++--- 4 files changed, 274 insertions(+), 194 deletions(-) diff --git a/adv_control/control.py b/adv_control/control.py index cf66378..f84e011 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -8,8 +8,10 @@ import comfy.model_management import comfy.model_detection import comfy.controlnet as comfy_cn from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, broadcast_image_to +from model_patcher import ModelPatcher from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper +from .control_lllite import LLLiteModule, LLLitePatch from .utils import (AdvancedControlBase, TimestepKeyframeGroup, LatentKeyframeGroup, ControlWeightType, ControlWeights, WeightTypeException, manual_cast_clean_groupnorm, disable_weight_init_clean_groupnorm, prepare_mask_batch, get_properly_arranged_t2i_weights, load_torch_file_with_dict_factory) from .logger import logger @@ -106,8 +108,6 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): 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) try: # if sub indexes present, replace original hint with subsection if self.sub_idxs is not None: @@ -168,59 +168,6 @@ class ControlLoraAdvanced(ControlLora, AdvancedControlBase): global_average_pooling=v.global_average_pooling, device=v.device) -class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): - # This ControlNet is more of an attention patch than a traditional controlnet - # So, the pre_run will be responsible for a lot of the functionality, - # while the usual get_control is mostly used to set some values - def __init__(self, timestep_keyframes: TimestepKeyframeGroup, device=None): - super().__init__(device) - AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite()) - self.already_patched = False - - def set_cond_hint(self, *args, **kwargs): - super().set_cond_hint(*args, **kwargs) - # cond hint for LLLite needs to be scaled between (-1, 1) instead of (0, 1) - self.cond_hint_original = self.cond_hint_original * 2.0 - 1.0 - - def pre_run_advanced(self, model, percent_to_timestep_function): - AdvancedControlBase.pre_run_advanced(self, model, percent_to_timestep_function) - logger.info(f"In ControlLLLiteAdvanced pre_run_advanced! {self.already_patched}") - # perform patches if not already patches - if not self.already_patched: - self.already_patched = True - - def get_control(self, x_noisy: Tensor, t, cond, batched_number): - logger.info("In ControlLLLiteAdvanced get_control!") - # prepare timestep and everything related - self.prepare_current_timestep(t=t, batched_number=batched_number) - # perform other controlnets - 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 - - def get_models(self): - logger.info(f"In ControlLLLiteAdvanced get_models!") - # get_models is called once at the start of every KSampler run - use to reset already_patched status - self.already_patched = False - out = super().get_models() - return out - - def copy(self): - c = ControlLLLiteAdvanced(self.timestep_keyframes) - self.copy_to(c) - self.copy_to_advanced(c) - return c - - def cleanup(self): - super().cleanup() - self.cleanup_advanced() - self.already_patched = False - - class SparseCtrlAdvanced(ControlNetAdvanced): def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, sparse_settings: SparseSettings=None, global_average_pooling=False, device=None, load_device=None, manual_cast_dtype=None): super().__init__(control_model=control_model, timestep_keyframes=timestep_keyframes, global_average_pooling=global_average_pooling, device=device, load_device=load_device, manual_cast_dtype=manual_cast_dtype) @@ -335,6 +282,72 @@ class SparseCtrlAdvanced(ControlNetAdvanced): return c +class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): + # This ControlNet is more of an attention patch than a traditional controlnet + def __init__(self, patch: LLLitePatch, timestep_keyframes: TimestepKeyframeGroup, device=None): + super().__init__(device) + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite(), require_model=True) + self.patch = patch.clone_with_control(self) + + def patch_model(self, model: ModelPatcher): + model.set_model_attn1_patch(self.patch) + model.set_model_attn2_patch(self.patch) + + def set_cond_hint(self, *args, **kwargs): + to_return = super().set_cond_hint(*args, **kwargs) + # cond hint for LLLite needs to be scaled between (-1, 1) instead of (0, 1) + self.cond_hint_original = self.cond_hint_original * 2.0 - 1.0 + return to_return + + def pre_run_advanced(self, *args, **kwargs): + AdvancedControlBase.pre_run_advanced(self, *args, **kwargs) + self.patch.control = self + + def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): + # normal ControlNet stuff + 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]: + return control_prev + + dtype = x_noisy.dtype + # prepare cond_hint + 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(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(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 + self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number) + # done preparing; model patches will take care of everything now. + # return normal controlnet stuff + return control_prev + + def cleanup_advanced(self): + super().cleanup_advanced() + self.patch.cleanup() + + def copy(self): + c = ControlLLLiteAdvanced(self.patch, self.timestep_keyframes) + self.copy_to(c) + self.copy_to_advanced(c) + return c + + # def get_models(self): + # # get_models is called once at the start of every KSampler run - use to reset already_patched status + # out = super().get_models() + # return out + + def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None): controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) control = None @@ -358,9 +371,7 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo if controlnet_type != ControlWeightType.DEFAULT: if controlnet_type == ControlWeightType.CONTROLLLLITE: - raise NotImplementedError("ControlLLLite has not been fully implemented yet!") - control = ControlLLLiteAdvanced(timestep_keyframes=timestep_keyframe) - # load Controll + control = load_controllllite(ckpt_path, controlnet_data=controlnet_data, timestep_keyframe=timestep_keyframe) elif controlnet_type == ControlWeightType.SPARSECTRL: control = load_sparsectrl(ckpt_path, controlnet_data=controlnet_data, timestep_keyframe=timestep_keyframe, model=model) # otherwise, load vanilla ControlNet @@ -542,3 +553,51 @@ def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, tim control = SparseCtrlAdvanced(control_model, timestep_keyframes=timestep_keyframe, sparse_settings=sparse_settings, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype) return control + + +def load_controllllite(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, timestep_keyframe: TimestepKeyframeGroup=None): + if controlnet_data is None: + controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) + # adapted from https://github.com/kohya-ss/ControlNet-LLLite-ComfyUI + # first, split weights for each module + module_weights = {} + for key, value in controlnet_data.items(): + fragments = key.split(".") + module_name = fragments[0] + weight_name = ".".join(fragments[1:]) + + if module_name not in module_weights: + module_weights[module_name] = {} + module_weights[module_name][weight_name] = value + + # next, load each module + modules = {} + for module_name, weights in module_weights.items(): + # kohya planned to do something about how these should be chosen, so I'm not touching this + # since I am not familiar with the logic for this + if "conditioning1.4.weight" in weights: + depth = 3 + elif weights["conditioning1.2.weight"].shape[-1] == 4: + depth = 2 + else: + depth = 1 + + module = LLLiteModule( + name=module_name, + is_conv2d=weights["down.0.weight"].ndim == 4, + in_dim=weights["down.0.weight"].shape[1], + depth=depth, + cond_emb_dim=weights["conditioning1.0.weight"].shape[0] * 2, + mlp_dim=weights["down.0.weight"].shape[0], + ) + # load weights into module + module.load_state_dict(weights) + modules[module_name] = module + if len(modules) == 1: + module.is_first = True + + logger.info(f"loaded {ckpt_path} successfully, {len(modules)} modules") + + patch = LLLitePatch(modules=modules) + control = ControlLLLiteAdvanced(patch=patch, timestep_keyframes=timestep_keyframe) + return control diff --git a/adv_control/control_lllite.py b/adv_control/control_lllite.py index ba327c6..e92a416 100644 --- a/adv_control/control_lllite.py +++ b/adv_control/control_lllite.py @@ -2,10 +2,16 @@ # basically, all the LLLite core code is from there, which I then combined with # Advanced-ControlNet features and QoL import math +from typing import Union +from torch import Tensor import torch import os -import comfy +import comfy.utils +from comfy.controlnet import ControlBase + +from .logger import logger +from .utils import AdvancedControlBase, prepare_mask_batch def extra_options_to_module_prefix(extra_options): @@ -27,97 +33,69 @@ def extra_options_to_module_prefix(extra_options): elif block[0] == "output": module_pfx = f"lllite_unet_output_blocks_{block[1]}_1_transformer_blocks_{block_index}" else: - raise Exception("invalid block name") + raise Exception(f"ControlLLLite: invalid block name '{block[0]}'. Expected 'input', 'middle', or 'output'.") return module_pfx -def load_control_net_lllite_patch(path, cond_image, multiplier, num_steps, start_percent, end_percent): - # calculate start and end step - start_step = math.floor(num_steps * start_percent * 0.01) if start_percent > 0 else 0 - end_step = math.floor(num_steps * end_percent * 0.01) if end_percent > 0 else num_steps +class LLLitePatch: + def __init__(self, modules: dict[str, 'LLLiteModule'], control: Union[AdvancedControlBase, ControlBase]=None): + self.modules = modules + self.control = control + + def __call__(self, q, k, v, extra_options): + # determine if have anything to run + if self.control.timestep_range is not None: + # it turns out comparing single-value tensors to floats is extremely slow + # a: Tensor = extra_options["sigmas"][0] + if self.control.t > self.control.timestep_range[0] or self.control.t < self.control.timestep_range[1]: + logger.info("Stopping short!!!") + return q, k, v - # load weights - ctrl_sd = comfy.utils.load_torch_file(path, safe_load=True) + module_pfx = extra_options_to_module_prefix(extra_options) - # split each weights for each module - module_weights = {} - for key, value in ctrl_sd.items(): - fragments = key.split(".") - module_name = fragments[0] - weight_name = ".".join(fragments[1:]) - - if module_name not in module_weights: - module_weights[module_name] = {} - module_weights[module_name][weight_name] = value - - # load each module - modules = {} - for module_name, weights in module_weights.items(): - # kohya planned to do something about how these should be chosen, so I'm not touching this - # since I am not familiar with the logic for this - if "conditioning1.4.weight" in weights: - depth = 3 - elif weights["conditioning1.2.weight"].shape[-1] == 4: - depth = 2 + is_attn1 = q.shape[-1] == k.shape[-1] # self attention + if is_attn1: + module_pfx = module_pfx + "_attn1" else: - depth = 1 + module_pfx = module_pfx + "_attn2" - module = LLLiteModule( - name=module_name, - is_conv2d=weights["down.0.weight"].ndim == 4, - in_dim=weights["down.0.weight"].shape[1], - depth=depth, - cond_emb_dim=weights["conditioning1.0.weight"].shape[0] * 2, - mlp_dim=weights["down.0.weight"].shape[0], - multiplier=multiplier, - num_steps=num_steps, - start_step=start_step, - end_step=end_step, - ) - info = module.load_state_dict(weights) - modules[module_name] = module - if len(modules) == 1: - module.is_first = True + module_pfx_to_q = module_pfx + "_to_q" + module_pfx_to_k = module_pfx + "_to_k" + module_pfx_to_v = module_pfx + "_to_v" - print(f"loaded {path} successfully, {len(modules)} modules") + # if masks present, get masks with same dims as attention + # if q.shape != k.shape or q.shape != v.shape: + # logger.warn(f"mismatch!!! q:{q.shape}, k:{k.shape}, v:{v.shape}") + #logger.warn(f"{q.shape}") - for module in modules.values(): - module.set_cond_image(cond_image) + if module_pfx_to_q in self.modules: + q = q + self.modules[module_pfx_to_q](q, self.control) + if module_pfx_to_k in self.modules: + k = k + self.modules[module_pfx_to_k](k, self.control) + if module_pfx_to_v in self.modules: + v = v + self.modules[module_pfx_to_v](v, self.control) - class control_net_lllite_patch: - def __init__(self, modules): - self.modules = modules + return q, k, v - def __call__(self, q, k, v, extra_options): - module_pfx = extra_options_to_module_prefix(extra_options) + def to(self, device): + for d in self.modules.keys(): + self.modules[d] = self.modules[d].to(device) + return self + + def set_control(self, control: Union[AdvancedControlBase, ControlBase]): + self.control = control - is_attn1 = q.shape[-1] == k.shape[-1] # self attention - if is_attn1: - module_pfx = module_pfx + "_attn1" - else: - module_pfx = module_pfx + "_attn2" + def clone_with_control(self, control: AdvancedControlBase): + return LLLitePatch(self.modules, control) - module_pfx_to_q = module_pfx + "_to_q" - module_pfx_to_k = module_pfx + "_to_k" - module_pfx_to_v = module_pfx + "_to_v" - - if module_pfx_to_q in self.modules: - q = q + self.modules[module_pfx_to_q](q) - if module_pfx_to_k in self.modules: - k = k + self.modules[module_pfx_to_k](k) - if module_pfx_to_v in self.modules: - v = v + self.modules[module_pfx_to_v](v) - - return q, k, v - - def to(self, device): - for d in self.modules.keys(): - self.modules[d] = self.modules[d].to(device) - return self - - return control_net_lllite_patch(modules) + def cleanup(self): + del self.control + self.control = None + for module in self.modules.values(): + module.cleanup() +# TODO: use comfy.ops to support fp8 properly class LLLiteModule(torch.nn.Module): def __init__( self, @@ -127,18 +105,10 @@ class LLLiteModule(torch.nn.Module): depth: int, cond_emb_dim: int, mlp_dim: int, - multiplier: int, - num_steps: int, - start_step: int, - end_step: int, ): super().__init__() self.name = name self.is_conv2d = is_conv2d - self.multiplier = multiplier - self.num_steps = num_steps - self.start_step = start_step - self.end_step = end_step self.is_first = False modules = [] @@ -184,48 +154,47 @@ class LLLiteModule(torch.nn.Module): ) self.depth = depth - self.cond_image = None self.cond_emb = None - self.current_step = 0 + self.cx_shape = None + self.prev_batch = 0 + self.prev_sub_idxs = None - # @torch.inference_mode() - def set_cond_image(self, cond_image): - # print("set_cond_image", self.name) - self.cond_image = cond_image + def cleanup(self): self.cond_emb = None - self.current_step = 0 + self.cx_shape = None + self.prev_batch = 0 + self.prev_sub_idxs = None - def forward(self, x): - if self.num_steps > 0: - if self.current_step < self.start_step: - self.current_step += 1 - return torch.zeros_like(x) - elif self.current_step >= self.end_step: - if self.is_first and self.current_step == self.end_step: - print(f"end LLLite: step {self.current_step}") - self.current_step += 1 - if self.current_step >= self.num_steps: - self.current_step = 0 # reset - return torch.zeros_like(x) - else: - if self.is_first and self.current_step == self.start_step: - print(f"start LLLite: step {self.current_step}") - self.current_step += 1 - if self.current_step >= self.num_steps: - self.current_step = 0 # reset - - if self.cond_emb is None: + def forward(self, x: Tensor, control: Union[AdvancedControlBase, ControlBase]): + mask = None + mask_tk = None + if self.cond_emb is None or control.sub_idxs != self.prev_sub_idxs or x.shape[0] != self.prev_batch: # print(f"cond_emb is None, {self.name}") - cx = self.conditioning1(self.cond_image.to(x.device, dtype=x.dtype)) + cx = self.conditioning1(control.cond_hint.to(x.device, dtype=x.dtype)) + self.cx_shape = cx.shape if not self.is_conv2d: # reshape / b,c,h,w -> b,h*w,c n, c, h, w = cx.shape cx = cx.view(n, c, h * w).permute(0, 2, 1) self.cond_emb = cx + # save prev values + self.prev_batch = x.shape[0] + self.prev_sub_idxs = control.sub_idxs cx: torch.Tensor = self.cond_emb # print(f"forward {self.name}, {cx.shape}, {x.shape}") + # TODO: make masks work for conv2d (could not find any ControlLLLites at this time that use them) + # create masks + if not self.is_conv2d: + n, c, h, w = self.cx_shape + if control.mask_cond_hint is not None: + mask = prepare_mask_batch(control.mask_cond_hint, (1, 1, h, w)).to(cx.dtype) + mask = mask.view(mask.shape[0], 1, h * w).permute(0, 2, 1) + if control.tk_mask_cond_hint is not None: + mask_tk = prepare_mask_batch(control.mask_cond_hint, (1, 1, h, w)).to(cx.dtype) + mask_tk = mask_tk.view(mask_tk.shape[0], 1, h * w).permute(0, 2, 1) + # x in uncond/cond doubles batch size if x.shape[0] != cx.shape[0]: if self.is_conv2d: @@ -233,8 +202,19 @@ class LLLiteModule(torch.nn.Module): else: # print("x.shape[0] != cx.shape[0]", x.shape[0], cx.shape[0]) cx = cx.repeat(x.shape[0] // cx.shape[0], 1, 1) + if mask is not None: + mask = mask.repeat(x.shape[0] // mask.shape[0], 1, 1) + if mask_tk is not None: + mask_tk = mask_tk.repeat(x.shape[0] // mask_tk.shape[0], 1, 1) + + if mask is None: + mask = 1.0 + elif mask_tk is not None: + mask = mask * mask_tk cx = torch.cat([cx, self.down(x)], dim=1 if self.is_conv2d else 2) cx = self.mid(cx) cx = self.up(cx) - return cx * self.multiplier \ No newline at end of file + if control.latent_keyframes is not None: + cx = cx * control.calc_latent_keyframe_mults(x=cx, batched_number=control.batched_number) + return cx * mask * control.strength * control.current_timestep_keyframe.strength diff --git a/adv_control/nodes.py b/adv_control/nodes.py index 9c6ed91..1cea1af 100644 --- a/adv_control/nodes.py +++ b/adv_control/nodes.py @@ -2,6 +2,7 @@ import numpy as np from torch import Tensor import folder_paths +from comfy.model_patcher import ModelPatcher from .control import load_controlnet, convert_to_advanced, is_advanced_controlnet from .utils import ControlWeights, ControlWeightType, LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup @@ -136,21 +137,24 @@ class AdvancedControlNetApply: "timestep_kf": ("TIMESTEP_KEYFRAME", ), "latent_kf_override": ("LATENT_KEYFRAME", ), "weights_override": ("CONTROL_NET_WEIGHTS", ), + "model_optional": ("MODEL",), } } - RETURN_TYPES = ("CONDITIONING","CONDITIONING") - RETURN_NAMES = ("positive", "negative") + RETURN_TYPES = ("CONDITIONING","CONDITIONING","MODEL",) + RETURN_NAMES = ("positive", "negative", "model_opt") FUNCTION = "apply_controlnet" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, - mask_optional: Tensor=None, + mask_optional: Tensor=None, model_optional: ModelPatcher=None, timestep_kf: TimestepKeyframeGroup=None, latent_kf_override: LatentKeyframeGroup=None, weights_override: ControlWeights=None): if strength == 0: - return (positive, negative) + return (positive, negative, model_optional) + if model_optional: + model_optional = model_optional.clone() control_hint = image.movedim(-1,1) cnets = {} @@ -168,6 +172,13 @@ class AdvancedControlNetApply: # 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): + # disarm node check + c_net.disarm() + # if model required, verify model is passed in, and if so patch it + if c_net.require_model: + if not model_optional: + raise Exception(f"Type '{type(c_net).__name__}' requires model_optional input, but got None.") + c_net.patch_model(model=model_optional) # apply optional parameters and overrides, if provided if timestep_kf is not None: c_net.set_timestep_keyframes(timestep_kf) @@ -192,7 +203,7 @@ class AdvancedControlNetApply: n = [t[0], d] c.append(n) out.append(c) - return (out[0], out[1]) + return (out[0], out[1], model_optional) # NODE MAPPING diff --git a/adv_control/utils.py b/adv_control/utils.py index be6b73b..606aaf6 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -2,10 +2,13 @@ from typing import Callable, Union import torch from torch import Tensor import torch.nn.functional as F + import comfy.ops import comfy.utils from comfy.controlnet import ControlBase, broadcast_image_to +from comfy.model_patcher import ModelPatcher +from .logger import logger def load_torch_file_with_dict_factory(controlnet_data: dict[str, Tensor], orig_load_torch_file: Callable): def load_torch_file_with_dict(*args, **kwargs): @@ -262,7 +265,7 @@ class WeightTypeException(TypeError): class AdvancedControlBase: - def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights): + def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights, require_model=False): self.base = base self.compatible_weights = [ControlWeightType.UNIVERSAL] self.add_compatible_weight(weights_default.weight_type) @@ -293,6 +296,14 @@ class AdvancedControlBase: self.control_merge = self.control_merge_inject#.__get__(self, type(self)) self.pre_run = self.pre_run_inject self.cleanup = self.cleanup_inject + self.set_previous_controlnet = self.set_previous_controlnet_inject + # require model to be passed into Apply Advanced ControlNet 🛂🅐🅒🅝 node + self.require_model = require_model + # disarm - when set to False, used to force usage of Apply Advanced ControlNet 🛂🅐🅒🅝 node (which will set it to True) + self.disarmed = not require_model + + def patch_model(self, model: ModelPatcher): + pass def add_compatible_weight(self, control_weight_type: str): self.compatible_weights.append(control_weight_type) @@ -322,10 +333,10 @@ class AdvancedControlBase: self.latent_keyframes = None def prepare_current_timestep(self, t: Tensor, batched_number: int): - self.t = t + self.t = float(t[0]) self.batched_number = batched_number # get current step percent - curr_t: float = t[0] + curr_t: float = self.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): @@ -395,23 +406,32 @@ class AdvancedControlBase: # clear variables self.cleanup_advanced() + def set_previous_controlnet_inject(self, *args, **kwargs): + to_return = self.base.set_previous_controlnet(*args, **kwargs) + if not self.disarmed: + raise Exception(f"Type '{type(self).__name__}' must be used with Apply Advanced ControlNet 🛂🅐🅒🅝 node (with model_optional passed in); otherwise, it will not work.") + return to_return + + def disarm(self): + self.disarmed = True + 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 + return self.default_control_actions(x_noisy, t, cond, batched_number) # 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 + return self.default_control_actions(x_noisy, t, cond, batched_number) + + def default_control_actions(self, x_noisy, 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) + return control_prev def calc_weight(self, idx: int, x: Tensor, layers: int) -> Union[float, Tensor]: if self.weights.weight_mask is not None: @@ -424,11 +444,12 @@ class AdvancedControlBase: 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): + def calc_latent_keyframe_mults(self, x: Tensor, batched_number: int) -> Tensor: # apply strengths, and get batch indeces to null out # AKA latents that should not be influenced by ControlNet - if self.latent_keyframes is not None: - latent_count = x.size(0)//batched_number + final_mults = [1.0] * x.shape[0] + if self.latent_keyframes: + latent_count = x.shape[0] // batched_number indeces_to_null = set(range(latent_count)) mapped_indeces = None # if expecting subdivision, will need to translate between subset and actual idx values @@ -459,13 +480,21 @@ class AdvancedControlBase: # 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 - + final_mults[(latent_count*b)+real_index] = keyframe.strength # 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] * self.current_timestep_keyframe.null_latent_kf_strength + final_mults[(latent_count*b)+batch_index] = self.current_timestep_keyframe.null_latent_kf_strength + # convert final_mults into tensor and match expected dimension count + final_tensor = torch.tensor(final_mults, dtype=x.dtype, device=x.device) + while len(final_tensor.shape) < len(x.shape): + final_tensor = final_tensor.unsqueeze(-1) + return final_tensor + + def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int): + if self.latent_keyframes is not None: + x[:] = x[:] * self.calc_latent_keyframe_mults(x=x, batched_number=batched_number) # apply masks, resizing mask to required dims if self.mask_cond_hint is not None: masks = prepare_mask_batch(self.mask_cond_hint, x.shape) @@ -602,3 +631,4 @@ class 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 + copied.disarmed = self.disarmed