Add support for uncond_multiplier in all Control Weights (identical to auto1111's 'controlnet is more important when set to 0.0)

This commit is contained in:
Jedrzej Kosinski
2024-05-16 04:15:36 -05:00
parent 2ed1fff6b9
commit db0bf14dac
3 changed files with 157 additions and 28 deletions
+3 -3
View File
@@ -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)
+35 -12
View File
@@ -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)))
+119 -13
View File
@@ -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