Scaled Soft Mask Weights added, fixed latent keyframe error when batch_index out of range, removed duplicate weight nodes, renamed TK outputs on weight nodes

This commit is contained in:
Jedrzej Kosinski
2023-11-29 11:39:23 -06:00
parent 1e2d9a5491
commit 583d2c59b7
3 changed files with 93 additions and 89 deletions
+41 -34
View File
@@ -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
+11 -11
View File
@@ -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] 🛂🅐🅒🅝"
+41 -44
View File
@@ -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"