From 74386d238012fdf94514c3b0d7efd76d8a37050a Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 12 Dec 2023 13:44:19 -0600 Subject: [PATCH 01/15] Initial work on ControlLLLite support --- control/control.py | 80 +++++++++++-- control/control_lllite.py | 239 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 310 insertions(+), 9 deletions(-) diff --git a/control/control.py b/control/control.py index 18cbffd..f9b972a 100644 --- a/control/control.py +++ b/control/control.py @@ -6,6 +6,7 @@ import comfy.utils import comfy.controlnet as comfy_cn from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, broadcast_image_to +from .logger import logger def get_properly_arranged_t2i_weights(initial_weights: list[float]): new_weights = [] @@ -708,20 +709,81 @@ class ControlLoraAdvanced(ControlLora, AdvancedControlBase): global_average_pooling=v.global_average_pooling, device=v.device) -class ControlLLLiteAdvanced(ControlNet, AdvancedControlBase): - def __init__(self, control_weights, timestep_keyframes: TimestepKeyframeGroup, device=None): +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 def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None): - 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) + control = None + # check if a non-vanilla ControlNet + controlnet_type = ControlWeightType.DEFAULT + for key in controlnet_data: + if "lllite" in key: + logger.info("ControlLLLite controlnet!") + controlnet_type = ControlWeightType.CONTROLLLLITE + break + if controlnet_type != ControlWeightType.DEFAULT: + if controlnet_type == ControlWeightType.CONTROLLLLITE: + control = ControlLLLiteAdvanced(timestep_keyframes=timestep_keyframe) + # load Controll + # otherwise, load vanilla ControlNet + else: + control = comfy_cn.load_controlnet(ckpt_path, model=model) + # from pathlib import Path + # with open(Path(__file__).parent.parent.parent / "controlnet_keys.txt", "w") as cfile: # controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) - # # check if lllite - # if "lllite_unet" in controlnet_data: - # pass + # for key in controlnet_data: + # cfile.write(f"{key}\n") return convert_to_advanced(control, timestep_keyframe=timestep_keyframe) diff --git a/control/control_lllite.py b/control/control_lllite.py index 8b13789..ba327c6 100644 --- a/control/control_lllite.py +++ b/control/control_lllite.py @@ -1 +1,240 @@ +# adapted from https://github.com/kohya-ss/ControlNet-LLLite-ComfyUI +# basically, all the LLLite core code is from there, which I then combined with +# Advanced-ControlNet features and QoL +import math +import torch +import os +import comfy + + +def extra_options_to_module_prefix(extra_options): + # extra_options = {'transformer_index': 2, 'block_index': 8, 'original_shape': [2, 4, 128, 128], 'block': ('input', 7), 'n_heads': 20, 'dim_head': 64} + + # block is: [('input', 4), ('input', 5), ('input', 7), ('input', 8), ('middle', 0), + # ('output', 0), ('output', 1), ('output', 2), ('output', 3), ('output', 4), ('output', 5)] + # transformer_index is: [0, 1, 2, 3, 4, 5, 6, 7, 8], for each block + # block_index is: 0-1 or 0-9, depends on the block + # input 7 and 8, middle has 10 blocks + + # make module name from extra_options + block = extra_options["block"] + block_index = extra_options["block_index"] + if block[0] == "input": + module_pfx = f"lllite_unet_input_blocks_{block[1]}_1_transformer_blocks_{block_index}" + elif block[0] == "middle": + module_pfx = f"lllite_unet_middle_block_1_transformer_blocks_{block_index}" + elif block[0] == "output": + module_pfx = f"lllite_unet_output_blocks_{block[1]}_1_transformer_blocks_{block_index}" + else: + raise Exception("invalid block name") + 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 + + # load weights + ctrl_sd = comfy.utils.load_torch_file(path, safe_load=True) + + # 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 + 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], + 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 + + print(f"loaded {path} successfully, {len(modules)} modules") + + for module in modules.values(): + module.set_cond_image(cond_image) + + class control_net_lllite_patch: + def __init__(self, modules): + self.modules = modules + + def __call__(self, q, k, v, extra_options): + module_pfx = extra_options_to_module_prefix(extra_options) + + is_attn1 = q.shape[-1] == k.shape[-1] # self attention + if is_attn1: + module_pfx = module_pfx + "_attn1" + else: + module_pfx = module_pfx + "_attn2" + + 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) + + +class LLLiteModule(torch.nn.Module): + def __init__( + self, + name: str, + is_conv2d: bool, + in_dim: int, + 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 = [] + modules.append(torch.nn.Conv2d(3, cond_emb_dim // 2, kernel_size=4, stride=4, padding=0)) # to latent (from VAE) size*2 + if depth == 1: + modules.append(torch.nn.ReLU(inplace=True)) + modules.append(torch.nn.Conv2d(cond_emb_dim // 2, cond_emb_dim, kernel_size=2, stride=2, padding=0)) + elif depth == 2: + modules.append(torch.nn.ReLU(inplace=True)) + modules.append(torch.nn.Conv2d(cond_emb_dim // 2, cond_emb_dim, kernel_size=4, stride=4, padding=0)) + elif depth == 3: + # kernel size 8 is too large, so set it to 4 + modules.append(torch.nn.ReLU(inplace=True)) + modules.append(torch.nn.Conv2d(cond_emb_dim // 2, cond_emb_dim // 2, kernel_size=4, stride=4, padding=0)) + modules.append(torch.nn.ReLU(inplace=True)) + modules.append(torch.nn.Conv2d(cond_emb_dim // 2, cond_emb_dim, kernel_size=2, stride=2, padding=0)) + + self.conditioning1 = torch.nn.Sequential(*modules) + + if self.is_conv2d: + self.down = torch.nn.Sequential( + torch.nn.Conv2d(in_dim, mlp_dim, kernel_size=1, stride=1, padding=0), + torch.nn.ReLU(inplace=True), + ) + self.mid = torch.nn.Sequential( + torch.nn.Conv2d(mlp_dim + cond_emb_dim, mlp_dim, kernel_size=1, stride=1, padding=0), + torch.nn.ReLU(inplace=True), + ) + self.up = torch.nn.Sequential( + torch.nn.Conv2d(mlp_dim, in_dim, kernel_size=1, stride=1, padding=0), + ) + else: + self.down = torch.nn.Sequential( + torch.nn.Linear(in_dim, mlp_dim), + torch.nn.ReLU(inplace=True), + ) + self.mid = torch.nn.Sequential( + torch.nn.Linear(mlp_dim + cond_emb_dim, mlp_dim), + torch.nn.ReLU(inplace=True), + ) + self.up = torch.nn.Sequential( + torch.nn.Linear(mlp_dim, in_dim), + ) + + self.depth = depth + self.cond_image = None + self.cond_emb = None + self.current_step = 0 + + # @torch.inference_mode() + def set_cond_image(self, cond_image): + # print("set_cond_image", self.name) + self.cond_image = cond_image + self.cond_emb = None + self.current_step = 0 + + 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: + # print(f"cond_emb is None, {self.name}") + cx = self.conditioning1(self.cond_image.to(x.device, dtype=x.dtype)) + 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 + + cx: torch.Tensor = self.cond_emb + # print(f"forward {self.name}, {cx.shape}, {x.shape}") + + # x in uncond/cond doubles batch size + if x.shape[0] != cx.shape[0]: + if self.is_conv2d: + cx = cx.repeat(x.shape[0] // cx.shape[0], 1, 1, 1) + 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) + + 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 From 88cc0ac149ee644927e2a87f08a0b0f3f9cd3731 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Mon, 18 Dec 2023 01:01:47 -0600 Subject: [PATCH 02/15] Progress on SparseCtrl support --- control/control.py | 413 +++---- control/control_sparsectrl.py | 1035 +++++++++++++++++ control/nodes.py | 14 +- ...eprecated_nodes.py => nodes_deprecated.py} | 2 +- ...rame_nodes.py => nodes_latent_keyframe.py} | 4 +- control/nodes_reference.py | 0 control/nodes_sparsectrl.py | 28 + control/{weight_nodes.py => nodes_weight.py} | 2 +- control/reference_nodes.py | 12 - control/utils.py | 260 +++++ 10 files changed, 1518 insertions(+), 252 deletions(-) create mode 100644 control/control_sparsectrl.py rename control/{deprecated_nodes.py => nodes_deprecated.py} (97%) rename control/{latent_keyframe_nodes.py => nodes_latent_keyframe.py} (99%) create mode 100644 control/nodes_reference.py create mode 100644 control/nodes_sparsectrl.py rename control/{weight_nodes.py => nodes_weight.py} (98%) delete mode 100644 control/reference_nodes.py create mode 100644 control/utils.py diff --git a/control/control.py b/control/control.py index f5e5495..7ed319f 100644 --- a/control/control.py +++ b/control/control.py @@ -1,223 +1,19 @@ -from typing import Union +from typing import Callable, Union from torch import Tensor import torch +import os import comfy.utils +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 .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper +from .utils import (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 -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) -> Union[float, Tensor]: - # 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 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: - 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: - 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 - self.strength = strength - - -# always maintain sorted state (by batch_index of LatentKeyframe) -class LatentKeyframeGroup: - def __init__(self) -> None: - self.keyframes: list[LatentKeyframe] = [] - - def add(self, keyframe: LatentKeyframe) -> None: - added = False - # replace existing keyframe if same batch_index - for i in range(len(self.keyframes)): - if self.keyframes[i].batch_index == keyframe.batch_index: - self.keyframes[i] = keyframe - added = True - break - if not added: - self.keyframes.append(keyframe) - self.keyframes.sort(key=lambda k: k.batch_index) - - def get_index(self, index: int) -> Union[LatentKeyframe, None]: - try: - return self.keyframes[index] - except IndexError: - return None - - def __getitem__(self, index) -> LatentKeyframe: - return self.keyframes[index] - - 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, - start_percent: float = 0.0, - strength: float = 1.0, - interpolation: str = StrengthInterpolation.NONE, - control_weights: ControlWeights = None, - latent_keyframes: LatentKeyframeGroup = None, - null_latent_kf_strength: float = 0.0, - inherit_missing: bool = True, - guarantee_usage: bool = True, - mask_hint_orig: Tensor = None) -> None: - self.start_percent = start_percent - self.start_t = 999999999.9 - self.strength = strength - self.interpolation = interpolation - self.control_weights = control_weights - self.latent_keyframes = latent_keyframes - 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 - - 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': - return cls(0.0) - - -# always maintain sorted state (by start_percent of TimestepKeyFrame) -class TimestepKeyframeGroup: - def __init__(self) -> None: - self.keyframes: list[TimestepKeyframe] = [] - self.keyframes.append(TimestepKeyframe.default()) - - def add(self, keyframe: TimestepKeyframe) -> None: - added = False - # replace existing keyframe if same start_percent - for i in range(len(self.keyframes)): - if self.keyframes[i].start_percent == keyframe.start_percent: - self.keyframes[i] = keyframe - added = True - break - if not added: - self.keyframes.append(keyframe) - self.keyframes.sort(key=lambda k: k.start_percent) - - def get_index(self, index: int) -> Union[TimestepKeyframe, None]: - try: - return self.keyframes[index] - 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 - - 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() - group.keyframes[0] = keyframe - return group - - -# used to inject ControlNetAdvanced and T2IAdapterAdvanced control_merge function - class AdvancedControlBase: def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights): @@ -765,23 +561,56 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): self.already_patched = False +class SparseCtrlAdvanced(ControlNetAdvanced): + def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, 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) + self.add_compatible_weight(ControlWeightType.SPARSECTRL) + + def copy(self): + c = SparseCtrlAdvanced(self.control_model, self.timestep_keyframes, self.global_average_pooling, self.device, self.load_device, self.manual_cast_dtype) + self.copy_to(c) + self.copy_to_advanced(c) + return c + + 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 # check if a non-vanilla ControlNet controlnet_type = ControlWeightType.DEFAULT + has_controlnet_key = False + has_motion_modules_key = False for key in controlnet_data: + # LLLLite check if "lllite" in key: logger.info("ControlLLLite controlnet!") controlnet_type = ControlWeightType.CONTROLLLLITE break + # SparseCtrl check + elif "motion_modules" in key: + has_motion_modules_key = True + elif "controlnet" in key: + has_controlnet_key = True + if has_controlnet_key and has_motion_modules_key: + controlnet_type = ControlWeightType.SPARSECTRL + 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 + elif controlnet_type == ControlWeightType.SPARSECTRL: + #raise NotImplementedError("SparseCtrl has not been fully implemented yet!") + control = load_sparsectrl(ckpt_path, controlnet_data=controlnet_data, timestep_keyframe=timestep_keyframe, model=model) # otherwise, load vanilla ControlNet else: - control = comfy_cn.load_controlnet(ckpt_path, model=model) + try: + # hacky way of getting load_torch_file in load_controlnet to use already-present controlnet_data and not redo loading + orig_load_torch_file = comfy.utils.load_torch_file + comfy.utils.load_torch_file = load_torch_file_with_dict_factory(controlnet_data, orig_load_torch_file) + control = comfy_cn.load_controlnet(ckpt_path, model=model) + finally: + comfy.utils.load_torch_file = orig_load_torch_file # from pathlib import Path # with open(Path(__file__).parent.parent.parent / "controlnet_keys.txt", "w") as cfile: # controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) @@ -811,25 +640,151 @@ def is_advanced_controlnet(input_object): return hasattr(input_object, "sub_idxs") -# adapted from comfy/sample.py -def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim1=False): - mask = mask.clone() - mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2]*multiplier, shape[3]*multiplier), mode="bilinear") - if match_dim1: - mask = torch.cat([mask] * shape[1], dim=1) - return mask +def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, timestep_keyframe: TimestepKeyframeGroup=None, model=None) -> SparseCtrlAdvanced: + if controlnet_data is None: + controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) + # first, separate out motion part from normal controlnet part and attempt to load that portion + motion_data = {} + for key in list(controlnet_data.keys()): + if "temporal" in key: + motion_data[key] = controlnet_data.pop(key) + motion_wrapper: SparseCtrlMotionWrapper = SparseCtrlMotionWrapper(motion_data).to(comfy.model_management.unet_dtype()) + missing, unexpected = motion_wrapper.load_state_dict(motion_data) + if len(missing) > 0 or len(unexpected) > 0: + logger.info(f"SparseCtrlMotionWrapper: {missing}, {unexpected}") + # now, load as if it was a normal controlnet - mostly copied from comfy load_controlnet function + controlnet_config = None + is_diffusers = False + use_simplified_conditioning_embedding = False + if "controlnet_cond_embedding.conv_in.weight" in controlnet_data: + is_diffusers = True + if "controlnet_cond_embedding.weight" in controlnet_data: + is_diffusers = True + use_simplified_conditioning_embedding = True + if is_diffusers: #diffusers format + unet_dtype = comfy.model_management.unet_dtype() + controlnet_config = comfy.model_detection.unet_config_from_diffusers_unet(controlnet_data, unet_dtype) + diffusers_keys = comfy.utils.unet_to_diffusers(controlnet_config) + diffusers_keys["controlnet_mid_block.weight"] = "middle_block_out.0.weight" + diffusers_keys["controlnet_mid_block.bias"] = "middle_block_out.0.bias" -# 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 + count = 0 + loop = True + while loop: + suffix = [".weight", ".bias"] + for s in suffix: + k_in = "controlnet_down_blocks.{}{}".format(count, s) + k_out = "zero_convs.{}.0{}".format(count, s) + if k_in not in controlnet_data: + loop = False + break + diffusers_keys[k_in] = k_out + count += 1 + # normal conditioning embedding + if not use_simplified_conditioning_embedding: + count = 0 + loop = True + while loop: + suffix = [".weight", ".bias"] + for s in suffix: + if count == 0: + k_in = "controlnet_cond_embedding.conv_in{}".format(s) + else: + k_in = "controlnet_cond_embedding.blocks.{}{}".format(count - 1, s) + k_out = "input_hint_block.{}{}".format(count * 2, s) + if k_in not in controlnet_data: + k_in = "controlnet_cond_embedding.conv_out{}".format(s) + loop = False + diffusers_keys[k_in] = k_out + count += 1 + # simplified conditioning embedding + else: + count = 0 + suffix = [".weight", ".bias"] + for s in suffix: + k_in = "controlnet_cond_embedding{}".format(s) + k_out = "input_hint_block.{}{}".format(count, s) + diffusers_keys[k_in] = k_out -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 + new_sd = {} + for k in diffusers_keys: + if k in controlnet_data: + new_sd[diffusers_keys[k]] = controlnet_data.pop(k) + leftover_keys = controlnet_data.keys() + if len(leftover_keys) > 0: + logger.info("leftover keys:", leftover_keys) + controlnet_data = new_sd -class WeightTypeException(TypeError): - "Raised when weight not compatible with AdvancedControlBase object" - pass + pth_key = 'control_model.zero_convs.0.0.weight' + pth = False + key = 'zero_convs.0.0.weight' + if pth_key in controlnet_data: + pth = True + key = pth_key + prefix = "control_model." + elif key in controlnet_data: + prefix = "" + else: + raise ValueError("The provided model is not a valid SparseCtrl model! [ErrorCode: HORSERADISH]") + + if controlnet_config is None: + unet_dtype = comfy.model_management.unet_dtype() + controlnet_config = comfy.model_detection.model_config_from_unet(controlnet_data, prefix, unet_dtype, True).unet_config + load_device = comfy.model_management.get_torch_device() + manual_cast_dtype = comfy.model_management.unet_manual_cast(unet_dtype, load_device) + if manual_cast_dtype is not None: + controlnet_config["operations"] = manual_cast_clean_groupnorm + else: + controlnet_config["operations"] = disable_weight_init_clean_groupnorm + controlnet_config.pop("out_channels") + # get proper hint channels + if use_simplified_conditioning_embedding: + controlnet_config["hint_channels"] = controlnet_data["{}input_hint_block.0.weight".format(prefix)].shape[1] + controlnet_config["use_simplified_conditioning_embedding"] = use_simplified_conditioning_embedding + else: + controlnet_config["hint_channels"] = controlnet_data["{}input_hint_block.0.weight".format(prefix)].shape[1] + controlnet_config["use_simplified_conditioning_embedding"] = use_simplified_conditioning_embedding + control_model = SparseControlNet(**controlnet_config) + + if pth: + if 'difference' in controlnet_data: + if model is not None: + comfy.model_management.load_models_gpu([model]) + model_sd = model.model_state_dict() + for x in controlnet_data: + c_m = "control_model." + if x.startswith(c_m): + sd_key = "diffusion_model.{}".format(x[len(c_m):]) + if sd_key in model_sd: + cd = controlnet_data[x] + cd += model_sd[sd_key].type(cd.dtype).to(cd.device) + else: + logger.warning("WARNING: Loaded a diff SparseCtrl without a model. It will very likely not work.") + + class WeightsLoader(torch.nn.Module): + pass + w = WeightsLoader() + w.control_model = control_model + missing, unexpected = w.load_state_dict(controlnet_data, strict=False) + else: + missing, unexpected = control_model.load_state_dict(controlnet_data, strict=False) + if len(missing) > 0 or len(unexpected) > 0: + logger.info(f"SparseCtrl ControlNet: {missing}, {unexpected}") + + global_average_pooling = False + filename = os.path.splitext(ckpt_path)[0] + if filename.endswith("_shuffle") or filename.endswith("_shuffle_fp16"): #TODO: smarter way of enabling global_average_pooling + global_average_pooling = True + + # both motion portion and controlnet portions are loaded; bring them together + motion_wrapper.inject(control_model) + + control = SparseCtrlAdvanced(control_model, timestep_keyframes=timestep_keyframe, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype) + new_state_dict = control_model.state_dict() + from pathlib import Path + with open(Path(__file__).parent.parent.parent / "sparcectrlstatedict.txt", "w") as cfile: + for key in new_state_dict: + cfile.write(f"{key}\n") + return control diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py new file mode 100644 index 0000000..8e1e15c --- /dev/null +++ b/control/control_sparsectrl.py @@ -0,0 +1,1035 @@ +#taken from: https://github.com/lllyasviel/ControlNet +#and modified +#and then taken from comfy/cldm/cldm.py and modified again + +import math +from typing import Iterable, Union +import torch +import torch as th +import torch.nn as nn +from torch import Tensor +from einops import rearrange, repeat + +from comfy.ldm.modules.diffusionmodules.util import ( + zero_module, + timestep_embedding, +) + +from comfy.cldm.cldm import ControlNet as ControlNet_cldm +from comfy.ldm.modules.attention import SpatialTransformer +from comfy.ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample +from comfy.ldm.util import exists +from comfy.ldm.modules.attention import default, optimized_attention +from comfy.ldm.modules.attention import FeedForward, SpatialTransformer +from comfy.controlnet import broadcast_image_to +from comfy.utils import repeat_to_batch_size +import comfy.ops + +from .utils import TimestepKeyframeGroup, disable_weight_init_clean_groupnorm, prepare_mask_batch + + +class SparseControlNet(ControlNet_cldm): + def __init__(self, *args,**kwargs): + super().__init__(*args, **kwargs) + hint_channels = kwargs.get("hint_channels") + operations: disable_weight_init_clean_groupnorm = kwargs.get("operations", disable_weight_init_clean_groupnorm) + device = kwargs.get("device", None) + use_simplified_conditioning_embedding = kwargs.get("use_simplified_conditioning_embedding", False) + if use_simplified_conditioning_embedding: + self.input_hint_block = TimestepEmbedSequential( + operations.conv_nd(self.dims, hint_channels, self.model_channels, 3, padding=1, dtype=self.dtype, device=device), + ) + + def forward(self, x, hint, timesteps, context, y=None, **kwargs): + t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype) + emb = self.time_embed(t_emb) + + x = torch.zeros_like(x) + + conditioning_mask1 = torch.ones_like(hint[:, :1]) + conditioning_mask2 = torch.zeros_like(hint[:, :1]) + conditioning_mask = conditioning_mask2 + conditioning_mask[0] = conditioning_mask1[0] + conditioning_mask[16] = conditioning_mask1[16] + #conditioning_mask[15] = conditioning_mask1[15] + #conditioning_mask[31] = conditioning_mask1[31] + modified_hint = torch.zeros_like(hint) + modified_hint[0] = hint[0] + modified_hint[16] = hint[16] + #modified_hint[15] = hint[15] + #modified_hint[31] = hint[31] + hint = torch.cat([modified_hint, conditioning_mask], dim=1) + guided_hint = self.input_hint_block(hint, emb, context) + + outs = [] + + hs = [] + if self.num_classes is not None: + assert y.shape[0] == x.shape[0] + emb = emb + self.label_emb(y) + + h = x + for module, zero_conv in zip(self.input_blocks, self.zero_convs): + if guided_hint is not None: + h = module(h, emb, context) + h += guided_hint + guided_hint = None + else: + h = module(h, emb, context) + outs.append(zero_conv(h, emb, context)) + + h = self.middle_block(h, emb, context) + outs.append(self.middle_block_out(h, emb, context)) + + return outs + + + +# main class for holding SparseControlNet +class SparseControlNetOld(nn.Module): + def __init__( + self, + image_size, + in_channels, + model_channels, + hint_channels, + num_res_blocks, + dropout=0, + channel_mult=(1, 2, 4, 8), + conv_resample=True, + dims=2, + num_classes=None, + use_checkpoint=False, + dtype=torch.float32, + num_heads=-1, + num_head_channels=-1, + num_heads_upsample=-1, + use_scale_shift_norm=False, + resblock_updown=False, + use_new_attention_order=False, + use_spatial_transformer=False, # custom transformer support + transformer_depth=1, # custom transformer support + context_dim=None, # custom transformer support + n_embed=None, # custom support for prediction of discrete ids into codebook of first stage vq model + legacy=True, + disable_self_attentions=None, + num_attention_blocks=None, + disable_middle_self_attn=False, + use_linear_in_transformer=False, + adm_in_channels=None, + transformer_depth_middle=None, + transformer_depth_output=None, + device=None, + operations=disable_weight_init_clean_groupnorm, + **kwargs, + ): + super().__init__() + assert use_spatial_transformer == True, "use_spatial_transformer has to be true" + if use_spatial_transformer: + assert context_dim is not None, 'Fool!! You forgot to include the dimension of your cross-attention conditioning...' + + if context_dim is not None: + assert use_spatial_transformer, 'Fool!! You forgot to use the spatial transformer for your cross-attention conditioning...' + # from omegaconf.listconfig import ListConfig + # if type(context_dim) == ListConfig: + # context_dim = list(context_dim) + if num_heads_upsample == -1: + num_heads_upsample = num_heads + + if num_heads == -1: + assert num_head_channels != -1, 'Either num_heads or num_head_channels has to be set' + + if num_head_channels == -1: + assert num_heads != -1, 'Either num_heads or num_head_channels has to be set' + + self.dims = dims + self.image_size = image_size + self.in_channels = in_channels + self.model_channels = model_channels + + if isinstance(num_res_blocks, int): + self.num_res_blocks = len(channel_mult) * [num_res_blocks] + else: + if len(num_res_blocks) != len(channel_mult): + raise ValueError("provide num_res_blocks either as an int (globally constant) or " + "as a list/tuple (per-level) with the same length as channel_mult") + self.num_res_blocks = num_res_blocks + + if disable_self_attentions is not None: + # should be a list of booleans, indicating whether to disable self-attention in TransformerBlocks or not + assert len(disable_self_attentions) == len(channel_mult) + if num_attention_blocks is not None: + assert len(num_attention_blocks) == len(self.num_res_blocks) + assert all(map(lambda i: self.num_res_blocks[i] >= num_attention_blocks[i], range(len(num_attention_blocks)))) + + transformer_depth = transformer_depth[:] + + self.dropout = dropout + self.channel_mult = channel_mult + self.conv_resample = conv_resample + self.num_classes = num_classes + self.use_checkpoint = use_checkpoint + self.dtype = dtype + self.num_heads = num_heads + self.num_head_channels = num_head_channels + self.num_heads_upsample = num_heads_upsample + self.predict_codebook_ids = n_embed is not None + + time_embed_dim = model_channels * 4 + self.time_embed = nn.Sequential( + operations.Linear(model_channels, time_embed_dim, dtype=self.dtype, device=device), + nn.SiLU(), + operations.Linear(time_embed_dim, time_embed_dim, dtype=self.dtype, device=device), + ) + + if self.num_classes is not None: + if isinstance(self.num_classes, int): + self.label_emb = nn.Embedding(num_classes, time_embed_dim) + elif self.num_classes == "continuous": + print("setting up linear c_adm embedding layer") + self.label_emb = nn.Linear(1, time_embed_dim) + elif self.num_classes == "sequential": + assert adm_in_channels is not None + self.label_emb = nn.Sequential( + nn.Sequential( + operations.Linear(adm_in_channels, time_embed_dim, dtype=self.dtype, device=device), + nn.SiLU(), + operations.Linear(time_embed_dim, time_embed_dim, dtype=self.dtype, device=device), + ) + ) + else: + raise ValueError() + + self.input_blocks = nn.ModuleList( + [ + TimestepEmbedSequential( + operations.conv_nd(dims, in_channels, model_channels, 3, padding=1, dtype=self.dtype, device=device) + ) + ] + ) + self.zero_convs = nn.ModuleList([self.make_zero_conv(model_channels, operations=operations, dtype=self.dtype, device=device)]) + + self.input_hint_block = TimestepEmbedSequential( + operations.conv_nd(dims, hint_channels, 16, 3, padding=1, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 16, 16, 3, padding=1, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 16, 32, 3, padding=1, stride=2, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 32, 32, 3, padding=1, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 32, 96, 3, padding=1, stride=2, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 96, 96, 3, padding=1, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 96, 256, 3, padding=1, stride=2, dtype=self.dtype, device=device), + nn.SiLU(), + operations.conv_nd(dims, 256, model_channels, 3, padding=1, dtype=self.dtype, device=device) + ) + + self._feature_size = model_channels + input_block_chans = [model_channels] + ch = model_channels + ds = 1 + for level, mult in enumerate(channel_mult): + for nr in range(self.num_res_blocks[level]): + layers = [ + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=mult * model_channels, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + dtype=self.dtype, + device=device, + operations=operations, + ) + ] + ch = mult * model_channels + num_transformers = transformer_depth.pop(0) + if num_transformers > 0: + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + if legacy: + #num_heads = 1 + dim_head = ch // num_heads if use_spatial_transformer else num_head_channels + if exists(disable_self_attentions): + disabled_sa = disable_self_attentions[level] + else: + disabled_sa = False + + if not exists(num_attention_blocks) or nr < num_attention_blocks[level]: + layers.append( + SpatialTransformer( + ch, num_heads, dim_head, depth=num_transformers, context_dim=context_dim, + disable_self_attn=disabled_sa, use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint, dtype=self.dtype, device=device, operations=operations + ) + ) + self.input_blocks.append(TimestepEmbedSequential(*layers)) + self.zero_convs.append(self.make_zero_conv(ch, operations=operations, dtype=self.dtype, device=device)) + self._feature_size += ch + input_block_chans.append(ch) + if level != len(channel_mult) - 1: + out_ch = ch + self.input_blocks.append( + TimestepEmbedSequential( + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=out_ch, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + down=True, + dtype=self.dtype, + device=device, + operations=operations + ) + if resblock_updown + else Downsample( + ch, conv_resample, dims=dims, out_channels=out_ch, dtype=self.dtype, device=device, operations=operations + ) + ) + ) + ch = out_ch + input_block_chans.append(ch) + self.zero_convs.append(self.make_zero_conv(ch, operations=operations, dtype=self.dtype, device=device)) + ds *= 2 + self._feature_size += ch + + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + if legacy: + #num_heads = 1 + dim_head = ch // num_heads if use_spatial_transformer else num_head_channels + mid_block = [ + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + dtype=self.dtype, + device=device, + operations=operations + )] + if transformer_depth_middle >= 0: + mid_block += [SpatialTransformer( # always uses a self-attn + ch, num_heads, dim_head, depth=transformer_depth_middle, context_dim=context_dim, + disable_self_attn=disable_middle_self_attn, use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint, dtype=self.dtype, device=device, operations=operations + ), + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + dtype=self.dtype, + device=device, + operations=operations + )] + self.middle_block = TimestepEmbedSequential(*mid_block) + self.middle_block_out = self.make_zero_conv(ch, operations=operations, dtype=self.dtype, device=device) + self._feature_size += ch + + #self._motion_wrapper: SparseCtrlMotionWrapper = None + + def make_zero_conv(self, channels, operations=None, dtype=None, device=None): + return TimestepEmbedSequential(operations.conv_nd(self.dims, channels, channels, 1, padding=0, dtype=dtype, device=device)) + + def forward(self, x, hint, timesteps, context, y=None, **kwargs): + t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype) + emb = self.time_embed(t_emb) + + x = torch.zeros_like(x) + + conditioning_mask1 = torch.ones_like(hint[:, :1]) + conditioning_mask2 = torch.zeros_like(hint[:, :1]) + conditioning_mask = conditioning_mask2 + conditioning_mask[0] = conditioning_mask1[0] + conditioning_mask[16] = conditioning_mask1[16] + #conditioning_mask[15] = conditioning_mask1[15] + #conditioning_mask[31] = conditioning_mask1[31] + modified_hint = torch.zeros_like(hint) + modified_hint[0] = hint[0] + modified_hint[16] = hint[16] + #modified_hint[15] = hint[15] + #modified_hint[31] = hint[31] + hint = torch.cat([modified_hint, conditioning_mask], dim=1) + guided_hint = self.input_hint_block(hint, emb, context) + + outs = [] + + hs = [] + if self.num_classes is not None: + assert y.shape[0] == x.shape[0] + emb = emb + self.label_emb(y) + + h = x + for module, zero_conv in zip(self.input_blocks, self.zero_convs): + if guided_hint is not None: + h = module(h, emb, context) + h += guided_hint + guided_hint = None + else: + h = module(h, emb, context) + outs.append(zero_conv(h, emb, context)) + + h = self.middle_block(h, emb, context) + outs.append(self.middle_block_out(h, emb, context)) + + return outs + + +# motion-related portion of controlnet +class BlockType: + UP = "up" + DOWN = "down" + MID = "mid" + +def get_down_block_max(mm_state_dict: dict[str, Tensor]) -> int: + return get_block_max(mm_state_dict, "down_blocks") + +def get_up_block_max(mm_state_dict: dict[str, Tensor]) -> int: + return get_block_max(mm_state_dict, "up_blocks") + +def get_block_max(mm_state_dict: dict[str, Tensor], block_name: str) -> int: + # keep track of biggest down_block count in module + biggest_block = -1 + for key in mm_state_dict.keys(): + if block_name in key: + try: + block_int = key.split(".")[1] + block_num = int(block_int) + if block_num > biggest_block: + biggest_block = block_num + except ValueError: + pass + return biggest_block + +def has_mid_block(mm_state_dict: dict[str, Tensor]): + # check if keys contain mid_block + for key in mm_state_dict.keys(): + if key.startswith("mid_block."): + return True + return False + +def get_position_encoding_max_len(mm_state_dict: dict[str, Tensor], mm_name: str=None) -> int: + # use pos_encoder.pe entries to determine max length - [1, {max_length}, {320|640|1280}] + for key in mm_state_dict.keys(): + if key.endswith("pos_encoder.pe"): + return mm_state_dict[key].size(1) # get middle dim + raise ValueError(f"No pos_encoder.pe found in SparseCtrl state_dict - {mm_name} is not a valid SparseCtrl model!") + + +class SparseCtrlMotionWrapper(nn.Module): + def __init__(self, mm_state_dict: dict[str, Tensor]): + super().__init__() + self.down_blocks: Iterable[MotionModule] = None + self.up_blocks: Iterable[MotionModule] = None + self.mid_block: MotionModule = None + self.encoding_max_len = get_position_encoding_max_len(mm_state_dict, "") + layer_channels = (320, 640, 1280, 1280) + if get_down_block_max(mm_state_dict) > -1: + self.down_blocks = nn.ModuleList([]) + for c in layer_channels: + self.down_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.DOWN)) + if get_up_block_max(mm_state_dict) > -1: + self.up_blocks = nn.ModuleList([]) + for c in reversed(layer_channels): + self.up_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.UP)) + if has_mid_block(mm_state_dict): + self.mid_block = MotionModule(1280, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.MID) + + def inject(self, unet: SparseControlNet): + # inject input (down) blocks + self._inject(unet.input_blocks, self.down_blocks) + # inject mid block, if present + if self.mid_block is not None: + self._inject([unet.middle_block], [self.mid_block]) + #unet._motion_wrapper = self + + def _inject(self, unet_blocks: nn.ModuleList, mm_blocks: nn.ModuleList): + # Rules for injection: + # For each component list in a unet block: + # if SpatialTransformer exists in list, place next block after last occurrence + # elif ResBlock exists in list, place next block after first occurrence + # else don't place block + injection_count = 0 + unet_idx = 0 + # details about blocks passed in + per_block = len(mm_blocks[0].motion_modules) + injection_goal = len(mm_blocks) * per_block + # only stop injecting when modules exhausted + while injection_count < injection_goal: + # figure out which VanillaTemporalModule from mm to inject + mm_blk_idx, mm_vtm_idx = injection_count // per_block, injection_count % per_block + # figure out layout of unet block components + st_idx = -1 # SpatialTransformer index + res_idx = -1 # first ResBlock index + # first, figure out indeces of relevant blocks + for idx, component in enumerate(unet_blocks[unet_idx]): + if type(component) == SpatialTransformer: + st_idx = idx + elif type(component).__name__ == "ResBlock" and res_idx < 0: + res_idx = idx + # if SpatialTransformer exists, inject right after + if st_idx >= 0: + unet_blocks[unet_idx].insert(st_idx+1, mm_blocks[mm_blk_idx].motion_modules[mm_vtm_idx]) + injection_count += 1 + # otherwise, if only ResBlock exists, inject right after + elif res_idx >= 0: + unet_blocks[unet_idx].insert(res_idx+1, mm_blocks[mm_blk_idx].motion_modules[mm_vtm_idx]) + injection_count += 1 + # increment unet_idx + unet_idx += 1 + + def eject(self, unet: SparseControlNet): + # remove from input blocks (downblocks) + self._eject(unet.input_blocks) + # remove from middle block (encapsulate in list to make compatible) + self._eject([unet.middle_block]) + #del unet._motion_wrapper + + def _eject(self, unet_blocks: nn.ModuleList): + # eject all VanillaTemporalModule objects from all blocks + for block in unet_blocks: + idx_to_pop = [] + for idx, component in enumerate(block): + if type(component) == VanillaTemporalModule: + idx_to_pop.append(idx) + # pop in backwards order, as to not disturb what the indeces refer to + for idx in sorted(idx_to_pop, reverse=True): + block.pop(idx) + + def set_video_length(self, video_length: int, full_length: int): + self.AD_video_length = video_length + for block in self.down_blocks: + block.set_video_length(video_length, full_length) + for block in self.up_blocks: + block.set_video_length(video_length, full_length) + if self.mid_block is not None: + self.mid_block.set_video_length(video_length, full_length) + + def set_scale_multiplier(self, multiplier: Union[float, None]): + for block in self.down_blocks: + block.set_scale_multiplier(multiplier) + for block in self.up_blocks: + block.set_scale_multiplier(multiplier) + if self.mid_block is not None: + self.mid_block.set_scale_multiplier(multiplier) + + def reset_temp_vars(self): + for block in self.down_blocks: + block.reset_temp_vars() + for block in self.up_blocks: + block.reset_temp_vars() + if self.mid_block is not None: + self.mid_block.reset_temp_vars() + + def reset_scale_multiplier(self): + self.set_scale_multiplier(None) + + def reset(self): + self.reset_scale_multiplier() + self.reset_temp_vars() + + +class MotionModule(nn.Module): + def __init__(self, in_channels, temporal_position_encoding_max_len=24, block_type: str=BlockType.DOWN): + super().__init__() + if block_type == BlockType.MID: + # mid blocks contain only a single VanillaTemporalModule + self.motion_modules: Iterable[VanillaTemporalModule] = nn.ModuleList([get_motion_module(in_channels, temporal_position_encoding_max_len)]) + else: + # down blocks contain two VanillaTemporalModules + self.motion_modules: Iterable[VanillaTemporalModule] = nn.ModuleList( + [ + get_motion_module(in_channels, temporal_position_encoding_max_len), + get_motion_module(in_channels, temporal_position_encoding_max_len) + ] + ) + # up blocks contain one additional VanillaTemporalModule + if block_type == BlockType.UP: + self.motion_modules.append(get_motion_module(in_channels, temporal_position_encoding_max_len)) + + def set_video_length(self, video_length: int, full_length: int): + for motion_module in self.motion_modules: + motion_module.set_video_length(video_length, full_length) + + def set_scale_multiplier(self, multiplier: Union[float, None]): + for motion_module in self.motion_modules: + motion_module.set_scale_multiplier(multiplier) + + def set_masks(self, masks: Tensor, min_val: float, max_val: float): + for motion_module in self.motion_modules: + motion_module.set_masks(masks, min_val, max_val) + + def set_sub_idxs(self, sub_idxs: list[int]): + for motion_module in self.motion_modules: + motion_module.set_sub_idxs(sub_idxs) + + def reset_temp_vars(self): + for motion_module in self.motion_modules: + motion_module.reset_temp_vars() + + +def get_motion_module(in_channels, temporal_position_encoding_max_len): + # unlike normal AD, there is only one attention block expected in SparseCtrl models + return VanillaTemporalModule(in_channels=in_channels, attention_block_types=("Temporal_Self",), temporal_position_encoding_max_len=temporal_position_encoding_max_len) + + +class VanillaTemporalModule(nn.Module): + def __init__( + self, + in_channels, + num_attention_heads=8, + num_transformer_block=1, + attention_block_types=("Temporal_Self", "Temporal_Self"), + cross_frame_attention_mode=None, + temporal_position_encoding=True, + temporal_position_encoding_max_len=24, + temporal_attention_dim_div=1, + zero_initialize=True, + ): + super().__init__() + + self.temporal_transformer = TemporalTransformer3DModel( + in_channels=in_channels, + num_attention_heads=num_attention_heads, + attention_head_dim=in_channels + // num_attention_heads + // temporal_attention_dim_div, + num_layers=num_transformer_block, + attention_block_types=attention_block_types, + cross_frame_attention_mode=cross_frame_attention_mode, + temporal_position_encoding=temporal_position_encoding, + temporal_position_encoding_max_len=temporal_position_encoding_max_len, + ) + + if zero_initialize: + self.temporal_transformer.proj_out = zero_module( + self.temporal_transformer.proj_out + ) + + def set_video_length(self, video_length: int, full_length: int): + self.temporal_transformer.set_video_length(video_length, full_length) + + def set_scale_multiplier(self, multiplier: Union[float, None]): + self.temporal_transformer.set_scale_multiplier(multiplier) + + def set_masks(self, masks: Tensor, min_val: float, max_val: float): + self.temporal_transformer.set_masks(masks, min_val, max_val) + + def set_sub_idxs(self, sub_idxs: list[int]): + self.temporal_transformer.set_sub_idxs(sub_idxs) + + def reset_temp_vars(self): + self.temporal_transformer.reset_temp_vars() + + def forward(self, input_tensor, encoder_hidden_states=None, attention_mask=None): + return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) + + +class TemporalTransformer3DModel(nn.Module): + def __init__( + self, + in_channels, + num_attention_heads, + attention_head_dim, + num_layers, + attention_block_types=( + "Temporal_Self", + "Temporal_Self", + ), + dropout=0.0, + norm_num_groups=32, + cross_attention_dim=768, + activation_fn="geglu", + attention_bias=False, + upcast_attention=False, + cross_frame_attention_mode=None, + temporal_position_encoding=False, + temporal_position_encoding_max_len=24, + ): + super().__init__() + self.video_length = 16 + self.full_length = 16 + self.scale_min = 1.0 + self.scale_max = 1.0 + self.raw_scale_mask: Union[Tensor, None] = None + self.temp_scale_mask: Union[Tensor, None] = None + self.sub_idxs: Union[list[int], None] = None + self.prev_hidden_states_batch = 0 + + + inner_dim = num_attention_heads * attention_head_dim + + self.norm = disable_weight_init_clean_groupnorm.GroupNorm( + num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True + ) + self.proj_in = nn.Linear(in_channels, inner_dim) + + self.transformer_blocks: Iterable[TemporalTransformerBlock] = nn.ModuleList( + [ + TemporalTransformerBlock( + dim=inner_dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + attention_block_types=attention_block_types, + dropout=dropout, + norm_num_groups=norm_num_groups, + cross_attention_dim=cross_attention_dim, + activation_fn=activation_fn, + attention_bias=attention_bias, + upcast_attention=upcast_attention, + cross_frame_attention_mode=cross_frame_attention_mode, + temporal_position_encoding=temporal_position_encoding, + temporal_position_encoding_max_len=temporal_position_encoding_max_len, + ) + for d in range(num_layers) + ] + ) + self.proj_out = nn.Linear(inner_dim, in_channels) + + def set_video_length(self, video_length: int, full_length: int): + self.video_length = video_length + self.full_length = full_length + + def set_scale_multiplier(self, multiplier: Union[float, None]): + for block in self.transformer_blocks: + block.set_scale_multiplier(multiplier) + + def set_masks(self, masks: Tensor, min_val: float, max_val: float): + self.scale_min = min_val + self.scale_max = max_val + self.raw_scale_mask = masks + + def set_sub_idxs(self, sub_idxs: list[int]): + self.sub_idxs = sub_idxs + for block in self.transformer_blocks: + block.set_sub_idxs(sub_idxs) + + def reset_temp_vars(self): + del self.temp_scale_mask + self.temp_scale_mask = None + self.prev_hidden_states_batch = 0 + + def get_scale_mask(self, hidden_states: Tensor) -> Union[Tensor, None]: + # if no raw mask, return None + if self.raw_scale_mask is None: + return None + shape = hidden_states.shape + batch, channel, height, width = shape + # if temp mask already calculated, return it + if self.temp_scale_mask != None: + # check if hidden_states batch matches + if batch == self.prev_hidden_states_batch: + if self.sub_idxs is not None: + return self.temp_scale_mask[:, self.sub_idxs, :] + return self.temp_scale_mask + # if does not match, reset cached temp_scale_mask and recalculate it + del self.temp_scale_mask + self.temp_scale_mask = None + # otherwise, calculate temp mask + self.prev_hidden_states_batch = batch + mask = prepare_mask_batch(self.raw_scale_mask, shape=(self.full_length, 1, height, width)) + mask = repeat_to_batch_size(mask, self.full_length) + # if mask not the same amount length as full length, make it match + if self.full_length != mask.shape[0]: + mask = broadcast_image_to(mask, self.full_length, 1) + # reshape mask to attention K shape (h*w, latent_count, 1) + batch, channel, height, width = mask.shape + # first, perform same operations as on hidden_states, + # turning (b, c, h, w) -> (b, h*w, c) + mask = mask.permute(0, 2, 3, 1).reshape(batch, height*width, channel) + # then, make it the same shape as attention's k, (h*w, b, c) + mask = mask.permute(1, 0, 2) + # make masks match the expected length of h*w + batched_number = shape[0] // self.video_length + if batched_number > 1: + mask = torch.cat([mask] * batched_number, dim=0) + # cache mask and set to proper device + self.temp_scale_mask = mask + # move temp_scale_mask to proper dtype + device + self.temp_scale_mask = self.temp_scale_mask.to(dtype=hidden_states.dtype, device=hidden_states.device) + # return subset of masks, if needed + if self.sub_idxs is not None: + return self.temp_scale_mask[:, self.sub_idxs, :] + return self.temp_scale_mask + + def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None): + batch, channel, height, width = hidden_states.shape + residual = hidden_states + scale_mask = self.get_scale_mask(hidden_states) + # add some casts for fp8 purposes - does not affect speed otherwise + hidden_states = self.norm(hidden_states).to(hidden_states.dtype) + inner_dim = hidden_states.shape[1] + hidden_states = hidden_states.permute(0, 2, 3, 1).reshape( + batch, height * width, inner_dim + ) + hidden_states = self.proj_in(hidden_states).to(hidden_states.dtype) + + # Transformer Blocks + for block in self.transformer_blocks: + hidden_states = block( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + attention_mask=attention_mask, + video_length=self.video_length, + scale_mask=scale_mask + ) + + # output + hidden_states = self.proj_out(hidden_states) + hidden_states = ( + hidden_states.reshape(batch, height, width, inner_dim) + .permute(0, 3, 1, 2) + .contiguous() + ) + + output = hidden_states + residual + + return output + + +class TemporalTransformerBlock(nn.Module): + def __init__( + self, + dim, + num_attention_heads, + attention_head_dim, + attention_block_types=( + "Temporal_Self", + "Temporal_Self", + ), + dropout=0.0, + norm_num_groups=32, + cross_attention_dim=768, + activation_fn="geglu", + attention_bias=False, + upcast_attention=False, + cross_frame_attention_mode=None, + temporal_position_encoding=False, + temporal_position_encoding_max_len=24, + ): + super().__init__() + + attention_blocks = [] + norms = [] + + for block_name in attention_block_types: + attention_blocks.append( + VersatileAttention( + attention_mode=block_name.split("_")[0], + context_dim=cross_attention_dim # called context_dim for ComfyUI impl + if block_name.endswith("_Cross") + else None, + query_dim=dim, + heads=num_attention_heads, + dim_head=attention_head_dim, + dropout=dropout, + #bias=attention_bias, # remove for Comfy CrossAttention + #upcast_attention=upcast_attention, # remove for Comfy CrossAttention + cross_frame_attention_mode=cross_frame_attention_mode, + temporal_position_encoding=temporal_position_encoding, + temporal_position_encoding_max_len=temporal_position_encoding_max_len, + ) + ) + norms.append(nn.LayerNorm(dim)) + + self.attention_blocks: Iterable[VersatileAttention] = nn.ModuleList(attention_blocks) + self.norms = nn.ModuleList(norms) + + self.ff = FeedForward(dim, dropout=dropout, glu=(activation_fn == "geglu")) + self.ff_norm = nn.LayerNorm(dim) + + def set_scale_multiplier(self, multiplier: Union[float, None]): + for block in self.attention_blocks: + block.set_scale_multiplier(multiplier) + + def set_sub_idxs(self, sub_idxs: list[int]): + for block in self.attention_blocks: + block.set_sub_idxs(sub_idxs) + + def forward( + self, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + video_length=None, + scale_mask=None + ): + for attention_block, norm in zip(self.attention_blocks, self.norms): + norm_hidden_states = norm(hidden_states).to(hidden_states.dtype) + hidden_states = ( + attention_block( + norm_hidden_states, + encoder_hidden_states=encoder_hidden_states + if attention_block.is_cross_attention + else None, + attention_mask=attention_mask, + video_length=video_length, + scale_mask=scale_mask + ) + + hidden_states + ) + + hidden_states = self.ff(self.ff_norm(hidden_states)) + hidden_states + + output = hidden_states + return output + + +class PositionalEncoding(nn.Module): + def __init__(self, d_model, dropout=0.0, max_len=24): + super().__init__() + self.dropout = nn.Dropout(p=dropout) + position = torch.arange(max_len).unsqueeze(1) + div_term = torch.exp( + torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model) + ) + pe = torch.zeros(1, max_len, d_model) + pe[0, :, 0::2] = torch.sin(position * div_term) + pe[0, :, 1::2] = torch.cos(position * div_term) + self.register_buffer("pe", pe) + self.sub_idxs = None + + def set_sub_idxs(self, sub_idxs: list[int]): + self.sub_idxs = sub_idxs + + def forward(self, x): + #if self.sub_idxs is not None: + # x = x + self.pe[:, self.sub_idxs] + #else: + x = x + self.pe[:, : x.size(1)] + return self.dropout(x) + + +class CrossAttentionMM(nn.Module): + def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0., dtype=None, device=None, + operations=comfy.ops.disable_weight_init): + super().__init__() + inner_dim = dim_head * heads + context_dim = default(context_dim, query_dim) + + self.heads = heads + self.dim_head = dim_head + self.scale = None + + self.to_q = operations.Linear(query_dim, inner_dim, bias=False, dtype=dtype, device=device) + self.to_k = operations.Linear(context_dim, inner_dim, bias=False, dtype=dtype, device=device) + self.to_v = operations.Linear(context_dim, inner_dim, bias=False, dtype=dtype, device=device) + + self.to_out = nn.Sequential(operations.Linear(inner_dim, query_dim, dtype=dtype, device=device), nn.Dropout(dropout)) + + def forward(self, x, context=None, value=None, mask=None, scale_mask=None): + q = self.to_q(x) + context = default(context, x) + k: Tensor = self.to_k(context) + if value is not None: + v = self.to_v(value) + del value + else: + v = self.to_v(context) + + # apply custom scale by multiplying k by scale factor + if self.scale is not None: + k *= self.scale + + # apply scale mask, if present + if scale_mask is not None: + k *= scale_mask + + out = optimized_attention(q, k, v, self.heads, mask) + return self.to_out(out) + + +class VersatileAttention(CrossAttentionMM): + def __init__( + self, + attention_mode=None, + cross_frame_attention_mode=None, + temporal_position_encoding=False, + temporal_position_encoding_max_len=24, + *args, + **kwargs, + ): + super().__init__(*args, **kwargs) + assert attention_mode == "Temporal" + + self.attention_mode = attention_mode + self.is_cross_attention = kwargs["context_dim"] is not None + + self.pos_encoder = ( + PositionalEncoding( + kwargs["query_dim"], + dropout=0.0, + max_len=temporal_position_encoding_max_len, + ) + if (temporal_position_encoding and attention_mode == "Temporal") + else None + ) + + def extra_repr(self): + return f"(Module Info) Attention_Mode: {self.attention_mode}, Is_Cross_Attention: {self.is_cross_attention}" + + def set_scale_multiplier(self, multiplier: Union[float, None]): + if multiplier is None or math.isclose(multiplier, 1.0): + self.scale = None + else: + self.scale = multiplier + + def set_sub_idxs(self, sub_idxs: list[int]): + if self.pos_encoder != None: + self.pos_encoder.set_sub_idxs(sub_idxs) + + def forward( + self, + hidden_states: Tensor, + encoder_hidden_states=None, + attention_mask=None, + video_length=None, + scale_mask=None, + ): + if self.attention_mode != "Temporal": + raise NotImplementedError + + d = hidden_states.shape[1] + hidden_states = rearrange( + hidden_states, "(b f) d c -> (b d) f c", f=video_length + ) + + if self.pos_encoder is not None: + hidden_states = self.pos_encoder(hidden_states).to(hidden_states.dtype) + + encoder_hidden_states = ( + repeat(encoder_hidden_states, "b n c -> (b d) n c", d=d) + if encoder_hidden_states is not None + else encoder_hidden_states + ) + + hidden_states = super().forward( + hidden_states, + encoder_hidden_states, + value=None, + mask=attention_mask, + scale_mask=scale_mask, + ) + + hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d) + + return hidden_states diff --git a/control/nodes.py b/control/nodes.py index 3794ac7..dd9ecde 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -3,13 +3,13 @@ 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, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, \ - SoftT2IAdapterWeights, CustomT2IAdapterWeights -from .latent_keyframe_nodes import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode -from .deprecated_nodes import LoadImagesFromDirectory +from .control import load_controlnet, convert_to_advanced, is_advanced_controlnet +from .utils import ControlWeights, ControlWeightType, LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup +from .utils import StrengthInterpolation as SI +from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, + SoftT2IAdapterWeights, CustomT2IAdapterWeights) +from .nodes_latent_keyframe import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode +from .nodes_deprecated import LoadImagesFromDirectory from .logger import logger diff --git a/control/deprecated_nodes.py b/control/nodes_deprecated.py similarity index 97% rename from control/deprecated_nodes.py rename to control/nodes_deprecated.py index a64ac9b..93ef08f 100644 --- a/control/deprecated_nodes.py +++ b/control/nodes_deprecated.py @@ -4,7 +4,7 @@ import torch import numpy as np from PIL import Image, ImageOps -from .control import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe +from .utils import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe from .logger import logger diff --git a/control/latent_keyframe_nodes.py b/control/nodes_latent_keyframe.py similarity index 99% rename from control/latent_keyframe_nodes.py rename to control/nodes_latent_keyframe.py index 2fde61e..2716295 100644 --- a/control/latent_keyframe_nodes.py +++ b/control/nodes_latent_keyframe.py @@ -2,8 +2,8 @@ from typing import Union import numpy as np from collections.abc import Iterable -from .control import LatentKeyframe, LatentKeyframeGroup -from .control import StrengthInterpolation as SI +from .utils import LatentKeyframe, LatentKeyframeGroup +from .utils import StrengthInterpolation as SI from .logger import logger diff --git a/control/nodes_reference.py b/control/nodes_reference.py new file mode 100644 index 0000000..e69de29 diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py new file mode 100644 index 0000000..2e62629 --- /dev/null +++ b/control/nodes_sparsectrl.py @@ -0,0 +1,28 @@ +import folder_paths + +from .utils import TimestepKeyframeGroup +from .control import load_sparsectrl + + +# node for SparseCtrl loading +class SparseCtrlLoaderAdvanced: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "control_net_name": (folder_paths.get_filename_list("controlnet"), ), + }, + "optional": { + "tk_optional": ("TIMESTEP_KEYFRAME", ), + } + } + + RETURN_TYPES = ("CONTROL_NET", ) + FUNCTION = "load_controlnet" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders" + + def load_controlnet(self, control_net_name: str, tk_optional: TimestepKeyframeGroup=None): + controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) + controlnet = load_sparsectrl(controlnet_path, timestep_keyframe=tk_optional) + return controlnet diff --git a/control/weight_nodes.py b/control/nodes_weight.py similarity index 98% rename from control/weight_nodes.py rename to control/nodes_weight.py index f80d607..35d0ffb 100644 --- a/control/weight_nodes.py +++ b/control/nodes_weight.py @@ -1,6 +1,6 @@ from torch import Tensor import torch -from .control import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights, linear_conversion +from .utils import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights, linear_conversion from .logger import logger diff --git a/control/reference_nodes.py b/control/reference_nodes.py deleted file mode 100644 index 6879f97..0000000 --- a/control/reference_nodes.py +++ /dev/null @@ -1,12 +0,0 @@ -class AnimateDiffLoaderWithContext: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "image": ("IMAGE",), - }, - } - - RETURN_TYPES = ("MODEL",) - CATEGORY = "" \ No newline at end of file diff --git a/control/utils.py b/control/utils.py new file mode 100644 index 0000000..7772601 --- /dev/null +++ b/control/utils.py @@ -0,0 +1,260 @@ +from typing import Callable, Union +import torch +from torch import Tensor +import torch.nn.functional as F +import comfy.ops +import comfy.utils + + +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 + comfy.utils.load_torch_file = orig_load_torch_file + return controlnet_data + return load_torch_file_with_dict + + +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" + SPARSECTRL = "sparsectrl" + + +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) -> Union[float, Tensor]: + # 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 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: + 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: + 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 + self.strength = strength + + +# always maintain sorted state (by batch_index of LatentKeyframe) +class LatentKeyframeGroup: + def __init__(self) -> None: + self.keyframes: list[LatentKeyframe] = [] + + def add(self, keyframe: LatentKeyframe) -> None: + added = False + # replace existing keyframe if same batch_index + for i in range(len(self.keyframes)): + if self.keyframes[i].batch_index == keyframe.batch_index: + self.keyframes[i] = keyframe + added = True + break + if not added: + self.keyframes.append(keyframe) + self.keyframes.sort(key=lambda k: k.batch_index) + + def get_index(self, index: int) -> Union[LatentKeyframe, None]: + try: + return self.keyframes[index] + except IndexError: + return None + + def __getitem__(self, index) -> LatentKeyframe: + return self.keyframes[index] + + 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, + start_percent: float = 0.0, + strength: float = 1.0, + interpolation: str = StrengthInterpolation.NONE, + control_weights: ControlWeights = None, + latent_keyframes: LatentKeyframeGroup = None, + null_latent_kf_strength: float = 0.0, + inherit_missing: bool = True, + guarantee_usage: bool = True, + mask_hint_orig: Tensor = None) -> None: + self.start_percent = start_percent + self.start_t = 999999999.9 + self.strength = strength + self.interpolation = interpolation + self.control_weights = control_weights + self.latent_keyframes = latent_keyframes + 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 + + 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': + return cls(0.0) + + +# always maintain sorted state (by start_percent of TimestepKeyFrame) +class TimestepKeyframeGroup: + def __init__(self) -> None: + self.keyframes: list[TimestepKeyframe] = [] + self.keyframes.append(TimestepKeyframe.default()) + + def add(self, keyframe: TimestepKeyframe) -> None: + added = False + # replace existing keyframe if same start_percent + for i in range(len(self.keyframes)): + if self.keyframes[i].start_percent == keyframe.start_percent: + self.keyframes[i] = keyframe + added = True + break + if not added: + self.keyframes.append(keyframe) + self.keyframes.sort(key=lambda k: k.start_percent) + + def get_index(self, index: int) -> Union[TimestepKeyframe, None]: + try: + return self.keyframes[index] + 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 + + 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() + group.keyframes[0] = keyframe + return group + + +# depending on model, AnimateDiff may inject into GroupNorm, so make sure GroupNorm will be clean +class disable_weight_init_clean_groupnorm(comfy.ops.disable_weight_init): + class GroupNorm(comfy.ops.disable_weight_init.GroupNorm): + def forward(self, input: Tensor) -> Tensor: + return F.group_norm( + input, self.num_groups, self.weight, self.bias, self.eps) +class manual_cast_clean_groupnorm(comfy.ops.manual_cast): + class GroupNorm(comfy.ops.manual_cast.GroupNorm): + def forward(self, input: Tensor) -> Tensor: + return F.group_norm( + input, self.num_groups, self.weight, self.bias, self.eps) + + +# adapted from comfy/sample.py +def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim1=False): + mask = mask.clone() + mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2]*multiplier, shape[3]*multiplier), mode="bilinear") + 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 + + +class WeightTypeException(TypeError): + "Raised when weight not compatible with AdvancedControlBase object" + pass From 3bf77b1f82f781a87a4f52d89682097d9e740520 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 19 Dec 2023 02:57:16 -0600 Subject: [PATCH 03/15] Scribble SparseCtrl working, RGB still in progress - added SparseCtrl-related nodes --- control/control.py | 439 +++++++--------------------------- control/control_sparsectrl.py | 379 +++++------------------------ control/nodes.py | 16 +- control/nodes_sparsectrl.py | 82 ++++++- control/utils.py | 341 ++++++++++++++++++++++++++ 5 files changed, 571 insertions(+), 686 deletions(-) diff --git a/control/control.py b/control/control.py index 7ed319f..8ce0a2e 100644 --- a/control/control.py +++ b/control/control.py @@ -9,352 +9,12 @@ import comfy.model_detection import comfy.controlnet as comfy_cn from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, broadcast_image_to -from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper -from .utils import (TimestepKeyframeGroup, LatentKeyframeGroup, ControlWeightType, ControlWeights, WeightTypeException, +from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod +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 -class AdvancedControlBase: - 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 - 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 - self.context_length = 0 - # 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 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 - 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] - 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 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: - 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 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 - # 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: - # 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: - return self.weights - - def set_cond_hint_mask(self, mask_hint): - self.mask_cond_hint_original = mask_hint - return self - - 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 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, 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 - # AKA latents that should not be influenced by ControlNet - if self.latent_keyframes is not None: - latent_count = x.size(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 - if self.sub_idxs: - mapped_indeces = {} - for i, actual in enumerate(self.sub_idxs): - mapped_indeces[actual] = i - for keyframe in self.latent_keyframes: - real_index = keyframe.batch_index - # if negative, count from end - if real_index < 0: - real_index += latent_count if self.sub_idxs is None else self.full_latent_length - - # if not mapping indeces, what you see is what you get - if mapped_indeces is None: - 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_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 - - # 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 - # 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) - 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 - - 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.calc_weight(i, x, len(control_input)) - 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.calc_weight(i, x, len(control_output)) - 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): - 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, 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, 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]: - 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 - 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: - 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() - 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 - # 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 - 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, load_device=None, manual_cast_dtype=None): super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, device=device, load_device=load_device, manual_cast_dtype=manual_cast_dtype) @@ -562,12 +222,88 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): class SparseCtrlAdvanced(ControlNetAdvanced): - def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None, load_device=None, manual_cast_dtype=None): + 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) self.add_compatible_weight(ControlWeightType.SPARSECTRL) + self.control_model: SparseControlNet = self.control_model # does nothing except help with IDE hints + self.sparse_settings = sparse_settings if sparse_settings is not None else SparseSettings.default() + 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]: + if control_prev is not None: + return control_prev + else: + return None + + dtype = self.control_model.dtype + if self.manual_cast_dtype is not None: + dtype = self.manual_cast_dtype + output_dtype = x_noisy.dtype + # prepare cond_hint, if needed + if self.sub_idxs is not None or self.cond_hint is None: + # clear out cond_hint and conditioning_mask + if self.cond_hint is not None: + del self.cond_hint + self.cond_hint = None + # first, figure out which cond idxs are relevant, and where they fit in + full_length = x_noisy.size(0)//batched_number if self.sub_idxs is None else self.full_latent_length + cond_idxs = self.sparse_settings.sparse_method.get_indeces(hint_length=self.cond_hint_original.size(0), full_length=full_length) + + range_idxs = list(range(full_length)) if self.sub_idxs is None else self.sub_idxs + hint_idxs = [] # idxs in cond_idxs + local_idxs = [] # idx to pun in final cond_hint + for i,cond_idx in enumerate(cond_idxs): + if cond_idx in range_idxs: + hint_idxs.append(i) + local_idxs.append(range_idxs.index(cond_idx)) + # sub_cond_hint now contains the hints relevant to current x_noisy + sub_cond_hint = self.cond_hint_original[hint_idxs].to(dtype).to(self.device) + + # scale cond_hints to match noisy input + if self.control_model.use_simplified_conditioning_embedding: + # RGB SparseCtrl; the inputs are latents - use bilinear to avoid blocky artifacts + sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3], x_noisy.shape[2], "bilinear", "center").to(dtype).to(self.device) + else: + # other SparseCtrl; inputs are typical images + sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) + # prepare cond_hint (b, c, h ,w) + cond_shape = list(sub_cond_hint.shape) + cond_shape[0] = len(range_idxs) + self.cond_hint = torch.zeros(cond_shape).to(dtype).to(self.device) + self.cond_hint[local_idxs] = sub_cond_hint[:] + # prepare cond_mask (b, 1, h, w) + cond_shape[1] = 1 + cond_mask = torch.zeros(cond_shape).to(dtype).to(self.device) + cond_mask[local_idxs] = 1.0 + # combine cond_hint and cond_mask into (b, c+1, h, w) + self.cond_hint = torch.cat([self.cond_hint, cond_mask], dim=1) + del sub_cond_hint + del cond_mask + # make cond_hint match x_noisy batch + 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=dtype) + + context = cond['c_crossattn'] + y = cond.get('y', None) + if y is not None: + y = y.to(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(dtype), hint=self.cond_hint, timesteps=timestep.float(), context=context.to(dtype), y=y) + return self.control_merge(None, control, control_prev, output_dtype) + def copy(self): - c = SparseCtrlAdvanced(self.control_model, self.timestep_keyframes, self.global_average_pooling, self.device, self.load_device, self.manual_cast_dtype) + c = SparseCtrlAdvanced(self.control_model, self.timestep_keyframes, self.sparse_settings, self.global_average_pooling, self.device, self.load_device, self.manual_cast_dtype) self.copy_to(c) self.copy_to_advanced(c) return c @@ -600,7 +336,6 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo control = ControlLLLiteAdvanced(timestep_keyframes=timestep_keyframe) # load Controll elif controlnet_type == ControlWeightType.SPARSECTRL: - #raise NotImplementedError("SparseCtrl has not been fully implemented yet!") control = load_sparsectrl(ckpt_path, controlnet_data=controlnet_data, timestep_keyframe=timestep_keyframe, model=model) # otherwise, load vanilla ControlNet else: @@ -611,11 +346,6 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo control = comfy_cn.load_controlnet(ckpt_path, model=model) finally: comfy.utils.load_torch_file = orig_load_torch_file - # from pathlib import Path - # with open(Path(__file__).parent.parent.parent / "controlnet_keys.txt", "w") as cfile: - # controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) - # for key in controlnet_data: - # cfile.write(f"{key}\n") return convert_to_advanced(control, timestep_keyframe=timestep_keyframe) @@ -640,7 +370,7 @@ def is_advanced_controlnet(input_object): return hasattr(input_object, "sub_idxs") -def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, timestep_keyframe: TimestepKeyframeGroup=None, model=None) -> SparseCtrlAdvanced: +def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, timestep_keyframe: TimestepKeyframeGroup=None, sparse_settings=None, model=None) -> SparseCtrlAdvanced: if controlnet_data is None: controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) # first, separate out motion part from normal controlnet part and attempt to load that portion @@ -781,10 +511,5 @@ def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, tim # both motion portion and controlnet portions are loaded; bring them together motion_wrapper.inject(control_model) - control = SparseCtrlAdvanced(control_model, timestep_keyframes=timestep_keyframe, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype) - new_state_dict = control_model.state_dict() - from pathlib import Path - with open(Path(__file__).parent.parent.parent / "sparcectrlstatedict.txt", "w") as cfile: - for key in new_state_dict: - cfile.write(f"{key}\n") + 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 diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index 8e1e15c..ee12807 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -2,6 +2,7 @@ #and modified #and then taken from comfy/cldm/cldm.py and modified again +from abc import ABC, abstractmethod import math from typing import Iterable, Union import torch @@ -15,7 +16,7 @@ from comfy.ldm.modules.diffusionmodules.util import ( timestep_embedding, ) -from comfy.cldm.cldm import ControlNet as ControlNet_cldm +from comfy.cldm.cldm import ControlNet as ControlNetCLDM from comfy.ldm.modules.attention import SpatialTransformer from comfy.ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample from comfy.ldm.util import exists @@ -28,37 +29,26 @@ import comfy.ops from .utils import TimestepKeyframeGroup, disable_weight_init_clean_groupnorm, prepare_mask_batch -class SparseControlNet(ControlNet_cldm): +class SparseControlNet(ControlNetCLDM): def __init__(self, *args,**kwargs): super().__init__(*args, **kwargs) hint_channels = kwargs.get("hint_channels") operations: disable_weight_init_clean_groupnorm = kwargs.get("operations", disable_weight_init_clean_groupnorm) device = kwargs.get("device", None) - use_simplified_conditioning_embedding = kwargs.get("use_simplified_conditioning_embedding", False) - if use_simplified_conditioning_embedding: + self.use_simplified_conditioning_embedding = kwargs.get("use_simplified_conditioning_embedding", False) + if self.use_simplified_conditioning_embedding: self.input_hint_block = TimestepEmbedSequential( - operations.conv_nd(self.dims, hint_channels, self.model_channels, 3, padding=1, dtype=self.dtype, device=device), + zero_module(operations.conv_nd(self.dims, hint_channels, self.model_channels, 3, padding=1, dtype=self.dtype, device=device)), + #zero_module(operations.conv_nd(self.dims, hint_channels, self.model_channels, 3, padding=1, dtype=self.dtype, device=device)), ) + self.motion_holder: MotionWrapperHolder = None - def forward(self, x, hint, timesteps, context, y=None, **kwargs): + def forward(self, x: Tensor, hint: Tensor, timesteps, context, y=None, **kwargs): t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype) emb = self.time_embed(t_emb) + # SparseCtrl sets noisy input to zeros x = torch.zeros_like(x) - - conditioning_mask1 = torch.ones_like(hint[:, :1]) - conditioning_mask2 = torch.zeros_like(hint[:, :1]) - conditioning_mask = conditioning_mask2 - conditioning_mask[0] = conditioning_mask1[0] - conditioning_mask[16] = conditioning_mask1[16] - #conditioning_mask[15] = conditioning_mask1[15] - #conditioning_mask[31] = conditioning_mask1[31] - modified_hint = torch.zeros_like(hint) - modified_hint[0] = hint[0] - modified_hint[16] = hint[16] - #modified_hint[15] = hint[15] - #modified_hint[31] = hint[31] - hint = torch.cat([modified_hint, conditioning_mask], dim=1) guided_hint = self.input_hint_block(hint, emb, context) outs = [] @@ -84,316 +74,59 @@ class SparseControlNet(ControlNet_cldm): return outs +class SparseSettings: + def __init__(self, sparse_method: 'SparseMethod'): + self.sparse_method = sparse_method + + @classmethod + def default(cls): + return cls(sparse_method=SparseSpreadMethod()) -# main class for holding SparseControlNet -class SparseControlNetOld(nn.Module): - def __init__( - self, - image_size, - in_channels, - model_channels, - hint_channels, - num_res_blocks, - dropout=0, - channel_mult=(1, 2, 4, 8), - conv_resample=True, - dims=2, - num_classes=None, - use_checkpoint=False, - dtype=torch.float32, - num_heads=-1, - num_head_channels=-1, - num_heads_upsample=-1, - use_scale_shift_norm=False, - resblock_updown=False, - use_new_attention_order=False, - use_spatial_transformer=False, # custom transformer support - transformer_depth=1, # custom transformer support - context_dim=None, # custom transformer support - n_embed=None, # custom support for prediction of discrete ids into codebook of first stage vq model - legacy=True, - disable_self_attentions=None, - num_attention_blocks=None, - disable_middle_self_attn=False, - use_linear_in_transformer=False, - adm_in_channels=None, - transformer_depth_middle=None, - transformer_depth_output=None, - device=None, - operations=disable_weight_init_clean_groupnorm, - **kwargs, - ): - super().__init__() - assert use_spatial_transformer == True, "use_spatial_transformer has to be true" - if use_spatial_transformer: - assert context_dim is not None, 'Fool!! You forgot to include the dimension of your cross-attention conditioning...' - if context_dim is not None: - assert use_spatial_transformer, 'Fool!! You forgot to use the spatial transformer for your cross-attention conditioning...' - # from omegaconf.listconfig import ListConfig - # if type(context_dim) == ListConfig: - # context_dim = list(context_dim) - if num_heads_upsample == -1: - num_heads_upsample = num_heads +class SparseMethod(ABC): + SPREAD = "spread" + INDEX = "index" + def __init__(self, method: str): + self.method = method - if num_heads == -1: - assert num_head_channels != -1, 'Either num_heads or num_head_channels has to be set' + @abstractmethod + def get_indeces(self, hint_length: int, full_length: int) -> list[int]: + pass - if num_head_channels == -1: - assert num_heads != -1, 'Either num_heads or num_head_channels has to be set' - self.dims = dims - self.image_size = image_size - self.in_channels = in_channels - self.model_channels = model_channels +class SparseSpreadMethod(SparseMethod): + def __init__(self, from_start=True): + super().__init__(self.SPREAD) + self.from_start = from_start - if isinstance(num_res_blocks, int): - self.num_res_blocks = len(channel_mult) * [num_res_blocks] - else: - if len(num_res_blocks) != len(channel_mult): - raise ValueError("provide num_res_blocks either as an int (globally constant) or " - "as a list/tuple (per-level) with the same length as channel_mult") - self.num_res_blocks = num_res_blocks + def get_indeces(self, hint_length: int, full_length: int) -> list[int]: + # handle special case of 1 hint image + if hint_length == 1: + return [0] if self.from_start else [full_length-1] + # handle special case of equal or more hint images than full length + if hint_length >= full_length: + return list(range(min(hint_length, full_length))) + if hint_length == 2: + return [0, full_length-1] + # TODO: other cases/modes - if disable_self_attentions is not None: - # should be a list of booleans, indicating whether to disable self-attention in TransformerBlocks or not - assert len(disable_self_attentions) == len(channel_mult) - if num_attention_blocks is not None: - assert len(num_attention_blocks) == len(self.num_res_blocks) - assert all(map(lambda i: self.num_res_blocks[i] >= num_attention_blocks[i], range(len(num_attention_blocks)))) - transformer_depth = transformer_depth[:] +class SparseIndexMethod(SparseMethod): + def __init__(self, idxs: list[int]): + super().__init__(self.INDEX) + self.idxs = idxs - self.dropout = dropout - self.channel_mult = channel_mult - self.conv_resample = conv_resample - self.num_classes = num_classes - self.use_checkpoint = use_checkpoint - self.dtype = dtype - self.num_heads = num_heads - self.num_head_channels = num_head_channels - self.num_heads_upsample = num_heads_upsample - self.predict_codebook_ids = n_embed is not None - - time_embed_dim = model_channels * 4 - self.time_embed = nn.Sequential( - operations.Linear(model_channels, time_embed_dim, dtype=self.dtype, device=device), - nn.SiLU(), - operations.Linear(time_embed_dim, time_embed_dim, dtype=self.dtype, device=device), - ) - - if self.num_classes is not None: - if isinstance(self.num_classes, int): - self.label_emb = nn.Embedding(num_classes, time_embed_dim) - elif self.num_classes == "continuous": - print("setting up linear c_adm embedding layer") - self.label_emb = nn.Linear(1, time_embed_dim) - elif self.num_classes == "sequential": - assert adm_in_channels is not None - self.label_emb = nn.Sequential( - nn.Sequential( - operations.Linear(adm_in_channels, time_embed_dim, dtype=self.dtype, device=device), - nn.SiLU(), - operations.Linear(time_embed_dim, time_embed_dim, dtype=self.dtype, device=device), - ) - ) + def get_indeces(self, hint_length: int, full_length: int) -> list[int]: + new_idxs = [] + for idx in self.idxs: + if idx < 0: + new_idxs.append(full_length+idx) else: - raise ValueError() - - self.input_blocks = nn.ModuleList( - [ - TimestepEmbedSequential( - operations.conv_nd(dims, in_channels, model_channels, 3, padding=1, dtype=self.dtype, device=device) - ) - ] - ) - self.zero_convs = nn.ModuleList([self.make_zero_conv(model_channels, operations=operations, dtype=self.dtype, device=device)]) - - self.input_hint_block = TimestepEmbedSequential( - operations.conv_nd(dims, hint_channels, 16, 3, padding=1, dtype=self.dtype, device=device), - nn.SiLU(), - operations.conv_nd(dims, 16, 16, 3, padding=1, dtype=self.dtype, device=device), - nn.SiLU(), - operations.conv_nd(dims, 16, 32, 3, padding=1, stride=2, dtype=self.dtype, device=device), - nn.SiLU(), - operations.conv_nd(dims, 32, 32, 3, padding=1, dtype=self.dtype, device=device), - nn.SiLU(), - operations.conv_nd(dims, 32, 96, 3, padding=1, stride=2, dtype=self.dtype, device=device), - nn.SiLU(), - operations.conv_nd(dims, 96, 96, 3, padding=1, dtype=self.dtype, device=device), - nn.SiLU(), - operations.conv_nd(dims, 96, 256, 3, padding=1, stride=2, dtype=self.dtype, device=device), - nn.SiLU(), - operations.conv_nd(dims, 256, model_channels, 3, padding=1, dtype=self.dtype, device=device) - ) - - self._feature_size = model_channels - input_block_chans = [model_channels] - ch = model_channels - ds = 1 - for level, mult in enumerate(channel_mult): - for nr in range(self.num_res_blocks[level]): - layers = [ - ResBlock( - ch, - time_embed_dim, - dropout, - out_channels=mult * model_channels, - dims=dims, - use_checkpoint=use_checkpoint, - use_scale_shift_norm=use_scale_shift_norm, - dtype=self.dtype, - device=device, - operations=operations, - ) - ] - ch = mult * model_channels - num_transformers = transformer_depth.pop(0) - if num_transformers > 0: - if num_head_channels == -1: - dim_head = ch // num_heads - else: - num_heads = ch // num_head_channels - dim_head = num_head_channels - if legacy: - #num_heads = 1 - dim_head = ch // num_heads if use_spatial_transformer else num_head_channels - if exists(disable_self_attentions): - disabled_sa = disable_self_attentions[level] - else: - disabled_sa = False - - if not exists(num_attention_blocks) or nr < num_attention_blocks[level]: - layers.append( - SpatialTransformer( - ch, num_heads, dim_head, depth=num_transformers, context_dim=context_dim, - disable_self_attn=disabled_sa, use_linear=use_linear_in_transformer, - use_checkpoint=use_checkpoint, dtype=self.dtype, device=device, operations=operations - ) - ) - self.input_blocks.append(TimestepEmbedSequential(*layers)) - self.zero_convs.append(self.make_zero_conv(ch, operations=operations, dtype=self.dtype, device=device)) - self._feature_size += ch - input_block_chans.append(ch) - if level != len(channel_mult) - 1: - out_ch = ch - self.input_blocks.append( - TimestepEmbedSequential( - ResBlock( - ch, - time_embed_dim, - dropout, - out_channels=out_ch, - dims=dims, - use_checkpoint=use_checkpoint, - use_scale_shift_norm=use_scale_shift_norm, - down=True, - dtype=self.dtype, - device=device, - operations=operations - ) - if resblock_updown - else Downsample( - ch, conv_resample, dims=dims, out_channels=out_ch, dtype=self.dtype, device=device, operations=operations - ) - ) - ) - ch = out_ch - input_block_chans.append(ch) - self.zero_convs.append(self.make_zero_conv(ch, operations=operations, dtype=self.dtype, device=device)) - ds *= 2 - self._feature_size += ch - - if num_head_channels == -1: - dim_head = ch // num_heads - else: - num_heads = ch // num_head_channels - dim_head = num_head_channels - if legacy: - #num_heads = 1 - dim_head = ch // num_heads if use_spatial_transformer else num_head_channels - mid_block = [ - ResBlock( - ch, - time_embed_dim, - dropout, - dims=dims, - use_checkpoint=use_checkpoint, - use_scale_shift_norm=use_scale_shift_norm, - dtype=self.dtype, - device=device, - operations=operations - )] - if transformer_depth_middle >= 0: - mid_block += [SpatialTransformer( # always uses a self-attn - ch, num_heads, dim_head, depth=transformer_depth_middle, context_dim=context_dim, - disable_self_attn=disable_middle_self_attn, use_linear=use_linear_in_transformer, - use_checkpoint=use_checkpoint, dtype=self.dtype, device=device, operations=operations - ), - ResBlock( - ch, - time_embed_dim, - dropout, - dims=dims, - use_checkpoint=use_checkpoint, - use_scale_shift_norm=use_scale_shift_norm, - dtype=self.dtype, - device=device, - operations=operations - )] - self.middle_block = TimestepEmbedSequential(*mid_block) - self.middle_block_out = self.make_zero_conv(ch, operations=operations, dtype=self.dtype, device=device) - self._feature_size += ch - - #self._motion_wrapper: SparseCtrlMotionWrapper = None - - def make_zero_conv(self, channels, operations=None, dtype=None, device=None): - return TimestepEmbedSequential(operations.conv_nd(self.dims, channels, channels, 1, padding=0, dtype=dtype, device=device)) - - def forward(self, x, hint, timesteps, context, y=None, **kwargs): - t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype) - emb = self.time_embed(t_emb) - - x = torch.zeros_like(x) - - conditioning_mask1 = torch.ones_like(hint[:, :1]) - conditioning_mask2 = torch.zeros_like(hint[:, :1]) - conditioning_mask = conditioning_mask2 - conditioning_mask[0] = conditioning_mask1[0] - conditioning_mask[16] = conditioning_mask1[16] - #conditioning_mask[15] = conditioning_mask1[15] - #conditioning_mask[31] = conditioning_mask1[31] - modified_hint = torch.zeros_like(hint) - modified_hint[0] = hint[0] - modified_hint[16] = hint[16] - #modified_hint[15] = hint[15] - #modified_hint[31] = hint[31] - hint = torch.cat([modified_hint, conditioning_mask], dim=1) - guided_hint = self.input_hint_block(hint, emb, context) - - outs = [] - - hs = [] - if self.num_classes is not None: - assert y.shape[0] == x.shape[0] - emb = emb + self.label_emb(y) - - h = x - for module, zero_conv in zip(self.input_blocks, self.zero_convs): - if guided_hint is not None: - h = module(h, emb, context) - h += guided_hint - guided_hint = None - else: - h = module(h, emb, context) - outs.append(zero_conv(h, emb, context)) - - h = self.middle_block(h, emb, context) - outs.append(self.middle_block_out(h, emb, context)) - - return outs + new_idxs.append(idx) + return new_idxs +######################################### # motion-related portion of controlnet class BlockType: UP = "up" @@ -435,6 +168,11 @@ def get_position_encoding_max_len(mm_state_dict: dict[str, Tensor], mm_name: str raise ValueError(f"No pos_encoder.pe found in SparseCtrl state_dict - {mm_name} is not a valid SparseCtrl model!") +class MotionWrapperHolder: + def __init__(self, motion_wrapper: 'SparseCtrlMotionWrapper'): + self.motion_wrapper = motion_wrapper + + class SparseCtrlMotionWrapper(nn.Module): def __init__(self, mm_state_dict: dict[str, Tensor]): super().__init__() @@ -460,7 +198,7 @@ class SparseCtrlMotionWrapper(nn.Module): # inject mid block, if present if self.mid_block is not None: self._inject([unet.middle_block], [self.mid_block]) - #unet._motion_wrapper = self + unet.motion_holder = MotionWrapperHolder(self) def _inject(self, unet_blocks: nn.ModuleList, mm_blocks: nn.ModuleList): # Rules for injection: @@ -502,7 +240,8 @@ class SparseCtrlMotionWrapper(nn.Module): self._eject(unet.input_blocks) # remove from middle block (encapsulate in list to make compatible) self._eject([unet.middle_block]) - #del unet._motion_wrapper + del unet.motion_holder + unet.motion_holder = None def _eject(self, unet_blocks: nn.ModuleList): # eject all VanillaTemporalModule objects from all blocks diff --git a/control/nodes.py b/control/nodes.py index dd9ecde..2ed68e1 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -9,7 +9,7 @@ from .utils import StrengthInterpolation as SI from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, SoftT2IAdapterWeights, CustomT2IAdapterWeights) from .nodes_latent_keyframe import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode -from .nodes_deprecated import LoadImagesFromDirectory +from .nodes_sparsectrl import SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, VAEEncodePreprocessor from .logger import logger @@ -214,8 +214,11 @@ NODE_CLASS_MAPPINGS = { "SoftT2IAdapterWeights": SoftT2IAdapterWeights, "CustomT2IAdapterWeights": CustomT2IAdapterWeights, "ACN_DefaultUniversalWeights": DefaultWeights, - # Image - "LoadImagesFromDirectory": LoadImagesFromDirectory + # SparseCtrl + "ACN_VAEEncodePreprocessor": VAEEncodePreprocessor, + "ACN_SparseCtrlLoaderAdvanced": SparseCtrlLoaderAdvanced, + "ACN_SparseCtrlIndexMethodNode": SparseIndexMethodNode, + "ACN_SparseCtrlSpreadMethodNode": SparseSpreadMethodNode, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -238,6 +241,9 @@ NODE_DISPLAY_NAME_MAPPINGS = { "SoftT2IAdapterWeights": "T2IAdapter Soft Weights 🛂🅐🅒🅝", "CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝", "ACN_DefaultUniversalWeights": "Force Default Weights 🛂🅐🅒🅝", - # Image - "LoadImagesFromDirectory": "Load Images [DEPRECATED] 🛂🅐🅒🅝" + # SparseCtrl + "ACN_VAEEncodePreprocessor": "RGB SparseCtrl 🛂🅐🅒🅝", + "ACN_SparseCtrlLoaderAdvanced": "Load SparseCtrl Model 🛂🅐🅒🅝", + "ACN_SparseCtrlIndexMethodNode": "SparseCtrl Index Method 🛂🅐🅒🅝", + "ACN_SparseCtrlSpreadMethodNode": "SparseCtrl Spread Method 🛂🅐🅒🅝", } diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index 2e62629..63e0ea0 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -1,6 +1,8 @@ import folder_paths +from nodes import VAEEncode from .utils import TimestepKeyframeGroup +from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod from .control import load_sparsectrl @@ -13,6 +15,7 @@ class SparseCtrlLoaderAdvanced: "control_net_name": (folder_paths.get_filename_list("controlnet"), ), }, "optional": { + "sparse_method": ("SPARSE_METHOD", ), "tk_optional": ("TIMESTEP_KEYFRAME", ), } } @@ -20,9 +23,80 @@ class SparseCtrlLoaderAdvanced: RETURN_TYPES = ("CONTROL_NET", ) FUNCTION = "load_controlnet" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" - def load_controlnet(self, control_net_name: str, tk_optional: TimestepKeyframeGroup=None): + def load_controlnet(self, control_net_name: str, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) - controlnet = load_sparsectrl(controlnet_path, timestep_keyframe=tk_optional) - return controlnet + sparse_settings = SparseSettings(sparse_method=sparse_method) + controlnet = load_sparsectrl(controlnet_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings) + return (controlnet,) + + +class SparseIndexMethodNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "indeces": ("STRING", {"default": "0"}), + } + } + + RETURN_TYPES = ("SPARSE_METHOD",) + FUNCTION = "get_method" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" + + def get_method(self, indeces: str): + idxs = [] + # get indeces from string + str_idxs = [x.strip() for x in indeces.strip().split(",")] + for str_idx in str_idxs: + try: + idx = int(str_idx) + idxs.append(idx) + except ValueError: + raise ValueError(f"'{str_idx}' is not a valid integer index.") + if len(idxs) == 0: + raise ValueError(f"No indeces were listed in Sparse Index Method.") + return (SparseIndexMethod(idxs),) + + +class SparseSpreadMethodNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "from_start": ("BOOLEAN", {"default": True}), + } + } + + RETURN_TYPES = ("SPARSE_METHOD",) + FUNCTION = "get_method" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" + + def get_method(self, from_start: bool): + return (SparseSpreadMethod(from_start=from_start),) + + +class VAEEncodePreprocessor: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE", ), + "vae": ("VAE", ) + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("latent_IMAGE",) + FUNCTION = "preprocess_images" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl/preprocess" + + def preprocess_images(self, vae, image): + image = VAEEncode.vae_encode_crop_pixels(image) + encoded = vae.encode(image[:,:,:,:3]) + encoded = encoded.movedim(1,-1) + return (encoded,) diff --git a/control/utils.py b/control/utils.py index 7772601..edd45a8 100644 --- a/control/utils.py +++ b/control/utils.py @@ -4,6 +4,7 @@ from torch import Tensor import torch.nn.functional as F import comfy.ops import comfy.utils +from comfy.controlnet import ControlBase, broadcast_image_to def load_torch_file_with_dict_factory(controlnet_data: dict[str, Tensor], orig_load_torch_file: Callable): @@ -258,3 +259,343 @@ def linear_conversion(x, x_min=0.0, x_max=1.0, new_min=0.0, new_max=1.0): class WeightTypeException(TypeError): "Raised when weight not compatible with AdvancedControlBase object" pass + + +class AdvancedControlBase: + 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 + 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 + self.context_length = 0 + # 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 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 + 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] + 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 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: + 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 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 + # 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: + # 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: + return self.weights + + def set_cond_hint_mask(self, mask_hint): + self.mask_cond_hint_original = mask_hint + return self + + 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 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, 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 + # AKA latents that should not be influenced by ControlNet + if self.latent_keyframes is not None: + latent_count = x.size(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 + if self.sub_idxs: + mapped_indeces = {} + for i, actual in enumerate(self.sub_idxs): + mapped_indeces[actual] = i + for keyframe in self.latent_keyframes: + real_index = keyframe.batch_index + # if negative, count from end + if real_index < 0: + real_index += latent_count if self.sub_idxs is None else self.full_latent_length + + # if not mapping indeces, what you see is what you get + if mapped_indeces is None: + 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_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 + + # 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 + # 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) + 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 + + 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.calc_weight(i, x, len(control_input)) + 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.calc_weight(i, x, len(control_output)) + 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): + 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, 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, 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]: + 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 + 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: + 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() + 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 + # 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 + copied.weights_override = self.weights_override + copied.latent_keyframe_override = self.latent_keyframe_override From 2757a4f17e8f7bd0ae3649b0dc357aba558f60b4 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 19 Dec 2023 10:59:45 -0600 Subject: [PATCH 04/15] Properly set video length in SparseCtrl motion module, account for potential batched_number changes --- control/control.py | 14 +++++++++----- control/control_sparsectrl.py | 14 ++++++++++---- 2 files changed, 19 insertions(+), 9 deletions(-) diff --git a/control/control.py b/control/control.py index 8ce0a2e..cc838eb 100644 --- a/control/control.py +++ b/control/control.py @@ -245,14 +245,18 @@ class SparseCtrlAdvanced(ControlNetAdvanced): if self.manual_cast_dtype is not None: dtype = self.manual_cast_dtype output_dtype = x_noisy.dtype + # set actual input length on motion model + actual_length = x_noisy.size(0)//batched_number + full_length = actual_length if self.sub_idxs is None else self.full_latent_length + self.control_model.set_actual_length(actual_length=actual_length, full_length=full_length) # prepare cond_hint, if needed - if self.sub_idxs is not None or self.cond_hint is None: + dim_mult = 1 if self.control_model.use_simplified_conditioning_embedding else 8 + if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[0] != self.cond_hint.shape[0] or x_noisy.shape[2]*dim_mult != self.cond_hint.shape[2] or x_noisy.shape[3]*dim_mult != self.cond_hint.shape[3]: # clear out cond_hint and conditioning_mask if self.cond_hint is not None: del self.cond_hint self.cond_hint = None # first, figure out which cond idxs are relevant, and where they fit in - full_length = x_noisy.size(0)//batched_number if self.sub_idxs is None else self.full_latent_length cond_idxs = self.sparse_settings.sparse_method.get_indeces(hint_length=self.cond_hint_original.size(0), full_length=full_length) range_idxs = list(range(full_length)) if self.sub_idxs is None else self.sub_idxs @@ -285,9 +289,9 @@ class SparseCtrlAdvanced(ControlNetAdvanced): self.cond_hint = torch.cat([self.cond_hint, cond_mask], dim=1) del sub_cond_hint del cond_mask - # make cond_hint match x_noisy batch - 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 cond_hint match x_noisy batch + 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=dtype) diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index ee12807..57bdec8 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -43,6 +43,10 @@ class SparseControlNet(ControlNetCLDM): ) self.motion_holder: MotionWrapperHolder = None + def set_actual_length(self, actual_length: int, full_length: int): + if self.motion_holder is not None: + self.motion_holder.motion_wrapper.set_video_length(video_length=actual_length, full_length=full_length) + def forward(self, x: Tensor, hint: Tensor, timesteps, context, y=None, **kwargs): t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype) emb = self.time_embed(t_emb) @@ -256,10 +260,12 @@ class SparseCtrlMotionWrapper(nn.Module): def set_video_length(self, video_length: int, full_length: int): self.AD_video_length = video_length - for block in self.down_blocks: - block.set_video_length(video_length, full_length) - for block in self.up_blocks: - block.set_video_length(video_length, full_length) + if self.down_blocks is not None: + for block in self.down_blocks: + block.set_video_length(video_length, full_length) + if self.up_blocks is not None: + for block in self.up_blocks: + block.set_video_length(video_length, full_length) if self.mid_block is not None: self.mid_block.set_video_length(video_length, full_length) From 1f66b7c5816106d2d9d614afdf781dc2ac8619dc Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 19 Dec 2023 11:03:36 -0600 Subject: [PATCH 05/15] Removed unnecessary shape[0] comparison - already handled --- control/control.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/control/control.py b/control/control.py index cc838eb..f81964f 100644 --- a/control/control.py +++ b/control/control.py @@ -251,7 +251,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced): self.control_model.set_actual_length(actual_length=actual_length, full_length=full_length) # prepare cond_hint, if needed dim_mult = 1 if self.control_model.use_simplified_conditioning_embedding else 8 - if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[0] != self.cond_hint.shape[0] or x_noisy.shape[2]*dim_mult != self.cond_hint.shape[2] or x_noisy.shape[3]*dim_mult != self.cond_hint.shape[3]: + if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2]*dim_mult != self.cond_hint.shape[2] or x_noisy.shape[3]*dim_mult != self.cond_hint.shape[3]: # clear out cond_hint and conditioning_mask if self.cond_hint is not None: del self.cond_hint From be0460ec04f1e41b4757f9257148280f4faccb66 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 20 Dec 2023 23:43:42 -0600 Subject: [PATCH 06/15] Made SparseIndex and SparseSpread much more robust, added usable spread types for SparseSpread --- control/control.py | 2 +- control/control_sparsectrl.py | 79 ++++++++++++++++++++++++++++------- control/nodes_sparsectrl.py | 18 ++++---- 3 files changed, 75 insertions(+), 24 deletions(-) diff --git a/control/control.py b/control/control.py index f81964f..c9bc138 100644 --- a/control/control.py +++ b/control/control.py @@ -257,7 +257,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced): del self.cond_hint self.cond_hint = None # first, figure out which cond idxs are relevant, and where they fit in - cond_idxs = self.sparse_settings.sparse_method.get_indeces(hint_length=self.cond_hint_original.size(0), full_length=full_length) + cond_idxs = self.sparse_settings.sparse_method.get_indexes(hint_length=self.cond_hint_original.size(0), full_length=full_length) range_idxs = list(range(full_length)) if self.sub_idxs is None else self.sub_idxs hint_idxs = [] # idxs in cond_idxs diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index 57bdec8..680a6d3 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -4,6 +4,7 @@ from abc import ABC, abstractmethod import math +import numpy as np from typing import Iterable, Union import torch import torch as th @@ -94,25 +95,53 @@ class SparseMethod(ABC): self.method = method @abstractmethod - def get_indeces(self, hint_length: int, full_length: int) -> list[int]: + def get_indexes(self, hint_length: int, full_length: int) -> list[int]: pass class SparseSpreadMethod(SparseMethod): - def __init__(self, from_start=True): - super().__init__(self.SPREAD) - self.from_start = from_start + UNIFORM = "uniform" + STARTING = "starting" + ENDING = "ending" + CENTER = "center" - def get_indeces(self, hint_length: int, full_length: int) -> list[int]: + LIST = [UNIFORM, STARTING, ENDING, CENTER] + + def __init__(self, spread=UNIFORM): + super().__init__(self.SPREAD) + self.spread = spread + + def get_indexes(self, hint_length: int, full_length: int) -> list[int]: + # if hint_length >= full_length, limit hints to full_length + if hint_length >= full_length: + return list(range(full_length)) # handle special case of 1 hint image if hint_length == 1: - return [0] if self.from_start else [full_length-1] - # handle special case of equal or more hint images than full length - if hint_length >= full_length: - return list(range(min(hint_length, full_length))) - if hint_length == 2: - return [0, full_length-1] - # TODO: other cases/modes + if self.spread in [self.UNIFORM, self.STARTING]: + return [0] + elif self.spread == self.ENDING: + return [full_length-1] + elif self.spread == self.CENTER: + # return second (of three) values as the center + return [np.linspace(0, full_length-1, 3, endpoint=True, dtype=int)[1]] + else: + raise ValueError(f"Unrecognized spread: {self.spread}") + # otherwise, handle other cases + if self.spread == self.UNIFORM: + return list(np.linspace(0, full_length-1, hint_length, endpoint=True, dtype=int)) + elif self.spread == self.STARTING: + # make split 1 larger, remove last element + return list(np.linspace(0, full_length-1, hint_length+1, endpoint=True, dtype=int))[:-1] + elif self.spread == self.ENDING: + # make split 1 larger, remove first element + return list(np.linspace(0, full_length-1, hint_length+1, endpoint=True, dtype=int))[1:] + elif self.spread == self.CENTER: + # if hint length is not 3 greater than full length, do STARTING behavior + if full_length-hint_length < 3: + return list(np.linspace(0, full_length-1, hint_length+1, endpoint=True, dtype=int))[:-1] + # otherwise, get linspace of 2 greater than needed, then cut off first and last + return list(np.linspace(0, full_length-1, hint_length, endpoint=True, dtype=int))[1:-1] + return ValueError(f"Unrecognized spread: {self.spread}") class SparseIndexMethod(SparseMethod): @@ -120,13 +149,31 @@ class SparseIndexMethod(SparseMethod): super().__init__(self.INDEX) self.idxs = idxs - def get_indeces(self, hint_length: int, full_length: int) -> list[int]: + def get_indexes(self, hint_length: int, full_length: int) -> list[int]: + orig_hint_length = hint_length + if hint_length > full_length: + hint_length = full_length + # if idxs is less than hint_length, throw error + if len(self.idxs) < hint_length: + err_msg = f"There are not enough indexes ({len(self.idxs)}) provided to fit the usable {hint_length} input images." + if orig_hint_length != hint_length: + err_msg = f"{err_msg} (original input images: {orig_hint_length})" + raise ValueError(err_msg) + # cap idxs to hint_length + idxs = self.idxs[:hint_length] new_idxs = [] - for idx in self.idxs: + real_idxs = set() + for idx in idxs: if idx < 0: - new_idxs.append(full_length+idx) + real_idx = full_length+idx + if real_idx in real_idxs: + raise ValueError(f"Index '{idx}' maps to '{real_idx}' and is duplicate - indexes in Sparse Index Method must be unique.") else: - new_idxs.append(idx) + real_idx = idx + if real_idx in real_idxs: + raise ValueError(f"Index '{idx}' is duplicate (or a negative index is equivalent) - indexes in Sparse Index Method must be unique.") + real_idxs.add(real_idx) + new_idxs.append(real_idx) return new_idxs diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index 63e0ea0..e396550 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -37,7 +37,7 @@ class SparseIndexMethodNode: def INPUT_TYPES(s): return { "required": { - "indeces": ("STRING", {"default": "0"}), + "indexes": ("STRING", {"default": "0"}), } } @@ -46,18 +46,22 @@ class SparseIndexMethodNode: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" - def get_method(self, indeces: str): + def get_method(self, indexes: str): idxs = [] + unique_idxs = set() # get indeces from string - str_idxs = [x.strip() for x in indeces.strip().split(",")] + str_idxs = [x.strip() for x in indexes.strip().split(",")] for str_idx in str_idxs: try: idx = int(str_idx) + if idx in unique_idxs: + raise ValueError(f"'{idx}' is duplicated; indexes must be unique.") idxs.append(idx) + unique_idxs.add(idx) except ValueError: raise ValueError(f"'{str_idx}' is not a valid integer index.") if len(idxs) == 0: - raise ValueError(f"No indeces were listed in Sparse Index Method.") + raise ValueError(f"No indexes were listed in Sparse Index Method.") return (SparseIndexMethod(idxs),) @@ -66,7 +70,7 @@ class SparseSpreadMethodNode: def INPUT_TYPES(s): return { "required": { - "from_start": ("BOOLEAN", {"default": True}), + "spread": (SparseSpreadMethod.LIST), } } @@ -75,8 +79,8 @@ class SparseSpreadMethodNode: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" - def get_method(self, from_start: bool): - return (SparseSpreadMethod(from_start=from_start),) + def get_method(self, spread: str): + return (SparseSpreadMethod(spread=spread),) class VAEEncodePreprocessor: From 2fc43f458b28f8b92803395384cc170b7a6f4e36 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 21 Dec 2023 01:24:13 -0600 Subject: [PATCH 07/15] Fixed Sparse Spread Method dropdown not rendering, added use_motion toggle --- control/control.py | 7 ++++--- control/control_sparsectrl.py | 5 +++-- control/nodes_sparsectrl.py | 7 ++++--- 3 files changed, 11 insertions(+), 8 deletions(-) diff --git a/control/control.py b/control/control.py index c9bc138..df53914 100644 --- a/control/control.py +++ b/control/control.py @@ -374,7 +374,7 @@ def is_advanced_controlnet(input_object): return hasattr(input_object, "sub_idxs") -def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, timestep_keyframe: TimestepKeyframeGroup=None, sparse_settings=None, model=None) -> SparseCtrlAdvanced: +def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, timestep_keyframe: TimestepKeyframeGroup=None, sparse_settings=SparseSettings.default(), model=None) -> SparseCtrlAdvanced: if controlnet_data is None: controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) # first, separate out motion part from normal controlnet part and attempt to load that portion @@ -512,8 +512,9 @@ def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, tim if filename.endswith("_shuffle") or filename.endswith("_shuffle_fp16"): #TODO: smarter way of enabling global_average_pooling global_average_pooling = True - # both motion portion and controlnet portions are loaded; bring them together - motion_wrapper.inject(control_model) + # both motion portion and controlnet portions are loaded; bring them together if using motion model + if sparse_settings.use_motion: + motion_wrapper.inject(control_model) 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 diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index 680a6d3..717e745 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -80,12 +80,13 @@ class SparseControlNet(ControlNetCLDM): class SparseSettings: - def __init__(self, sparse_method: 'SparseMethod'): + def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True): self.sparse_method = sparse_method + self.use_motion = use_motion @classmethod def default(cls): - return cls(sparse_method=SparseSpreadMethod()) + return SparseSettings(sparse_method=SparseSpreadMethod(), use_motion=True) class SparseMethod(ABC): diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index e396550..7759013 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -13,6 +13,7 @@ class SparseCtrlLoaderAdvanced: return { "required": { "control_net_name": (folder_paths.get_filename_list("controlnet"), ), + "use_motion": ("BOOLEAN", {"default": True}, ), }, "optional": { "sparse_method": ("SPARSE_METHOD", ), @@ -25,9 +26,9 @@ class SparseCtrlLoaderAdvanced: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" - def load_controlnet(self, control_net_name: str, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): + def load_controlnet(self, control_net_name: str, use_motion: bool, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) - sparse_settings = SparseSettings(sparse_method=sparse_method) + sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion) controlnet = load_sparsectrl(controlnet_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings) return (controlnet,) @@ -70,7 +71,7 @@ class SparseSpreadMethodNode: def INPUT_TYPES(s): return { "required": { - "spread": (SparseSpreadMethod.LIST), + "spread": (SparseSpreadMethod.LIST,), } } From 2ca56093ff23dfaf16c75ea6bb471564567f2b45 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 21 Dec 2023 23:35:36 -0600 Subject: [PATCH 08/15] Add motion_strength and motion_scale to Load SparseCtrl node --- control/control.py | 7 +++++ control/control_sparsectrl.py | 53 +++++++++++++++++++++++++++-------- control/nodes_sparsectrl.py | 6 ++-- 3 files changed, 53 insertions(+), 13 deletions(-) diff --git a/control/control.py b/control/control.py index df53914..b473b3e 100644 --- a/control/control.py +++ b/control/control.py @@ -306,6 +306,13 @@ class SparseCtrlAdvanced(ControlNetAdvanced): control = self.control_model(x=x_noisy.to(dtype), hint=self.cond_hint, timesteps=timestep.float(), context=context.to(dtype), y=y) return self.control_merge(None, control, control_prev, output_dtype) + def pre_run_advanced(self, model, percent_to_timestep_function): + super().pre_run_advanced(model, percent_to_timestep_function) + if self.control_model.motion_holder is not None: + self.control_model.motion_holder.motion_wrapper.reset() + self.control_model.motion_holder.motion_wrapper.set_strength(self.sparse_settings.motion_strength) + self.control_model.motion_holder.motion_wrapper.set_scale_multiplier(self.sparse_settings.motion_scale) + def copy(self): c = SparseCtrlAdvanced(self.control_model, self.timestep_keyframes, self.sparse_settings, self.global_average_pooling, self.device, self.load_device, self.manual_cast_dtype) self.copy_to(c) diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index 717e745..1abb8cd 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -80,9 +80,11 @@ class SparseControlNet(ControlNetCLDM): class SparseSettings: - def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True): + def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True, motion_strength=1.0, motion_scale=1.0): self.sparse_method = sparse_method self.use_motion = use_motion + self.motion_strength = motion_strength + self.motion_scale = motion_scale @classmethod def default(cls): @@ -318,18 +320,32 @@ class SparseCtrlMotionWrapper(nn.Module): self.mid_block.set_video_length(video_length, full_length) def set_scale_multiplier(self, multiplier: Union[float, None]): - for block in self.down_blocks: - block.set_scale_multiplier(multiplier) - for block in self.up_blocks: - block.set_scale_multiplier(multiplier) + if self.down_blocks is not None: + for block in self.down_blocks: + block.set_scale_multiplier(multiplier) + if self.up_blocks is not None: + for block in self.up_blocks: + block.set_scale_multiplier(multiplier) if self.mid_block is not None: self.mid_block.set_scale_multiplier(multiplier) + def set_strength(self, strength: float): + if self.down_blocks is not None: + for block in self.down_blocks: + block.set_strength(strength) + if self.up_blocks is not None: + for block in self.up_blocks: + block.set_strength(strength) + if self.mid_block is not None: + self.mid_block.set_strength(strength) + def reset_temp_vars(self): - for block in self.down_blocks: - block.reset_temp_vars() - for block in self.up_blocks: - block.reset_temp_vars() + if self.down_blocks is not None: + for block in self.down_blocks: + block.reset_temp_vars() + if self.up_blocks is not None: + for block in self.up_blocks: + block.reset_temp_vars() if self.mid_block is not None: self.mid_block.reset_temp_vars() @@ -375,6 +391,10 @@ class MotionModule(nn.Module): for motion_module in self.motion_modules: motion_module.set_sub_idxs(sub_idxs) + def set_strength(self, strength: float): + for motion_module in self.motion_modules: + motion_module.set_strength(strength) + def reset_temp_vars(self): for motion_module in self.motion_modules: motion_module.reset_temp_vars() @@ -399,7 +419,7 @@ class VanillaTemporalModule(nn.Module): zero_initialize=True, ): super().__init__() - + self.strength = 1.0 self.temporal_transformer = TemporalTransformer3DModel( in_channels=in_channels, num_attention_heads=num_attention_heads, @@ -430,11 +450,22 @@ class VanillaTemporalModule(nn.Module): def set_sub_idxs(self, sub_idxs: list[int]): self.temporal_transformer.set_sub_idxs(sub_idxs) + def set_strength(self, strength: float): + self.strength = strength + def reset_temp_vars(self): + self.set_strength(1.0) self.temporal_transformer.reset_temp_vars() def forward(self, input_tensor, encoder_hidden_states=None, attention_mask=None): - return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) + if math.isclose(self.strength, 1.0): + return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) + elif math.isclose(self.strength, 0.0): + return input_tensor + elif self.strength > 1.0: + return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)*self.strength + else: + return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)*self.strength + input_tensor*(1.0-self.strength) class TemporalTransformer3DModel(nn.Module): diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index 7759013..2faae66 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -14,6 +14,8 @@ class SparseCtrlLoaderAdvanced: "required": { "control_net_name": (folder_paths.get_filename_list("controlnet"), ), "use_motion": ("BOOLEAN", {"default": True}, ), + "motion_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), }, "optional": { "sparse_method": ("SPARSE_METHOD", ), @@ -26,9 +28,9 @@ class SparseCtrlLoaderAdvanced: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" - def load_controlnet(self, control_net_name: str, use_motion: bool, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): + def load_controlnet(self, control_net_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) - sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion) + sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale) controlnet = load_sparsectrl(controlnet_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings) return (controlnet,) From bcc7fa385ef55e91b2121ef7af0979b0f493ae37 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 22 Dec 2023 00:31:06 -0600 Subject: [PATCH 09/15] Add experimental Load Merged SparseCtrl Model node --- control/control.py | 3 +- control/control_sparsectrl.py | 3 +- control/nodes.py | 4 ++- control/nodes_sparsectrl.py | 52 +++++++++++++++++++++++++++++++---- 4 files changed, 54 insertions(+), 8 deletions(-) diff --git a/control/control.py b/control/control.py index b473b3e..32b4912 100644 --- a/control/control.py +++ b/control/control.py @@ -286,7 +286,8 @@ class SparseCtrlAdvanced(ControlNetAdvanced): cond_mask = torch.zeros(cond_shape).to(dtype).to(self.device) cond_mask[local_idxs] = 1.0 # combine cond_hint and cond_mask into (b, c+1, h, w) - self.cond_hint = torch.cat([self.cond_hint, cond_mask], dim=1) + if not self.sparse_settings.merged: + self.cond_hint = torch.cat([self.cond_hint, cond_mask], dim=1) del sub_cond_hint del cond_mask # make cond_hint match x_noisy batch diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index 1abb8cd..5328cac 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -80,11 +80,12 @@ class SparseControlNet(ControlNetCLDM): class SparseSettings: - def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True, motion_strength=1.0, motion_scale=1.0): + def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True, motion_strength=1.0, motion_scale=1.0, merged=False): self.sparse_method = sparse_method self.use_motion = use_motion self.motion_strength = motion_strength self.motion_scale = motion_scale + self.merged = merged @classmethod def default(cls): diff --git a/control/nodes.py b/control/nodes.py index 2ed68e1..5076062 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -9,7 +9,7 @@ from .utils import StrengthInterpolation as SI from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, SoftT2IAdapterWeights, CustomT2IAdapterWeights) from .nodes_latent_keyframe import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode -from .nodes_sparsectrl import SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, VAEEncodePreprocessor +from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, VAEEncodePreprocessor from .logger import logger @@ -217,6 +217,7 @@ NODE_CLASS_MAPPINGS = { # SparseCtrl "ACN_VAEEncodePreprocessor": VAEEncodePreprocessor, "ACN_SparseCtrlLoaderAdvanced": SparseCtrlLoaderAdvanced, + "ACN_SparseCtrlMergedLoaderAdvanced": SparseCtrlMergedLoaderAdvanced, "ACN_SparseCtrlIndexMethodNode": SparseIndexMethodNode, "ACN_SparseCtrlSpreadMethodNode": SparseSpreadMethodNode, } @@ -244,6 +245,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { # SparseCtrl "ACN_VAEEncodePreprocessor": "RGB SparseCtrl 🛂🅐🅒🅝", "ACN_SparseCtrlLoaderAdvanced": "Load SparseCtrl Model 🛂🅐🅒🅝", + "ACN_SparseCtrlMergedLoaderAdvanced": "Load Merged SparseCtrl Model 🛂🅐🅒🅝", "ACN_SparseCtrlIndexMethodNode": "SparseCtrl Index Method 🛂🅐🅒🅝", "ACN_SparseCtrlSpreadMethodNode": "SparseCtrl Spread Method 🛂🅐🅒🅝", } diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index 2faae66..8a724a7 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -3,7 +3,7 @@ from nodes import VAEEncode from .utils import TimestepKeyframeGroup from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod -from .control import load_sparsectrl +from .control import load_sparsectrl, load_controlnet, ControlNetAdvanced, SparseCtrlAdvanced # node for SparseCtrl loading @@ -12,6 +12,35 @@ class SparseCtrlLoaderAdvanced: def INPUT_TYPES(s): return { "required": { + "sparsectrl_name": (folder_paths.get_filename_list("controlnet"), ), + "use_motion": ("BOOLEAN", {"default": True}, ), + "motion_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + }, + "optional": { + "sparse_method": ("SPARSE_METHOD", ), + "tk_optional": ("TIMESTEP_KEYFRAME", ), + } + } + + RETURN_TYPES = ("CONTROL_NET", ) + FUNCTION = "load_controlnet" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" + + def load_controlnet(self, sparsectrl_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): + sparsectrl_path = folder_paths.get_full_path("controlnet", sparsectrl_name) + sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale) + sparsectrl = load_sparsectrl(sparsectrl_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings) + return (sparsectrl,) + + +class SparseCtrlMergedLoaderAdvanced: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sparsectrl_name": (folder_paths.get_filename_list("controlnet"), ), "control_net_name": (folder_paths.get_filename_list("controlnet"), ), "use_motion": ("BOOLEAN", {"default": True}, ), "motion_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), @@ -28,11 +57,24 @@ class SparseCtrlLoaderAdvanced: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" - def load_controlnet(self, control_net_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): + def load_controlnet(self, sparsectrl_name: str, control_net_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): + sparsectrl_path = folder_paths.get_full_path("controlnet", sparsectrl_name) controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) - sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale) - controlnet = load_sparsectrl(controlnet_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings) - return (controlnet,) + sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale, merged=True) + # first, load normal controlnet + controlnet = load_controlnet(controlnet_path, timestep_keyframe=tk_optional) + # confirm that controlnet is ControlNetAdvanced + if controlnet is None or type(controlnet) != ControlNetAdvanced: + raise ValueError(f"controlnet_path must point to a normal ControlNet, but instead: {type(controlnet).__name__}") + # next, load sparsectrl, making sure to load motion portion + sparsectrl = load_sparsectrl(sparsectrl_path, timestep_keyframe=tk_optional, sparse_settings=SparseSettings.default()) + # now, combine state dicts + new_state_dict = controlnet.control_model.state_dict() + for key, value in sparsectrl.control_model.motion_holder.motion_wrapper.state_dict().items(): + new_state_dict[key] = value + # now, reload sparsectrl with real settings + sparsectrl = load_sparsectrl(sparsectrl_path, controlnet_data=new_state_dict, timestep_keyframe=tk_optional, sparse_settings=sparse_settings) + return (sparsectrl,) class SparseIndexMethodNode: From 2eebf83ff6ed630ba0dead46a591105e3eee6985 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 22 Dec 2023 00:47:38 -0600 Subject: [PATCH 10/15] Fix center Spread logic to work as intended instead of just returning one index --- control/control_sparsectrl.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index 5328cac..c3bb217 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -144,7 +144,7 @@ class SparseSpreadMethod(SparseMethod): if full_length-hint_length < 3: return list(np.linspace(0, full_length-1, hint_length+1, endpoint=True, dtype=int))[:-1] # otherwise, get linspace of 2 greater than needed, then cut off first and last - return list(np.linspace(0, full_length-1, hint_length, endpoint=True, dtype=int))[1:-1] + return list(np.linspace(0, full_length-1, hint_length+2, endpoint=True, dtype=int))[1:-1] return ValueError(f"Unrecognized spread: {self.spread}") From 491045b8f3ef3a17e2f3c4ef17b82a4908da2c6e Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 22 Dec 2023 03:00:24 -0600 Subject: [PATCH 11/15] Apply missing scale factor via model's LatentFormat for RGB SparseCtrl --- control/control.py | 9 +++++++++ control/nodes_sparsectrl.py | 2 +- 2 files changed, 10 insertions(+), 1 deletion(-) diff --git a/control/control.py b/control/control.py index 32b4912..d211751 100644 --- a/control/control.py +++ b/control/control.py @@ -227,6 +227,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced): self.add_compatible_weight(ControlWeightType.SPARSECTRL) self.control_model: SparseControlNet = self.control_model # does nothing except help with IDE hints self.sparse_settings = sparse_settings if sparse_settings is not None else SparseSettings.default() + self.latent_format = None def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): # normal ControlNet stuff @@ -272,6 +273,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced): # scale cond_hints to match noisy input if self.control_model.use_simplified_conditioning_embedding: # RGB SparseCtrl; the inputs are latents - use bilinear to avoid blocky artifacts + sub_cond_hint = self.latent_format.process_in(sub_cond_hint) # multiplies by model scale factor sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3], x_noisy.shape[2], "bilinear", "center").to(dtype).to(self.device) else: # other SparseCtrl; inputs are typical images @@ -309,11 +311,18 @@ class SparseCtrlAdvanced(ControlNetAdvanced): def pre_run_advanced(self, model, percent_to_timestep_function): super().pre_run_advanced(model, percent_to_timestep_function) + self.latent_format = model.latent_format # LatentFormat object, used to process_in latent cond hint if self.control_model.motion_holder is not None: self.control_model.motion_holder.motion_wrapper.reset() self.control_model.motion_holder.motion_wrapper.set_strength(self.sparse_settings.motion_strength) self.control_model.motion_holder.motion_wrapper.set_scale_multiplier(self.sparse_settings.motion_scale) + def cleanup_advanced(self): + super().cleanup_advanced() + if self.latent_format is not None: + del self.latent_format + self.latent_format = None + def copy(self): c = SparseCtrlAdvanced(self.control_model, self.timestep_keyframes, self.sparse_settings, self.global_average_pooling, self.device, self.load_device, self.manual_cast_dtype) self.copy_to(c) diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index 8a724a7..0e713fe 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -134,7 +134,7 @@ class VAEEncodePreprocessor: return { "required": { "image": ("IMAGE", ), - "vae": ("VAE", ) + "vae": ("VAE", ), } } From e674b995d0c75d1783d05d9c8080f02848f9d204 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 22 Dec 2023 03:20:19 -0600 Subject: [PATCH 12/15] Change latent upscaling to nearest-exact for now - will come up with a better solution after getting sleep --- control/control.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/control/control.py b/control/control.py index d211751..c5812ed 100644 --- a/control/control.py +++ b/control/control.py @@ -274,7 +274,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced): if self.control_model.use_simplified_conditioning_embedding: # RGB SparseCtrl; the inputs are latents - use bilinear to avoid blocky artifacts sub_cond_hint = self.latent_format.process_in(sub_cond_hint) # multiplies by model scale factor - sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3], x_noisy.shape[2], "bilinear", "center").to(dtype).to(self.device) + sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3], x_noisy.shape[2], "nearest-exact", "center").to(dtype).to(self.device) else: # other SparseCtrl; inputs are typical images sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) From fa0a3dba4c00f82d45c1fab53a8c40cfd59f4733 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 22 Dec 2023 11:12:23 -0600 Subject: [PATCH 13/15] Automatically scale image before RGB Preproc does vae encoding, make sure RGB Preproc image isn't attempted to be used outside of Apply ControlNet nodes --- control/control.py | 5 ++++- control/control_sparsectrl.py | 11 +++++++++++ control/nodes.py | 6 +++--- control/nodes_sparsectrl.py | 22 +++++++++++++++------- 4 files changed, 33 insertions(+), 11 deletions(-) diff --git a/control/control.py b/control/control.py index c5812ed..6500af7 100644 --- a/control/control.py +++ b/control/control.py @@ -9,7 +9,7 @@ import comfy.model_detection import comfy.controlnet as comfy_cn from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, broadcast_image_to -from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod +from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper 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 @@ -228,6 +228,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced): self.control_model: SparseControlNet = self.control_model # does nothing except help with IDE hints self.sparse_settings = sparse_settings if sparse_settings is not None else SparseSettings.default() self.latent_format = None + self.preprocessed = False def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): # normal ControlNet stuff @@ -311,6 +312,8 @@ class SparseCtrlAdvanced(ControlNetAdvanced): def pre_run_advanced(self, model, percent_to_timestep_function): super().pre_run_advanced(model, percent_to_timestep_function) + if type(self.cond_hint_original) == PreprocSparseRGBWrapper: + self.cond_hint_original = self.cond_hint_original.condhint self.latent_format = model.latent_format # LatentFormat object, used to process_in latent cond hint if self.control_model.motion_holder is not None: self.control_model.motion_holder.motion_wrapper.reset() diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index c3bb217..41698b4 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -79,6 +79,17 @@ class SparseControlNet(ControlNetCLDM): return outs +class PreprocSparseRGBWrapper: + def __init__(self, condhint: Tensor): + self.condhint = condhint + + def movedim(self, *args, **kwargs): + return self + + def __getattr__(self, name): + raise AttributeError("Invalid use of RGB SparseCtrl output. The output of RGB SparseCtrl preprocessor is NOT a usual image, but a latent pretending to be an image - you must connect the output directly to an Apply ControlNet node (advanced or otherwise).") + + class SparseSettings: def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True, motion_strength=1.0, motion_scale=1.0, merged=False): self.sparse_method = sparse_method diff --git a/control/nodes.py b/control/nodes.py index 5076062..3a071c9 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -9,7 +9,7 @@ from .utils import StrengthInterpolation as SI from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, SoftT2IAdapterWeights, CustomT2IAdapterWeights) from .nodes_latent_keyframe import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode -from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, VAEEncodePreprocessor +from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, RgbSparseCtrlPreprocessor from .logger import logger @@ -215,7 +215,7 @@ NODE_CLASS_MAPPINGS = { "CustomT2IAdapterWeights": CustomT2IAdapterWeights, "ACN_DefaultUniversalWeights": DefaultWeights, # SparseCtrl - "ACN_VAEEncodePreprocessor": VAEEncodePreprocessor, + "ACN_SparseCtrlRGBPreprocessor": RgbSparseCtrlPreprocessor, "ACN_SparseCtrlLoaderAdvanced": SparseCtrlLoaderAdvanced, "ACN_SparseCtrlMergedLoaderAdvanced": SparseCtrlMergedLoaderAdvanced, "ACN_SparseCtrlIndexMethodNode": SparseIndexMethodNode, @@ -243,7 +243,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝", "ACN_DefaultUniversalWeights": "Force Default Weights 🛂🅐🅒🅝", # SparseCtrl - "ACN_VAEEncodePreprocessor": "RGB SparseCtrl 🛂🅐🅒🅝", + "ACN_SparseCtrlRGBPreprocessor": "RGB SparseCtrl 🛂🅐🅒🅝", "ACN_SparseCtrlLoaderAdvanced": "Load SparseCtrl Model 🛂🅐🅒🅝", "ACN_SparseCtrlMergedLoaderAdvanced": "Load Merged SparseCtrl Model 🛂🅐🅒🅝", "ACN_SparseCtrlIndexMethodNode": "SparseCtrl Index Method 🛂🅐🅒🅝", diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index 0e713fe..dca2e3f 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -1,8 +1,11 @@ +from torch import Tensor + import folder_paths from nodes import VAEEncode +import comfy.utils from .utils import TimestepKeyframeGroup -from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod +from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper from .control import load_sparsectrl, load_controlnet, ControlNetAdvanced, SparseCtrlAdvanced @@ -55,7 +58,7 @@ class SparseCtrlMergedLoaderAdvanced: RETURN_TYPES = ("CONTROL_NET", ) FUNCTION = "load_controlnet" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl/experimental" def load_controlnet(self, sparsectrl_name: str, control_net_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): sparsectrl_path = folder_paths.get_full_path("controlnet", sparsectrl_name) @@ -128,24 +131,29 @@ class SparseSpreadMethodNode: return (SparseSpreadMethod(spread=spread),) -class VAEEncodePreprocessor: +class RgbSparseCtrlPreprocessor: @classmethod def INPUT_TYPES(s): return { "required": { "image": ("IMAGE", ), "vae": ("VAE", ), + "latent_size": ("LATENT", ), } } RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("latent_IMAGE",) + RETURN_NAMES = ("proc_IMAGE",) FUNCTION = "preprocess_images" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl/preprocess" - def preprocess_images(self, vae, image): + def preprocess_images(self, vae, image: Tensor, latent_size: Tensor): + # first, resize image to match latents + image = image.movedim(-1,1) + image = comfy.utils.common_upscale(image, latent_size["samples"].shape[3] * 8, latent_size["samples"].shape[2] * 8, 'nearest-exact', "center") + image = image.movedim(1,-1) + # then, vae encode image = VAEEncode.vae_encode_crop_pixels(image) encoded = vae.encode(image[:,:,:,:3]) - encoded = encoded.movedim(1,-1) - return (encoded,) + return (PreprocSparseRGBWrapper(condhint=encoded),) From db0493a34b0615f8e2b269dabc88b3d09fcf819f Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 22 Dec 2023 11:44:57 -0600 Subject: [PATCH 14/15] Make it clear Sparse RGB preproc output cannot be used for anything but Apply ControlNet nodes --- control/control_sparsectrl.py | 25 +++++++++++++++++++++++-- 1 file changed, 23 insertions(+), 2 deletions(-) diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index 41698b4..d423872 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -80,14 +80,35 @@ class SparseControlNet(ControlNetCLDM): class PreprocSparseRGBWrapper: + error_msg = "Invalid use of RGB SparseCtrl output. The output of RGB SparseCtrl preprocessor is NOT a usual image, but a latent pretending to be an image - you must connect the output directly to an Apply ControlNet node (advanced or otherwise). It cannot be used for anything else that accepts IMAGE input." def __init__(self, condhint: Tensor): self.condhint = condhint def movedim(self, *args, **kwargs): return self - def __getattr__(self, name): - raise AttributeError("Invalid use of RGB SparseCtrl output. The output of RGB SparseCtrl preprocessor is NOT a usual image, but a latent pretending to be an image - you must connect the output directly to an Apply ControlNet node (advanced or otherwise).") + def __getattr__(self, *args, **kwargs): + raise AttributeError(self.error_msg) + + def __setattr__(self, name, value): + if name != "condhint": + raise AttributeError(self.error_msg) + super().__setattr__(name, value) + + def __iter__(self, *args, **kwargs): + raise AttributeError(self.error_msg) + + def __next__(self, *args, **kwargs): + raise AttributeError(self.error_msg) + + def __len__(self, *args, **kwargs): + raise AttributeError(self.error_msg) + + def __getitem__(self, *args, **kwargs): + raise AttributeError(self.error_msg) + + def __setitem__(self, *args, **kwargs): + raise AttributeError(self.error_msg) class SparseSettings: From d984aa23cfbeaca280168a22d566e942ee6e0b08 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 22 Dec 2023 12:08:56 -0600 Subject: [PATCH 15/15] Add some helpful error messages when not loading SparseCtrl model or trying to use preprocessed RGB images with incompatible SparseCtrl --- control/control.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/control/control.py b/control/control.py index 6500af7..cf66378 100644 --- a/control/control.py +++ b/control/control.py @@ -313,6 +313,8 @@ class SparseCtrlAdvanced(ControlNetAdvanced): def pre_run_advanced(self, model, percent_to_timestep_function): super().pre_run_advanced(model, percent_to_timestep_function) if type(self.cond_hint_original) == PreprocSparseRGBWrapper: + if not self.control_model.use_simplified_conditioning_embedding: + raise ValueError("Any model besides RGB SparseCtrl should NOT have its images go through the RGB SparseCtrl preprocessor.") self.cond_hint_original = self.cond_hint_original.condhint self.latent_format = model.latent_format # LatentFormat object, used to process_in latent cond hint if self.control_model.motion_holder is not None: @@ -402,6 +404,8 @@ def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, tim for key in list(controlnet_data.keys()): if "temporal" in key: motion_data[key] = controlnet_data.pop(key) + if len(motion_data) == 0: + raise ValueError(f"No motion-related keys in '{ckpt_path}'; not a valid SparseCtrl model!") motion_wrapper: SparseCtrlMotionWrapper = SparseCtrlMotionWrapper(motion_data).to(comfy.model_management.unet_dtype()) missing, unexpected = motion_wrapper.load_state_dict(motion_data) if len(missing) > 0 or len(unexpected) > 0: