Started some initial scaffolding for prompt scheduling

This commit is contained in:
Jedrzej Kosinski
2024-08-11 22:42:15 -05:00
parent 3c44784a36
commit f909507bd9
2 changed files with 45 additions and 25 deletions
+11 -8
View File
@@ -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
View File
@@ -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] = []