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:
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user