mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-30 12:47:10 +08:00
Added initial hook scheduling nodes, small renaming/refactoring
This commit is contained in:
parent
a5034df6db
commit
5a9aa5817c
@ -1,6 +1,7 @@
|
|||||||
from typing import TYPE_CHECKING, List, Dict, Tuple
|
from typing import TYPE_CHECKING, List, Dict, Tuple
|
||||||
import enum
|
import enum
|
||||||
import torch
|
import torch
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from comfy.model_patcher import ModelPatcher
|
from comfy.model_patcher import ModelPatcher
|
||||||
@ -14,13 +15,41 @@ class EnumHookMode(enum.Enum):
|
|||||||
MinVram = "minvram"
|
MinVram = "minvram"
|
||||||
MaxSpeed = "maxspeed"
|
MaxSpeed = "maxspeed"
|
||||||
|
|
||||||
|
class InterpolationMethod:
|
||||||
|
LINEAR = "linear"
|
||||||
|
EASE_IN = "ease_in"
|
||||||
|
EASE_OUT = "ease_out"
|
||||||
|
EASE_IN_OUT = "ease_in_out"
|
||||||
|
|
||||||
|
_LIST = [LINEAR, EASE_IN, EASE_OUT, EASE_IN_OUT]
|
||||||
|
|
||||||
|
@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 HookRef:
|
class HookRef:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
class Hook:
|
class Hook:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.hook_ref = HookRef()
|
self.hook_ref = HookRef()
|
||||||
self.hook_keyframe = HookWeightKeyframeGroup()
|
self.hook_keyframe = HookKeyframeGroup()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def strength(self):
|
def strength(self):
|
||||||
@ -68,7 +97,7 @@ class HookGroup:
|
|||||||
c.add(hook.clone())
|
c.add(hook.clone())
|
||||||
return c
|
return c
|
||||||
|
|
||||||
def set_keyframes_on_hooks(self, hook_kf: 'HookWeightKeyframeGroup'):
|
def set_keyframes_on_hooks(self, hook_kf: 'HookKeyframeGroup'):
|
||||||
hook_kf = hook_kf.clone()
|
hook_kf = hook_kf.clone()
|
||||||
for hook in self.hooks:
|
for hook in self.hooks:
|
||||||
hook.hook_keyframe = hook_kf
|
hook.hook_keyframe = hook_kf
|
||||||
@ -92,7 +121,7 @@ class HookGroup:
|
|||||||
final_hook = final_hook.clone_and_combine(hook)
|
final_hook = final_hook.clone_and_combine(hook)
|
||||||
return final_hook
|
return final_hook
|
||||||
|
|
||||||
class HookWeightKeyframe:
|
class HookKeyframe:
|
||||||
def __init__(self, strength: float, start_percent=0.0, guarantee_steps=1):
|
def __init__(self, strength: float, start_percent=0.0, guarantee_steps=1):
|
||||||
self.strength = strength
|
self.strength = strength
|
||||||
# scheduling
|
# scheduling
|
||||||
@ -101,15 +130,15 @@ class HookWeightKeyframe:
|
|||||||
self.guarantee_steps = guarantee_steps
|
self.guarantee_steps = guarantee_steps
|
||||||
|
|
||||||
def clone(self):
|
def clone(self):
|
||||||
c = HookWeightKeyframe(strength=self.strength,
|
c = HookKeyframe(strength=self.strength,
|
||||||
start_percent=self.start_percent, guarantee_steps=self.guarantee_steps)
|
start_percent=self.start_percent, guarantee_steps=self.guarantee_steps)
|
||||||
c.start_t = self.start_t
|
c.start_t = self.start_t
|
||||||
return c
|
return c
|
||||||
|
|
||||||
class HookWeightKeyframeGroup:
|
class HookKeyframeGroup:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.keyframes: List[HookWeightKeyframe] = []
|
self.keyframes: List[HookKeyframe] = []
|
||||||
self._current_keyframe: HookWeightKeyframe = None
|
self._current_keyframe: HookKeyframe = None
|
||||||
self._current_used_steps = 0
|
self._current_used_steps = 0
|
||||||
self._current_index = 0
|
self._current_index = 0
|
||||||
self._curr_t = -1.
|
self._curr_t = -1.
|
||||||
@ -126,8 +155,9 @@ class HookWeightKeyframeGroup:
|
|||||||
self._current_used_steps = 0
|
self._current_used_steps = 0
|
||||||
self._current_index = 0
|
self._current_index = 0
|
||||||
self.curr_t = -1.
|
self.curr_t = -1.
|
||||||
|
self._set_first_as_current()
|
||||||
|
|
||||||
def add(self, keyframe: HookWeightKeyframe):
|
def add(self, keyframe: HookKeyframe):
|
||||||
# add to end of list, then sort
|
# add to end of list, then sort
|
||||||
self.keyframes.append(keyframe)
|
self.keyframes.append(keyframe)
|
||||||
self.keyframes = get_sorted_list_via_attr(self.keyframes, "start_percent")
|
self.keyframes = get_sorted_list_via_attr(self.keyframes, "start_percent")
|
||||||
@ -146,7 +176,7 @@ class HookWeightKeyframeGroup:
|
|||||||
return len(self.keyframes) == 0
|
return len(self.keyframes) == 0
|
||||||
|
|
||||||
def clone(self):
|
def clone(self):
|
||||||
c = HookWeightKeyframeGroup()
|
c = HookKeyframeGroup()
|
||||||
for keyframe in self.keyframes:
|
for keyframe in self.keyframes:
|
||||||
c.keyframes.append(keyframe)
|
c.keyframes.append(keyframe)
|
||||||
c._set_first_as_current()
|
c._set_first_as_current()
|
||||||
|
|||||||
@ -1,5 +1,6 @@
|
|||||||
from typing import TYPE_CHECKING, Dict, List, Tuple
|
from typing import TYPE_CHECKING, Dict, List, Tuple, Union
|
||||||
import torch
|
import torch
|
||||||
|
from collections.abc import Iterable
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from comfy.model_patcher import ModelPatcher
|
from comfy.model_patcher import ModelPatcher
|
||||||
@ -159,6 +160,8 @@ class SetClipHooks:
|
|||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"clip": ("CLIP",),
|
"clip": ("CLIP",),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
"hooks": ("HOOKS",),
|
"hooks": ("HOOKS",),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@ -167,10 +170,30 @@ class SetClipHooks:
|
|||||||
CATEGORY = "advanced/hooks/clip"
|
CATEGORY = "advanced/hooks/clip"
|
||||||
FUNCTION = "apply_hooks"
|
FUNCTION = "apply_hooks"
|
||||||
|
|
||||||
def apply_hooks(self, clip: 'CLIP', hooks: comfy.hooks.HookGroup):
|
def apply_hooks(self, clip: 'CLIP', hooks: comfy.hooks.HookGroup=None):
|
||||||
clip = clip.clone()
|
if hooks is not None:
|
||||||
clip.patcher.forced_hooks = hooks
|
clip = clip.clone()
|
||||||
|
clip.patcher.forced_hooks = hooks
|
||||||
return (clip,)
|
return (clip,)
|
||||||
|
|
||||||
|
class ConditioningTimestepsRange:
|
||||||
|
NodeId = 'ConditioningTimestepsRange'
|
||||||
|
NodeName = 'Timesteps Range'
|
||||||
|
@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})
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("TIMESTEPS_RANGE",)
|
||||||
|
CATEGORY = "advanced/hooks"
|
||||||
|
FUNCTION = "create_range"
|
||||||
|
|
||||||
|
def create_range(self, start_percent: float, end_percent: float):
|
||||||
|
return ((start_percent, end_percent),)
|
||||||
#------------------------------------------
|
#------------------------------------------
|
||||||
###########################################
|
###########################################
|
||||||
|
|
||||||
@ -313,6 +336,105 @@ class RegisterHookModelAsLoraModelOnly:
|
|||||||
###########################################
|
###########################################
|
||||||
# Schedule Hooks
|
# Schedule Hooks
|
||||||
#------------------------------------------
|
#------------------------------------------
|
||||||
|
class SetHookKeyframes:
|
||||||
|
NodeId = 'SetHookKeyframes'
|
||||||
|
NodeName = 'Set Hook Keyframes'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"hooks": ("HOOKS",),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"hook_kf": ("HOOK_KEYFRAMES",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("HOOKS",)
|
||||||
|
CATEGORY = "advanced/hooks/scheduling"
|
||||||
|
FUNCTION = "set_hook_keyframes"
|
||||||
|
|
||||||
|
def set_hook_keyframes(self, hooks: comfy.hooks.HookGroup, hook_kf: comfy.hooks.HookKeyframeGroup=None):
|
||||||
|
if hook_kf is not None:
|
||||||
|
hooks = hooks.clone()
|
||||||
|
hooks.set_keyframes_on_hooks(hook_kf=hook_kf)
|
||||||
|
return (hooks,)
|
||||||
|
|
||||||
|
class CreateHookKeyframe:
|
||||||
|
NodeId = 'CreateHookKeyframe'
|
||||||
|
NodeName = 'Create Hook Keyframe'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"strength_mult": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||||
|
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"prev_hook_kf": ("HOOK_KEYFRAMES",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("HOOK_KEYFRAMES",)
|
||||||
|
RETURN_NAMES = ("HOOK_KF",)
|
||||||
|
CATEGORY = "advanced/hooks/scheduling"
|
||||||
|
FUNCTION = "create_hook_keyframe"
|
||||||
|
|
||||||
|
def create_hook_keyframe(self, strength_mult: float, start_percent: float, prev_hook_kf: comfy.hooks.HookKeyframeGroup=None):
|
||||||
|
if prev_hook_kf is None:
|
||||||
|
prev_hook_kf = comfy.hooks.HookKeyframeGroup()
|
||||||
|
prev_hook_kf = prev_hook_kf.clone()
|
||||||
|
keyframe = comfy.hooks.HookKeyframe(strength=strength_mult, start_percent=start_percent)
|
||||||
|
prev_hook_kf.add(keyframe)
|
||||||
|
return (prev_hook_kf,)
|
||||||
|
|
||||||
|
class CreateHookKeyframesFromFloats:
|
||||||
|
NodeId = 'CreateHookKeyframesFromFloats'
|
||||||
|
NodeName = 'Create Hook Keyframes From Floats'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"floats_strength": ("FLOATS", {"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}),
|
||||||
|
"print_keyframes": ("BOOLEAN", {"default": False}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"prev_hook_kf": ("HOOK_KEYFRAMES",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("HOOK_KEYFRAMES",)
|
||||||
|
RETURN_NAMES = ("HOOK_KF",)
|
||||||
|
CATEGORY = "advanced/hooks/scheduling"
|
||||||
|
FUNCTION = "create_hook_keyframes"
|
||||||
|
|
||||||
|
def create_hook_keyframes(self, floats_strength: Union[float, List[float]],
|
||||||
|
start_percent: float, end_percent: float,
|
||||||
|
prev_hook_kf: comfy.hooks.HookKeyframeGroup=None, print_keyframes=False):
|
||||||
|
if prev_hook_kf is None:
|
||||||
|
prev_hook_kf = comfy.hooks.HookKeyframeGroup()
|
||||||
|
prev_hook_kf = prev_hook_kf.clone()
|
||||||
|
if type(floats_strength) in (float, int):
|
||||||
|
floats_strength = [float(floats_strength)]
|
||||||
|
elif isinstance(floats_strength, Iterable):
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
raise Exception(f"floats_strength must be either an iterable input or a float, but was{type(floats_strength).__repr__}.")
|
||||||
|
percents = comfy.hooks.InterpolationMethod.get_weights(num_from=start_percent, num_to=end_percent, length=len(floats_strength),
|
||||||
|
method=comfy.hooks.InterpolationMethod.LINEAR)
|
||||||
|
|
||||||
|
is_first = True
|
||||||
|
for percent, strength in zip(percents, floats_strength):
|
||||||
|
guarantee_steps = 0
|
||||||
|
if is_first:
|
||||||
|
guarantee_steps = 1
|
||||||
|
is_first = False
|
||||||
|
prev_hook_kf.add(comfy.hooks.HookKeyframe(strength=strength, start_percent=percent, guarantee_steps=guarantee_steps))
|
||||||
|
if print_keyframes:
|
||||||
|
print(f"Hook Keyframe - start_percent:{percent} = {strength}")
|
||||||
|
return (prev_hook_kf,)
|
||||||
#------------------------------------------
|
#------------------------------------------
|
||||||
###########################################
|
###########################################
|
||||||
|
|
||||||
@ -434,6 +556,10 @@ node_list = [
|
|||||||
RegisterHookLoraModelOnly,
|
RegisterHookLoraModelOnly,
|
||||||
RegisterHookModelAsLora,
|
RegisterHookModelAsLora,
|
||||||
RegisterHookModelAsLoraModelOnly,
|
RegisterHookModelAsLoraModelOnly,
|
||||||
|
# Scheduling
|
||||||
|
SetHookKeyframes,
|
||||||
|
CreateHookKeyframe,
|
||||||
|
CreateHookKeyframesFromFloats,
|
||||||
# Combine
|
# Combine
|
||||||
CombineHooks,
|
CombineHooks,
|
||||||
CombineHooksFour,
|
CombineHooksFour,
|
||||||
@ -445,6 +571,8 @@ node_list = [
|
|||||||
PairConditioningSetDefaultAndCombine,
|
PairConditioningSetDefaultAndCombine,
|
||||||
PairConditioningCombine,
|
PairConditioningCombine,
|
||||||
SetClipHooks,
|
SetClipHooks,
|
||||||
|
# Other
|
||||||
|
ConditioningTimestepsRange,
|
||||||
]
|
]
|
||||||
NODE_CLASS_MAPPINGS = {}
|
NODE_CLASS_MAPPINGS = {}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user