diff --git a/prompt_control/prompts.py b/prompt_control/prompts.py index 018726e..9f34e3c 100644 --- a/prompt_control/prompts.py +++ b/prompt_control/prompts.py @@ -5,7 +5,7 @@ from functools import partial from comfy_extras.nodes_mask import FeatherMask, MaskComposite from nodes import ConditioningAverage -from .utils import safe_float, get_function, parse_floats, smarter_split +from .utils import safe_float, get_function, split_by_function, parse_floats, smarter_split from .adv_encode import advanced_encode_from_tokens from .cutoff import process_cuts from .parser import parse_cuts @@ -231,15 +231,13 @@ def encode_prompt_segment( # Chunks to ConditioningAverage: - text, averages = get_function(text, "AVG", ["0.5"], return_dict=True) - prev = 0 + text, averages = split_by_function(text, "AVG", ["0.5"]) prompts_to_avg = [] for avg in averages: w = safe_float(avg["args"][0], 0.5) - p = text[prev : avg["position"]], w - prompts_to_avg.append(p) - prev = avg["position"] - prompts_to_avg.append((text[prev:], 1.0)) + prompts_to_avg.append(text, w) + text = avg["text"] + prompts_to_avg.append((text, 1.0)) conds_to_avg = [] for prompt, weight in prompts_to_avg: diff --git a/prompt_control/utils.py b/prompt_control/utils.py index 32e58c0..7650fb5 100644 --- a/prompt_control/utils.py +++ b/prompt_control/utils.py @@ -116,14 +116,30 @@ def get_function(text, func, defaults, return_func_name=False, placeholder="", r instances.append(args) if placeholder: - text = text[:start] + f"\0{placeholder}{count}\0" + text[end + 1 :] + text = text[:start] + f"\0{placeholder}{count}\0" + text[end:] else: - text = text[:start] + text[end + 1 :] + text = text[:start] + text[end:] match = rex.search(text) count += 1 return text, instances +def split_by_function(text, func, defaults=None): + """ + Splits a string by function calls, returning the text preceding the first call and a list of dictionaries with a "text" key with the prompt before the next split or until hthe end of the text. + """ + text, functions = get_function(text, func, defaults, return_dict=True) + chunks = [] + prev = 0 + for f in functions: + chunks.append(text[prev : f["position"]]) + prev = f["position"] + chunks.append(text[prev:]) + for i, f in enumerate(functions): + f["text"] = chunks[i + 1] + return chunks[0], functions + + def parse_args(strings, arg_spec, strip=True): args = [s[1] for s in arg_spec] for i, spec in list(enumerate(arg_spec))[: len(strings)]: