diff --git a/animatediff/context_extras.py b/animatediff/context_extras.py index 3af7df6..2f19856 100644 --- a/animatediff/context_extras.py +++ b/animatediff/context_extras.py @@ -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) diff --git a/animatediff/nodes.py b/animatediff/nodes.py index e8a0a57..69a5496 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -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 diff --git a/animatediff/nodes_context_extras.py b/animatediff/nodes_context_extras.py index e9860df..8a28e8d 100644 --- a/animatediff/nodes_context_extras.py +++ b/animatediff/nodes_context_extras.py @@ -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): diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index 8425f02..501a7f6 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -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