Started some initial scaffolding for prompt scheduling
This commit is contained in:
@@ -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)
|
||||
|
||||
+34
-17
@@ -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] = []
|
||||
|
||||
Reference in New Issue
Block a user