diff --git a/control/control.py b/control/control.py index 180c9e9..7607112 100644 --- a/control/control.py +++ b/control/control.py @@ -35,11 +35,7 @@ class ControlWeights: self.weights.reverse() self.weight_mask = weight_mask - def get(self, idx: int, shape: Tensor) -> Union[float, Tensor]: - # if weight_mask present, return normalized mask - if self.base_multiplier != 1.0 and self.weight_mask is not None: - # TODO: fill out - pass + def get(self, idx: int) -> Union[float, Tensor]: # if weights is not none, return index if self.weights is not None: return self.weights[idx] @@ -53,6 +49,10 @@ class ControlWeights: def universal(cls, base_multiplier: float, flip_weights: bool=False): return cls(ControlWeightType.UNIVERSAL, base_multiplier=base_multiplier, flip_weights=flip_weights) + @classmethod + def universal_mask(cls, weight_mask: Tensor): + return cls(ControlWeightType.UNIVERSAL, weight_mask=weight_mask) + @classmethod def t2iadapter(cls, weights: list[float]=None, flip_weights: bool=False): if weights is None: @@ -268,7 +268,6 @@ class AdvancedControlBase: self.batched_number = batched_number # get current step percent curr_t: float = t[0] - print(f"$$$$ {curr_t}") prev_index = self.current_timestep_index # if has next index, loop through and see if need to switch if self.timestep_keyframes.has_index(self.current_timestep_index+1): @@ -278,7 +277,8 @@ class AdvancedControlBase: if eval_tk.start_t >= curr_t: self.current_timestep_index = i self.current_timestep_keyframe = eval_tk - # keep track of weights and control weights, accounting for inherit_missing + # keep track of control weights, latent keyframes, and masks, + # accounting for inherit_missing if self.current_timestep_keyframe.has_control_weights(): self.weights = self.current_timestep_keyframe.control_weights elif not self.current_timestep_keyframe.inherit_missing: @@ -314,6 +314,9 @@ class AdvancedControlBase: if self.weights is None or self.weights.weight_type == ControlWeightType.DEFAULT: self.weights = self.weights_default elif self.weights.weight_type == ControlWeightType.UNIVERSAL: + # if universal and weight_mask present, no need to convert + if self.weights.weight_mask is not None: + return self.weights = self.get_universal_weights() def get_universal_weights(self) -> ControlWeights: @@ -352,6 +355,14 @@ class AdvancedControlBase: def get_control_advanced(self, x_noisy, t, cond, batched_number): pass + def calc_weight(self, idx: int, x: Tensor, layers: int) -> Union[float, Tensor]: + if self.weights.weight_mask is not None: + # prepare weight mask + self.prepare_weight_mask_cond_hint(x, self.batched_number) + # adjust mask for current layer and return + return torch.pow(self.weight_mask_cond_hint, (layers-1)-idx) + return self.weights.get(idx=idx) + def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int): # apply strengths, and get batch indeces to null out # AKA latents that should not be influenced by ControlNet @@ -381,6 +392,10 @@ class AdvancedControlBase: continue indeces_to_null.remove(real_index) + # if real_index is outside the bounds of latents, don't apply + if real_index >= latent_count or real_index < 0: + continue + # apply strength for each batched cond/uncond for b in range(batched_number): x[(latent_count*b)+real_index] = x[(latent_count*b)+real_index] * keyframe.strength @@ -411,7 +426,7 @@ class AdvancedControlBase: if x is not None: self.apply_advanced_strengths_and_masks(x, self.batched_number) - x *= self.strength * self.weights.get(i, x.shape) + x *= self.strength * self.calc_weight(i, x, len(control_input)) if x.dtype != output_dtype: x = x.to(output_dtype) out[key].insert(0, x) @@ -431,7 +446,7 @@ class AdvancedControlBase: if self.global_average_pooling: x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3]) - x *= self.strength * self.weights.get(i, x.shape) + x *= self.strength * self.calc_weight(i, x, len(control_output)) if x.dtype != output_dtype: x = x.to(output_dtype) @@ -451,36 +466,17 @@ class AdvancedControlBase: return out def prepare_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): - # make mask appropriate dimensions, if present - if self.mask_cond_hint_original is not None: - if self.sub_idxs is not None or self.mask_cond_hint is None or x_noisy.shape[2] * 8 != self.mask_cond_hint.shape[1] or x_noisy.shape[3] * 8 != self.mask_cond_hint.shape[2]: - if self.mask_cond_hint is not None: - del self.mask_cond_hint - self.mask_cond_hint = None - # TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM - # resize mask and match batch count - self.mask_cond_hint = prepare_mask_batch(self.mask_cond_hint_original, x_noisy.shape, multiplier=8) - actual_latent_length = x_noisy.shape[0] // batched_number - self.mask_cond_hint = comfy.utils.repeat_to_batch_size(self.mask_cond_hint, actual_latent_length if self.sub_idxs is None else self.full_latent_length) - if self.sub_idxs is not None: - self.mask_cond_hint = self.mask_cond_hint[self.sub_idxs] - # make cond_hint_mask length match x_noise - if x_noisy.shape[0] != self.mask_cond_hint.shape[0]: - self.mask_cond_hint = broadcast_image_to(self.mask_cond_hint, x_noisy.shape[0], batched_number) - # default dtype to be same as x_noisy - if dtype is None: - dtype = x_noisy.dtype - self.mask_cond_hint = self.mask_cond_hint.to(dtype=dtype).to(self.device) - # prepare other masks + self._prepare_mask("mask_cond_hint", self.mask_cond_hint_original, x_noisy, t, cond, batched_number, dtype) self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype) def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype) - def prepare_weight_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): - return self._prepare_mask("weight_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype) + def prepare_weight_mask_cond_hint(self, x_noisy: Tensor, batched_number, dtype=None): + return self._prepare_mask("weight_mask_cond_hint", self.weights.weight_mask, x_noisy, t=None, cond=None, batched_number=batched_number, dtype=dtype, direct_attn=True) - def _prepare_mask(self, attr_name, orig_mask: Tensor, x_noisy: Tensor, t, cond, batched_number, dtype=None): + def _prepare_mask(self, attr_name, orig_mask: Tensor, x_noisy: Tensor, t, cond, batched_number, dtype=None, direct_attn=False): + # make mask appropriate dimensions, if present if orig_mask is not None: out_mask = getattr(self, attr_name) if self.sub_idxs is not None or out_mask is None or x_noisy.shape[2] * 8 != out_mask.shape[1] or x_noisy.shape[3] * 8 != out_mask.shape[2]: @@ -488,7 +484,8 @@ class AdvancedControlBase: del out_mask # TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM # resize mask and match batch count - out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=8) + multiplier = 1 if direct_attn else 8 + out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=multiplier) actual_latent_length = x_noisy.shape[0] // batched_number out_mask = comfy.utils.repeat_to_batch_size(out_mask, actual_latent_length if self.sub_idxs is None else self.full_latent_length) if self.sub_idxs is not None: @@ -734,3 +731,13 @@ def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim if match_dim1: mask = torch.cat([mask] * shape[1], dim=1) return mask + + +# applies min-max normalization, from: +# https://stackoverflow.com/questions/68791508/min-max-normalization-of-a-tensor-in-pytorch +def normalize_min_max(x: Tensor, new_min = 0.0, new_max = 1.0): + x_min, x_max = x.min(), x.max() + return (((x - x_min)/(x_max - x_min)) * (new_max - new_min)) + new_min + +def linear_conversion(x, x_min=0.0, x_max=1.0, new_min=0.0, new_max=1.0): + return (((x - x_min)/(x_max - x_min)) * (new_max - new_min)) + new_min diff --git a/control/nodes.py b/control/nodes.py index 3d03faf..f4c194c 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -1,11 +1,12 @@ import numpy as np +from torch import Tensor import folder_paths from .control import load_controlnet, convert_to_advanced, ControlWeights, ControlWeightType,\ LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup, is_advanced_controlnet from .control import StrengthInterpolation as SI -from .weight_nodes import DefaultWeights, ScaledSoftControlLoraWeights, ScaledSoftControlNetWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, \ +from .weight_nodes import DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, \ SoftT2IAdapterWeights, CustomT2IAdapterWeights from .latent_keyframe_nodes import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode from .deprecated_nodes import LoadImagesFromDirectory @@ -124,7 +125,7 @@ class AdvancedControlNetApply: "mask_optional": ("MASK", ), "timestep_kf": ("TIMESTEP_KEYFRAME", ), "latent_kf_override": ("LATENT_KEYFRAME", ), - "cn_weights_override": ("CONTROL_NET_WEIGHTS", ), + "weights_override": ("CONTROL_NET_WEIGHTS", ), } } @@ -135,7 +136,7 @@ class AdvancedControlNetApply: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝" def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, - mask_optional=None, + mask_optional: Tensor=None, timestep_kf: TimestepKeyframeGroup=None, latent_kf_override: LatentKeyframeGroup=None, weights_override: ControlWeights=None): if strength == 0: @@ -166,6 +167,7 @@ class AdvancedControlNetApply: c_net.weights_override = weights_override # set cond hint mask if mask_optional is not None: + mask_optional = mask_optional.clone() # if not in the form of a batch, make it so if len(mask_optional.shape) < 3: mask_optional = mask_optional.unsqueeze(0) @@ -195,13 +197,12 @@ NODE_CLASS_MAPPINGS = { "ControlNetLoaderAdvanced": ControlNetLoaderAdvanced, "DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced, # Weights - "ScaledSoftUniversalWeights": ScaledSoftUniversalWeights, - "ScaledSoftControlNetWeights": ScaledSoftControlNetWeights, + "ScaledSoftControlNetWeights": ScaledSoftUniversalWeights, + "ScaledSoftMaskedUniversalWeights": ScaledSoftMaskedUniversalWeights, "SoftControlNetWeights": SoftControlNetWeights, "CustomControlNetWeights": CustomControlNetWeights, "SoftT2IAdapterWeights": SoftT2IAdapterWeights, "CustomT2IAdapterWeights": CustomT2IAdapterWeights, - "ScaledSoftControlLoraWeights": ScaledSoftControlLoraWeights, "ACN_DefaultUniversalWeights": DefaultWeights, # Image "LoadImagesFromDirectory": LoadImagesFromDirectory @@ -217,16 +218,15 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Conditioning "ACN_AdvancedControlNetApply": "Apply Advanced ControlNet 🛂🅐🅒🅝", # Loaders - "ControlNetLoaderAdvanced": "Load ControlNet Model (Advanced) 🛂🅐🅒🅝", - "DiffControlNetLoaderAdvanced": "Load ControlNet Model (diff Advanced) 🛂🅐🅒🅝", + "ControlNetLoaderAdvanced": "Load Advanced ControlNet Model 🛂🅐🅒🅝", + "DiffControlNetLoaderAdvanced": "Load Advanced ControlNet Model (diff) 🛂🅐🅒🅝", # Weights - "ScaledSoftUniversalWeights": "Scaled Soft Weights 🛂🅐🅒🅝", - "ScaledSoftControlNetWeights": "ControlNet Scaled Soft Weights 🛂🅐🅒🅝", + "ScaledSoftControlNetWeights": "Scaled Soft Weights 🛂🅐🅒🅝", + "ScaledSoftMaskedUniversalWeights": "Scaled Soft Masked Weights 🛂🅐🅒🅝", "SoftControlNetWeights": "ControlNet Soft Weights 🛂🅐🅒🅝", "CustomControlNetWeights": "ControlNet Custom Weights 🛂🅐🅒🅝", "SoftT2IAdapterWeights": "T2IAdapter Soft Weights 🛂🅐🅒🅝", "CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝", - "ScaledSoftControlLoraWeights": "ControlLora Scaled Soft Weights 🛂🅐🅒🅝", "ACN_DefaultUniversalWeights": "Force Default Weights 🛂🅐🅒🅝", # Image "LoadImagesFromDirectory": "Load Images [DEPRECATED] 🛂🅐🅒🅝" diff --git a/control/weight_nodes.py b/control/weight_nodes.py index 8d95593..1245397 100644 --- a/control/weight_nodes.py +++ b/control/weight_nodes.py @@ -1,7 +1,11 @@ -from .control import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights +from torch import Tensor +from .control import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights, linear_conversion from .logger import logger +WEIGHTS_RETURN_NAMES = ("CONTROL_NET_WEIGHTS", "TK_SHORTCUT") + + class DefaultWeights: @classmethod def INPUT_TYPES(s): @@ -9,6 +13,7 @@ class DefaultWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" @@ -18,17 +23,47 @@ class DefaultWeights: return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) +class ScaledSoftMaskedUniversalWeights: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "mask": ("MASK", ), + "min_base_multiplier": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), + "max_base_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}, ), + #"lock_min": ("BOOLEAN", {"default": False}, ), + #"lock_max": ("BOOLEAN", {"default": False}, ), + }, + } + + RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES + FUNCTION = "load_weights" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" + + def load_weights(self, mask: Tensor, min_base_multiplier: float, max_base_multiplier: float, lock_min=False, lock_max=False): + # normalize mask + mask = mask.clone() + x_min = 0.0 if lock_min else mask.min() + x_max = 1.0 if lock_max else mask.max() + mask = linear_conversion(mask, x_min, x_max, min_base_multiplier, max_base_multiplier) + weights = ControlWeights.universal_mask(weight_mask=mask) + return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) + + class ScaledSoftUniversalWeights: @classmethod def INPUT_TYPES(s): return { "required": { - "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 1.0, "step": 0.001}, ), "flip_weights": ("BOOLEAN", {"default": False}), }, } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" @@ -38,48 +73,6 @@ class ScaledSoftUniversalWeights: return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) -class ScaledSoftControlLoraWeights: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "flip_weights": ("BOOLEAN", {"default": False}), - }, - } - - RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) - FUNCTION = "load_weights" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlLoRA" - - def load_weights(self, base_multiplier, flip_weights): - weights = [(base_multiplier ** float(9 - i)) for i in range(10)] - weights = ControlWeights.controllora(weights, flip_weights=flip_weights) - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) - - -class ScaledSoftControlNetWeights: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ), - "flip_weights": ("BOOLEAN", {"default": False}), - }, - } - - RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) - FUNCTION = "load_weights" - - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" - - def load_weights(self, base_multiplier, flip_weights): - weights = [(base_multiplier ** float(12 - i)) for i in range(13)] - weights = ControlWeights.controlnet(weights, flip_weights=flip_weights) - return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) - - class SoftControlNetWeights: @classmethod def INPUT_TYPES(s): @@ -103,6 +96,7 @@ class SoftControlNetWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" @@ -138,6 +132,7 @@ class CustomControlNetWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" @@ -164,6 +159,7 @@ class SoftT2IAdapterWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter" @@ -189,6 +185,7 @@ class CustomT2IAdapterWeights: } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) + RETURN_NAMES = WEIGHTS_RETURN_NAMES FUNCTION = "load_weights" CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter"