attempted modification for updated keyframed lib

This commit is contained in:
David
2023-12-07 23:28:42 -08:00
parent 9c954acce9
commit 4d08ef2298
+159 -29
View File
@@ -1,35 +1,84 @@
import keyframed as kf import keyframed as kf
#from keyframed.interpolation import bisect_left_keyframe, bisect_right_keyframe
from functools import total_ordering from functools import total_ordering
from sortedcontainers import SortedDict, SortedList from sortedcontainers import SortedDict, SortedList
from .core import CATEGORY as RootCategory from .core import CATEGORY as RootCategory
from numbers import Number
import torch import torch
from copy import deepcopy from copy import deepcopy
import logging import logging
logging.basicConfig(level=logging.INFO, logging.basicConfig(level=logging.DEBUG,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s') format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
CATEGORY=RootCategory + "/schedule" CATEGORY=RootCategory + "/schedule"
@total_ordering
class ScheduleKeyframe(kf.Keyframe):
def __lt__(self, other):
return self.t < other
def update_schedule(schedule, keyframe): # @total_ordering
bl_idx = schedule.bisect_left(keyframe.t) # class ScheduleKeyframe(kf.Keyframe):
try: # def __lt__(self, other):
if schedule[bl_idx].t == keyframe.t: # return self.t < other
#del schedule[bl_idx]
schedule.pop(bl_idx)
except IndexError: # def update_schedule(schedule, keyframe):
pass # bl_idx = schedule.bisect_left(keyframe.t)
schedule.add(keyframe) # try:
return schedule # 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() # schedule = SortedList()
@@ -67,17 +116,34 @@ class KfKeyframedCondition:
"required": { "required": {
"conditioning": ("CONDITIONING", {}), "conditioning": ("CONDITIONING", {}),
"time": ("FLOAT", {"default": 0}), "time": ("FLOAT", {"default": 0}),
"weight": ("FLOAT", {"default": 1}), # maybe i should hide this attribute #"weight": ("FLOAT", {"default": 1}), # maybe i should hide this attribute
"interpolation_method": (list(kf.interpolation.INTERPOLATORS.keys()),), "interpolation_method": (list(kf.interpolation.INTERPOLATORS.keys()),),
}, },
} }
def main(self, conditioning, time, weight, interpolation_method): def main(self, conditioning, time, interpolation_method):
#keyframe = kf.Keyframe(t=time, value=weight, interpolation_method=interpolation_method) #keyframe = kf.Keyframe(t=time, value=weight, interpolation_method=interpolation_method)
#keyframe = ScheduleKeyframe(t=time, value=weight, interpolation_method=interpolation_method) #keyframe = ScheduleKeyframe(t=time, value=weight, interpolation_method=interpolation_method)
#return (keyframe, conditioning) #return (keyframe, conditioning)
kf_cond = ScheduleKeyframe(t=time, value=conditioning, interpolation_method=interpolation_method) #kf_cond = ScheduleKeyframe(t=time, value=conditioning, interpolation_method=interpolation_method)
return (kf_cond,) #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: class KfSetKeyframe:
@@ -97,14 +163,78 @@ class KfSetKeyframe:
} }
def main(self, keyframed_condition, schedule=None): def main(self, keyframed_condition, schedule=None):
#keyframe, kf_condition = keyframed_condition #keyframe, kf_condition = keyframed_condition
cond_dict = keyframed_condition.pop("cond_dict")
if schedule is None: if schedule is None:
schedule = SortedList() #SortedDict curve_tokenized = kf.Curve(keyframed_condition["kf_cond_t"], label="kf_cond_t")
schedule = update_schedule(schedule, keyframed_condition) 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)
#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,) return (schedule,)
def evaluate_schedule_at_time(schedule, time): def evaluate_schedule_at_time(schedule, time):
schedule, cond_dict = schedule
cond_dict = deepcopy(cond_dict)
values = schedule[time]
kf_cond_t = values.pop("kf_cond_t")
kf_cond_pooled = values.pop("kf_cond_pooled")
if kf_cond_pooled is not None:
cond_dict["pooled"] = kf_cond_pooled.clone()
return (kf_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) 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, left_cond = schedule[bl_idx]
left_kf = schedule[bl_idx] left_kf = schedule[bl_idx]
left_cond = left_kf.value left_cond = left_kf.value
@@ -117,8 +247,8 @@ def evaluate_schedule_at_time(schedule, time):
#right_kf, right_cond = schedule[bl_idx+1] #right_kf, right_cond = schedule[bl_idx+1]
right_kf = schedule[bl_idx+1] right_kf = schedule[bl_idx+1]
right_cond = right_kf.value right_cond = right_kf.value
logger.info("type(right_kf):{type(right_kf)}") logger.info(f"type(right_kf):{type(right_kf)}")
logger.info("type(right_cond):{type(right_cond)}") logger.info(f"type(right_cond):{type(right_cond)}")
start, end = left_kf.t, right_kf.t start, end = left_kf.t, right_kf.t
interval_length = end - start interval_length = end - start
elapsed = time-start elapsed = time-start
@@ -130,13 +260,13 @@ def evaluate_schedule_at_time(schedule, time):
#lerped_cond = perc_complete * right_cond + (1-perc_complete)*left_cond #lerped_cond = perc_complete * right_cond + (1-perc_complete)*left_cond
right_tokenized, right_dict = right_cond right_tokenized, right_dict = right_cond
right_pooled = right_dict.get["pooled_output"] right_pooled = right_dict.get["pooled_output"]
logger.info("type(right_tokenized):{type(right_tokenized)}") logger.info(f"type(right_tokenized):{type(right_tokenized)}")
logger.info("type(right_pooled):{type(right_pooled)}") logger.info(f"type(right_pooled):{type(right_pooled)}")
left_tokenized, left_dict = left_cond left_tokenized, left_dict = left_cond
left_pooled = left_dict.get["pooled_output"] left_pooled = left_dict.get["pooled_output"]
logger.info("type(left_tokenized):{type(left_tokenized)}") logger.info(f"type(left_tokenized):{type(left_tokenized)}")
logger.info("type(left_pooled):{type(left_pooled)}") logger.info(f"type(left_pooled):{type(left_pooled)}")
lerped_tokenized = perc_complete * right_tokenized + (1-perc_complete)*left_tokenized lerped_tokenized = perc_complete * right_tokenized + (1-perc_complete)*left_tokenized
@@ -148,7 +278,7 @@ def evaluate_schedule_at_time(schedule, time):
lerped_pooled = perc_complete * right_pooled lerped_pooled = perc_complete * right_pooled
elif left_pooled is not None: elif left_pooled is not None:
lerped_pooled = (1-perc_complete) * left_pooled lerped_pooled = (1-perc_complete) * left_pooled
logger.info("type(lerped_pooled):{type(lerped_pooled)}") logger.info(f"type(lerped_pooled):{type(lerped_pooled)}")
out_dict = deepcopy(left_dict) out_dict = deepcopy(left_dict)
if lerped_pooled is not None: if lerped_pooled is not None:
out_dict['pooled_output'] = lerped_pooled out_dict['pooled_output'] = lerped_pooled