fleshing out scheduling
This commit is contained in:
+101
-19
@@ -3,6 +3,13 @@ from functools import total_ordering
|
||||
from sortedcontainers import SortedDict, SortedList
|
||||
from .core import CATEGORY as RootCategory
|
||||
|
||||
import torch
|
||||
from copy import deepcopy
|
||||
import logging
|
||||
|
||||
logging.basicConfig(level=logging.INFO,
|
||||
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
CATEGORY=RootCategory + "/schedule"
|
||||
|
||||
@@ -66,8 +73,11 @@ class KfKeyframedCondition:
|
||||
}
|
||||
|
||||
def main(self, conditioning, time, weight, interpolation_method):
|
||||
keyframe = kf.Keyframe(t=time, value=weight, interpolation_method=interpolation_method)
|
||||
return (keyframe, conditioning)
|
||||
#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,)
|
||||
|
||||
|
||||
class KfSetKeyframe:
|
||||
@@ -85,14 +95,69 @@ class KfSetKeyframe:
|
||||
"schedule": ("SCHEDULE", {}),
|
||||
}
|
||||
}
|
||||
def main(keyframed_condition, schedule=None):
|
||||
def main(self, keyframed_condition, schedule=None):
|
||||
#keyframe, kf_condition = keyframed_condition
|
||||
if schedule is None:
|
||||
schedule = SortedDict
|
||||
schedule = SortedList() #SortedDict
|
||||
schedule = update_schedule(schedule, keyframed_condition)
|
||||
return (schedule,)
|
||||
|
||||
|
||||
def evaluate_schedule_at_time(schedule, time):
|
||||
bl_idx = schedule.bisect_left(time)
|
||||
#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("type(right_kf):{type(right_kf)}")
|
||||
logger.info("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("type(right_tokenized):{type(right_tokenized)}")
|
||||
logger.info("type(right_pooled):{type(right_pooled)}")
|
||||
|
||||
left_tokenized, left_dict = left_cond
|
||||
left_pooled = left_dict.get["pooled_output"]
|
||||
logger.info("type(left_tokenized):{type(left_tokenized)}")
|
||||
logger.info("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("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'
|
||||
@@ -108,24 +173,41 @@ class KfGetScheduleConditionAtTime:
|
||||
}
|
||||
|
||||
def main(self, schedule, time):
|
||||
bl_idx = schedule.bisect_left(time)
|
||||
left_kf, left_cond = schedule[bl_idx]
|
||||
if left_kf.t == time: # hit time exactly, return
|
||||
return (left_kf, left_cond)
|
||||
if bl_idx == len(schedule): # there's nothing to our right, return
|
||||
return (left_kf, left_cond)
|
||||
right_kf, right_cond = schedule[bl_idx+1]
|
||||
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
|
||||
lerped_cond = perc_complete * right_cond + (1-perc_complete)*left_cond
|
||||
# TODO: we could also interpolate and return an associated weight
|
||||
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}),
|
||||
"n": ("INTEGER", {"default":1}),
|
||||
"endpoint": ("BOOL", {"default":True})
|
||||
}
|
||||
}
|
||||
|
||||
def main(self, schedule, start, stop, n, endpoint):
|
||||
times = np.linspace(start=start, stop=stop, num=n, endpoint=endpoint)
|
||||
conds = [evaluate_schedule_at_time(schedule, time) 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)
|
||||
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)
|
||||
return (lerped_conds, out_dict)
|
||||
|
||||
###################################################################
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"KfKeyframedCondition": KfKeyframedCondition,
|
||||
"KfSetKeyframe": KfSetKeyframe,
|
||||
|
||||
Reference in New Issue
Block a user