diff --git a/nodes/schedule.py b/nodes/schedule.py index 64a95f3..9c4dc10 100644 --- a/nodes/schedule.py +++ b/nodes/schedule.py @@ -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,