Refactored Timestep Keyframes to use guarantee_steps instead of guarantee_usage, Timestep Keyframes with same start_percent will no longer overwrite one another, allowing guarantee_steps to be used to schedule on a per-step basis without worrying about unique start_percents, added Timestep Keyframe Interpolation and From List nodes to ease creation of Timestep Keyframe schedules, small cleanup/upkeep

This commit is contained in:
Jedrzej Kosinski
2024-05-14 04:05:18 -05:00
parent 33d9884b76
commit cdfbdc0bbe
6 changed files with 307 additions and 162 deletions
+1 -1
View File
@@ -244,4 +244,4 @@ class LLLiteModule(torch.nn.Module):
cx = self.up(cx)
if control.latent_keyframes is not None:
cx = cx * control.calc_latent_keyframe_mults(x=cx, batched_number=control.batched_number)
return cx * mask * control.strength * control.current_timestep_keyframe.strength
return cx * mask * control.strength * control._current_timestep_keyframe.strength
+4 -4
View File
@@ -243,8 +243,8 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase):
def get_effective_strength(self):
effective_strength = self.strength
if self.current_timestep_keyframe is not None:
effective_strength = effective_strength * self.current_timestep_keyframe.strength
if self._current_timestep_keyframe is not None:
effective_strength = effective_strength * self._current_timestep_keyframe.strength
return effective_strength
def get_effective_attn_mask_or_float(self, x: Tensor, channels: int, is_mid: bool):
@@ -328,8 +328,8 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase):
self.cond_hint = self.latent_format.process_in(self.cond_hint)
self.cond_hint = ref_noise_latents(self.cond_hint, sigma=t, noise=None)
timestep = self.model_sampling_current.timestep(t)
self.should_apply_attn_effective_strength = not (math.isclose(self.strength, 1.0) and math.isclose(self.current_timestep_keyframe.strength, 1.0) and math.isclose(self.ref_opts.attn_strength, 1.0))
self.should_apply_adain_effective_strength = not (math.isclose(self.strength, 1.0) and math.isclose(self.current_timestep_keyframe.strength, 1.0) and math.isclose(self.ref_opts.adain_strength, 1.0))
self.should_apply_attn_effective_strength = not (math.isclose(self.strength, 1.0) and math.isclose(self._current_timestep_keyframe.strength, 1.0) and math.isclose(self.ref_opts.attn_strength, 1.0))
self.should_apply_adain_effective_strength = not (math.isclose(self.strength, 1.0) and math.isclose(self._current_timestep_keyframe.strength, 1.0) and math.isclose(self.ref_opts.adain_strength, 1.0))
# prepare mask - use direct_attn, so the mask dims will match source latents (and be smaller)
self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, direct_attn=True)
self.should_apply_effective_masks = self.latent_keyframes is not None or self.mask_cond_hint is not None or self.tk_mask_cond_hint is not None
+11 -57
View File
@@ -5,11 +5,11 @@ import folder_paths
from comfy.model_patcher import ModelPatcher
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 .utils import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, BIGMAX
from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights,
SoftT2IAdapterWeights, CustomT2IAdapterWeights)
from .nodes_latent_keyframe import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode
from .nodes_keyframes import (LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode,
TimestepKeyframeNode, TimestepKeyframeInterpolationNode, TimestepKeyframeFromStrengthListNode)
from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, RgbSparseCtrlPreprocessor
from .nodes_reference import ReferenceControlNetNode, ReferenceControlFinetune, ReferencePreprocessorNode
from .nodes_loosecontrol import ControlNetLoaderWithLoraAdvanced
@@ -17,56 +17,6 @@ from .nodes_deprecated import LoadImagesFromDirectory
from .logger import logger
class TimestepKeyframeNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
},
"optional": {
"prev_timestep_kf": ("TIMESTEP_KEYFRAME", ),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
"cn_weights": ("CONTROL_NET_WEIGHTS", ),
"latent_keyframe": ("LATENT_KEYFRAME", ),
"null_latent_kf_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
"inherit_missing": ("BOOLEAN", {"default": True}, ),
"guarantee_usage": ("BOOLEAN", {"default": True}, ),
"mask_optional": ("MASK", ),
#"interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT, SI.NONE], {"default": SI.NONE}, ),
}
}
RETURN_NAMES = ("TIMESTEP_KF", )
RETURN_TYPES = ("TIMESTEP_KEYFRAME", )
FUNCTION = "load_keyframe"
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes"
def load_keyframe(self,
start_percent: float,
strength: float=1.0,
cn_weights: ControlWeights=None, control_net_weights: ControlWeights=None, # old name
latent_keyframe: LatentKeyframeGroup=None,
prev_timestep_kf: TimestepKeyframeGroup=None, prev_timestep_keyframe: TimestepKeyframeGroup=None, # old name
null_latent_kf_strength: float=0.0,
inherit_missing=True,
guarantee_usage=True,
mask_optional=None,
interpolation: str=SI.NONE,):
control_net_weights = control_net_weights if control_net_weights else cn_weights
prev_timestep_keyframe = prev_timestep_keyframe if prev_timestep_keyframe else prev_timestep_kf
if not prev_timestep_keyframe:
prev_timestep_keyframe = TimestepKeyframeGroup()
else:
prev_timestep_keyframe = prev_timestep_keyframe.clone()
keyframe = TimestepKeyframe(start_percent=start_percent, strength=strength, interpolation=interpolation, null_latent_kf_strength=null_latent_kf_strength,
control_weights=control_net_weights, latent_keyframes=latent_keyframe, inherit_missing=inherit_missing, guarantee_usage=guarantee_usage,
mask_hint_orig=mask_optional)
prev_timestep_keyframe.add(keyframe)
return (prev_timestep_keyframe,)
class ControlNetLoaderAdvanced:
@classmethod
def INPUT_TYPES(s):
@@ -211,10 +161,12 @@ class AdvancedControlNetApply:
NODE_CLASS_MAPPINGS = {
# Keyframes
"TimestepKeyframe": TimestepKeyframeNode,
"ACN_TimestepKeyframeInterpolation": TimestepKeyframeInterpolationNode,
"ACN_TimestepKeyframeFromStrengthList": TimestepKeyframeFromStrengthListNode,
"LatentKeyframe": LatentKeyframeNode,
"LatentKeyframeGroup": LatentKeyframeGroupNode,
"LatentKeyframeBatchedGroup": LatentKeyframeBatchedGroupNode,
"LatentKeyframeTiming": LatentKeyframeInterpolationNode,
"LatentKeyframeBatchedGroup": LatentKeyframeBatchedGroupNode,
"LatentKeyframeGroup": LatentKeyframeGroupNode,
# Conditioning
"ACN_AdvancedControlNetApply": AdvancedControlNetApply,
# Loaders
@@ -247,10 +199,12 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
# Keyframes
"TimestepKeyframe": "Timestep Keyframe 🛂🅐🅒🅝",
"ACN_TimestepKeyframeInterpolation": "Timestep Keyframe Interpolation 🛂🅐🅒🅝",
"ACN_TimestepKeyframeFromStrengthList": "Timestep Keyframe From List 🛂🅐🅒🅝",
"LatentKeyframe": "Latent Keyframe 🛂🅐🅒🅝",
"LatentKeyframeGroup": "Latent Keyframe Group 🛂🅐🅒🅝",
"LatentKeyframeBatchedGroup": "Latent Keyframe Batched Group 🛂🅐🅒🅝",
"LatentKeyframeTiming": "Latent Keyframe Interpolation 🛂🅐🅒🅝",
"LatentKeyframeBatchedGroup": "Latent Keyframe From List 🛂🅐🅒🅝",
"LatentKeyframeGroup": "Latent Keyframe Group 🛂🅐🅒🅝",
# Conditioning
"ACN_AdvancedControlNetApply": "Apply Advanced ControlNet 🛂🅐🅒🅝",
# Loaders
+2 -34
View File
@@ -4,7 +4,7 @@ import torch
import numpy as np
from PIL import Image, ImageOps
from .utils import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe, BIGMAX
from .utils import BIGMAX
from .logger import logger
@@ -24,7 +24,7 @@ class LoadImagesFromDirectory:
RETURN_TYPES = ("IMAGE", "MASK", "INT")
FUNCTION = "load_images"
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/deprecated"
CATEGORY = ""
def load_images(self, directory: str, image_load_cap: int = 0, start_index: int = 0):
if not os.path.isdir(directory):
@@ -69,35 +69,3 @@ class LoadImagesFromDirectory:
raise FileNotFoundError(f"No images could be loaded from directory '{directory}'.")
return (torch.cat(images, dim=0), torch.stack(masks, dim=0), image_count)
class TimestepKeyframeNodeDeprecated:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
},
"optional": {
"control_net_weights": ("CONTROL_NET_WEIGHTS", ),
"t2i_adapter_weights": ("T2I_ADAPTER_WEIGHTS", ),
"latent_keyframe": ("LATENT_KEYFRAME", ),
"prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ),
}
}
RETURN_TYPES = ("TIMESTEP_KEYFRAME", )
FUNCTION = "load_keyframe"
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes"
def load_keyframe(self,
start_percent: float,
control_net_weights: ControlWeights=None,
latent_keyframe: LatentKeyframeGroup=None,
prev_timestep_keyframe: TimestepKeyframeGroup=None):
if not prev_timestep_keyframe:
prev_timestep_keyframe = TimestepKeyframeGroup()
keyframe = TimestepKeyframe(start_percent, control_net_weights, latent_keyframe)
prev_timestep_keyframe.add(keyframe)
return (prev_timestep_keyframe,)
@@ -2,11 +2,189 @@ from typing import Union
import numpy as np
from collections.abc import Iterable
from .utils import LatentKeyframe, LatentKeyframeGroup, BIGMIN, BIGMAX
from .utils import ControlWeights, TimestepKeyframe, TimestepKeyframeGroup, LatentKeyframe, LatentKeyframeGroup, BIGMIN, BIGMAX
from .utils import StrengthInterpolation as SI
from .logger import logger
class TimestepKeyframeNode:
OUTDATED_DUMMY = -39
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
},
"optional": {
"prev_timestep_kf": ("TIMESTEP_KEYFRAME", ),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
"cn_weights": ("CONTROL_NET_WEIGHTS", ),
"latent_keyframe": ("LATENT_KEYFRAME", ),
"null_latent_kf_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
"inherit_missing": ("BOOLEAN", {"default": True}, ),
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
"mask_optional": ("MASK", ),
}
}
RETURN_NAMES = ("TIMESTEP_KF", )
RETURN_TYPES = ("TIMESTEP_KEYFRAME", )
FUNCTION = "load_keyframe"
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes"
def load_keyframe(self,
start_percent: float,
strength: float=1.0,
cn_weights: ControlWeights=None, control_net_weights: ControlWeights=None, # old name
latent_keyframe: LatentKeyframeGroup=None,
prev_timestep_kf: TimestepKeyframeGroup=None, prev_timestep_keyframe: TimestepKeyframeGroup=None, # old name
null_latent_kf_strength: float=0.0,
inherit_missing=True,
guarantee_steps=OUTDATED_DUMMY,
guarantee_usage=True, # old input
mask_optional=None,):
# if using outdated dummy value, means node on workflow is outdated and should appropriately convert behavior
if guarantee_steps == self.OUTDATED_DUMMY:
guarantee_steps = int(guarantee_usage)
control_net_weights = control_net_weights if control_net_weights else cn_weights
prev_timestep_keyframe = prev_timestep_keyframe if prev_timestep_keyframe else prev_timestep_kf
if not prev_timestep_keyframe:
prev_timestep_keyframe = TimestepKeyframeGroup()
else:
prev_timestep_keyframe = prev_timestep_keyframe.clone()
keyframe = TimestepKeyframe(start_percent=start_percent, strength=strength, null_latent_kf_strength=null_latent_kf_strength,
control_weights=control_net_weights, latent_keyframes=latent_keyframe, inherit_missing=inherit_missing,
guarantee_steps=guarantee_steps, mask_hint_orig=mask_optional)
prev_timestep_keyframe.add(keyframe)
return (prev_timestep_keyframe,)
class TimestepKeyframeInterpolationNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001},),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"strength_start": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001},),
"strength_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001},),
"interpolation": (SI._LIST, ),
"intervals": ("INT", {"default": 50, "min": 2, "max": 100, "step": 1}),
},
"optional": {
"prev_timestep_kf": ("TIMESTEP_KEYFRAME", ),
"cn_weights": ("CONTROL_NET_WEIGHTS", ),
"latent_keyframe": ("LATENT_KEYFRAME", ),
"null_latent_kf_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001},),
"inherit_missing": ("BOOLEAN", {"default": True},),
"mask_optional": ("MASK", ),
"print_keyframes": ("BOOLEAN", {"default": False}),
}
}
RETURN_NAMES = ("TIMESTEP_KF", )
RETURN_TYPES = ("TIMESTEP_KEYFRAME", )
FUNCTION = "load_keyframes"
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes"
def load_keyframe(self,
start_percent: float, end_percent: float,
strength_start: float, strength_end: float, interpolation: str, intervals: int,
cn_weights: ControlWeights=None,
latent_keyframe: LatentKeyframeGroup=None,
prev_timestep_kf: TimestepKeyframeGroup=None,
null_latent_kf_strength: float=0.0,
inherit_missing=True,
guarantee_steps=1,
mask_optional=None, print_keyframes=False):
if not prev_timestep_kf:
prev_timestep_kf = TimestepKeyframeGroup()
else:
prev_timestep_kf = prev_timestep_kf.clone()
percents = SI.get_weights(num_from=start_percent, num_to=end_percent, length=intervals, method=SI.LINEAR)
strengths = SI.get_weights(num_from=strength_start, num_to=strength_end, length=intervals, method=interpolation)
is_first = True
for percent, strength in zip(percents, strengths):
guarantee_steps = 0
if is_first:
guarantee_steps = 1
is_first = False
prev_timestep_kf.add(TimestepKeyframe(start_percent=percent, strength=strength, null_latent_kf_strength=null_latent_kf_strength,
control_weights=cn_weights, latent_keyframes=latent_keyframe, inherit_missing=inherit_missing,
guarantee_steps=guarantee_steps, mask_hint_orig=mask_optional))
if print_keyframes:
logger.info(f"TimestepKeyframe - start_percent:{percent} = {strength}")
return (prev_timestep_kf,)
class TimestepKeyframeFromStrengthListNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"float_strengths": ("FLOAT", {"default": -1, "min": -1, "step": 0.001, "forceInput": True}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001},),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
},
"optional": {
"prev_timestep_kf": ("TIMESTEP_KEYFRAME", ),
"cn_weights": ("CONTROL_NET_WEIGHTS", ),
"latent_keyframe": ("LATENT_KEYFRAME", ),
"null_latent_kf_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001},),
"inherit_missing": ("BOOLEAN", {"default": True},),
"mask_optional": ("MASK", ),
"print_keyframes": ("BOOLEAN", {"default": False}),
}
}
RETURN_NAMES = ("TIMESTEP_KF", )
RETURN_TYPES = ("TIMESTEP_KEYFRAME", )
FUNCTION = "load_keyframes"
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/keyframes"
def load_keyframe(self,
start_percent: float, end_percent: float,
float_strengths: float,
cn_weights: ControlWeights=None,
latent_keyframe: LatentKeyframeGroup=None,
prev_timestep_kf: TimestepKeyframeGroup=None,
null_latent_kf_strength: float=0.0,
inherit_missing=True,
guarantee_steps=1,
mask_optional=None, print_keyframes=False):
if not prev_timestep_kf:
prev_timestep_kf = TimestepKeyframeGroup()
else:
prev_timestep_kf = prev_timestep_kf.clone()
if type(float_strengths) in (float, int):
float_strengths = [float(float_strengths)]
elif isinstance(float_strengths, Iterable):
pass
else:
raise Exception(f"strengths_float must be either an iterable input or a float, but was {type(float_strengths).__repr__}.")
percents = SI.get_weights(num_from=start_percent, num_to=end_percent, length=len(float_strengths), method=SI.LINEAR)
is_first = True
for percent, strength in zip(percents, float_strengths):
guarantee_steps = 0
if is_first:
guarantee_steps = 1
is_first = False
prev_timestep_kf.add(TimestepKeyframe(start_percent=percent, strength=strength, null_latent_kf_strength=null_latent_kf_strength,
control_weights=cn_weights, latent_keyframes=latent_keyframe, inherit_missing=inherit_missing,
guarantee_steps=guarantee_steps, mask_hint_orig=mask_optional))
if print_keyframes:
logger.info(f"TimestepKeyframe - start_percent:{percent} = {strength}")
return (prev_timestep_kf,)
class LatentKeyframeNode:
@classmethod
def INPUT_TYPES(s):
@@ -149,7 +327,7 @@ class LatentKeyframeGroupNode:
if print_keyframes:
for keyframe in curr_latent_keyframe.keyframes:
logger.info(f"keyframe {keyframe.batch_index}:{keyframe.strength}")
logger.info(f"LatentKeyframe {keyframe.batch_index}={keyframe.strength}")
# replace values with prev_latent_keyframes
for latent_keyframe in prev_latent_keyframe.keyframes:
@@ -167,7 +345,7 @@ class LatentKeyframeInterpolationNode:
"batch_index_to_excl": ("INT", {"default": 0, "min": BIGMIN, "max": BIGMAX, "step": 1}),
"strength_from": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
"strength_to": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
"interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT], ),
"interpolation": (SI._LIST, ),
},
"optional": {
"prev_latent_kf": ("LATENT_KEYFRAME", ),
@@ -223,7 +401,7 @@ class LatentKeyframeInterpolationNode:
if print_keyframes:
for keyframe in curr_latent_keyframe.keyframes:
logger.info(f"keyframe {keyframe.batch_index}:{keyframe.strength}")
logger.info(f"LatentKeyframe {keyframe.batch_index}={keyframe.strength}")
# replace values with prev_latent_keyframes
for latent_keyframe in prev_latent_keyframe.keyframes:
@@ -274,7 +452,7 @@ class LatentKeyframeBatchedGroupNode:
if print_keyframes:
for keyframe in curr_latent_keyframe.keyframes:
logger.info(f"keyframe {keyframe.batch_index}:{keyframe.strength}")
logger.info(f"LatentKeyframe {keyframe.batch_index}={keyframe.strength}")
# replace values with prev_latent_keyframes
for latent_keyframe in prev_latent_keyframe.keyframes:
+106 -61
View File
@@ -3,6 +3,7 @@ from typing import Callable, Union
import torch
from torch import Tensor
import torch.nn.functional
import numpy as np
import math
import comfy.ops
@@ -107,6 +108,29 @@ class StrengthInterpolation:
EASE_IN_OUT = "ease-in-out"
NONE = "none"
_LIST = [LINEAR, EASE_IN, EASE_OUT, EASE_IN_OUT]
_LIST_WITH_NONE = [LINEAR, EASE_IN, EASE_OUT, EASE_IN_OUT, NONE]
@classmethod
def get_weights(cls, num_from: float, num_to: float, length: int, method: str, reverse=False):
diff = num_to - num_from
if method == cls.LINEAR:
weights = torch.linspace(num_from, num_to, length)
elif method == cls.EASE_IN:
index = torch.linspace(0, 1, length)
weights = diff * np.power(index, 2) + num_from
elif method == cls.EASE_OUT:
index = torch.linspace(0, 1, length)
weights = diff * (1 - np.power(1 - index, 2)) + num_from
elif method == cls.EASE_IN_OUT:
index = torch.linspace(0, 1, length)
weights = diff * ((1 - np.cos(index * np.pi)) / 2) + num_from
else:
raise ValueError(f"Unrecognized interpolation method '{method}'.")
if reverse:
weights = weights.flip(dims=(0,))
return weights
class LatentKeyframe:
def __init__(self, batch_index: int, strength: float) -> None:
@@ -154,22 +178,20 @@ 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,
guarantee_steps: int = 1,
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.guarantee_steps = guarantee_steps
self.mask_hint_orig = mask_hint_orig
def has_control_weights(self):
@@ -182,9 +204,9 @@ class TimestepKeyframe:
return self.mask_hint_orig is not None
@classmethod
def default(cls) -> 'TimestepKeyframe':
return cls(0.0)
@staticmethod
def default() -> 'TimestepKeyframe':
return TimestepKeyframe(start_percent=0.0, guarantee_steps=0)
# always maintain sorted state (by start_percent of TimestepKeyFrame)
@@ -194,16 +216,9 @@ class TimestepKeyframeGroup:
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)
# add to end of list, then sort
self.keyframes.append(keyframe)
self.keyframes = get_sorted_list_via_attr(self.keyframes, attr="start_percent")
def get_index(self, index: int) -> Union[TimestepKeyframe, None]:
try:
@@ -225,8 +240,9 @@ class TimestepKeyframeGroup:
def clone(self) -> 'TimestepKeyframeGroup':
cloned = TimestepKeyframeGroup()
# already sorted, so don't use add function to make cloning quicker
for tk in self.keyframes:
cloned.add(tk)
cloned.keyframes.append(tk)
return cloned
@classmethod
@@ -362,6 +378,30 @@ def deepcopy_with_sharing(obj, shared_attribute_names, memo=None):
return clone
def get_sorted_list_via_attr(objects: list, attr: str) -> list:
if not objects:
return objects
elif len(objects) <= 1:
return [x for x in objects]
# now that we know we have to sort, do it following these rules:
# a) if objects have same value of attribute, maintain their relative order
# b) perform sorting of the groups of objects with same attributes
unique_attrs = {}
for o in objects:
val_attr = getattr(o, attr)
attr_list: list = unique_attrs.get(val_attr, list())
attr_list.append(o)
if val_attr not in unique_attrs:
unique_attrs[val_attr] = attr_list
# now that we have the unique attr values grouped together in relative order, sort them by key
sorted_attrs = dict(sorted(unique_attrs.items()))
# now flatten out the dict into a list to return
sorted_list = []
for object_list in sorted_attrs.values():
sorted_list.extend(object_list)
return sorted_list
class WeightTypeException(TypeError):
"Raised when weight not compatible with AdvancedControlBase object"
pass
@@ -396,7 +436,7 @@ class AdvancedControlBase:
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.control_merge = self.control_merge_inject
self.pre_run = self.pre_run_inject
self.cleanup = self.cleanup_inject
self.set_previous_controlnet = self.set_previous_controlnet_inject
@@ -429,9 +469,9 @@ class AdvancedControlBase:
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._current_timestep_keyframe = None
self._current_timestep_index = -1
self._current_used_steps = 0
self.weights = None
self.latent_keyframes = None
@@ -440,39 +480,44 @@ class AdvancedControlBase:
self.batched_number = batched_number
# get current step percent
curr_t: float = self.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):
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:
prev_index = self._current_timestep_index
# if met guaranteed steps (or no current keyframe), look for next keyframe in case need to switch
if self._current_timestep_keyframe is None or self._current_used_steps >= self._current_timestep_keyframe.guarantee_steps:
# 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
self._current_used_steps = 0
# 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_steps greater than zero, stop searching for other keyframes
if self._current_timestep_keyframe.guarantee_steps > 0:
break
# if eval_tk is outside of percent range, stop looking further
else:
break
# if eval_tk is outside of percent range, stop looking further
else:
break
# update steps current keyframe is used
self._current_used_steps += 1
# if index changed, apply overrides
if prev_index != self.current_timestep_index:
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:
@@ -519,7 +564,7 @@ class AdvancedControlBase:
self.disarmed = True
def should_run(self):
if math.isclose(self.strength, 0.0) or math.isclose(self.current_timestep_keyframe.strength, 0.0):
if math.isclose(self.strength, 0.0) or math.isclose(self._current_timestep_keyframe.strength, 0.0):
return False
if self.timestep_range is not None:
if self.t > self.timestep_range[0] or self.t < self.timestep_range[1]:
@@ -530,7 +575,7 @@ class AdvancedControlBase:
# 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:
if self.strength == 0.0 or self._current_timestep_keyframe.strength == 0.0:
return self.default_control_actions(x_noisy, t, cond, batched_number)
# otherwise, perform normal function
return self.get_control_advanced(x_noisy, t, cond, batched_number)
@@ -596,7 +641,7 @@ class AdvancedControlBase:
for batch_index in indeces_to_null:
# apply null for each batched cond/uncond
for b in range(batched_number):
final_mults[(latent_count*b)+batch_index] = self.current_timestep_keyframe.null_latent_kf_strength
final_mults[(latent_count*b)+batch_index] = self._current_timestep_keyframe.null_latent_kf_strength
# convert final_mults into tensor and match expected dimension count
final_tensor = torch.tensor(final_mults, dtype=x.dtype, device=x.device)
while len(final_tensor.shape) < len(x.shape):
@@ -614,8 +659,8 @@ class AdvancedControlBase:
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
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': []}
@@ -674,7 +719,7 @@ class AdvancedControlBase:
self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype, direct_attn=direct_attn)
def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None, direct_attn=False):
return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype, direct_attn=direct_attn)
return self._prepare_mask("tk_mask_cond_hint", self._current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype, direct_attn=direct_attn)
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)
@@ -721,9 +766,9 @@ class AdvancedControlBase:
self.weights = None
self.latent_keyframes = None
# timestep stuff
self.current_timestep_keyframe = None
self.next_timestep_keyframe = None
self.current_timestep_index = -1
self._current_timestep_keyframe = None
self._current_timestep_index = -1
self._current_used_steps = 0
# clear mask hints
if self.mask_cond_hint is not None:
del self.mask_cond_hint