From 77e92f5f8fb7af4bf5305350323591e334b92271 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 6 Feb 2024 01:19:07 -0600 Subject: [PATCH 1/2] Fixed ControlLLLite error with certain latent sizes --- adv_control/control.py | 59 ++++++++++++++++++++++++++++------- adv_control/control_lllite.py | 41 +++++++++++++++++++----- adv_control/utils.py | 39 +++++++++++++++++++++++ 3 files changed, 121 insertions(+), 18 deletions(-) diff --git a/adv_control/control.py b/adv_control/control.py index a5c1dd2..a140e5c 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -284,14 +284,18 @@ class SparseCtrlAdvanced(ControlNetAdvanced): 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): + def __init__(self, patch_attn1: LLLitePatch, patch_attn2: 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) + self.patch_attn1 = patch_attn1.clone_with_control(self) + self.patch_attn2 = patch_attn2.clone_with_control(self) + self.latent_dims_div2 = None + self.latent_dims_div4 = None + def patch_model(self, model: ModelPatcher): - model.set_model_attn1_patch(self.patch) - model.set_model_attn2_patch(self.patch) + model.set_model_attn1_patch(self.patch_attn1) + model.set_model_attn2_patch(self.patch_attn2) def set_cond_hint(self, *args, **kwargs): to_return = super().set_cond_hint(*args, **kwargs) @@ -301,7 +305,9 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): def pre_run_advanced(self, *args, **kwargs): AdvancedControlBase.pre_run_advanced(self, *args, **kwargs) - self.patch.set_control(self) + #logger.error(f"in cn: {id(self.patch_attn1)},{id(self.patch_attn2)}") + self.patch_attn1.set_control(self) + self.patch_attn2.set_control(self) #logger.warn(f"in pre_run_advanced: {id(self)}") def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): @@ -327,6 +333,31 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): 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) + # some special logic here compared to other controlnets: + # * The cond_emb in attn patches will divide latent dims by 2 or 4, integer + # * Due to this loss, the cond_emb will become smaller than x input if latent dims are not divisble by 2 or 4 + divisible_by_2_h = x_noisy.shape[2]%2==0 + divisible_by_2_w = x_noisy.shape[3]%2==0 + if not (divisible_by_2_h and divisible_by_2_w): + #logger.warn(f"{x_noisy.shape} not divisible by 2!") + new_h = (x_noisy.shape[2]//2)*2 + new_w = (x_noisy.shape[3]//2)*2 + if not divisible_by_2_h: + new_h += 2 + if not divisible_by_2_w: + new_w += 2 + self.latent_dims_div2 = (new_h, new_w) + divisible_by_4_h = x_noisy.shape[2]%4==0 + divisible_by_4_w = x_noisy.shape[3]%4==0 + if not (divisible_by_4_h and divisible_by_4_w): + #logger.warn(f"{x_noisy.shape} not divisible by 4!") + new_h = (x_noisy.shape[2]//4)*4 + new_w = (x_noisy.shape[3]//4)*4 + if not divisible_by_4_h: + new_h += 4 + if not divisible_by_4_w: + new_w += 4 + self.latent_dims_div4 = (new_h, new_w) # 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. @@ -335,21 +366,26 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): def cleanup_advanced(self): super().cleanup_advanced() - self.patch.cleanup() + self.patch_attn1.cleanup() + self.patch_attn2.cleanup() + self.latent_dims_div2 = None + self.latent_dims_div4 = None def copy(self): - c = ControlLLLiteAdvanced(self.patch, self.timestep_keyframes) + c = ControlLLLiteAdvanced(self.patch_attn1, self.patch_attn2, self.timestep_keyframes) self.copy_to(c) self.copy_to_advanced(c) return c # deepcopy needs to properly keep track of objects to work between model.clone calls! - def __deepcopy__(self, *args, **kwargs): - return self + # def __deepcopy__(self, *args, **kwargs): + # self.cleanup_advanced() + # return self # 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() + # logger.error(f"in get_models! {id(self)}") # return out @@ -602,6 +638,7 @@ def load_controllllite(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, #logger.info(f"loaded {ckpt_path} successfully, {len(modules)} modules") - patch = LLLitePatch(modules=modules) - control = ControlLLLiteAdvanced(patch=patch, timestep_keyframes=timestep_keyframe) + patch_attn1 = LLLitePatch(modules=modules, patch_type=LLLitePatch.ATTN1) + patch_attn2 = LLLitePatch(modules=modules, patch_type=LLLitePatch.ATTN2) + control = ControlLLLiteAdvanced(patch_attn1=patch_attn1, patch_attn2=patch_attn2, timestep_keyframes=timestep_keyframe) return control diff --git a/adv_control/control_lllite.py b/adv_control/control_lllite.py index 5298e9c..0cdd01c 100644 --- a/adv_control/control_lllite.py +++ b/adv_control/control_lllite.py @@ -11,7 +11,7 @@ import comfy.utils from comfy.controlnet import ControlBase from .logger import logger -from .utils import AdvancedControlBase, prepare_mask_batch +from .utils import AdvancedControlBase, deepcopy_with_sharing, prepare_mask_batch def extra_options_to_module_prefix(extra_options): @@ -38,12 +38,16 @@ def extra_options_to_module_prefix(extra_options): class LLLitePatch: - def __init__(self, modules: dict[str, 'LLLiteModule'], control: Union[AdvancedControlBase, ControlBase]=None): + ATTN1 = "attn1" + ATTN2 = "attn2" + def __init__(self, modules: dict[str, 'LLLiteModule'], patch_type: str, control: Union[AdvancedControlBase, ControlBase]=None): self.modules = modules self.control = control + self.patch_type = patch_type #logger.error(f"create LLLitePatch: {id(self)},{control}") def __call__(self, q, k, v, extra_options): + #logger.error(f"in __call__: {id(self)}") # 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 @@ -80,19 +84,34 @@ class LLLitePatch: def set_control(self, control: Union[AdvancedControlBase, ControlBase]): self.control = control - #logger.error(f"set control for LLLitePatch: {id(self)},{id(control)}") + #logger.error(f"set control for LLLitePatch: {id(self)}, cn: {id(control)}") def clone_with_control(self, control: AdvancedControlBase): #logger.error(f"clone-set control for LLLitePatch: {id(self)},{id(control)}") - return LLLitePatch(self.modules, control) + return LLLitePatch(self.modules, self.patch_type, control) def cleanup(self): - #del self.control - #self.control = None + #total_cleaned = 0 for module in self.modules.values(): module.cleanup() + # total_cleaned += 1 + #logger.info(f"cleaned modules: {total_cleaned}, {id(self)}") #logger.error(f"cleanup LLLitePatch: {id(self)}") + # make sure deepcopy does not copy control, and deepcopied LLLitePatch should be assigned to control + def __deepcopy__(self, memo): + self.cleanup() + to_return: LLLitePatch = deepcopy_with_sharing(self, shared_attribute_names = ['control'], memo=memo) + #logger.warn(f"patch {id(self)} turned into {id(to_return)}") + try: + if self.patch_type == self.ATTN1: + to_return.control.patch_attn1 = to_return + elif self.patch_type == self.ATTN2: + to_return.control.patch_attn2 = to_return + except Exception: + pass + return to_return + # TODO: use comfy.ops to support fp8 properly class LLLiteModule(torch.nn.Module): @@ -159,6 +178,7 @@ class LLLiteModule(torch.nn.Module): self.prev_sub_idxs = None def cleanup(self): + del self.cond_emb self.cond_emb = None self.cx_shape = None self.prev_batch = 0 @@ -167,9 +187,15 @@ class LLLiteModule(torch.nn.Module): def forward(self, x: Tensor, control: Union[AdvancedControlBase, ControlBase]): mask = None mask_tk = None + #logger.info(x.shape) 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(control.cond_hint.to(x.device, dtype=x.dtype)) + cond_hint = control.cond_hint.to(x.device, dtype=x.dtype) + if control.latent_dims_div2 is not None and x.shape[-1] != 1280: + cond_hint = comfy.utils.common_upscale(cond_hint, control.latent_dims_div2[0] * 8, control.latent_dims_div2[1] * 8, 'nearest-exact', "center").to(x.device, dtype=x.dtype) + elif control.latent_dims_div4 is not None and x.shape[-1] == 1280: + cond_hint = comfy.utils.common_upscale(cond_hint, control.latent_dims_div4[0] * 8, control.latent_dims_div4[1] * 8, 'nearest-exact', "center").to(x.device, dtype=x.dtype) + cx = self.conditioning1(cond_hint) self.cx_shape = cx.shape if not self.is_conv2d: # reshape / b,c,h,w -> b,h*w,c @@ -211,6 +237,7 @@ class LLLiteModule(torch.nn.Module): elif mask_tk is not None: mask = mask * mask_tk + #logger.info(f"cs: {cx.shape}, x: {x.shape}, is_conv2d: {self.is_conv2d}") cx = torch.cat([cx, self.down(x)], dim=1 if self.is_conv2d else 2) cx = self.mid(cx) cx = self.up(cx) diff --git a/adv_control/utils.py b/adv_control/utils.py index 606aaf6..4338574 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -1,3 +1,4 @@ +from copy import deepcopy from typing import Callable, Union import torch from torch import Tensor @@ -259,6 +260,44 @@ 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 +# from https://stackoverflow.com/a/24621200 +def deepcopy_with_sharing(obj, shared_attribute_names, memo=None): + ''' + Deepcopy an object, except for a given list of attributes, which should + be shared between the original object and its copy. + + obj is some object + shared_attribute_names: A list of strings identifying the attributes that + should be shared between the original and its copy. + memo is the dictionary passed into __deepcopy__. Ignore this argument if + not calling from within __deepcopy__. + ''' + assert isinstance(shared_attribute_names, (list, tuple)) + + shared_attributes = {k: getattr(obj, k) for k in shared_attribute_names} + + if hasattr(obj, '__deepcopy__'): + # Do hack to prevent infinite recursion in call to deepcopy + deepcopy_method = obj.__deepcopy__ + obj.__deepcopy__ = None + + for attr in shared_attribute_names: + del obj.__dict__[attr] + + clone = deepcopy(obj) + + for attr, val in shared_attributes.items(): + setattr(obj, attr, val) + setattr(clone, attr, val) + + if hasattr(obj, '__deepcopy__'): + # Undo hack + obj.__deepcopy__ = deepcopy_method + del clone.__deepcopy__ + + return clone + + class WeightTypeException(TypeError): "Raised when weight not compatible with AdvancedControlBase object" pass From 7a34d00c6abfdd2165f25be61b11a746d3910cad Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 6 Feb 2024 01:34:47 -0600 Subject: [PATCH 2/2] Added BIGMIN and BIGMAX to use as placeholder limits for INT widgets that don't care about limits --- adv_control/nodes_deprecated.py | 6 +++--- adv_control/nodes_latent_keyframe.py | 8 ++++---- adv_control/utils.py | 5 ++++- 3 files changed, 11 insertions(+), 8 deletions(-) diff --git a/adv_control/nodes_deprecated.py b/adv_control/nodes_deprecated.py index 93ef08f..b08b5b5 100644 --- a/adv_control/nodes_deprecated.py +++ b/adv_control/nodes_deprecated.py @@ -4,7 +4,7 @@ import torch import numpy as np from PIL import Image, ImageOps -from .utils import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe +from .utils import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe, BIGMAX from .logger import logger @@ -16,8 +16,8 @@ class LoadImagesFromDirectory: "directory": ("STRING", {"default": ""}), }, "optional": { - "image_load_cap": ("INT", {"default": 0, "min": 0, "step": 1}), - "start_index": ("INT", {"default": 0, "min": 0, "step": 1}), + "image_load_cap": ("INT", {"default": 0, "min": 0, "max": BIGMAX, "step": 1}), + "start_index": ("INT", {"default": 0, "min": 0, "max": BIGMAX, "step": 1}), } } diff --git a/adv_control/nodes_latent_keyframe.py b/adv_control/nodes_latent_keyframe.py index 2716295..668036d 100644 --- a/adv_control/nodes_latent_keyframe.py +++ b/adv_control/nodes_latent_keyframe.py @@ -2,7 +2,7 @@ from typing import Union import numpy as np from collections.abc import Iterable -from .utils import LatentKeyframe, LatentKeyframeGroup +from .utils import LatentKeyframe, LatentKeyframeGroup, BIGMIN, BIGMAX from .utils import StrengthInterpolation as SI from .logger import logger @@ -12,7 +12,7 @@ class LatentKeyframeNode: def INPUT_TYPES(s): return { "required": { - "batch_index": ("INT", {"default": 0, "min": -1000, "max": 1000, "step": 1}), + "batch_index": ("INT", {"default": 0, "min": BIGMIN, "max": BIGMAX, "step": 1}), "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), }, "optional": { @@ -163,8 +163,8 @@ class LatentKeyframeInterpolationNode: def INPUT_TYPES(s): return { "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}), + "batch_index_from": ("INT", {"default": 0, "min": BIGMIN, "max": BIGMAX, "step": 1}), + "batch_index_to_excl": ("INT", {"default": 0, "min": BIGMIN, "max": BIGMAX, "step": 1}), "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], ), diff --git a/adv_control/utils.py b/adv_control/utils.py index 4338574..3ac0c42 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -11,6 +11,9 @@ from comfy.model_patcher import ModelPatcher from .logger import logger +BIGMIN = -(2**63-1) +BIGMAX = (2**63-1) + 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): # immediately restore load_torch_file to original version @@ -273,7 +276,7 @@ def deepcopy_with_sharing(obj, shared_attribute_names, memo=None): not calling from within __deepcopy__. ''' assert isinstance(shared_attribute_names, (list, tuple)) - + shared_attributes = {k: getattr(obj, k) for k in shared_attribute_names} if hasattr(obj, '__deepcopy__'):