Added basic NaiveReuse Keyframe nodes

This commit is contained in:
Jedrzej Kosinski
2024-08-05 04:32:22 -05:00
parent f051ca9944
commit 6167eee9f4
4 changed files with 252 additions and 29 deletions
+128 -11
View File
@@ -1,8 +1,12 @@
from typing import Union
import math
import torch
from torch import Tensor
from comfy.model_base import BaseModel
from .utils_motion import prepare_mask_batch, extend_to_batch_size, get_combined_multival
from .utils_motion import (prepare_mask_batch, extend_to_batch_size, get_combined_multival, resize_multival,
get_sorted_list_via_attr)
class ContextExtra:
@@ -93,34 +97,147 @@ class ContextRef(ContextExtra):
################################
# NaiveReuse
# NaiveReuse
class NaiveReuseKeyframe:
def __init__(self, mult_multival: Union[float, Tensor], start_percent=0.0, guarantee_steps=1):
self.mult_multival = mult_multival
# scheduling
self.start_percent = float(start_percent)
self.start_t = 999999999.9
self.guarantee_steps = guarantee_steps
def clone(self):
c = NaiveReuseKeyframe(mult_multival=self.mult_multival,
start_percent=self.start_percent, guarantee_steps=self.guarantee_steps)
c.start_t = self.start_t
return c
class NaiveReuseKeyframeGroup:
def __init__(self):
self.keyframes: list[NaiveReuseKeyframe] = []
self._current_keyframe: NaiveReuseKeyframe = None
self._current_used_steps: int = 0
self._current_index: int = 0
self._previous_t = -1
def reset(self):
self._current_keyframe = None
self._current_used_steps = 0
self._current_index = 0
self._set_first_as_current()
def add(self, keyframe: NaiveReuseKeyframe):
# add to end of list, then sort
self.keyframes.append(keyframe)
self.keyframes = get_sorted_list_via_attr(self.keyframes, "start_percent")
self._set_first_as_current()
def _set_first_as_current(self):
if len(self.keyframes) > 0:
self._current_keyframe = self.keyframes[0]
else:
self._current_keyframe = None
def has_index(self, index: int) -> int:
return index >=0 and index < len(self.keyframes)
def is_empty(self) -> bool:
return len(self.keyframes) == 0
def clone(self):
cloned = NaiveReuseKeyframeGroup()
for keyframe in self.keyframes:
cloned.keyframes.append(keyframe)
cloned._set_first_as_current()
return cloned
def initialize_timesteps(self, model: BaseModel):
for keyframe in self.keyframes:
keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent)
def prepare_current_keyframe(self, t: Tensor):
curr_t: float = t[0]
# if curr_t same as before, do nothing as step already accounted for
if curr_t == self._previous_t:
return
prev_index = self._current_index
# if met guaranteed steps, look for next keyframe in case need to switch
if self._current_used_steps >= self._current_keyframe.guarantee_steps:
# if has next index, loop through and see if need t oswitch
if self.has_index(self._current_index+1):
for i in range(self._current_index+1, len(self.keyframes)):
eval_c = self.keyframes[i]
# check if start_t is greater or equal to curr_t
# NOTE: t is in terms of sigmas, not percent, so bigger number = earlier step in sampling
if eval_c.start_t >= curr_t:
self._current_index = i
self._current_keyframe = eval_c
self._current_used_steps = 0
# if guarantee_steps greater than zero, stop searching for other keyframes
if self._current_keyframe.guarantee_steps > 0:
break
# if eval_c is outside the percent range, stop looking further
else: break
# update steps current context is used
self._current_used_steps += 1
# update previous_t
self._previous_t = curr_t
# properties shadow those of NaiveReuseKeyframe
@property
def mult_multival(self):
if self._current_keyframe != None:
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: Tensor=None):
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)
self.weighted_mean = weighted_mean
self.orig_multival = multival_opt
self.mask: Tensor = None
self.keyframe = naivereuse_kf if naivereuse_kf else NaiveReuseKeyframeGroup()
self._prev_keyframe = None
def cleanup(self):
super().cleanup()
del self.mask
self.mask = None
self._prev_keyframe = None
self.keyframe.reset()
def initialize_timesteps(self, model: BaseModel):
super().initialize_timesteps(model)
self.keyframe.initialize_timesteps(model)
def prepare_current(self, t: Tensor):
super().prepare_current(t)
self.keyframe.prepare_current_keyframe(t)
def get_effective_weighted_mean(self, x: Tensor, idxs: list[int]):
if self.orig_multival is None:
if self.orig_multival is None and self.keyframe.mult_multival is None:
return self.weighted_mean
# otherwise, is Tensor and should be extended to match dims and size of x;
# see if needs to be recalculated
if type(self.orig_multival) != Tensor:
return self.weighted_mean * self.orig_multival
elif self.mask is None or self.mask.shape[0] != x.shape[0] or self.mask.shape[-1] != x.shape[-1] or self.mask.shape[-2] != x.shape[-2]:
# check if keyframe changed
keyframe_changed = False
if self.keyframe._current_keyframe != self._prev_keyframe:
keyframe_changed = True
self._prev_keyframe = self.keyframe._current_keyframe
if type(self.orig_multival) != Tensor and type(self.keyframe.mult_multival) != Tensor:
return self.weighted_mean * get_combined_multival(self.orig_multival, self.keyframe.mult_multival)
if self.mask is None or keyframe_changed or self.mask.shape[0] != x.shape[0] or self.mask.shape[-1] != x.shape[-1] or self.mask.shape[-2] != x.shape[-2]:
del self.mask
self.mask = prepare_mask_batch(self.orig_multival, x.shape)
self.mask = extend_to_batch_size(self.mask, x.shape[0])
real_mult_multival = resize_multival(self.keyframe.mult_multival, batch_size=x.shape[0], height=x.shape[-1], width=x.shape[-2])
self.mask = resize_multival(self.orig_multival, batch_size=x.shape[0], height=x.shape[-1], width=x.shape[-2])
self.mask = get_combined_multival(self.mask, real_mult_multival)
return self.weighted_mean * self.mask[idxs].to(dtype=x.dtype, device=x.device)
def should_run(self):
to_return = super().should_run()
# if keyframe has 0.0 val, should not run
if self.keyframe.mult_multival is not None and type(self.keyframe.mult_multival) != Tensor and math.isclose(self.keyframe.mult_multival, 0.0):
return False
# if weighted_mean is 0.0, then reuse will take no effect anyway
return to_return and self.weighted_mean > 0.0
#--------------------------------
+8 -3
View File
@@ -31,7 +31,8 @@ from .nodes_context import (LegacyLoopedUniformContextOptionsNode, LoopedUniform
VisualizeContextOptionsK, VisualizeContextOptionsKAdv, VisualizeContextOptionsSCustom,
SetContextExtrasOnContextOptions, ContextExtras_NaiveReuse, ContextExtras_ContextRef,
ContextRef_ModeFirst, ContextRef_ModeSliding, ContextRef_ModeIndexes,
ContextRef_TuneAttn, ContextRef_TuneAttnAdain)
ContextRef_TuneAttn, ContextRef_TuneAttnAdain,
NaiveReuse_KeyframeNode, NaiveReuse_KeyframeMultivalNode)
from .nodes_ad_settings import (AnimateDiffSettingsNode, ManualAdjustPENode, SweetspotStretchPENode, FullStretchPENode,
WeightAdjustAllAddNode, WeightAdjustAllMultNode, WeightAdjustIndivAddNode, WeightAdjustIndivMultNode,
WeightAdjustIndivAttnAddNode, WeightAdjustIndivAttnMultNode)
@@ -76,12 +77,14 @@ NODE_CLASS_MAPPINGS = {
# Context Extras
"ADE_ContextExtras_Set": SetContextExtrasOnContextOptions,
"ADE_ContextExtras_ContextRef": ContextExtras_ContextRef,
"ADE_ContextExtras_NaiveReuse": ContextExtras_NaiveReuse,
"ADE_ContextExtras_ContextRef_ModeFirst": ContextRef_ModeFirst,
"ADE_ContextExtras_ContextRef_ModeSliding": ContextRef_ModeSliding,
"ADE_ContextExtras_ContextRef_ModeIndexes": ContextRef_ModeIndexes,
"ADE_ContextExtras_ContextRef_TuneAttn": ContextRef_TuneAttn,
"ADE_ContextExtras_ContextRef_TuneAttnAdain": ContextRef_TuneAttnAdain,
"ADE_ContextExtras_NaiveReuse": ContextExtras_NaiveReuse,
"ADE_ContextExtras_NaiveReuse_Keyframe": NaiveReuse_KeyframeNode,
"ADE_ContextExtras_NaiveReuse_KeyframeMultival": NaiveReuse_KeyframeMultivalNode,
#------------------------------------------------------------------------------
###############################################################################
# Iteration Opts
@@ -219,12 +222,14 @@ NODE_DISPLAY_NAME_MAPPINGS = {
# Context Extras
"ADE_ContextExtras_Set": "Set Context Extras 🎭🅐🅓",
"ADE_ContextExtras_ContextRef": "Context Extras◆ContextRef 🎭🅐🅓",
"ADE_ContextExtras_NaiveReuse": "Context Extras◆NaiveReuse 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_ModeFirst": "ContextRef Mode◆First 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_ModeSliding": "ContextRef Mode◆Sliding 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_ModeIndexes": "ContextRef Mode◆Indexes 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_TuneAttn": "ContextRef Tune◆Attn 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_TuneAttnAdain": "ContextRef Tune◆Attn+Adain 🎭🅐🅓",
"ADE_ContextExtras_NaiveReuse": "Context Extras◆NaiveReuse 🎭🅐🅓",
"ADE_ContextExtras_NaiveReuse_Keyframe": "NaiveReuse Keyframe 🎭🅐🅓",
"ADE_ContextExtras_NaiveReuse_KeyframeMultival": "NaiveReuse Keyframe [Multival] 🎭🅐🅓",
#------------------------------------------------------------------------------
###############################################################################
# Iteration Opts
+93 -8
View File
@@ -7,7 +7,8 @@ from comfy.model_patcher import ModelPatcher
from .context import (ContextFuseMethod, ContextOptions, ContextOptionsGroup, ContextSchedules,
generate_context_visualization)
from .context_extras import ContextExtrasGroup, ContextRef, ContextRefParams, ContextRefMode, NaiveReuse
from .context_extras import (ContextExtrasGroup, ContextRef, ContextRefParams, ContextRefMode, NaiveReuse,
NaiveReuseKeyframe, NaiveReuseKeyframeGroup)
from .utils_model import BIGMAX, MAX_RESOLUTION
from .utils_scheduling import convert_str_to_indexes
@@ -480,6 +481,7 @@ class ContextExtras_NaiveReuse:
"optional": {
"prev_extras": ("CONTEXT_EXTRAS",),
"strength_multival": ("MULTIVAL",),
"naivereuse_kf": ("NAIVEREUSE_KEYFRAME",),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_percent": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001}),
"weighted_mean": ("FLOAT", {"default": 0.95, "min": 0.0, "max": 1.0, "step": 0.001}),
@@ -492,16 +494,73 @@ class ContextExtras_NaiveReuse:
FUNCTION = "create_context_extra"
def create_context_extra(self, start_percent=0.0, end_percent=0.1, weighted_mean=0.95, strength_multival: Union[float, Tensor]=None,
prev_extras: ContextExtrasGroup=None):
naivereuse_kf: NaiveReuseKeyframeGroup=None, prev_extras: ContextExtrasGroup=None):
if prev_extras is None:
prev_extras = prev_extras = ContextExtrasGroup()
prev_extras = prev_extras.clone()
# create extra
naive_reuse = NaiveReuse(start_percent=start_percent, end_percent=end_percent, weighted_mean=weighted_mean, multival_opt=strength_multival)
naive_reuse = NaiveReuse(start_percent=start_percent, end_percent=end_percent, weighted_mean=weighted_mean, multival_opt=strength_multival,
naivereuse_kf=naivereuse_kf)
prev_extras.add(naive_reuse)
return (prev_extras,)
class NaiveReuse_KeyframeMultivalNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mult_multival": ("MULTIVAL",),
},
"optional": {
"prev_kf": ("NAIVEREUSE_KEYFRAME",),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
"autosize": ("ADEAUTOSIZE", {"padding": 80}),
}
}
RETURN_TYPES = ("NAIVEREUSE_KEYFRAME",)
RETURN_NAMES = ("NAIVEREUSE_KF",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/naivereuse"
FUNCTION = "create_keyframe"
def create_keyframe(self, prev_kf=None, mult_multival=1.0,
start_percent=0.0, guarantee_steps=1):
if prev_kf is None:
prev_kf = NaiveReuseKeyframeGroup()
prev_kf = prev_kf.clone()
kf = NaiveReuseKeyframe(mult_multival=mult_multival, start_percent=start_percent, guarantee_steps=guarantee_steps)
prev_kf.add(kf)
return (prev_kf,)
class NaiveReuse_KeyframeNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
},
"optional": {
"prev_kf": ("NAIVEREUSE_KEYFRAME",),
"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}),
"autosize": ("ADEAUTOSIZE", {"padding": 10}),
}
}
RETURN_TYPES = ("NAIVEREUSE_KEYFRAME",)
RETURN_NAMES = ("NAIVEREUSE_KF",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/naivereuse"
FUNCTION = "create_keyframe"
def create_keyframe(self, prev_kf=None, mult=1.0,
start_percent=0.0, guarantee_steps=1):
return NaiveReuse_KeyframeMultivalNode.create_keyframe(self, prev_kf=prev_kf, mult_multival=float(mult),
start_percent=start_percent, guarantee_steps=guarantee_steps)
class ContextExtras_ContextRef:
@classmethod
def INPUT_TYPES(s):
@@ -541,6 +600,32 @@ class ContextExtras_ContextRef:
return (prev_extras,)
class ContextRef_KeyframeNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
},
"optional": {
"prev_keyframe": ("CONTEXTREF_KEYFRAME",),
"mult_multival": ("MULTIVAL",),
"mode_replace": ("CONTEXTREF_MODE",),
"tune_replace": ("CONTEXTREF_TUNE",),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
}
}
RETURN_TYPES = ("CONTEXTREF_KEYFRAME",)
RETURN_NAMES = ("CONTEXTREF_KF",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
FUNCTION = "create_keyframe"
def create_keyframe(self, prev_keyframe=None, mult_multival=1.0, mode_replace=None, tune_replace=None,
start_percent=1.0, guarantee_steps=1):
pass
class ContextRef_ModeFirst:
@classmethod
def INPUT_TYPES(s):
@@ -553,7 +638,7 @@ class ContextRef_ModeFirst:
}
RETURN_TYPES = ("CONTEXTREF_MODE",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
FUNCTION = "create_contextref_mode"
def create_contextref_mode(self):
@@ -574,7 +659,7 @@ class ContextRef_ModeSliding:
}
RETURN_TYPES = ("CONTEXTREF_MODE",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
FUNCTION = "create_contextref_mode"
def create_contextref_mode(self, sliding_width):
@@ -596,7 +681,7 @@ class ContextRef_ModeIndexes:
}
RETURN_TYPES = ("CONTEXTREF_MODE",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
FUNCTION = "create_contextref_mode"
def create_contextref_mode(self, switch_on_idxs: str, always_include_0: bool):
@@ -625,7 +710,7 @@ class ContextRef_TuneAttnAdain:
}
RETURN_TYPES = ("CONTEXTREF_TUNE",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
FUNCTION = "create_contextref_tune"
def create_contextref_tune(self, attn_style_fidelity=1.0, attn_ref_weight=1.0, attn_strength=1.0,
@@ -651,7 +736,7 @@ class ContextRef_TuneAttn:
}
RETURN_TYPES = ("CONTEXTREF_TUNE",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/contextref"
FUNCTION = "create_contextref_tune"
def create_contextref_tune(self, attn_style_fidelity=1.0, attn_ref_weight=1.0, attn_strength=1.0):
+23 -7
View File
@@ -275,7 +275,7 @@ def create_multival_combo(float_val: Union[float, list[float]], mask_optional: T
return mask_optional
def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[float, Tensor]) -> Union[float, Tensor]:
def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[float, Tensor], force_leader_A=False) -> Union[float, Tensor]:
# if one is None, use the other
if multivalA == None:
return multivalB
@@ -284,14 +284,18 @@ def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[floa
# both have a value - combine them based on type
# if both are Tensors, make dims match before multiplying
if type(multivalA) == Tensor and type(multivalB) == Tensor:
areaA = multivalA.shape[1]*multivalA.shape[2]
areaB = multivalB.shape[1]*multivalB.shape[2]
# match height/width to mask with larger area
leader,follower = (multivalA,multivalB) if areaA >= areaB else (multivalB,multivalA)
batch_size = multivalA.shape[0] if multivalA.shape[0] >= multivalB.shape[0] else multivalB.shape[0]
if force_leader_A:
leader,follower = (multivalA,multivalB)
batch_size = multivalA.shape[0]
else:
areaA = multivalA.shape[1]*multivalA.shape[2]
areaB = multivalB.shape[1]*multivalB.shape[2]
# match height/width to mask with larger area
leader,follower = (multivalA,multivalB) if areaA >= areaB else (multivalB,multivalA)
batch_size = multivalA.shape[0] if multivalA.shape[0] >= multivalB.shape[0] else multivalB.shape[0]
# make follower same dimensions as leader
follower = torch.unsqueeze(follower, 1)
follower = comfy.utils.common_upscale(follower, leader.shape[2], leader.shape[1], "bilinear", "center")
follower = comfy.utils.common_upscale(follower, leader.shape[-1], leader.shape[-2], "bilinear", "center")
follower = torch.squeeze(follower, 1)
# make sure batch size will match
leader = extend_to_batch_size(leader, batch_size)
@@ -301,6 +305,18 @@ def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[floa
return multivalA * multivalB
def resize_multival(multival: Union[float, Tensor], batch_size: int, height: int, width: int):
if multival == None:
return 1.0
if type(multival) != Tensor:
return multival
multival = torch.unsqueeze(multival, 1)
multival = comfy.utils.common_upscale(multival, height, width, "bilinear", "center")
multival = torch.squeeze(multival, 1)
multival = extend_to_batch_size(multival, batch_size)
return multival
def get_combined_input(inputA: Union[InputPIA, None], inputB: Union[InputPIA, None], x: Tensor):
if inputA is None:
inputA = InputPIA_Multival(1.0)