Added ContextRef Keyframe From List node, added NaiveReuse Keyframe Interp. and From List nodes, added inherit_missing to NaiveReuse Keyframes

This commit is contained in:
Jedrzej Kosinski
2024-08-06 21:57:25 -05:00
parent 0ecd0ac82c
commit fd7f0ffa24
4 changed files with 204 additions and 16 deletions
+26 -5
View File
@@ -106,9 +106,6 @@ class ContextRefKeyframe:
def clone(self):
c = ContextRefKeyframe(mult=self.mult, mult_multival=self.orig_mult_multival, tune_replace=self.orig_tune_replace, mode_replace=self.orig_mode_replace,
start_percent=self.start_percent, guarantee_steps=self.guarantee_steps, inherit_missing=self.inherit_missing)
c.mult_multival = self.mult_multival
c.tune_replace = self.tune_replace
c.mode_replace = self.mode_replace
return c
@@ -253,13 +250,15 @@ class ContextRef(ContextExtra):
################################
# NaiveReuse
class NaiveReuseKeyframe:
def __init__(self, mult=1.0, mult_multival: Union[float, Tensor]=1.0, start_percent=0.0, guarantee_steps=1):
def __init__(self, mult=1.0, mult_multival: Union[float, Tensor]=None, start_percent=0.0, guarantee_steps=1, inherit_missing=True):
self.mult = mult
self.orig_mult_multival = mult_multival
self.mult_multival = mult_multival
# scheduling
self.start_percent = float(start_percent)
self.start_t = 999999999.9
self.guarantee_steps = guarantee_steps
self.inherit_missing = inherit_missing
def clone(self):
c = NaiveReuseKeyframe(mult=self.mult, mult_multival=self.mult_multival,
@@ -267,6 +266,7 @@ class NaiveReuseKeyframe:
c.start_t = self.start_t
return c
class NaiveReuseKeyframeGroup:
def __init__(self):
self.keyframes: list[NaiveReuseKeyframe] = []
@@ -286,13 +286,32 @@ class NaiveReuseKeyframeGroup:
self.keyframes.append(keyframe)
self.keyframes = get_sorted_list_via_attr(self.keyframes, "start_percent")
self._set_first_as_current()
self._prepare_all_keyframe_vals()
def _set_first_as_current(self):
if len(self.keyframes) > 0:
self._current_keyframe = self.keyframes[0]
else:
self._current_keyframe = None
def _prepare_all_keyframe_vals(self):
if self.is_empty():
return
multival = None
for kf in self.keyframes:
# if shouldn't inherit, clear cache
if not kf.inherit_missing:
multival = None
# assign cached values, if origs were None
# Mult #################
if kf.orig_mult_multival is None:
kf.mult_multival = multival
else:
kf.mult_multival = kf.orig_mult_multival
# save new caches, in case next keyframe inherits missing
if kf.mult_multival is not None:
multival = kf.mult_multival
def has_index(self, index: int) -> int:
return index >=0 and index < len(self.keyframes)
@@ -304,6 +323,7 @@ class NaiveReuseKeyframeGroup:
for keyframe in self.keyframes:
cloned.keyframes.append(keyframe)
cloned._set_first_as_current()
cloned._prepare_all_keyframe_vals()
return cloned
def initialize_timesteps(self, model: BaseModel):
@@ -353,6 +373,7 @@ class NaiveReuseKeyframeGroup:
return self._current_keyframe.mult_multival
return None
class NaiveReuse(ContextExtra):
def __init__(self, start_percent: float, end_percent: float, weighted_mean: float, multival_opt: Union[float, Tensor]=None, naivereuse_kf: NaiveReuseKeyframeGroup=None):
super().__init__(start_percent=start_percent, end_percent=end_percent)
+10 -3
View File
@@ -31,8 +31,9 @@ from .nodes_context import (LegacyLoopedUniformContextOptionsNode, LoopedUniform
VisualizeContextOptionsK, VisualizeContextOptionsKAdv, VisualizeContextOptionsSCustom)
from .nodes_context_extras import (SetContextExtrasOnContextOptions, ContextExtras_NaiveReuse, ContextExtras_ContextRef,
ContextRef_ModeFirst, ContextRef_ModeSliding, ContextRef_ModeIndexes,
ContextRef_TuneAttn, ContextRef_TuneAttnAdain, ContextRef_KeyframeMultivalNode, ContextRef_KeyframeInterpolationNode,
NaiveReuse_KeyframeMultivalNode)
ContextRef_TuneAttn, ContextRef_TuneAttnAdain,
ContextRef_KeyframeMultivalNode, ContextRef_KeyframeInterpolationNode, ContextRef_KeyframeFromListNode,
NaiveReuse_KeyframeMultivalNode, NaiveReuse_KeyframeInterpolationNode, NaiveReuse_KeyframeFromListNode)
from .nodes_ad_settings import (AnimateDiffSettingsNode, ManualAdjustPENode, SweetspotStretchPENode, FullStretchPENode,
WeightAdjustAllAddNode, WeightAdjustAllMultNode, WeightAdjustIndivAddNode, WeightAdjustIndivMultNode,
WeightAdjustIndivAttnAddNode, WeightAdjustIndivAttnMultNode)
@@ -84,8 +85,11 @@ NODE_CLASS_MAPPINGS = {
"ADE_ContextExtras_ContextRef_TuneAttnAdain": ContextRef_TuneAttnAdain,
"ADE_ContextExtras_ContextRef_Keyframe": ContextRef_KeyframeMultivalNode,
"ADE_ContextExtras_ContextRef_KeyframeInterpolation": ContextRef_KeyframeInterpolationNode,
"ADE_ContextExtras_ContextRef_KeyframeFromList": ContextRef_KeyframeFromListNode,
"ADE_ContextExtras_NaiveReuse": ContextExtras_NaiveReuse,
"ADE_ContextExtras_NaiveReuse_Keyframe": NaiveReuse_KeyframeMultivalNode,
"ADE_ContextExtras_NaiveReuse_KeyframeInterpolation": NaiveReuse_KeyframeInterpolationNode,
"ADE_ContextExtras_NaiveReuse_KeyframeFromList": NaiveReuse_KeyframeFromListNode,
#------------------------------------------------------------------------------
###############################################################################
# Iteration Opts
@@ -229,9 +233,12 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ADE_ContextExtras_ContextRef_TuneAttn": "ContextRef Tune◆Attn 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_TuneAttnAdain": "ContextRef Tune◆Attn+Adain 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_Keyframe": "ContextRef Keyframe 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_KeyframeInterpolation": "ContextRef Keyframe Interp. 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_KeyframeInterpolation": "ContextRef Keyframes Interp. 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_KeyframeFromList": "ContextRef Keyframes From List 🎭🅐🅓",
"ADE_ContextExtras_NaiveReuse": "Context Extras◆NaiveReuse 🎭🅐🅓",
"ADE_ContextExtras_NaiveReuse_Keyframe": "NaiveReuse Keyframe 🎭🅐🅓",
"ADE_ContextExtras_NaiveReuse_KeyframeInterpolation": "NaiveReuse Keyframes Interp. 🎭🅐🅓",
"ADE_ContextExtras_NaiveReuse_KeyframeFromList": "NaiveReuse Keyframes From List 🎭🅐🅓",
#------------------------------------------------------------------------------
###############################################################################
# Iteration Opts
+163 -5
View File
@@ -1,5 +1,6 @@
from torch import Tensor
from typing import Union
from collections.abc import Iterable
from .context import (ContextOptionsGroup)
from .context_extras import (ContextExtrasGroup,
@@ -81,6 +82,7 @@ class NaiveReuse_KeyframeMultivalNode:
"mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
"inherit_missing": ("BOOLEAN", {"default": True}, ),
"autosize": ("ADEAUTOSIZE", {"padding": 0}),
}
}
@@ -90,14 +92,117 @@ class NaiveReuse_KeyframeMultivalNode:
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/naivereuse"
FUNCTION = "create_keyframe"
def create_keyframe(self, prev_kf=None, mult=1.0, mult_multival=1.0,
start_percent=0.0, guarantee_steps=1):
def create_keyframe(self, prev_kf=None, mult=1.0, mult_multival=None,
start_percent=0.0, guarantee_steps=1, inherit_missing=True):
if prev_kf is None:
prev_kf = NaiveReuseKeyframeGroup()
prev_kf = prev_kf.clone()
kf = NaiveReuseKeyframe(mult=mult, mult_multival=mult_multival, start_percent=start_percent, guarantee_steps=guarantee_steps)
kf = NaiveReuseKeyframe(mult=mult, mult_multival=mult_multival,
start_percent=start_percent, guarantee_steps=guarantee_steps, inherit_missing=inherit_missing)
prev_kf.add(kf)
return (prev_kf,)
class NaiveReuse_KeyframeInterpolationNode:
@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}),
"mult_start": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"mult_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"interpolation": (InterpolationMethod._LIST, ),
"intervals": ("INT", {"default": 50, "min": 2, "max": 100, "step": 1}),
"inherit_missing": ("BOOLEAN", {"default": True}),
"print_keyframes": ("BOOLEAN", {"default": False}),
},
"optional": {
"prev_kf": ("NAIVEREUSE_KEYFRAME",),
"mult_multival": ("MULTIVAL",),
"autosize": ("ADEAUTOSIZE", {"padding": 50}),
}
}
RETURN_TYPES = ("NAIVEREUSE_KEYFRAME",)
RETURN_NAMES = ("NAIVEREUSE_KF",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/naivereuse"
FUNCTION = "create_keyframe"
def create_keyframe(self,
start_percent: float, end_percent: float,
mult_start: float, mult_end: float, interpolation: str, intervals: int,
inherit_missing=True, prev_kf: NaiveReuseKeyframeGroup=None,
mult_multival=None, print_keyframes=False):
if prev_kf is None:
prev_kf = NaiveReuseKeyframeGroup()
prev_kf = prev_kf.clone()
prev_kf = prev_kf.clone()
percents = InterpolationMethod.get_weights(num_from=start_percent, num_to=end_percent, length=intervals, method=InterpolationMethod.LINEAR)
mults = InterpolationMethod.get_weights(num_from=mult_start, num_to=mult_end, length=intervals, method=interpolation)
is_first = True
for percent, mult in zip(percents, mults):
guarantee_steps = 0
if is_first:
guarantee_steps = 1
is_first = False
prev_kf.add(NaiveReuseKeyframe(mult=mult, mult_multival=mult_multival,
start_percent=percent, guarantee_steps=guarantee_steps, inherit_missing=inherit_missing))
if print_keyframes:
logger.info(f"NaiveReuseKeyframe - start_percent:{percent} = {mult}")
return (prev_kf,)
class NaiveReuse_KeyframeFromListNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mults_float": ("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}),
"inherit_missing": ("BOOLEAN", {"default": True}),
"print_keyframes": ("BOOLEAN", {"default": False}),
},
"optional": {
"prev_kf": ("NAIVEREUSE_KEYFRAME",),
"mult_multival": ("MULTIVAL",),
"autosize": ("ADEAUTOSIZE", {"padding": 0}),
}
}
RETURN_TYPES = ("NAIVEREUSE_KEYFRAME",)
RETURN_NAMES = ("NAIVEREUSE_KF",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/naivereuse"
FUNCTION = "create_keyframe"
def create_keyframe(self, mults_float: Union[float, list[float]],
start_percent: float, end_percent: float,
inherit_missing=True, prev_kf: NaiveReuseKeyframeGroup=None,
mult_multival=None, print_keyframes=False):
if prev_kf is None:
prev_kf = NaiveReuseKeyframeGroup()
prev_kf = prev_kf.clone()
if type(mults_float) in (float, int):
mults_float = [float(mults_float)]
elif isinstance(mults_float, Iterable):
pass
else:
raise Exception(f"strengths_float must be either an interable input or a float, but was {type(mults_float).__repr__}.")
percents = InterpolationMethod.get_weights(num_from=start_percent, num_to=end_percent, length=len(mults_float), method=InterpolationMethod.LINEAR)
is_first = True
for percent, mult in zip(percents, mults_float):
guarantee_steps = 0
if is_first:
guarantee_steps = 1
is_first = False
prev_kf.add(NaiveReuseKeyframe(mult=mult, mult_multival=mult_multival,
start_percent=percent, guarantee_steps=guarantee_steps, inherit_missing=inherit_missing))
if print_keyframes:
logger.info(f"NaiveReuseKeyframe - start_percent:{percent} = {mult}")
return (prev_kf,)
#----------------------------------------
#########################################
@@ -170,7 +275,7 @@ class ContextRef_KeyframeMultivalNode:
FUNCTION = "create_keyframe"
def create_keyframe(self, prev_kf: ContextRefKeyframeGroup=None,
mult=1.0, mult_multival=1.0, mode_replace=None, tune_replace=None,
mult=1.0, mult_multival=None, mode_replace=None, tune_replace=None,
start_percent=1.0, guarantee_steps=1, inherit_missing=True):
if prev_kf is None:
prev_kf = ContextRefKeyframeGroup()
@@ -213,7 +318,7 @@ class ContextRef_KeyframeInterpolationNode:
start_percent: float, end_percent: float,
mult_start: float, mult_end: float, interpolation: str, intervals: int,
inherit_missing=True, prev_kf: ContextRefKeyframeGroup=None,
mult_multival=1.0, mode_replace=None, tune_replace=None, print_keyframes=False):
mult_multival=None, mode_replace=None, tune_replace=None, print_keyframes=False):
if prev_kf is None:
prev_kf = ContextRefKeyframeGroup()
prev_kf = prev_kf.clone()
@@ -233,6 +338,59 @@ class ContextRef_KeyframeInterpolationNode:
return (prev_kf,)
class ContextRef_KeyframeFromListNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mults_float": ("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}),
"inherit_missing": ("BOOLEAN", {"default": True}),
"print_keyframes": ("BOOLEAN", {"default": False}),
},
"optional": {
"prev_kf": ("CONTEXTREF_KEYFRAME",),
"mult_multival": ("MULTIVAL",),
"mode_replace": ("CONTEXTREF_MODE",),
"tune_replace": ("CONTEXTREF_TUNE",),
"autosize": ("ADEAUTOSIZE", {"padding": 50}),
}
}
RETURN_TYPES = ("CONTEXTREF_KEYFRAME",)
RETURN_NAMES = ("CONTEXTREF_KF",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
FUNCTION = "create_keyframe"
def create_keyframe(self, mults_float: Union[float, list[float]],
start_percent: float, end_percent: float,
inherit_missing=True, prev_kf: ContextRefKeyframeGroup=None,
mult_multival=None, mode_replace=None, tune_replace=None, print_keyframes=False):
if prev_kf is None:
prev_kf = ContextRefKeyframeGroup()
prev_kf = prev_kf.clone()
if type(mults_float) in (float, int):
mults_float = [float(mults_float)]
elif isinstance(mults_float, Iterable):
pass
else:
raise Exception(f"strengths_float must be either an interable input or a float, but was {type(mults_float).__repr__}.")
percents = InterpolationMethod.get_weights(num_from=start_percent, num_to=end_percent, length=len(mults_float), method=InterpolationMethod.LINEAR)
is_first = True
for percent, mult in zip(percents, mults_float):
guarantee_steps = 0
if is_first:
guarantee_steps = 1
is_first = False
prev_kf.add(ContextRefKeyframe(mult=mult, mult_multival=mult_multival, tune_replace=tune_replace, mode_replace=mode_replace,
start_percent=percent, guarantee_steps=guarantee_steps, inherit_missing=inherit_missing))
if print_keyframes:
logger.info(f"ContextRefKeyframe - start_percent:{percent} = {mult}")
return (prev_kf,)
class ContextRef_ModeFirst:
@classmethod
def INPUT_TYPES(s):
+5 -3
View File
@@ -276,10 +276,12 @@ def create_multival_combo(float_val: Union[float, list[float]], mask_optional: T
def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[float, Tensor], force_leader_A=False) -> Union[float, Tensor]:
if multivalA is None and multivalB is None:
return 1.0
# if one is None, use the other
if multivalA == None:
if multivalA is None:
return multivalB
elif multivalB == None:
elif multivalB is None:
return multivalA
# both have a value - combine them based on type
# if both are Tensors, make dims match before multiplying
@@ -306,7 +308,7 @@ def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[floa
def resize_multival(multival: Union[float, Tensor], batch_size: int, height: int, width: int):
if multival == None:
if multival is None:
return 1.0
if type(multival) != Tensor:
return multival