diff --git a/README.md b/README.md index 066588d..bb004f2 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,5 @@ # ComfyUI-Advanced-ControlNet -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, ControlLoRAs, and SparseCtrls. Kohya Controllllite support coming soon. +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, ControlLoRAs, ControlLLLite, and SparseCtrls. Custom weights allow replication of the "My prompt is more important" feature of Auto1111's sd-webui ControlNet extension. @@ -10,6 +10,7 @@ ControlNet preprocessors are available through [comfyui_controlnet_aux](https:// - 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. +- ControlLLLite support (requires model_optional to be passed into and out of Apply Advanced ControlNet node) - SparseCtrl support ## Table of Contents: 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