diff --git a/adv_control/control.py b/adv_control/control.py index dfca88f..2ea2586 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -25,7 +25,7 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase): def get_universal_weights(self) -> ControlWeights: raw_weights = [(self.weights.base_multiplier ** float(12 - i)) for i in range(13)] - return ControlWeights.controlnet(raw_weights, self.weights.flip_weights) + return self.weights.copy_with_new_weights(raw_weights) def get_control_advanced(self, x_noisy, t, cond, batched_number): # perform special version of get_control that supports sliding context and masks @@ -99,7 +99,7 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): raw_weights = [(self.weights.base_multiplier ** float(7 - i)) for i in range(8)] raw_weights = [raw_weights[-8], raw_weights[-3], raw_weights[-2], raw_weights[-1]] raw_weights = get_properly_arranged_t2i_weights(raw_weights) - return ControlWeights.t2iadapter(raw_weights, self.weights.flip_weights) + return self.weights.copy_with_new_weights(raw_weights) def get_calc_pow(self, idx: int, layers: int) -> int: # match how T2IAdapterAdvanced deals with universal weights @@ -152,7 +152,7 @@ class ControlLoraAdvanced(ControlLora, AdvancedControlBase): def get_universal_weights(self) -> ControlWeights: raw_weights = [(self.weights.base_multiplier ** float(9 - i)) for i in range(10)] - return ControlWeights.controllora(raw_weights, self.weights.flip_weights) + return self.weights.copy_with_new_weights(raw_weights) def copy(self): c = ControlLoraAdvanced(self.control_weights, self.timestep_keyframes, global_average_pooling=self.global_average_pooling) diff --git a/adv_control/nodes_weight.py b/adv_control/nodes_weight.py index 35d0ffb..7eb58d4 100644 --- a/adv_control/nodes_weight.py +++ b/adv_control/nodes_weight.py @@ -35,6 +35,9 @@ class ScaledSoftMaskedUniversalWeights: #"lock_min": ("BOOLEAN", {"default": False}, ), #"lock_max": ("BOOLEAN", {"default": False}, ), }, + "optional": { + "uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ), + } } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) @@ -43,7 +46,8 @@ class ScaledSoftMaskedUniversalWeights: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" - def load_weights(self, mask: Tensor, min_base_multiplier: float, max_base_multiplier: float, lock_min=False, lock_max=False): + def load_weights(self, mask: Tensor, min_base_multiplier: float, max_base_multiplier: float, lock_min=False, lock_max=False, + uncond_multiplier: float=1.0): # normalize mask mask = mask.clone() x_min = 0.0 if lock_min else mask.min() @@ -52,7 +56,7 @@ class ScaledSoftMaskedUniversalWeights: mask = torch.ones_like(mask) * max_base_multiplier else: mask = linear_conversion(mask, x_min, x_max, min_base_multiplier, max_base_multiplier) - weights = ControlWeights.universal_mask(weight_mask=mask) + weights = ControlWeights.universal_mask(weight_mask=mask, uncond_multiplier=uncond_multiplier) return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) @@ -64,6 +68,9 @@ class ScaledSoftUniversalWeights: "base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 1.0, "step": 0.001}, ), "flip_weights": ("BOOLEAN", {"default": False}), }, + "optional": { + "uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ), + } } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) @@ -72,8 +79,8 @@ class ScaledSoftUniversalWeights: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights" - def load_weights(self, base_multiplier, flip_weights): - weights = ControlWeights.universal(base_multiplier=base_multiplier, flip_weights=flip_weights) + def load_weights(self, base_multiplier, flip_weights, uncond_multiplier: float=1.0): + weights = ControlWeights.universal(base_multiplier=base_multiplier, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier) return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) @@ -97,6 +104,9 @@ class SoftControlNetWeights: "weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "flip_weights": ("BOOLEAN", {"default": False}), }, + "optional": { + "uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ), + } } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) @@ -106,10 +116,11 @@ class SoftControlNetWeights: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, - weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights): + weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights, + uncond_multiplier: float=1.0): weights = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, weight_07, weight_08, weight_09, weight_10, weight_11, weight_12] - weights = ControlWeights.controlnet(weights, flip_weights=flip_weights) + weights = ControlWeights.controlnet(weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier) return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) @@ -132,6 +143,9 @@ class CustomControlNetWeights: "weight_11": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "flip_weights": ("BOOLEAN", {"default": False}), + }, + "optional": { + "uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ), } } @@ -142,10 +156,11 @@ class CustomControlNetWeights: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet" def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, - weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights): + weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights, + uncond_multiplier: float=1.0): weights = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06, weight_07, weight_08, weight_09, weight_10, weight_11, weight_12] - weights = ControlWeights.controlnet(weights, flip_weights=flip_weights) + weights = ControlWeights.controlnet(weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier) return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) @@ -160,6 +175,9 @@ class SoftT2IAdapterWeights: "weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "flip_weights": ("BOOLEAN", {"default": False}), }, + "optional": { + "uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ), + } } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) @@ -168,10 +186,11 @@ class SoftT2IAdapterWeights: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter" - def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights): + def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights, + uncond_multiplier: float=1.0): weights = [weight_00, weight_01, weight_02, weight_03] weights = get_properly_arranged_t2i_weights(weights) - weights = ControlWeights.t2iadapter(weights, flip_weights=flip_weights) + weights = ControlWeights.t2iadapter(weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier) return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) @@ -186,6 +205,9 @@ class CustomT2IAdapterWeights: "weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "flip_weights": ("BOOLEAN", {"default": False}), }, + "optional": { + "uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ), + } } RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",) @@ -194,8 +216,9 @@ class CustomT2IAdapterWeights: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter" - def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights): + def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights, + uncond_multiplier: float=1.0): weights = [weight_00, weight_01, weight_02, weight_03] weights = get_properly_arranged_t2i_weights(weights) - weights = ControlWeights.t2iadapter(weights, flip_weights=flip_weights) + weights = ControlWeights.t2iadapter(weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier) return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights))) diff --git a/adv_control/utils.py b/adv_control/utils.py index 4048b89..4257b79 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -8,6 +8,8 @@ import math import comfy.ops import comfy.utils +import comfy.sample +import comfy.samplers from comfy.controlnet import ControlBase, broadcast_image_to from comfy.model_patcher import ModelPatcher @@ -23,6 +25,92 @@ def load_torch_file_with_dict_factory(controlnet_data: dict[str, Tensor], orig_l return controlnet_data return load_torch_file_with_dict +# wrapping len function so that it will save the thing len is trying to get the length of; +# this will be assumed to be the cond_or_uncond variable; +# automatically restores len to original function after running +def wrapper_len_factory(orig_len: Callable) -> Callable: + def wrapper_len(*args, **kwargs): + cond_or_uncond = args[0] + if type(cond_or_uncond) == list and (0 in cond_or_uncond or 1 in cond_or_uncond): + try: + to_return = IntWithCondOrUncond(orig_len(*args, **kwargs)) + setattr(to_return, "cond_or_uncond", cond_or_uncond) + return to_return + finally: + __builtins__["len"] = orig_len + else: + return orig_len(*args, **kwargs) + return wrapper_len + +# wrapping cond_cat function so that it will wrap around len function to get cond_or_uncond variable value +# from comfy.samplers.calc_conds_batch +def wrapper_cond_cat_factory(orig_cond_cat: Callable): + def wrapper_cond_cat(*args, **kwargs): + __builtins__["len"] = wrapper_len_factory(__builtins__["len"]) + return orig_cond_cat(*args, **kwargs) + return wrapper_cond_cat +orig_cond_cat = comfy.samplers.cond_cat +comfy.samplers.cond_cat = wrapper_cond_cat_factory(orig_cond_cat) + + +def uncond_multiplier_check_cn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callable: + def contains_uncond_multiplier(control: Union[ControlBase, 'AdvancedControlBase']): + if control is None: + return False + if not isinstance(control, AdvancedControlBase): + return contains_uncond_multiplier(control.previous_controlnet) + # check if weights_override has an uncond_multiplier + if control.weights_override is not None and control.weights_override.has_uncond_multiplier: + return True + # check if any timestep_keyframes have an uncond_multiplier on their weights + if control.timestep_keyframes is not None: + for tk in control.timestep_keyframes.keyframes: + if tk.has_control_weights() and tk.control_weights.has_uncond_multiplier: + return True + return False + + # check if positive or negative conds contain Adv. Cns that use multiply_negative on weights + def uncond_multiplier_check_cn_sample(model: ModelPatcher, *args, **kwargs): + positive = args[-3] + negative = args[-2] + has_uncond_multiplier = False + if positive is not None: + for cond in positive: + if "control" in cond[1]: + has_uncond_multiplier = contains_uncond_multiplier(cond[1]["control"]) + if has_uncond_multiplier: + break + if negative is not None and not has_uncond_multiplier: + for cond in negative: + if "control" in cond[1]: + has_uncond_multiplier = contains_uncond_multiplier(cond[1]["control"]) + if has_uncond_multiplier: + break + # if uncond_multiplier found, continue to use wrapped version of function + if has_uncond_multiplier: + return orig_comfy_sample(model, *args, **kwargs) + # otherwise, use original version of function to prevent even the smallest of slowdowns (0.XX%) + try: + wrapped_cond_cat = comfy.samplers.cond_cat + comfy.samplers.cond_cat = orig_cond_cat + return orig_comfy_sample(model, *args, **kwargs) + finally: + comfy.samplers.cond_cat = wrapped_cond_cat + return uncond_multiplier_check_cn_sample +# inject sample functions +comfy.sample.sample = uncond_multiplier_check_cn_sample_factory(comfy.sample.sample) +comfy.sample.sample_custom = uncond_multiplier_check_cn_sample_factory(comfy.sample.sample_custom, is_custom=True) + + +class IntWithCondOrUncond(int): + def __new__(cls, *args, **kwargs): + return super(IntWithCondOrUncond, cls).__new__(cls, *args, **kwargs) + + def __init__(self, *args, **kwargs): + super().__init__() + self.cond_or_uncond = None + + def get_properly_arranged_t2i_weights(initial_weights: list[float]): new_weights = [] @@ -45,7 +133,8 @@ class ControlWeightType: 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): + def __init__(self, weight_type: str, base_multiplier: float=1.0, flip_weights: bool=False, weights: list[float]=None, weight_mask: Tensor=None, + uncond_multiplier=1.0): self.weight_type = weight_type self.base_multiplier = base_multiplier self.flip_weights = flip_weights @@ -53,6 +142,8 @@ class ControlWeights: if self.weights is not None and self.flip_weights: self.weights.reverse() self.weight_mask = weight_mask + self.uncond_multiplier = float(uncond_multiplier) + self.has_uncond_multiplier = not math.isclose(self.uncond_multiplier, 1.0) def get(self, idx: int, default=1.0) -> Union[float, Tensor]: # if weights is not none, return index @@ -63,42 +154,46 @@ class ControlWeights: return self.weights[idx] return 1.0 + def copy_with_new_weights(self, new_weights: list[float]): + return ControlWeights(weight_type=self.weight_type, base_multiplier=self.base_multiplier, flip_weights=self.flip_weights, + weights=new_weights, weight_mask=self.weight_mask, uncond_multiplier=self.uncond_multiplier) + @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) + def universal(cls, base_multiplier: float, flip_weights: bool=False, uncond_multiplier: float=1.0): + return cls(ControlWeightType.UNIVERSAL, base_multiplier=base_multiplier, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier) @classmethod - def universal_mask(cls, weight_mask: Tensor): - return cls(ControlWeightType.UNIVERSAL, weight_mask=weight_mask) + def universal_mask(cls, weight_mask: Tensor, uncond_multiplier: float=1.0): + return cls(ControlWeightType.UNIVERSAL, weight_mask=weight_mask, uncond_multiplier=uncond_multiplier) @classmethod - def t2iadapter(cls, weights: list[float]=None, flip_weights: bool=False): + def t2iadapter(cls, weights: list[float]=None, flip_weights: bool=False, uncond_multiplier: float=1.0): if weights is None: weights = [1.0]*12 - return cls(ControlWeightType.T2IADAPTER, weights=weights,flip_weights=flip_weights) + return cls(ControlWeightType.T2IADAPTER, weights=weights,flip_weights=flip_weights, uncond_multiplier=uncond_multiplier) @classmethod - def controlnet(cls, weights: list[float]=None, flip_weights: bool=False): + def controlnet(cls, weights: list[float]=None, flip_weights: bool=False, uncond_multiplier: float=1.0): if weights is None: weights = [1.0]*13 - return cls(ControlWeightType.CONTROLNET, weights=weights, flip_weights=flip_weights) + return cls(ControlWeightType.CONTROLNET, weights=weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier) @classmethod - def controllora(cls, weights: list[float]=None, flip_weights: bool=False): + def controllora(cls, weights: list[float]=None, flip_weights: bool=False, uncond_multiplier: float=1.0): if weights is None: weights = [1.0]*10 - return cls(ControlWeightType.CONTROLLORA, weights=weights, flip_weights=flip_weights) + return cls(ControlWeightType.CONTROLLORA, weights=weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier) @classmethod - def controllllite(cls, weights: list[float]=None, flip_weights: bool=False): + def controllllite(cls, weights: list[float]=None, flip_weights: bool=False, uncond_multiplier: float=1.0): 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) + return cls(ControlWeightType.CONTROLLLLITE, weights=weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier) class StrengthInterpolation: @@ -572,6 +667,8 @@ class AdvancedControlBase: return True def get_control_inject(self, x_noisy, t, cond, batched_number): + if type(batched_number) != IntWithCondOrUncond: + logger.warn(f"not IntWithCondOrUncond! {type(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 @@ -649,6 +746,15 @@ class AdvancedControlBase: return final_tensor def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int): + # handle weight's uncond_multiplier, if applicable + if self.weights.has_uncond_multiplier: + cond_or_uncond = self.batched_number.cond_or_uncond + actual_length = x.size(0) // batched_number + for idx, cond_type in enumerate(cond_or_uncond): + # if uncond, set to weight's uncond_multiplier + if cond_type == 1: + x[actual_length*idx:actual_length*(idx+1)] *= self.weights.uncond_multiplier + if self.latent_keyframes is not None: x[:] = x[:] * self.calc_latent_keyframe_mults(x=x, batched_number=batched_number) # apply masks, resizing mask to required dims