diff --git a/__init__.py b/__init__.py index af13844..7a6face 100644 --- a/__init__.py +++ b/__init__.py @@ -29,11 +29,4 @@ from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS print(os.environ.get('COMFYUI_DEBUG_MODE')) -from .debug import NODE_CLASS_MAPPINGS as ncm0, NODE_DISPLAY_NAME_MAPPINGS as ndnm0 - -# there's probably a cleaner, more-dummy-proof way to do this. -# feels like an accident waiting to happen. low risk though. -NODE_CLASS_MAPPINGS.update(ncm0) -NODE_DISPLAY_NAME_MAPPINGS.update(ndnm0) - __all__ =["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/nodes/__init__.py b/nodes/__init__.py new file mode 100644 index 0000000..6f0169a --- /dev/null +++ b/nodes/__init__.py @@ -0,0 +1,18 @@ + +from .core import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +from .debug import NODE_CLASS_MAPPINGS as ncm0, NODE_DISPLAY_NAME_MAPPINGS as ndnm0 +from .entangled import NODE_CLASS_MAPPINGS as ncm1, NODE_DISPLAY_NAME_MAPPINGS as ndnm1 +from .schedule import NODE_CLASS_MAPPINGS as ncm2, NODE_DISPLAY_NAME_MAPPINGS as ndnm2 + + +# there's probably a cleaner, more-dummy-proof way to do this. +# feels like an accident waiting to happen. low risk though. +NODE_CLASS_MAPPINGS.update(ncm0) +NODE_CLASS_MAPPINGS.update(ncm1) +NODE_CLASS_MAPPINGS.update(ncm2) +NODE_DISPLAY_NAME_MAPPINGS.update(ndnm0) +NODE_DISPLAY_NAME_MAPPINGS.update(ndnm1) +NODE_DISPLAY_NAME_MAPPINGS.update(ndnm2) + + +__all__ =["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/nodes.py b/nodes/core.py similarity index 58% rename from nodes.py rename to nodes/core.py index 42bc7fb..8ab167e 100644 --- a/nodes.py +++ b/nodes/core.py @@ -69,7 +69,7 @@ label: foo""" class KfEvaluateCurveAtT: - CATEGORY=CATEGORY + CATEGORY=CATEGORY # TODO: create a "utils" group FUNCTION = 'main' RETURN_TYPES = ("FLOAT","INT") @@ -313,6 +313,8 @@ class KfCurvesMultiply: return (curve_1 * curve_2, ) +## This seems to not be working properly. I think the issue is upstream in Keyframed +# TODO: set as experimental? class KfCurvesDivide: CATEGORY = CATEGORY FUNCTION = "main" @@ -331,9 +333,26 @@ class KfCurvesDivide: return (curve_1 / curve_2, ) +class KfCurveConstant: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("KEYFRAMED_CURVE",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "value": ("FLOAT", {"forceInput": True,}) + }} + + def main(self, value): + curve = kf.Curve(value) + return (curve,) + + ################################################################## -#### Working with parameter groups +### TODO: Working with parameter groups # Label curve @@ -351,6 +370,257 @@ class KfCurvesDivide: # extract a time slice from the parameter group +################################################################## + +### Sinusoidal + +class KfSinusoidalWithFrequency: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("KEYFRAMED_CURVE", "SINUSOIDAL_CURVE") + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "frequency": ("FLOAT",{ + "default": 1/12, + "step": 0.01, + }), + "phase": ("FLOAT", { + "default": 0.0, + #"min": 0.0, + #"max": 6.28318530718, # 2*pi + "step": 0.1308996939, # pi/24 + }), + "amplitude": ("FLOAT",{ + "default": 1, + "step": 0.01, + }), + }, + } + + def main(self, frequency, phase, amplitude): + curve = kf.SinusoidalCurve(frequency=frequency, phase=phase, amplitude=amplitude) + return (curve, curve) + +class KfSinusoidalWithWavelength: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("KEYFRAMED_CURVE", "SINUSOIDAL_CURVE") + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "wavelength": ("FLOAT",{ + "default": 12, + "step": 0.5, + }), + "phase": ("FLOAT", { + "default": 0.0, + #"min": 0.0, + #"max": 6.28318530718, # 2*pi + "step": 0.1308996939, # pi/24 + }), + "amplitude": ("FLOAT",{ + "default": 1, + "step": 0.01, + }), + }, + } + + def main(self, wavelength, phase, amplitude): + curve = kf.SinusoidalCurve(wavelength=wavelength, phase=phase, amplitude=amplitude) + return (curve, curve) + + +### # ### # ### # ### # ### # ### # ### # ### # ### # + + +### # ### # ### # ### # ### # ### # ### # ### # ### # + +class KfSinusoidalAdjustWavelength: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("KEYFRAMED_CURVE", "SINUSOIDAL_CURVE") + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve": ("SINUSOIDAL_CURVE",{"forceInput": True,}), + "adjustment": ("FLOAT",{ + "default": 0.0, + "step": 0.5, + }), + }} + + def main(self, curve, adjustment): + wavelength, phase, amplitude = curve.wavelength, curve.phase, curve.amplitude + wavelength += adjustment + curve = kf.SinusoidalCurve(wavelength=wavelength, phase=phase, amplitude=amplitude) + return (curve, curve) + + +class KfSinusoidalAdjustPhase: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("KEYFRAMED_CURVE", "SINUSOIDAL_CURVE") + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve": ("SINUSOIDAL_CURVE",{"forceInput": True,}), + "adjustment": ("FLOAT", { + "default": 0.0, + "step": 0.1308996939, # pi/24 + }), + }} + + def main(self, curve, adjustment): + wavelength, phase, amplitude = curve.wavelength, curve.phase, curve.amplitude + phase += adjustment + curve = kf.SinusoidalCurve(wavelength=wavelength, phase=phase, amplitude=amplitude) + return (curve, curve) + + +class KfSinusoidalAdjustFrequency: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("KEYFRAMED_CURVE", "SINUSOIDAL_CURVE") + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve": ("SINUSOIDAL_CURVE",{"forceInput": True,}), + "adjustment": ("FLOAT",{ + "default": 0, + "step": 0.01, + }), + }} + + def main(self, curve, adjustment): + wavelength, phase, amplitude = curve.wavelength, curve.phase, curve.amplitude + frequency = 1/wavelength + frequency += adjustment + wavelength = 1/frequency + curve = kf.SinusoidalCurve(wavelength=wavelength, phase=phase, amplitude=amplitude) + return (curve, curve) + + +class KfSinusoidalAdjustAmplitude: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("KEYFRAMED_CURVE", "SINUSOIDAL_CURVE") + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve": ("SINUSOIDAL_CURVE",{"forceInput": True,}), + "adjustment": ("FLOAT",{ + "default": 0, + "step": 0.01, + }), + }} + + def main(self, curve, adjustment): + wavelength, phase, amplitude = curve.wavelength, curve.phase, curve.amplitude + amplitude += adjustment + curve = kf.SinusoidalCurve(wavelength=wavelength, phase=phase, amplitude=amplitude) + return (curve, curve) + + +### # ### # ### # ### # ### # ### # ### # ### # ### # + + +class KfSinusoidalGetWavelength: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("FLOAT",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve": ("SINUSOIDAL_CURVE",{"forceInput": True,}), + }} + + def main(self, curve): + return (curve.wavelength,) + + +class KfSinusoidalGetFrequency: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("FLOAT",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve": ("SINUSOIDAL_CURVE",{"forceInput": True,}), + }} + + def main(self, curve): + return (1/curve.wavelength,) + + +class KfSinusoidalGetPhase: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("FLOAT",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve": ("SINUSOIDAL_CURVE",{"forceInput": True,}), + }} + + def main(self, curve): + return (curve.phase,) + + +class KfSinusoidalGetAmplitude: + CATEGORY = CATEGORY + FUNCTION = "main" + RETURN_TYPES = ("FLOAT",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "curve": ("SINUSOIDAL_CURVE",{"forceInput": True,}), + }} + + def main(self, curve): + return (curve.amplitude,) + +################################################################## + + +# KfScheduleConditions: +# """ +# Carry curve and cond together for Simplicity +# """ +# CATEGORY = CATEGORY +# FUNCTION = "main" +# RETURN_TYPES = ("COND_SCHEDULE",) + + +# KfCombineWeightedConditions: + +################################################################## + +# TODO: 0-1 curves (low frequency oscillators) +# --> "1-X" operator + +# TODO: pre-entangled curves + ################################################################## NODE_CLASS_MAPPINGS = { @@ -359,13 +629,26 @@ NODE_CLASS_MAPPINGS = { "KfEvaluateCurveAtT": KfEvaluateCurveAtT, "KfApplyCurveToCond": KfApplyCurveToCond, "KfConditioningAdd": KfConditioningAdd, + #"KfCurveToAcnLatentKeyframe": KfCurveToAcnLatentKeyframe, + ####################################### #"KfCurveInverse": KfCurveInverse, "KfCurveDraw": KfCurveDraw, "KfCurvesAdd": KfCurvesAdd, "KfCurvesSubtract": KfCurvesSubtract, "KfCurvesMultiply": KfCurvesMultiply, "KfCurvesDivide": KfCurvesDivide, - #"KfCurveToAcnLatentKeyframe": KfCurveToAcnLatentKeyframe, + "KfCurveConstant": KfCurveConstant, + ######################### + "KfSinusoidalWithFrequency": KfSinusoidalWithFrequency, + "KfSinusoidalWithWavelength": KfSinusoidalWithWavelength, + "KfSinusoidalAdjustWavelength": KfSinusoidalAdjustWavelength, + "KfSinusoidalAdjustPhase": KfSinusoidalAdjustPhase, + "KfSinusoidalAdjustFrequency": KfSinusoidalAdjustFrequency, + "KfSinusoidalAdjustAmplitude": KfSinusoidalAdjustAmplitude, + "KfSinusoidalGetWavelength": KfSinusoidalGetWavelength, + "KfSinusoidalGetPhase": KfSinusoidalGetPhase, + "KfSinusoidalGetAmplitude": KfSinusoidalGetAmplitude, + "KfSinusoidalGetFrequency": KfSinusoidalGetFrequency, } # A dictionary that contains the friendly/humanly readable titles for the nodes diff --git a/debug.py b/nodes/debug.py similarity index 100% rename from debug.py rename to nodes/debug.py diff --git a/nodes/entangled.py b/nodes/entangled.py new file mode 100644 index 0000000..6050a0c --- /dev/null +++ b/nodes/entangled.py @@ -0,0 +1,172 @@ +from .core import CATEGORY +import keyframed as kf +import numpy as np + +class KfSinusoidalEntangledZeroOne: + CATEGORY = CATEGORY + "/entangled [0-1]" + FUNCTION = "main" + #RETURN_TYPES = ("KEYFRAMED_CURVE", "SINUSOIDAL_CURVE") + + def main(self, n, **kargs): + tau = 2*np.pi + floor=.001 # fully zeroing out conditions creates "sharp edges" in the traversal + a = 1/n + amplitude = a - floor/2 + + + #return [a+kf.SinusoidalCurve(phase=i/tau, amplitude=a, **kargs) for i in range(n)] + return [amplitude+kf.SinusoidalCurve(phase=tau*(n-i-1)/n, amplitude=amplitude, **kargs) for i in range(n)] + + +class KfSinusoidalEntangledZeroOneFromWavelength(KfSinusoidalEntangledZeroOne): + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "wavelength": ("FLOAT",{ + "default": 12, + "step": 0.5, + }), + } + } + + +class KfSinusoidalEntangledZeroOneFromWavelengthx2(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*2 + def main(self, wavelength): + return super().main(n=2, wavelength=wavelength) + +class KfSinusoidalEntangledZeroOneFromWavelengthx3(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*3 + def main(self, wavelength): + return super().main(n=3, wavelength=wavelength) + +class KfSinusoidalEntangledZeroOneFromWavelengthx4(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*4 + def main(self, wavelength): + return super().main(n=4, wavelength=wavelength) + +class KfSinusoidalEntangledZeroOneFromWavelengthx5(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*5 + def main(self, wavelength): + return super().main(n=5, wavelength=wavelength) + +class KfSinusoidalEntangledZeroOneFromWavelengthx6(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*6 + def main(self, wavelength): + return super().main(n=6, wavelength=wavelength) + +class KfSinusoidalEntangledZeroOneFromWavelengthx7(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*7 + def main(self, wavelength): + return super().main(n=7, wavelength=wavelength) + +class KfSinusoidalEntangledZeroOneFromWavelengthx8(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*8 + def main(self, wavelength): + return super().main(n=8, wavelength=wavelength) + +class KfSinusoidalEntangledZeroOneFromWavelengthx9(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*9 + def main(self, wavelength): + return super().main(n=9, wavelength=wavelength) + + +############################################################################################### + +class KfSinusoidalEntangledZeroOneFromFrequency(KfSinusoidalEntangledZeroOne): + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "frequency": ("FLOAT",{ + "default": 1/12, + "step": 0.01, + }), + } + } + + +class KfSinusoidalEntangledZeroOneFromFrequencyx2(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*2 + def main(self, frequency): + return super().main(n=2, frequency=frequency) + +class KfSinusoidalEntangledZeroOneFromFrequencyx3(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*3 + def main(self, frequency): + return super().main(n=3, frequency=frequency) + +class KfSinusoidalEntangledZeroOneFromFrequencyx4(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*4 + def main(self, frequency): + return super().main(n=4, frequency=frequency) + +class KfSinusoidalEntangledZeroOneFromFrequencyx5(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*5 + def main(self, frequency): + return super().main(n=5, frequency=frequency) + +class KfSinusoidalEntangledZeroOneFromFrequencyx6(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*6 + def main(self, frequency): + return super().main(n=6, frequency=frequency) + +class KfSinusoidalEntangledZeroOneFromFrequencyx7(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*7 + def main(self, frequency): + return super().main(n=7, frequency=frequency) + +class KfSinusoidalEntangledZeroOneFromFrequencyx8(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*8 + def main(self, frequency): + return super().main(n=8, frequency=frequency) + +class KfSinusoidalEntangledZeroOneFromFrequencyx9(KfSinusoidalEntangledZeroOneFromWavelength): + RETURN_TYPES = ("KEYFRAMED_CURVE",)*9 + def main(self, frequency): + return super().main(n=9, frequency=frequency) + + +############################################################################################### + +NODE_CLASS_MAPPINGS = { + "KfSinusoidalEntangledZeroOneFromWavelengthx2": KfSinusoidalEntangledZeroOneFromWavelengthx2, + "KfSinusoidalEntangledZeroOneFromWavelengthx3": KfSinusoidalEntangledZeroOneFromWavelengthx3, + "KfSinusoidalEntangledZeroOneFromWavelengthx4": KfSinusoidalEntangledZeroOneFromWavelengthx4, + "KfSinusoidalEntangledZeroOneFromWavelengthx5": KfSinusoidalEntangledZeroOneFromWavelengthx5, + "KfSinusoidalEntangledZeroOneFromWavelengthx6": KfSinusoidalEntangledZeroOneFromWavelengthx6, + "KfSinusoidalEntangledZeroOneFromWavelengthx7": KfSinusoidalEntangledZeroOneFromWavelengthx7, + "KfSinusoidalEntangledZeroOneFromWavelengthx8": KfSinusoidalEntangledZeroOneFromWavelengthx8, + "KfSinusoidalEntangledZeroOneFromWavelengthx9": KfSinusoidalEntangledZeroOneFromWavelengthx9, + + "KfSinusoidalEntangledZeroOneFromFrequencyx2": KfSinusoidalEntangledZeroOneFromFrequencyx2, + "KfSinusoidalEntangledZeroOneFromFrequencyx3": KfSinusoidalEntangledZeroOneFromFrequencyx3, + "KfSinusoidalEntangledZeroOneFromFrequencyx4": KfSinusoidalEntangledZeroOneFromFrequencyx4, + "KfSinusoidalEntangledZeroOneFromFrequencyx5": KfSinusoidalEntangledZeroOneFromFrequencyx5, + "KfSinusoidalEntangledZeroOneFromFrequencyx6": KfSinusoidalEntangledZeroOneFromFrequencyx6, + "KfSinusoidalEntangledZeroOneFromFrequencyx7": KfSinusoidalEntangledZeroOneFromFrequencyx7, + "KfSinusoidalEntangledZeroOneFromFrequencyx8": KfSinusoidalEntangledZeroOneFromFrequencyx8, + "KfSinusoidalEntangledZeroOneFromFrequencyx9": KfSinusoidalEntangledZeroOneFromFrequencyx9, +} + +# A dictionary that contains the friendly/humanly readable titles for the nodes +NODE_DISPLAY_NAME_MAPPINGS = { + "KfSinusoidalEntangledZeroOneFromWavelengthx2": "2x Entangled Curves [0,1] (Wavelength)", + "KfSinusoidalEntangledZeroOneFromWavelengthx3": "3x Entangled Curves [0,1] (Wavelength)", + "KfSinusoidalEntangledZeroOneFromWavelengthx4": "4x Entangled Curves [0,1] (Wavelength)", + "KfSinusoidalEntangledZeroOneFromWavelengthx5": "5x Entangled Curves [0,1] (Wavelength)", + "KfSinusoidalEntangledZeroOneFromWavelengthx6": "6x Entangled Curves [0,1] (Wavelength)", + "KfSinusoidalEntangledZeroOneFromWavelengthx7": "7x Entangled Curves [0,1] (Wavelength)", + "KfSinusoidalEntangledZeroOneFromWavelengthx8": "8x Entangled Curves [0,1] (Wavelength)", + "KfSinusoidalEntangledZeroOneFromWavelengthx9": "9x Entangled Curves [0,1] (Wavelength)", + + "KfSinusoidalEntangledZeroOneFromFrequencyx2": "2x Entangled Curves [0,1] (Frequency)", + "KfSinusoidalEntangledZeroOneFromFrequencyx3": "3x Entangled Curves [0,1] (Frequency)", + "KfSinusoidalEntangledZeroOneFromFrequencyx4": "4x Entangled Curves [0,1] (Frequency)", + "KfSinusoidalEntangledZeroOneFromFrequencyx5": "5x Entangled Curves [0,1] (Frequency)", + "KfSinusoidalEntangledZeroOneFromFrequencyx6": "6x Entangled Curves [0,1] (Frequency)", + "KfSinusoidalEntangledZeroOneFromFrequencyx7": "7x Entangled Curves [0,1] (Frequency)", + "KfSinusoidalEntangledZeroOneFromFrequencyx8": "8x Entangled Curves [0,1] (Frequency)", + "KfSinusoidalEntangledZeroOneFromFrequencyx9": "9x Entangled Curves [0,1] (Frequency)", +} \ No newline at end of file diff --git a/nodes/schedule.py b/nodes/schedule.py new file mode 100644 index 0000000..46abc16 --- /dev/null +++ b/nodes/schedule.py @@ -0,0 +1,406 @@ +import keyframed as kf +#from keyframed.interpolation import bisect_left_keyframe, bisect_right_keyframe + +from functools import total_ordering +from sortedcontainers import SortedDict, SortedList +from .core import CATEGORY as RootCategory + +import numpy as np +from numbers import Number + +import torch +from copy import deepcopy +import logging + +logging.basicConfig(level=logging.DEBUG, + format='%(asctime)s - %(name)s - %(levelname)s - %(message)s') +logger = logging.getLogger(__name__) + +CATEGORY=RootCategory + "/schedule" + + + + +# @total_ordering +# class ScheduleKeyframe(kf.Keyframe): +# def __lt__(self, other): +# return self.t < other + + +# def update_schedule(schedule, keyframe): +# bl_idx = schedule.bisect_left(keyframe.t) +# try: +# if schedule[bl_idx].t == keyframe.t: +# #del schedule[bl_idx] +# schedule.pop(bl_idx) +# except IndexError: +# pass +# schedule.add(keyframe) +# return schedule + + + + + +# # scavenged from Keyframed... + +# def bisect_left_keyframe(k: Number, curve:SortedList, *args, **kargs) -> ScheduleKeyframe: +# """ +# finds the value of the keyframe in a sorted dictionary to the left of a given key, i.e. performs "previous" interpolation +# """ +# right_index = curve.bisect_right(k) +# left_index = right_index - 1 +# #if right_index > 0: +# if right_index >= 0: +# #_, left_value = self._data.peekitem(left_index) +# left_value = curve[left_index] +# else: +# raise RuntimeError( +# "The return value of bisect_right should always be greater than zero, " +# f"however self._data.bisect_right({k}) returned {right_index}." +# "You should never see this error. Please report the circumstances to the library issue tracker on github." +# ) +# return left_value + + +# def bisect_right_keyframe(k: Number, curve:SortedList, *args, **kargs) -> ScheduleKeyframe: +# """ +# finds the value of the keyframe in a sorted dictionary to the right of a given key, i.e. performs "next" interpolation +# """ +# right_index = curve.bisect_right(k) +# #if right_index > 0: +# if right_index >= 0: +# #_, right_value = curve.peekitem(right_index) +# right_value = curve[right_index] +# else: +# raise RuntimeError( +# "The return value of bisect_right should always be greater than zero, " +# f"however self._data.bisect_right({k}) returned {right_index}." +# "You should never see this error. Please report the circumstances to the library issue tracker on github." +# ) +# return right_value + + + +# schedule = SortedList() + +# x0 = ScheduleKeyframe(t=0, value="a") +# x1 = ScheduleKeyframe(t=5, value="b") +# x2 = ScheduleKeyframe(t=5, value="c") +# x3 = ScheduleKeyframe(t=6, value="d") + + + +# schedule = update_schedule(schedule, x0) +# schedule = update_schedule(schedule, x3) # +# schedule = update_schedule(schedule, x2) +# schedule = update_schedule(schedule, x1) +# schedule + + + +################################################################################### + + + +class KfKeyframedCondition: + """ + Attaches a condition to a keyframe + """ + CATEGORY=CATEGORY + FUNCTION = 'main' + RETURN_TYPES = ("KEYFRAMED_CONDITION",) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "conditioning": ("CONDITIONING", {}), + "time": ("FLOAT", {"default": 0}), + #"weight": ("FLOAT", {"default": 1}), # maybe i should hide this attribute + "interpolation_method": (list(kf.interpolation.INTERPOLATORS.keys()),), + }, + } + + def main(self, conditioning, time, interpolation_method): + #keyframe = kf.Keyframe(t=time, value=weight, interpolation_method=interpolation_method) + #keyframe = ScheduleKeyframe(t=time, value=weight, interpolation_method=interpolation_method) + #return (keyframe, conditioning) + #kf_cond = ScheduleKeyframe(t=time, value=conditioning, interpolation_method=interpolation_method) + #return (kf_cond,) + ########################### + + # separately create keyframes for the parts that need interpolating, and carry around anything else + cond_tensor, cond_dict = conditioning[0] # uh... i have NO idea what to do if there are multiple condition entries here... map over them i guess? + #cond_tensor = deepcopy(cond_tensor) + cond_tensor = cond_tensor.clone() + kf_cond_t = kf.Keyframe(t=time, value=cond_tensor, interpolation_method=interpolation_method) + + cond_pooled = cond_dict.get("pooled_output") + cond_dict = deepcopy(cond_dict) + kf_cond_pooled = None + if cond_pooled is not None: + cond_pooled = cond_pooled.clone() + kf_cond_pooled = kf.Keyframe(t=time, value=cond_pooled, interpolation_method=interpolation_method) + cond_dict["pooled_output"] = cond_pooled + + return ({"kf_cond_t":kf_cond_t, "kf_cond_pooled":kf_cond_pooled, "cond_dict":cond_dict},) + + +class KfSetKeyframe: + CATEGORY=CATEGORY + FUNCTION = 'main' + RETURN_TYPES = ("SCHEDULE",) + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "keyframed_condition": ("KEYFRAMED_CONDITION", {}), + }, + "optional": { + "schedule": ("SCHEDULE", {}), + } + } + def main(self, keyframed_condition, schedule=None): + #keyframe, kf_condition = keyframed_condition + cond_dict = keyframed_condition.pop("cond_dict") + cond_dict = deepcopy(cond_dict) + + if schedule is None: + curve_tokenized = kf.Curve([keyframed_condition["kf_cond_t"]], label="kf_cond_t") + curves = [curve_tokenized] + if keyframed_condition["kf_cond_pooled"] is not None: + curve_pooled = kf.Curve([keyframed_condition["kf_cond_pooled"]], label="kf_cond_pooled") + curves.append(curve_pooled) + schedule = (kf.ParameterGroup(curves), cond_dict) + #schedule = kf.ParameterGroup(keyframed_condition) #parameters + #schedule = SortedList() #SortedDict + #schedule = kf.Curve([keyframed_condition]) + #schedule = kf.Curve({keyframed_condition.t: keyframed_condition}) + else: + schedule, old_cond_dict = schedule + for k, v in keyframed_condition.items(): + if (v is not None): + # for now, assume we already have a schedule for k. + # Not sure how to handle new conditioning type appearing. + schedule.parameters[k][v.t] = v + #schedule[keyframed_condition.t] = keyframed_condition + #schedule._data[keyframed_condition.t] = keyframed_condition + old_cond_dict.update(cond_dict) # NB: mutating this is probably bad + schedule = (schedule, old_cond_dict) + + #schedule = update_schedule(schedule, keyframed_condition) + return (schedule,) + +def evaluate_schedule_at_time(schedule, time): + #schedule, cond_dict = schedule + schedule, cond_dict = schedule + #cond_dict = deepcopy(cond_dict) + values = schedule[time] + cond_t = values.get("kf_cond_t") + cond_pooled = values.get("kf_cond_pooled") + if cond_pooled is not None: + #cond_dict = deepcopy(cond_dict) + cond_dict["pooled_output"] = cond_pooled.clone() + logger.debug(f"type(cond_t):{type(cond_t)}") + logger.debug(f"type(cond_pooled):{type(cond_pooled)}") + logger.debug(f"type(cond_dict):{type(cond_dict)}") + return [(cond_t.clone(), cond_dict)] + + +# def evaluate_schedule_at_time__OLD2(schedule, time): +# kf_cond_left, kf_cond_right = bisect_left_keyframe(time, schedule), bisect_right_keyframe(time, schedule) +# logger.debug(f"kf_cond_left: {kf_cond_left}") +# #return (kf_cond_left.value,) + +# kf_tokenized_left = deepcopy(kf_cond_left) +# #kf_pooled_left = deepcopy(kf_cond_left) +# kf_tokenized_left.value = kf_cond_left.value[0] +# #kf_pooled_left.value = kf_cond_left.value[1].get("pooled_output") + +# kf_tokenized_right = deepcopy(kf_cond_right) +# #kf_pooled_right = deepcopy(kf_cond_right) +# kf_tokenized_right.value = kf_cond_right.value[0] +# #kf_pooled_right.value = kf_cond_right.value[1].get("pooled_output") + +# curve_tokenized = kf.Curve([kf_tokenized_left, kf_tokenized_right]) +# #curve_pooled = kf.Curve([kf_pooled_left, kf_pooled_right]) + +# lerped_tokenized = curve_tokenized[time] +# logger.debug(lerped_tokenized) +# #lerped_pooled = curve_pooled[time] + +# #out_dict = deepcopy(kf_cond_left.value[1]) +# #out_dict["pooled_output"] = lerped_pooled +# out_dict={} + +# return (lerped_tokenized, out_dict) + + +def evaluate_schedule_at_time__OLD(schedule, time): + bl_idx = schedule.bisect_left(time) + logger.debug(f"bl_idx:{bl_idx}") + print(f"bl_idx:{bl_idx}") + #left_kf, left_cond = schedule[bl_idx] + left_kf = schedule[bl_idx] + left_cond = left_kf.value + if left_kf.t == time: # hit time exactly, return + #return (left_kf, left_cond) + return left_cond + if bl_idx == len(schedule): # there's nothing to our right, return + #return (left_kf, left_cond) + return left_cond + #right_kf, right_cond = schedule[bl_idx+1] + right_kf = schedule[bl_idx+1] + right_cond = right_kf.value + logger.info(f"type(right_kf):{type(right_kf)}") + logger.info(f"type(right_cond):{type(right_cond)}") + start, end = left_kf.t, right_kf.t + interval_length = end - start + elapsed = time-start + perc_complete = elapsed / interval_length + # TODO: use interpolation method on keyframe to compute transition weight + # For now, simple lerp + + # TODO: This isn't a proper cond object. need to separately lerp the cond and the pooled output + #lerped_cond = perc_complete * right_cond + (1-perc_complete)*left_cond + right_tokenized, right_dict = right_cond + right_pooled = right_dict.get["pooled_output"] + logger.info(f"type(right_tokenized):{type(right_tokenized)}") + logger.info(f"type(right_pooled):{type(right_pooled)}") + + left_tokenized, left_dict = left_cond + left_pooled = left_dict.get["pooled_output"] + logger.info(f"type(left_tokenized):{type(left_tokenized)}") + logger.info(f"type(left_pooled):{type(left_pooled)}") + + lerped_tokenized = perc_complete * right_tokenized + (1-perc_complete)*left_tokenized + + # TODO: simplify this + if (right_pooled is not None) and (left_pooled is not None): + lerped_pooled = perc_complete * right_pooled + (1-perc_complete)*left_pooled + else: + if right_pooled is not None: + lerped_pooled = perc_complete * right_pooled + elif left_pooled is not None: + lerped_pooled = (1-perc_complete) * left_pooled + logger.info(f"type(lerped_pooled):{type(lerped_pooled)}") + out_dict = deepcopy(left_dict) + if lerped_pooled is not None: + out_dict['pooled_output'] = lerped_pooled + + logger.info("type(lerped_tokenized):{type(lerped_tokenized)}") + # TODO: we could also interpolate and return an associated weight + return (lerped_tokenized, out_dict) + + +class KfGetScheduleConditionAtTime: + CATEGORY=CATEGORY + FUNCTION = 'main' + RETURN_TYPES = ("CONDITIONING",) + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "schedule": ("SCHEDULE",{}), + "time": ("FLOAT",{}), + } + } + + def main(self, schedule, time): + lerped_cond = evaluate_schedule_at_time(schedule, time) + return (lerped_cond,) + + + +class KfGetScheduleConditionSlice: + CATEGORY=CATEGORY + FUNCTION = 'main' + RETURN_TYPES = ("CONDITIONING",) + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "schedule": ("SCHEDULE",{}), + "start": ("FLOAT",{"default":0}), + #"stop": ("FLOAT",{"default":0}), + "step": ("FLOAT",{"default":1}), + "n": ("INT", {"default":24}), + #"endpoint": ("BOOL", {"default":True}) + } + } + + #def main(self, schedule, start, stop, n, endpoint): + def main(self, schedule, start, step, n): + stop = start+n*step + times = np.linspace(start=start, stop=stop, num=n, endpoint=True) + conds = [evaluate_schedule_at_time(schedule, time)[0] for time in times] + lerped_tokenized = [c[0] for c in conds] + lerped_pooled = [c[1]["pooled_output"] for c in conds] + lerped_tokenized_t = torch.cat(lerped_tokenized, dim=0) + logger.info(f"lerped_tokenized_t.shape: {lerped_tokenized_t.shape}") + out_dict = deepcopy(conds[0][1]) + if isinstance(lerped_pooled[0], torch.Tensor) and isinstance(lerped_pooled[-1], torch.Tensor): + out_dict['pooled_output'] = torch.cat(lerped_pooled, dim=0) + return [[(lerped_tokenized_t, out_dict)]] # uh... wrap it in lists until it doesn't complain? + +################################################################### + +NODE_CLASS_MAPPINGS = { + "KfKeyframedCondition": KfKeyframedCondition, + "KfSetKeyframe": KfSetKeyframe, + "KfGetScheduleConditionAtTime": KfGetScheduleConditionAtTime, + "KfGetScheduleConditionSlice": KfGetScheduleConditionSlice, +} + +NODE_DISPLAY_NAME_MAPPINGS = {} + + +################################################################################### + +# class KfSetKeyframe: +# CATEGORY=CATEGORY +# FUNCTION = 'main' +# RETURN_TYPES = ("SCHEDULE",) + +# @classmethod +# def INPUT_TYPES(cls): +# return { +# "required": { +# "keyframed_condition": ("KEYFRAMED_CONDITION", {}), +# }, +# "optional": { +# "schedule": ("SCHEDULE", {}), +# } +# } +# def main(keyframed_condition, schedule=None): +# keyframe, kf_condition = keyframed_condition +# if schedule is None: +# schedule = SortedDict +# schedule[keyframe.t] = keyframed_condition +# return (schedule,) + +# class KfGetScheduleConditionAtTime: +# CATEGORY=CATEGORY +# FUNCTION = 'main' +# RETURN_TYPES = ("KEYFRAME",) + +# @classmethod +# def INPUT_TYPES(cls): +# return { +# "required": { +# "schedule": ("SCHEDULE",{}), +# "time": ("FLOAT",{}), +# } +# } + +# def main(self, schedule, time): + +# # right_index = self._data.bisect_right(k) +# # left_index = right_index - 1 +# # if right_index > 0: +# # _, left_value = self._data.peekitem(left_index) +# # else: \ No newline at end of file