diff --git a/animatediff/nodes_scheduling.py b/animatediff/nodes_scheduling.py index 1a8d76b..fd62a5c 100644 --- a/animatediff/nodes_scheduling.py +++ b/animatediff/nodes_scheduling.py @@ -9,14 +9,17 @@ class PromptSchedulingLatentsNode: "required": { "prompts": ("STRING", {"multiline": True, "default": ""}), "latent": ("LATENT",), - } + }, + "optional": { + "print_schedule": ("BOOLEAN", {"default": False}), + }, } RETURN_TYPES = ("LATENT",) CATEGORY = "Animate Diff 🎭🅐🅓/scheduling" FUNCTION = "create_schedule" - def create_schedule(self, prompts: str, latent: dict): + def create_schedule(self, prompts: str, latent: dict, print_schedule=False): evaluate_prompt_schedule(prompts, latent["samples"].size(0)) return (latent) @@ -30,7 +33,7 @@ class ValueSchedulingLatentsNode: "latent": ("LATENT",), }, "optional": { - "print_values": ("BOOLEAN", {"default": False}), + "print_schedule": ("BOOLEAN", {"default": False}), }, } @@ -38,10 +41,10 @@ class ValueSchedulingLatentsNode: CATEGORY = "Animate Diff 🎭🅐🅓/scheduling" FUNCTION = "create_schedule" - def create_schedule(self, values: str, latent: dict, print_values=False): + def create_schedule(self, values: str, latent: dict, print_schedule=False): float_vals = evaluate_value_schedule(values, latent["samples"].size(0)) int_vals = [round(x) for x in float_vals] - if print_values: + if print_schedule: for i, val in enumerate(float_vals): logger.info(f"ValueScheduling: {i} = {val}") return (float_vals, float_vals, int_vals, int_vals) @@ -56,7 +59,7 @@ class ValueSchedulingNode: "max_length": ("INT", {"default": 0, "min": 0, "max": BIGMAX, "step": 1}), }, "optional": { - "print_values": ("BOOLEAN", {"default": False}), + "print_schedule": ("BOOLEAN", {"default": False}), }, } @@ -64,10 +67,10 @@ class ValueSchedulingNode: CATEGORY = "Animate Diff 🎭🅐🅓/scheduling" FUNCTION = "create_schedule" - def create_schedule(self, values: str, max_length: int, print_values=False): + def create_schedule(self, values: str, max_length: int, print_schedule=False): float_vals = evaluate_value_schedule(values, max_length) int_vals = [round(x) for x in float_vals] - if print_values: + if print_schedule: for i, val in enumerate(float_vals): logger.info(f"ValueScheduling: {i} = {val}") return (float_vals, float_vals, int_vals, int_vals) diff --git a/animatediff/scheduling.py b/animatediff/scheduling.py index e6bb9f0..08366e1 100644 --- a/animatediff/scheduling.py +++ b/animatediff/scheduling.py @@ -1,5 +1,6 @@ import re import math +from typing import Union from collections import namedtuple from dataclasses import dataclass @@ -65,9 +66,8 @@ class RegexErrorReport: def evaluate_prompt_schedule(text: str, length: int): text = strip_input(text) - # TODO: handle case of no text provided if len(text) == 0: - pass + raise Exception("No text provided to Prompt Scheduling.") # prioritize formats based on best guess to minimize redo's if text.startswith('"'): formats = [SFormat.JSON, SFormat.PYTH] @@ -114,8 +114,25 @@ def evaluate_prompt_schedule(text: str, length: int): error_msg = "\n".join(error_msg_list) raise Exception(error_msg) + def parse_prompt_groups(groups: tuple, length: int): - pass + pairs: list[InputPair] + errors: list[ParseErrorReport] + # perform first parse, to get idea of indexes to handle + pairs, errors = handle_group_idxs(groups, length) + if len(errors) == 0: + # do next step + raise Exception("Looks good.") + if len(errors) > 0: + error_msg_list = [] + issues_formatted = f"{len(errors)} issue{'s' if len(errors)> 1 else ''}" + error_msg_list.append(f"Found {issues_formatted} with idxs:") + for error in errors: + error_msg_list.append(f"{error.idx_str}: {error.reason}") + error_msg = "\n".join(error_msg_list) + raise Exception(error_msg) + final_vals = [] + return final_vals def evaluate_value_schedule(text: str, length: int): @@ -173,20 +190,6 @@ def evaluate_value_schedule(text: str, length: int): raise Exception(error_msg) -@dataclass -class InputPair: - idx: int - val: int - hold: bool = False - end: bool = False - -@dataclass -class ParseErrorReport: - idx_str: str - val_str: str - reason: str - - def parse_value_groups(groups: tuple, length: int): #logger.info(groups) pairs: list[InputPair] @@ -224,6 +227,20 @@ def handle_float_vals(groups: list[tuple]): return actual_pairs, errors +@dataclass +class InputPair: + idx: int + val: Union[int, str] + hold: bool = False + end: bool = False + +@dataclass +class ParseErrorReport: + idx_str: str + val_str: str + reason: str + + def handle_group_idxs(pairs: list[InputPair], length: int): actual_pairs: list[InputPair] = [] errors: list[ParseErrorReport] = []