* Added combinatorial generation.

* Updated pylint settings.
* Fixed a bug with wildcard or extranetwork mappings folders in a different drive than the extension.

Co-authored-by: Copilot <copilot@github.com>
This commit is contained in:
Antonio Cordero Balcazar
2026-05-01 23:02:57 +02:00
co-authored by Copilot
parent ad5cad30ff
commit a31070b986
22 changed files with 1076 additions and 587 deletions
+3
View File
@@ -1,4 +1,5 @@
**/__pycache__
.venv
.vscode/**/*
!.vscode/settings.json
@@ -7,3 +8,5 @@
tests/tests_local.py
tests/local_wildcards
tests/logs
scripts/last_prompts.txt
-12
View File
@@ -59,16 +59,6 @@ single-line-class-stmt=no
single-line-if-stmt=no
[IMPORTS]
allow-any-import-level=
allow-reexport-from-package=no
allow-wildcard-with-all=no
deprecated-modules=
ext-import-graph=
import-graph=
int-import-graph=
known-standard-library=
known-third-party=enchant
preferred-modules=
[LOGGING]
logging-format-style=new
@@ -94,9 +84,7 @@ disable=raw-checker-failed,
missing-module-docstring,
missing-class-docstring,
logging-fstring-interpolation,
import-outside-toplevel,
consider-iterating-dictionary,
wrong-import-position,
unnecessary-lambda,
consider-using-dict-items,
dangerous-default-value,
+1
View File
@@ -0,0 +1 @@
3.10.11
+7
View File
@@ -12,5 +12,12 @@
"python.analysis.typeCheckingMode": "off",
"black-formatter.args": [
"--line-length=120"
],
"python-envs.pythonProjects": [
{
"path": ".",
"envManager": "ms-python.python:venv",
"packageManager": "ms-python.python:pip"
}
]
}
+6 -1
View File
@@ -17,6 +17,7 @@ These are some features:
* Filter content based on the loaded SD model/variant or a variable.
* Map extranetworks (LoRAs) depending on conditions (like the loaded model variant). This allows you to add "virtual" loras to the prompt that will be translated to the correct one.
* Clean up the prompt of unnecessary separators or spaces.
* Combinatorial mode.
Note: when used in an *A1111* compatible webui, the extension must be loaded after any other extension that modifies the prompt (like another wildcards extension). Usually extensions load by their folder name in alphanumeric order, so if the extensions are not loading in the correct order just rename this extension's folder so the ordering works out. When in doubt, just rename this extension's folder with a "z" in front (for example) so that it is the last one to load, or manually set such folder name when installing it.
@@ -75,10 +76,14 @@ See the [syntax documentation](docs/SYNTAX.md).
See the [cookbook](docs/COOKBOOK.md) for interesting usages.
## Contributing
To develop, I suggest creating a virtual environment just for the extension, so the tests work and can be debugged properly.
## License
MIT
## Contact
If you have any questions or concerns, please start a thread in the discussions.
If you have any questions or concerns, please start a thread in the discussions. For bug reports and feature requests open an issue.
+6
View File
@@ -29,6 +29,8 @@ Inputs:
* **process_wildcards**: Activates the wildcard processing.
* **do_cleanup**: Activates the cleanup processing.
* **cleanup_variables**: Do a cleanup of the output variables (depends on do_cleanup).
* **do_combinatorial**: Activates combinatorial mode, where the output are all the combinations of choices/wildcards of the prompt.
* **combinatorial_limit**: Limit for the number of generated combinations.
* **wc_options**: Connection to a Wildcards options node.
* **stn_options**: Connection to a Send-To-Negative options node.
* **cup_options**: Connection to a Cleanup options node.
@@ -42,6 +44,8 @@ Outputs:
* **neg_prompt**: the resulting negative prompt
* **variables**: the dictionary of variables set or echoed.
The outputs are lists, and in combinatorial mode there will be multiple elements that ComfyUI will process sequentially.
### ACB PPP Select Variable node
Lets you extract the variables used from the output (or just one of them). You can use this to send only part of the prompt to, for example, a detailer node. For example:
@@ -127,6 +131,8 @@ Options for extranetworks mapping, in case you want to change them from the defa
* **Unlink seed**: Uses the specified seed for the prompt generation instead of the one from the image. This seed is only used for wildcards and choices.
* **Prompt seed**: The seed to use for the prompt generation. If -1 a random one will be used.
* **Incremental seed**: When using a batch you can use this to set the rest of the prompt seeds with consecutive values.
* **Combinatorial mode**: Generate all possible prompt combinations (from choices and wildcards) and cycle through them to fill the batch.
* **Combinations limit**: Maximum number of combinations to generate (0 = no limit). When generating a batch the limit is automatically raised to at least the batch size.
### General settings
+22
View File
@@ -0,0 +1,22 @@
{
"folders": [
{
"name": "sd-webui-prompt-postprocessor",
"path": "."
}
],
"settings": {
"[python]": {
"editor.defaultFormatter": "ms-python.black-formatter"
},
"yaml.schemaStore.enable": false,
"cSpell.enabledFiletypes": [
"lark"
],
"powershell.cwd": "sd-webui-prompt-postprocessor"
},
"launch": {
"version": "0.2.0",
"configurations": []
}
}
+143 -110
View File
@@ -8,7 +8,7 @@ import lark
import numpy as np
import yaml
from pydantic import ValidationError # pylint: disable=import-error
from pydantic import ValidationError
from ppp_classes import (
FindInFilenamePattern,
HostConfig,
@@ -21,7 +21,7 @@ from ppp_classes import (
PPPInterrupt,
PPPState,
PPPStateOptions,
) # pylint: disable=import-error
)
from ppp_variables import VariableRepository
from ppp_logging import DEBUG_LEVEL, log
from ppp_tree import TreeProcessor
@@ -83,6 +83,8 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
DEFAULT_CUP_MERGE_ATTENTION = defopt["cup_merge_attention"]
DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS = defopt["cup_remove_extranetwork_tags"]
DEFAULT_STRICT_OPERATORS = defopt["strict_operators"]
DEFAULT_DO_COMBINATORIAL = defopt["do_combinatorial"]
DEFAULT_COMBINATORIAL_LIMIT = defopt["combinatorial_limit"]
WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK '
WILDCARD_STOP = "INVALID WILDCARD! {0}\nBREAK "
UNPROCESSED_STOP = "UNPROCESSED CONSTRUCTS!\nBREAK "
@@ -138,7 +140,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
if user_config_file == "":
if self.env_info.get("app", "") == SUPPORTED_APPS.comfyui.value:
try:
import folder_paths # pylint: disable=import-error # type: ignore
import folder_paths # type: ignore
user_dir = folder_paths.get_user_directory()
if user_dir and os.path.isdir(user_dir):
@@ -772,7 +774,115 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
"simple parser without new constructs",
)
def __processprompts(self, rng, prompt, negative_prompt):
def __postprocess_result(
self,
result: tuple[str, list[tuple[str, bool]], tuple[dict[str, str | None], dict[str, str | None]]],
) -> tuple[str, str, dict[str, str | None]]:
variables = {}
unified_prompt, rem_wildcards, (_, echoed_variables_snapshot) = result
# Split the unified prompt back into prompt and negative prompt
split_parts = unified_prompt.split("\x1d", 1)
prompt = split_parts[0]
negative_prompt = split_parts[1] if len(split_parts) > 1 else ""
# Clean up
prompt = self.__cleanup(prompt, 1)
negative_prompt = self.__cleanup(negative_prompt, -1)
self.log(logging.INFO, f"Result prompt: {prompt}")
self.log(logging.INFO, f"Result negative_prompt: {negative_prompt}")
try:
# Get and clean variables
var_keys = sorted(echoed_variables_snapshot.keys())
for k in var_keys:
ev = echoed_variables_snapshot.get(k)
variables[k] = self.__cleanup(ev, 0) if self.state.options.cup_cleanup_variables else ev
self.log(logging.DEBUG, f"Result variables: {variables}")
# Result checks
warnings = []
# Check for special character sequences that should not be in the result
compound_prompt = prompt + "\n" + negative_prompt
found_sequences = re.findall(r"::|\$\$|\$\{|[{}]", compound_prompt)
if found_sequences:
s = ", ".join(map(lambda x: '"' + x + '"', set(found_sequences)))
warnings.append(f"Probably invalid character sequences: {s}.")
# Check for correctly nested parentheses and brackets
stack = []
prev_char = ""
for char in compound_prompt:
if prev_char != "\\":
if char in "([": # opening characters
stack.append(char)
elif char in ")]": # closing characters
if not stack:
warnings.append(f"Unmatched '{char}' character.")
break
last_open = stack.pop()
if (last_open == "(" and char != ")") or (last_open == "[" and char != "]"):
warnings.append(f"Mismatched '{last_open}' and '{char}' characters.")
break
prev_char = char
else:
prev_char = "" # reset prev_char to avoid treating escaped characters as escapes
if stack:
warnings.append(f"Unmatched '{''.join(stack)}' characters.")
if warnings:
self.log(
logging.WARNING,
"Found some weird things in the result. Something might be wrong!\n"
+ "\n".join(f" - {w}" for w in warnings),
)
# Check for wildcards not processed
if rem_wildcards:
w_found_p = [wc for wc, n in rem_wildcards if not n]
w_found_n = [wc for wc, n in rem_wildcards if n]
if self.state.options.if_wildcards == IFWILDCARDS_CHOICES.stop:
self.log(logging.ERROR, "Found unprocessed wildcards!")
else:
self.log(logging.INFO, "Found unprocessed wildcards.")
ppwl = ", ".join(w_found_p)
npwl = ", ".join(w_found_n)
if ppwl:
self.log(logging.ERROR, f"In the prompt: {ppwl}")
if npwl:
self.log(logging.ERROR, f"In the negative prompt: {npwl}")
if self.state.options.if_wildcards == IFWILDCARDS_CHOICES.warn:
prompt = self.WILDCARD_WARNING + prompt
elif self.state.options.if_wildcards == IFWILDCARDS_CHOICES.stop:
raise PPPInterrupt(
"Found unprocessed wildcards!",
self.WILDCARD_STOP.format(ppwl) if ppwl else "",
self.WILDCARD_STOP.format(npwl) if npwl else "",
)
# Check for constructs not processed due to parsing problems
ppp_in_prompt = prompt.find("<ppp:") >= 0
ppp_in_negative_prompt = negative_prompt.find("<ppp:") >= 0
if ppp_in_prompt or ppp_in_negative_prompt:
raise PPPInterrupt(
"Found unprocessed constructs!",
self.UNPROCESSED_STOP if ppp_in_prompt else "",
self.UNPROCESSED_STOP if ppp_in_negative_prompt else "",
)
except PPPInterrupt as e:
self.log(logging.ERROR, e.message)
if e.pos_prefix:
prompt = e.pos_prefix + prompt
if e.neg_prefix:
negative_prompt = e.neg_prefix + negative_prompt
self.log(logging.ERROR, "Interrupting!")
self.interrupt()
v = self.state.variables.get_all_system()
v.update(variables)
return prompt, negative_prompt, v
def __processprompts(self, rng, prompt, negative_prompt) -> list[tuple[str, str, dict[str, str | None]]]:
"""
Process the prompt and negative prompt.
@@ -782,16 +892,15 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
negative_prompt (str): The negative prompt.
Returns:
tuple: A tuple containing the processed prompt and negative prompt.
list: A list of tuples, each containing the processed prompt, negative prompt, and all variables.
"""
self.state.variables.clear_user()
self.state.variables.clear_echoed()
all_variables = self.state.variables.get_all_system()
# Parse both prompts
processor = TreeProcessor(self.state, rng)
unified_prompt = prompt + "\x1D" + negative_prompt
(prompt_parser, parser_description) = self.__get_best_parser(unified_prompt)
unified_prompt = prompt + "\x1d" + negative_prompt
prompt_parser, parser_description = self.__get_best_parser(unified_prompt)
self.log(logging.DEBUG, f"Using {parser_description} for prompt")
parsed = parse_prompt(
self.state,
@@ -801,100 +910,28 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
)
# Process the unified prompt
unified_prompt, rem_wildcards = processor.start_visit(parsed)
t1 = time.monotonic_ns()
try:
results = processor.start_visit(parsed)
except PPPInterrupt as e:
self.log(logging.ERROR, e.message)
if e.pos_prefix:
prompt = e.pos_prefix + prompt
if e.neg_prefix:
negative_prompt = e.neg_prefix + negative_prompt
self.log(logging.ERROR, "Interrupting!")
self.interrupt()
t2 = time.monotonic_ns()
self.log(logging.INFO, f"Visit time: {(t2 - t1) / 1_000_000_000:.3f} seconds")
# Complete variables
var_keys = sorted(self.state.variables.all_user_or_echoed_keys())
for k in var_keys:
ev = self.state.variables.get_echoed_value(k)
if ev is None:
ev = self.state.variables.get_user(k)
if ev is None or not isinstance(ev, str):
self.log(logging.DEBUG, f"Completing variable: {k}")
ev = processor.get_final_variable(k)
all_variables[k] = self.__cleanup(ev, 0) if self.state.options.cup_cleanup_variables else ev
self.log(logging.DEBUG, f"All variables: {all_variables}")
# Split the unified prompt back into prompt and negative prompt
split_parts = unified_prompt.split("\x1D", 1)
prompt = split_parts[0]
negative_prompt = split_parts[1] if len(split_parts) > 1 else ""
# Clean up
prompt = self.__cleanup(prompt, 1)
negative_prompt = self.__cleanup(negative_prompt, -1)
# Result checks
warnings = []
# Check for special character sequences that should not be in the result
compound_prompt = prompt + "\n" + negative_prompt
found_sequences = re.findall(r"::|\$\$|\$\{|[{}]", compound_prompt)
if found_sequences:
warnings.append(
f"Probably invalid character sequences: {', '.join(map(lambda x: '"' + x + '"', set(found_sequences)))}."
)
# Check for correctly nested parentheses and brackets
stack = []
prev_char = ""
for char in compound_prompt:
if prev_char != "\\":
if char in "([": # opening characters
stack.append(char)
elif char in ")]": # closing characters
if not stack:
warnings.append(f"Unmatched '{char}' character.")
break
last_open = stack.pop()
if (last_open == "(" and char != ")") or (last_open == "[" and char != "]"):
warnings.append(f"Mismatched '{last_open}' and '{char}' characters.")
break
prev_char = char
else:
prev_char = "" # reset prev_char to avoid treating escaped characters as escapes
if stack:
warnings.append(f"Unmatched '{''.join(stack)}' characters.")
if warnings:
self.log(
logging.WARNING,
"Found some weird things in the result. Something might be wrong!\n"
+ "\n".join(f" - {w}" for w in warnings),
)
# Check for wildcards not processed
if rem_wildcards:
w_found_p = [wc for wc, n in rem_wildcards if not n]
w_found_n = [wc for wc, n in rem_wildcards if n]
if self.state.options.if_wildcards == IFWILDCARDS_CHOICES.stop:
self.log(logging.ERROR, "Found unprocessed wildcards!")
else:
self.log(logging.INFO, "Found unprocessed wildcards.")
ppwl = ", ".join(w_found_p)
npwl = ", ".join(w_found_n)
if ppwl:
self.log(logging.ERROR, f"In the prompt: {ppwl}")
if npwl:
self.log(logging.ERROR, f"In the negative prompt: {npwl}")
if self.state.options.if_wildcards == IFWILDCARDS_CHOICES.warn:
prompt = self.WILDCARD_WARNING + prompt
elif self.state.options.if_wildcards == IFWILDCARDS_CHOICES.stop:
raise PPPInterrupt(
"Found unprocessed wildcards!",
self.WILDCARD_STOP.format(ppwl) if ppwl else "",
self.WILDCARD_STOP.format(npwl) if npwl else "",
)
# Check for constructs not processed due to parsing problems
ppp_in_prompt = prompt.find("<ppp:") >= 0
ppp_in_negative_prompt = negative_prompt.find("<ppp:") >= 0
if ppp_in_prompt or ppp_in_negative_prompt:
raise PPPInterrupt(
"Found unprocessed constructs!",
self.UNPROCESSED_STOP if ppp_in_prompt else "",
self.UNPROCESSED_STOP if ppp_in_negative_prompt else "",
)
return prompt, negative_prompt, all_variables
final_results = []
for i, r in enumerate(results):
if self.state.options.do_combinatorial:
self.log(logging.INFO, f"Combination {i + 1}:")
final_results.append(self.__postprocess_result(r))
if self.state.options.do_combinatorial:
self.log(logging.INFO, f"Total combinations: {len(final_results)}")
return final_results
def process_prompt(
self,
@@ -913,7 +950,6 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
Returns:
tuple: A tuple containing the processed prompt, negative prompt and all the prompt variables.
"""
all_variables = {}
try:
if seed == -1:
seed = np.random.randint(0, 2**32, dtype=np.int64)
@@ -923,16 +959,13 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
self.log(logging.INFO, f"Input seed: {seed}")
self.log(logging.INFO, f"Input prompt: {prompt}")
self.log(logging.INFO, f"Input negative_prompt: {negative_prompt}")
self.log(logging.INFO, f"Combinatorial: {self.state.options.do_combinatorial}")
t1 = time.monotonic_ns()
prompt, negative_prompt, all_variables = self.__processprompts(
np.random.default_rng(seed & 0xFFFFFFFF), prompt, negative_prompt
)
results = self.__processprompts(np.random.default_rng(seed & 0xFFFFFFFF), prompt, negative_prompt)
t2 = time.monotonic_ns()
self.log(logging.INFO, f"Result prompt: {prompt}")
self.log(logging.INFO, f"Result negative_prompt: {negative_prompt}")
self.log(logging.INFO, f"Process prompt pair time: {(t2 - t1) / 1_000_000_000:.3f} seconds")
# self.log(logging.DEBUG,f"Wildcards memory usage: {self.state.wildcards_obj.__sizeof__()}")
return prompt, negative_prompt, all_variables
return results
except PPPInterrupt as e:
self.log(logging.ERROR, e.message)
if e.pos_prefix:
@@ -941,7 +974,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
negative_prompt = e.neg_prefix + negative_prompt
self.log(logging.ERROR, "Interrupting!")
self.interrupt()
return prompt, negative_prompt, all_variables
return [prompt, negative_prompt, {}]
except Exception: # pylint: disable=broad-exception-caught
self.log(logging.ERROR, "Unexpected error", exc_info=True)
return original_prompt, original_negative_prompt, all_variables
return [original_prompt, original_negative_prompt, {}]
+17
View File
@@ -201,7 +201,24 @@ class PPPStateOptions:
cup_merge_attention: bool = True
cup_remove_extranetwork_tags: bool = False
strict_operators: bool = True
do_combinatorial: bool = False
combinatorial_limit: int = 100 # 0 = no limit
def __post_init__(self):
if not self.cup_do_cleanup:
object.__setattr__(self, "cup_cleanup_variables", False)
object.__setattr__(self, "cup_extra_spaces", False)
object.__setattr__(self, "cup_empty_constructs", False)
object.__setattr__(self, "cup_extra_separators", False)
object.__setattr__(self, "cup_extra_separators2", False)
object.__setattr__(self, "cup_extra_separators_include_eol", False)
object.__setattr__(self, "cup_breaks", False)
object.__setattr__(self, "cup_breaks_eol", False)
object.__setattr__(self, "cup_ands", False)
object.__setattr__(self, "cup_ands_eol", False)
object.__setattr__(self, "cup_extranetwork_tags", False)
object.__setattr__(self, "cup_merge_attention", False)
object.__setattr__(self, "cup_remove_extranetwork_tags", False)
@dataclass(frozen=True)
class PPPState:
+40 -11
View File
@@ -1,8 +1,8 @@
import logging
import os
import folder_paths # pylint: disable=import-error # type: ignore
import nodes # pylint: disable=import-error # type: ignore
import folder_paths # type: ignore
import nodes # type: ignore
from ppp import PromptPostProcessor
from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, SUPPORTED_APPS, PPPStateOptions
@@ -180,6 +180,22 @@ class PromptPostProcessorComfyUINode:
"label_off": "No",
},
),
"do_combinatorial": (
"BOOLEAN",
{
"default": PromptPostProcessor.DEFAULT_DO_COMBINATORIAL,
"tooltip": "Enable combinatorial mode",
"label_on": "Yes",
"label_off": "No",
},
),
"combinatorial_limit": (
"INT",
{
"default": PromptPostProcessor.DEFAULT_COMBINATORIAL_LIMIT,
"tooltip": "Limit for combinatorial mode",
},
),
"wc_options": (
"PPP_OPTIONS_WC",
{
@@ -229,6 +245,11 @@ class PromptPostProcessorComfyUINode:
"STRING",
"PPP_DICT",
)
OUTPUT_IS_LIST = (
True,
True,
True,
)
RETURN_NAMES = (
"pos_prompt",
"neg_prompt",
@@ -255,6 +276,8 @@ class PromptPostProcessorComfyUINode:
process_wildcards,
do_cleanup,
cleanup_variables,
do_combinatorial,
combinatorial_limit,
wc_options=None,
stn_options=None,
cup_options=None,
@@ -292,7 +315,9 @@ class PromptPostProcessorComfyUINode:
options = PPPStateOptions(
debug_level=DEBUG_LEVEL(debug_level),
on_warning=ONWARNING_CHOICES(on_warnings) if on_warnings else PromptPostProcessor.DEFAULT_ON_WARNING,
strict_operators=strict_operators if strict_operators is not None else PromptPostProcessor.DEFAULT_STRICT_OPERATORS,
strict_operators=(
strict_operators if strict_operators is not None else PromptPostProcessor.DEFAULT_STRICT_OPERATORS
),
process_wildcards=process_wildcards,
if_wildcards=(wc_options["wc_if_wildcards"] if wc_options else IFWILDCARDS_CHOICES.stop.value),
choice_separator=(
@@ -345,6 +370,8 @@ class PromptPostProcessorComfyUINode:
if cup_options
else PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS
),
do_combinatorial=do_combinatorial,
combinatorial_limit=combinatorial_limit,
)
self.wildcards_obj.refresh_wildcards(
options.debug_level,
@@ -365,12 +392,14 @@ class PromptPostProcessorComfyUINode:
self.wildcards_obj,
self.extranetwork_mappings_obj,
)
pos_prompt, neg_prompt, variables = ppp.process_prompt(pos_prompt, neg_prompt, seed if seed is not None else 1)
return (
pos_prompt,
neg_prompt,
variables,
)
results = ppp.process_prompt(pos_prompt, neg_prompt, seed if seed is not None else 1)
# pos_prompt, neg_prompt, variables = results[0]
# return (
# pos_prompt,
# neg_prompt,
# variables,
# )
return tuple(zip(*results)) # unzip the list of tuples into tuple of lists
def interrupt(self):
nodes.interrupt_processing(True)
@@ -917,8 +946,8 @@ class PromptPostProcessorWildcardConcatComfyUINode:
try:
from server import PromptServer # type: ignore # pylint: disable=import-error
from aiohttp import web as _aiohttp_web # type: ignore # pylint: disable=import-error
from server import PromptServer # type: ignore
from aiohttp import web as _aiohttp_web # type: ignore
@PromptServer.instance.routes.get("/acb_ppp/wildcards")
async def _acb_ppp_get_wildcards(request):
+2 -1
View File
@@ -1,4 +1,5 @@
import os
from pathlib import Path
from typing import Optional
import logging
import yaml
@@ -102,7 +103,7 @@ class PPPExtraNetworkMappings:
if fullpath != self.LOCALINPUT_FILENAME:
path = os.path.dirname(fullpath)
if not os.path.exists(fullpath) or not any(
os.path.commonpath([path, folder]) == folder for folder in self.__enmappings_folders
Path(path).is_relative_to(folder) for folder in self.__enmappings_folders
):
self.__remove_extranetwork_mappings_from_path(fullpath)
elif enmappings_input is None:
+272 -140
View File
@@ -1,10 +1,11 @@
from collections import namedtuple
from itertools import combinations, combinations_with_replacement, permutations, product
import logging
import math
import re
import textwrap
import time
from typing import Optional
from typing import Any, Optional
import lark
import numpy as np
@@ -31,7 +32,7 @@ class TreeProcessor(lark.visitors.Interpreter):
result (str): The final processed prompt.
"""
NEGATIVE_SEP = "\x1D"
NEGATIVE_SEP = "\x1d"
AccumulatedShell = namedtuple("AccumulatedShell", ["type", "data"])
NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"])
@@ -46,10 +47,12 @@ class TreeProcessor(lark.visitors.Interpreter):
self.__is_negative = False
self.__wildcard_filters = {}
self.__seen_wildcards: list[str] = []
self.add_at: dict[str, list] = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []}
self.insertion_at: list[tuple[int, int]] = [None for _ in range(10)]
self.detectedWildcards: list[tuple[str,bool]] = []
self.result = ""
self.__add_at: dict[str, list] = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []}
self.__insertion_at: list[tuple[int, int]] = [None for _ in range(10)]
self.__detectedWildcards: list[tuple[str, bool]] = []
self.__result = ""
self.__comb_forced_path: list[int] = []
self.__comb_trace: list[int] = []
def log(self, kind, message: str, min_level: DEBUG_LEVEL | None = None):
log(self.state.logger, self.state.options.debug_level, kind, message, min_level)
@@ -57,10 +60,25 @@ class TreeProcessor(lark.visitors.Interpreter):
def warn_or_stop(self, message: str, e: Exception = None):
warn_or_stop(self.state, self.__is_negative, message, e)
def __reset_run_state(self):
"""Reset all per-run mutable state for a fresh combinatorial pass."""
self.__shell = []
self.__negtags = []
self.__already_processed = []
self.__is_negative = False
self.__wildcard_filters = {}
self.__seen_wildcards = []
self.__add_at = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []}
self.__insertion_at = [None for _ in range(10)]
self.__detectedWildcards = []
self.__result = ""
if self.state.extranetwork_mappings_obj is not None:
self.state.extranetwork_mappings_obj.cached_mappings.clear()
def start_visit(
self,
parsed: lark.Tree,
) -> tuple[str, list[tuple[str,bool]]]:
) -> list[tuple[str, list[tuple[str, bool]], tuple[dict[str, Any], dict[str, str]]]]:
"""
Process the positive and negative prompts in a unified way using the same processor.
STN insertions are applied to the negative result directly inside this processor.
@@ -69,17 +87,78 @@ class TreeProcessor(lark.visitors.Interpreter):
parsed (Tree): The parsed unified prompt.
Returns:
tuple[str, list[tuple[str,bool]]]: The processed prompt and its detected wildcards.
list[tuple[str, list[tuple[str,bool]], tuple[dict[str, Any], dict[str, str]]]]: A list of
(processed prompt, detected wildcards, variables snapshot) triples — one entry per
combination in combinatorial mode, or a single entry otherwise. The variables snapshot
is the return value of ``state.variables.backup_user_and_echoed()`` captured after processing.
"""
t1 = time.monotonic_ns()
self.log(logging.INFO, f"Processing prompt...")
self.detectedWildcards = []
self.log(logging.INFO, "Processing prompt...")
self.__detectedWildcards = []
self.__is_negative = False
self.result = ""
self.visit(parsed)
t2 = time.monotonic_ns()
self.log(logging.INFO, f"Process prompt time: {(t2 - t1) / 1_000_000_000:.3f} seconds")
return self.result, self.detectedWildcards
self.__result = ""
if not self.state.options.do_combinatorial:
self.visit(parsed)
self.__finalize_echoed_variables()
return [(self.__result, self.__detectedWildcards, self.state.variables.backup_user_and_echoed())]
# Combinatorial mode: explore every possible path through choices and wildcards via DFS.
# __comb_forced_path drives which option is selected at each decision point;
# __comb_trace records how many options were available at each point so the DFS can
# correctly enumerate unexplored branches after each run.
initial_vars = self.state.variables.backup_user_and_echoed()
results: list[tuple[str, list[tuple[str, bool]], tuple]] = []
limit = self.state.options.combinatorial_limit
def _run(forced_path: tuple[int, ...]) -> tuple[int, ...]:
self.log(logging.DEBUG, f"Running combinatorial path: {forced_path}")
self.__comb_forced_path = list(forced_path)
self.__comb_trace = []
self.__reset_run_state()
self.state.variables.restore_user_and_echoed(initial_vars)
self.visit(parsed)
self.__finalize_echoed_variables()
results.append(
(self.__result, self.__detectedWildcards.copy(), self.state.variables.backup_user_and_echoed())
)
return tuple(self.__comb_trace)
def _dfs(forced_path: tuple[int, ...]):
if 0 < limit <= len(results):
return
trace = _run(forced_path)
# For each decision that was reached but not forced, spawn branches for all
# options beyond the default (index 0).
# Iterate in reverse so later decisions vary fastest, producing lexicographic order.
for i in range(len(trace) - 1, len(forced_path) - 1, -1):
if 0 < limit <= len(results):
return
num_options = trace[i]
for opt in range(1, num_options):
if 0 < limit <= len(results):
return
# Pad with zeros for intermediate decisions so they keep the default.
new_path = forced_path + (0,) * (i - len(forced_path)) + (opt,)
_dfs(new_path)
_dfs(())
if 0 < limit <= len(results):
self.log(
logging.WARNING, f"Combinatorial limit of {limit} reached; some combinations may have been skipped."
)
return results
def __finalize_echoed_variables(self):
var_keys = self.state.variables.all_user_or_echoed_keys()
for k in var_keys:
ev = self.state.variables.get_echoed_value(k)
if ev is None:
ev = self.state.variables.get_user(k)
if ev is None or ev.__class__ != str: # strict check to avoid problems with Tokens
self.log(logging.DEBUG, f"Completing variable: {k}")
ev = self.get_final_variable(k)
self.state.variables.echo(k, ev) # ensure all variables are echoed so they are included in the snapshot
def __visit(
self,
@@ -98,16 +177,16 @@ class TreeProcessor(lark.visitors.Interpreter):
Returns:
str: The result of the visit.
"""
backup_result = self.result
backup_result = self.__result
# self.log(logging.DEBUG, f"Visiting node {node}.")
if restore_state:
# self.log(logging.DEBUG, "Backing up state before visiting.")
backup_shell = self.__shell.copy()
backup_negtags = self.__negtags.copy()
backup_already_processed = self.__already_processed.copy()
backup_add_at = self.add_at.copy()
backup_insertion_at = self.insertion_at.copy()
backup_detectedwildcards = self.detectedWildcards.copy()
backup_add_at = self.__add_at.copy()
backup_insertion_at = self.__insertion_at.copy()
backup_detectedwildcards = self.__detectedWildcards.copy()
backup_vars = self.state.variables.backup_user_and_echoed()
if node is not None:
if isinstance(node, list):
@@ -116,22 +195,22 @@ class TreeProcessor(lark.visitors.Interpreter):
elif isinstance(node, lark.Tree):
self.visit(node)
elif isinstance(node, lark.Token):
self.result += node
self.__result += node
len_backup = len(backup_result)
# if self.result[:len_backup] == backup_result: # this is only necessary if we call parse_prompt with a parser from "start", because it resets the result
added_result = self.result[len_backup:]
added_result = self.__result[len_backup:]
# else:
# added_result = self.result
if discard_content or restore_state:
self.result = backup_result
self.__result = backup_result
if restore_state:
# self.log(logging.DEBUG, "Restoring state after visiting.")
self.__shell = backup_shell
self.__negtags = backup_negtags
self.__already_processed = backup_already_processed
self.add_at = backup_add_at
self.insertion_at = backup_insertion_at
self.detectedWildcards = backup_detectedwildcards
self.__add_at = backup_add_at
self.__insertion_at = backup_insertion_at
self.__detectedWildcards = backup_detectedwildcards
self.state.variables.restore_user_and_echoed(backup_vars)
return added_result
@@ -162,7 +241,7 @@ class TreeProcessor(lark.visitors.Interpreter):
return None, None, None
if specifier == "#": # special value to indicate length of the array variable
return None, None, True
if specifier.startswith("&"): # special value to indicate a separator
if specifier.startswith("&"): # special value to indicate a separator
return None, specifier[2:-1], False
if not specifier.isdecimal():
# bare identifier: resolve as variable
@@ -199,7 +278,7 @@ class TreeProcessor(lark.visitors.Interpreter):
elif isinstance(v, lark.Token):
v = str(v)
if visit and not visited:
self.result += v
self.__result += v
return v
v = self.state.variables.get(name)
@@ -212,7 +291,7 @@ class TreeProcessor(lark.visitors.Interpreter):
if cnt:
v = len(v)
if visit:
self.result += str(v)
self.__result += str(v)
elif idx is not None:
if 0 <= idx < len(v):
v = visit_value(v[idx])
@@ -225,7 +304,7 @@ class TreeProcessor(lark.visitors.Interpreter):
for i, item in enumerate(v):
v2.append(visit_value(item))
if visit and i < len(v) - 1:
self.result += sep
self.__result += sep
v = v2
else:
v = None # error
@@ -283,7 +362,7 @@ class TreeProcessor(lark.visitors.Interpreter):
"""
if self.__debug_level == DEBUG_LEVEL.full:
info = f"({info}) " if info is not None and info != "" else ""
output = self.result[len(start_result) :]
output = self.__result[len(start_result) :]
if output != "":
output = f" >> '{escape_single_quotes(output)}'"
self.log(logging.DEBUG, f"TreeProcessor.{construct} {info}({duration / 1_000_000_000:.3f} seconds){output}")
@@ -333,7 +412,9 @@ class TreeProcessor(lark.visitors.Interpreter):
(bool, list),
]
if not any(isinstance(operand1, t1) and isinstance(operand2, t2) for t1, t2 in compatible_types):
self.warn_or_stop(f"Mixed type values ({type(operand1).__name__}, {type(operand2).__name__}) used in comparison: '{escape_single_quotes(desc)}'")
self.warn_or_stop(
f"Mixed type values ({type(operand1).__name__}, {type(operand2).__name__}) used in comparison: '{escape_single_quotes(desc)}'"
)
return False
return operation(operand1, operand2)
@@ -696,10 +777,10 @@ class TreeProcessor(lark.visitors.Interpreter):
"""
Process a negative prompt separator in the tree.
"""
start_result = self.result
start_result = self.__result
t1 = time.monotonic_ns()
x = tree.children[0]
self.result += x.value
self.__result += x.value
self.__is_negative = True
t2 = time.monotonic_ns()
self.__debug_end("negative_sep", start_result, t2 - t1)
@@ -708,7 +789,7 @@ class TreeProcessor(lark.visitors.Interpreter):
"""
Process a prompt composition construct in the tree.
"""
start_result = self.result
start_result = self.__result
t1 = time.monotonic_ns()
self.__visit(tree.children[0])
and_processing = self.state.host_config.and_
@@ -719,11 +800,11 @@ class TreeProcessor(lark.visitors.Interpreter):
"remove": ("removed", " "),
}
if tree.children[1] is not None:
self.result += f":{tree.children[1]}"
self.__result += f":{tree.children[1]}"
for i in range(2, len(tree.children), 3):
if and_processing in and_replacements.keys():
self.result = (
self.result.rstrip()
self.__result = (
self.__result.rstrip()
+ and_replacements[and_processing][1]
+ self.__visit(tree.children[i + 1], False, True).lstrip()
)
@@ -732,18 +813,20 @@ class TreeProcessor(lark.visitors.Interpreter):
self.warn_or_stop("AND constructs are not allowed!")
else: # and_processing == "ok":
if self.state.options.cup_ands:
self.result = re.sub(r"[, ]+$", "\n" if self.state.options.cup_ands_eol else " ", self.result)
if self.result[-1:].isalnum(): # add space if needed
self.result += " "
self.result += "AND"
self.__result = re.sub(
r"[, ]+$", "\n" if self.state.options.cup_ands_eol else " ", self.__result
)
if self.__result[-1:].isalnum(): # add space if needed
self.__result += " "
self.__result += "AND"
added_result = self.__visit(tree.children[i + 1], False, True)
if self.state.options.cup_ands:
added_result = re.sub(r"^[, ]+", " ", added_result)
if added_result[0:1].isalnum(): # add space if needed
added_result = " " + added_result
self.result += added_result
self.__result += added_result
if tree.children[i + 2] is not None:
self.result += f":{tree.children[i+2]}"
self.__result += f":{tree.children[i+2]}"
t2 = time.monotonic_ns()
self.__debug_end("promptcomp", start_result, t2 - t1)
@@ -751,7 +834,7 @@ class TreeProcessor(lark.visitors.Interpreter):
"""
Process a scheduling construct in the tree and add it to the accumulated shell.
"""
start_result = self.result
start_result = self.__result
t1 = time.monotonic_ns()
before = tree.children[0]
after = tree.children[-2]
@@ -780,7 +863,7 @@ class TreeProcessor(lark.visitors.Interpreter):
self.warn_or_stop("Scheduling constructs are not allowed!")
else: # scheduling_processing == "ok"
# self.__shell.append(TreeProcessor.AccumulatedShell("sc", pos))
self.result += "["
self.__result += "["
if before is not None:
self.log(logging.DEBUG, f"Shell scheduled before with position {pos}")
self.__shell.append(TreeProcessor.AccumulatedShell("scb", pos))
@@ -788,15 +871,15 @@ class TreeProcessor(lark.visitors.Interpreter):
self.__shell.pop()
self.log(logging.DEBUG, f"Shell scheduled after with position {pos}")
self.__shell.append(TreeProcessor.AccumulatedShell("sca", pos))
self.result += ":"
self.__result += ":"
self.__visit(after)
self.__shell.pop()
if self.state.options.cup_empty_constructs and re.fullmatch(
re.escape(start_result) + r"\[:\s*", self.result
re.escape(start_result) + r"\[:\s*", self.__result
):
self.result = start_result
self.__result = start_result
else:
self.result += f":{pos_str}]"
self.__result += f":{pos_str}]"
# self.__shell.pop()
t2 = time.monotonic_ns()
self.__debug_end("scheduled", start_result, t2 - t1, pos_str)
@@ -805,7 +888,7 @@ class TreeProcessor(lark.visitors.Interpreter):
"""
Process an alternation construct in the tree and add it to the accumulated shell.
"""
start_result = self.result
start_result = self.__result
t1 = time.monotonic_ns()
alternation_processing = self.state.host_config.alternation
if alternation_processing == "first":
@@ -817,19 +900,19 @@ class TreeProcessor(lark.visitors.Interpreter):
self.warn_or_stop("Alternation constructs are not allowed!")
else: # alternation_processing == "ok"
# self.__shell.append(TreeProcessor.AccumulatedShell("al", len(tree.children)))
self.result += "["
self.__result += "["
for i, opt in enumerate(tree.children):
self.log(logging.DEBUG, f"Shell alternate option {i+1}")
self.__shell.append(TreeProcessor.AccumulatedShell("alo", {"pos": i + 1, "len": len(tree.children)}))
if i > 0:
self.result += "|"
self.__result += "|"
self.__visit(opt)
self.__shell.pop()
self.result += "]"
self.__result += "]"
if self.state.options.cup_empty_constructs and re.fullmatch(
re.escape(start_result) + r"\[\s*\]", self.result
re.escape(start_result) + r"\[\s*\]", self.__result
):
self.result = start_result
self.__result = start_result
# self.__shell.pop()
t2 = time.monotonic_ns()
self.__debug_end("alternate", start_result, t2 - t1)
@@ -838,7 +921,7 @@ class TreeProcessor(lark.visitors.Interpreter):
"""
Process a attention change construct in the tree and add it to the accumulated shell.
"""
start_result = self.result
start_result = self.__result
t1 = time.monotonic_ns()
# weight_kind: -1: remove, 0=none, 1=decrease, 2=increase, 3=specific
if len(tree.children) == 2:
@@ -902,25 +985,25 @@ class TreeProcessor(lark.visitors.Interpreter):
self.__shell.append(TreeProcessor.AccumulatedShell("at", (weight_kind, weight_str)))
if weight_kind == 1:
starttag = "["
self.result += starttag
self.__result += starttag
self.__visit(current_tree)
endtag = "]"
elif weight_kind == 2:
starttag = "("
self.result += starttag
self.__result += starttag
self.__visit(current_tree)
endtag = ")"
else: # weight_kind == 3
starttag = "("
self.result += starttag
self.__result += starttag
self.__visit(current_tree)
endtag = f":{weight_str})"
if self.state.options.cup_empty_constructs and re.fullmatch(
re.escape(start_result + starttag) + r"\s*", self.result
re.escape(start_result + starttag) + r"\s*", self.__result
):
self.result = start_result
self.__result = start_result
else:
self.result += endtag
self.__result += endtag
self.__shell.pop()
t2 = time.monotonic_ns()
self.__debug_end("attention", start_result, t2 - t1, weight_str)
@@ -929,7 +1012,7 @@ class TreeProcessor(lark.visitors.Interpreter):
"""
Process a send to negative command in the tree and add it to the list of negative tags.
"""
start_result = self.result
start_result = self.__result
info = None
t1 = time.monotonic_ns()
if not self.__is_negative:
@@ -940,7 +1023,7 @@ class TreeProcessor(lark.visitors.Interpreter):
parameters = ""
content = self.__visit(tree.children[1::], False, True)
self.__negtags.append(
TreeProcessor.NegTag(len(self.result), len(self.result), content, parameters, self.__shell.copy())
TreeProcessor.NegTag(len(self.__result), len(self.__result), content, parameters, self.__shell.copy())
)
info = f"with {escape_single_quotes(parameters) or 'no parameters'} : {escape_single_quotes(content)}"
else:
@@ -953,7 +1036,7 @@ class TreeProcessor(lark.visitors.Interpreter):
"""
Process a send to negative insertion point command in the tree and add it to the list of negative tags.
"""
start_result = self.result
start_result = self.__result
info = None
t1 = time.monotonic_ns()
if self.__is_negative:
@@ -962,7 +1045,9 @@ class TreeProcessor(lark.visitors.Interpreter):
parameters = str(negtagparameters)
else:
parameters = ""
self.__negtags.append(TreeProcessor.NegTag(len(self.result), len(self.result), "", parameters, self.__shell.copy()))
self.__negtags.append(
TreeProcessor.NegTag(len(self.__result), len(self.__result), "", parameters, self.__shell.copy())
)
info = f"with {parameters or 'no parameters'}"
else:
self.warn_or_stop("Ignored negative insertion point command in positive prompt")
@@ -981,7 +1066,7 @@ class TreeProcessor(lark.visitors.Interpreter):
Process a generic set command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
start_result = self.__result
if self.state.variables.name_is_system(variable_name):
self.warn_or_stop(
f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! System variables cannot be set."
@@ -1069,7 +1154,9 @@ class TreeProcessor(lark.visitors.Interpreter):
if is_starred:
if access_full_array and isinstance(newvalue.children[0], lark.Tree):
if newvalue.children[0].data == "vardescriptor_get":
vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor(newvalue.children[0])
vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor(
newvalue.children[0]
)
if vardescriptor_specifier is not None:
newvalue = None
else:
@@ -1079,9 +1166,9 @@ class TreeProcessor(lark.visitors.Interpreter):
self.__resolve_operand(c) for c in self.__get_cond_operand(newvalue.children[0])
)
elif newvalue.children[0].data == "wildcard":
backup_result = self.result
backup_result = self.__result
newvalue = self.__process_wildcard(newvalue.children[0])
self.result = backup_result
self.__result = backup_result
else:
newvalue = None
else:
@@ -1150,7 +1237,7 @@ class TreeProcessor(lark.visitors.Interpreter):
Process a generic echo command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
start_result = self.__result
default_value = None
# if default is not None:
# default_value = self.__visit(default, True) # for log
@@ -1161,7 +1248,7 @@ class TreeProcessor(lark.visitors.Interpreter):
if default is not None:
self.log(logging.DEBUG, f"Variable '{escape_single_quotes(vname)}' not found, using default value")
value = self.__visit(default, False, True)
self.result += value
self.__result += value
default_value = value
else:
self.warn_or_stop(f"Unknown variable {escape_single_quotes(vname)}")
@@ -1206,7 +1293,7 @@ class TreeProcessor(lark.visitors.Interpreter):
Process an if command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
start_result = self.__result
for i, n in enumerate(tree.children):
content = n.children[-1]
if len(n.children) == 2: # its not an else
@@ -1229,7 +1316,7 @@ class TreeProcessor(lark.visitors.Interpreter):
Process an extranetwork command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
start_result = self.__result
extnet = "(ignored)"
if not self.state.options.cup_remove_extranetwork_tags:
extnet_type: str = (tree.children[0].children[0] or "") + str(tree.children[0].children[1])
@@ -1290,12 +1377,23 @@ class TreeProcessor(lark.visitors.Interpreter):
else:
else_mapping = v
if found_mappings:
found = found_mappings[
self.__rng.choice(
len(found_mappings),
p=[v.weight or 1 for v in found_mappings],
if self.state.options.do_combinatorial:
N = len(found_mappings)
decision_idx = len(self.__comb_trace)
self.__comb_trace.append(N)
chosen_idx = (
min(self.__comb_forced_path[decision_idx], N - 1)
if decision_idx < len(self.__comb_forced_path)
else 0
)
]
found = found_mappings[chosen_idx]
else:
found = found_mappings[
self.__rng.choice(
len(found_mappings),
p=[v.weight or 1 for v in found_mappings],
)
]
else:
found = else_mapping
self.state.extranetwork_mappings_obj.cached_mappings[extnet_id] = found
@@ -1348,23 +1446,23 @@ class TreeProcessor(lark.visitors.Interpreter):
self.warn_or_stop(f"Extranetwork mapping '{escape_single_quotes(extnet_id)}' not found!")
if extnet_id:
extnet = f"<{extnet_id}:{parameters}>"
self.result += extnet
self.__result += extnet
elif triggers or compiled_extra_triggers:
extnet = "(only triggers)"
if triggers or compiled_extra_triggers:
if extnet_id:
if not self.state.options.cup_extranetwork_tags:
self.result += " "
self.__result += " "
else:
self.result += ", "
self.__result += ", "
if triggers:
self.result += self.__visit(triggers, True, True)
self.__result += self.__visit(triggers, True, True)
if compiled_extra_triggers:
if triggers:
self.result += ", "
self.result += self.__visit(compiled_extra_triggers, True, True)
self.__result += ", "
self.__result += self.__visit(compiled_extra_triggers, True, True)
if triggers or compiled_extra_triggers:
self.result += ", "
self.__result += ", "
t2 = time.monotonic_ns()
self.__debug_end("commandext", start_result, t2 - t1, extnet)
@@ -1373,7 +1471,7 @@ class TreeProcessor(lark.visitors.Interpreter):
Process a setwcdeffilter (Set Wildcard Default Filter) command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
start_result = self.__result
wildcard_key: str = self.__visit(tree.children[0].children[1], False, True)
selected_wildcards = [x.key for x in self.state.wildcards_obj.get_wildcards(wildcard_key)]
if not selected_wildcards:
@@ -1397,11 +1495,11 @@ class TreeProcessor(lark.visitors.Interpreter):
Process an extra network construct in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
start_result = self.__result
if not self.state.options.cup_remove_extranetwork_tags:
self.result += f"<{tree.children[0]}"
self.__result += f"<{tree.children[0]}"
self.__visit(tree.children[1])
self.result += ">"
self.__result += ">"
t2 = time.monotonic_ns()
self.__debug_end("extranetworktag", start_result, t2 - t1)
@@ -1450,7 +1548,7 @@ class TreeProcessor(lark.visitors.Interpreter):
for i, c in enumerate(filtered_choice_values):
if c.get("command", False):
content_text = self.__visit(c.get("content", ""), False, True).strip()
(cmd, cmd_args) = content_text.split()
cmd, cmd_args = content_text.split()
if cmd == "include":
wcs = self.state.wildcards_obj.get_wildcards(cmd_args)
if not wcs:
@@ -1467,7 +1565,7 @@ class TreeProcessor(lark.visitors.Interpreter):
self.__seen_wildcards.append(wc.key)
self.log(logging.DEBUG, f"Seen wildcard '{escape_single_quotes(wc.key)}'")
self.log(logging.DEBUG, f"Including choices from wildcard '{escape_single_quotes(wc.key)}'")
(_, choice_values) = self.__check_wildcard_initialization(wc)
_, choice_values = self.__check_wildcard_initialization(wc)
if choice_values is not None:
ch_values = self.__get_choices_internal_get(choice_values, None, wc.key)
for cv in ch_values:
@@ -1550,9 +1648,40 @@ class TreeProcessor(lark.visitors.Interpreter):
to_value = 1
elif (to_value > len(available_choices) and not repeating) or from_value > to_value:
to_value = len(available_choices)
num_choices = (
self.__rng.integers(from_value, to_value, endpoint=True) if from_value < to_value else from_value
)
comb_chosen_selection: Optional[list[dict]] = None
if self.state.options.do_combinatorial:
# Enumerate every distinct selection of choices, accounting for count range and repetition.
all_selections: list[tuple] = []
# When keep_choices_order is False the output depends on the selection order,
# so we must enumerate ordered sequences (permutations / product).
# When keep_choices_order is True selections are sorted afterward, so all
# orderings of the same items produce identical output and we only need
# unordered iterators (combinations / combinations_with_replacement).
for k in range(from_value, to_value + 1):
if repeating:
if self.state.options.keep_choices_order:
all_selections.extend(combinations_with_replacement(available_choices, k))
else:
all_selections.extend(product(available_choices, repeat=k))
else:
if self.state.options.keep_choices_order:
all_selections.extend(combinations(available_choices, k))
else:
all_selections.extend(permutations(available_choices, k))
num_selections = len(all_selections)
decision_idx = len(self.__comb_trace)
self.__comb_trace.append(num_selections)
chosen_idx = (
min(self.__comb_forced_path[decision_idx], num_selections - 1)
if decision_idx < len(self.__comb_forced_path)
else 0
)
comb_chosen_selection = list(all_selections[chosen_idx])
num_choices = len(comb_chosen_selection)
else:
num_choices = (
self.__rng.integers(from_value, to_value, endpoint=True) if from_value < to_value else from_value
)
else:
num_choices = 0
if not optional and from_value > 0:
@@ -1566,11 +1695,14 @@ class TreeProcessor(lark.visitors.Interpreter):
+ (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else ""),
)
if num_choices > 0:
selected_choices: list[dict] = (
list(self.__rng.choice(available_choices, size=num_choices, p=weights, replace=repeating))
if available_choices
else []
)
if self.state.options.do_combinatorial and comb_chosen_selection is not None:
selected_choices: list[dict] = comb_chosen_selection
else:
selected_choices: list[dict] = (
list(self.__rng.choice(available_choices, size=num_choices, p=weights, replace=repeating))
if available_choices
else []
)
if self.state.options.keep_choices_order:
selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"])
selected_choices_text = []
@@ -1829,7 +1961,7 @@ class TreeProcessor(lark.visitors.Interpreter):
"""
t1 = time.monotonic_ns()
chosen_choices = []
start_result = self.result
start_result = self.__result
seen_wildcards_len = len(self.__seen_wildcards)
applied_options = self.__clean_wildcard_options(self.__convert_choices_options(tree.children[0], False))
wildcard_key: str = self.__visit(tree.children[1], False, True)
@@ -1838,8 +1970,8 @@ class TreeProcessor(lark.visitors.Interpreter):
self.log(logging.DEBUG, f"Processing wildcard: {wildcard_key}")
selected_wildcards = self.state.wildcards_obj.get_wildcards(wildcard_key)
if not selected_wildcards:
self.detectedWildcards.append((wc, self.__is_negative))
self.result += wc
self.__detectedWildcards.append((wc, self.__is_negative))
self.__result += wc
t2 = time.monotonic_ns()
self.__debug_end("wildcard", start_result, t2 - t1, wc)
return []
@@ -1891,8 +2023,8 @@ class TreeProcessor(lark.visitors.Interpreter):
choice_values_all = []
for wildcard in selected_wildcards:
if wildcard is None:
self.detectedWildcards.append((wc, self.__is_negative))
self.result += wc
self.__detectedWildcards.append((wc, self.__is_negative))
self.__result += wc
t2 = time.monotonic_ns()
self.__debug_end("wildcard", start_result, t2 - t1, wc)
return []
@@ -1903,7 +2035,7 @@ class TreeProcessor(lark.visitors.Interpreter):
continue
self.__seen_wildcards.append(wildcard.key)
self.log(logging.DEBUG, f"Seen wildcard '{escape_single_quotes(wildcard.key)}'")
(options, choice_values) = self.__check_wildcard_initialization(wildcard)
options, choice_values = self.__check_wildcard_initialization(wildcard)
if options is not None:
if applied_options is None:
applied_options = options
@@ -1916,7 +2048,7 @@ class TreeProcessor(lark.visitors.Interpreter):
applied_options, choice_values_all, filter_specifier, wildcard_key
)
if chosen_choices:
self.result += prefix + separator.join(chosen_choices) + suffix
self.__result += prefix + separator.join(chosen_choices) + suffix
if wildcard_key in self.__wildcard_filters:
del self.__wildcard_filters[wildcard_key]
if variablename is not None:
@@ -1924,8 +2056,8 @@ class TreeProcessor(lark.visitors.Interpreter):
if variablebackup is not None:
self.state.variables.set_user(variablename, variablebackup)
elif self.state.options.if_wildcards != IFWILDCARDS_CHOICES.remove:
self.detectedWildcards.append((wc, self.__is_negative))
self.result += wc
self.__detectedWildcards.append((wc, self.__is_negative))
self.__result += wc
if self.__debug_level == DEBUG_LEVEL.full:
list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]
self.log(logging.DEBUG, f"Unseen wildcards: {', '.join(list_unseen)}")
@@ -1959,7 +2091,7 @@ class TreeProcessor(lark.visitors.Interpreter):
Process a choices construct in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
start_result = self.__result
options = self.__convert_choices_options(tree.children[0], False)
choice_values = [self.__convert_choice(c) for c in tree.children[1::]]
ch = self.__get_original_node_content(tree, "?{...}")
@@ -1967,22 +2099,22 @@ class TreeProcessor(lark.visitors.Interpreter):
self.log(logging.DEBUG, "Processing choices:")
prefix, chosen_choices, separator, suffix = self.__get_choices_select(options, choice_values)
if chosen_choices:
self.result += prefix + separator.join(chosen_choices) + suffix
self.__result += prefix + separator.join(chosen_choices) + suffix
elif self.state.options.if_wildcards != IFWILDCARDS_CHOICES.remove:
self.detectedWildcards.append((ch, self.__is_negative))
self.result += ch
self.__detectedWildcards.append((ch, self.__is_negative))
self.__result += ch
t2 = time.monotonic_ns()
self.__debug_end("choices", start_result, t2 - t1, f"'{escape_single_quotes(ch)}'")
def __default__(self, tree):
t1 = time.monotonic_ns()
start_result = self.result
start_result = self.__result
self.__visit(tree.children)
t2 = time.monotonic_ns()
self.__debug_end(tree.data.value, start_result, t2 - t1)
def __process_negtags(self):
# process the found negative tags
# process the found negative tags
for negtag in self.__negtags:
if self.state.options.cup_merge_attention:
# join consecutive attention elements
@@ -2032,49 +2164,49 @@ class TreeProcessor(lark.visitors.Interpreter):
position = negtag.parameters or "s"
if position.startswith("i"):
n = int(position[1])
self.insertion_at[n] = [negtag.start, negtag.end]
self.__insertion_at[n] = [negtag.start, negtag.end]
elif len(content) > 0:
if content not in self.__already_processed:
if self.state.options.stn_ignore_repeats:
self.__already_processed.append(content)
self.log(logging.DEBUG, f"Adding content at position {position}: {content}")
if position == "e":
self.add_at["end"].append(content)
self.__add_at["end"].append(content)
elif position.startswith("p"):
n = int(position[1])
self.add_at["insertion_point"][n].append(content)
self.__add_at["insertion_point"][n].append(content)
else: # position == "s" or invalid
self.add_at["start"].append(content)
self.__add_at["start"].append(content)
else:
self.log(logging.WARNING, f"Ignoring repeated content: {content}")
self.__negtags = []
self.__negtags = []
def __apply_stn_insertions(self):
"""
Apply all accumulated STN content from add_at to self.result using the recorded
insertion_at positions, then reset both so ppp.py does not re-apply them.
"""
pos, neg = self.result.split(self.NEGATIVE_SEP, 1)
pos, neg = self.__result.split(self.NEGATIVE_SEP, 1)
neg_start = len(pos) + len(self.NEGATIVE_SEP)
stn_sep = self.state.options.stn_separator
self.log(logging.DEBUG, f"Applying STN additions to negative: {self.add_at}")
self.log(logging.DEBUG, f"Applying STN indexes: {self.insertion_at}")
self.log(logging.DEBUG, f"Applying STN additions to negative: {self.__add_at}")
self.log(logging.DEBUG, f"Applying STN indexes: {self.__insertion_at}")
ordered_range = sorted(
range(10),
key=lambda x: self.insertion_at[x][0] if self.insertion_at[x] is not None else float("-inf"),
key=lambda x: self.__insertion_at[x][0] if self.__insertion_at[x] is not None else float("-inf"),
reverse=True,
)
for n in ordered_range:
if self.insertion_at[n] is not None:
insertion_point_n: list[str] = self.add_at["insertion_point"][n]
ipp = self.insertion_at[n][0] - neg_start
ipl = self.insertion_at[n][1] - self.insertion_at[n][0]
if self.__insertion_at[n] is not None:
insertion_point_n: list[str] = self.__add_at["insertion_point"][n]
ipp = self.__insertion_at[n][0] - neg_start
ipl = self.__insertion_at[n][1] - self.__insertion_at[n][0]
if neg[ipp - len(stn_sep) : ipp] == stn_sep:
ipp -= len(stn_sep) # adjust for existing start separator
ipp -= len(stn_sep) # adjust for existing start separator
ipl += len(stn_sep)
insertion_point_n.insert(0, neg[:ipp])
if neg[ipp + ipl : ipp + ipl + len(stn_sep)] == stn_sep:
ipl += len(stn_sep) # adjust for existing end separator
ipl += len(stn_sep) # adjust for existing end separator
end_part = neg[ipp + ipl :]
if len(end_part) > 0:
insertion_point_n.append(end_part)
@@ -2083,30 +2215,30 @@ class TreeProcessor(lark.visitors.Interpreter):
ipp = 0
if neg.startswith(stn_sep):
ipp = len(stn_sep)
self.add_at["insertion_point"][n].append(neg[ipp:])
neg = stn_sep.join(self.add_at["insertion_point"][n])
if self.add_at["start"]:
add_at_start = self.add_at["start"]
self.__add_at["insertion_point"][n].append(neg[ipp:])
neg = stn_sep.join(self.__add_at["insertion_point"][n])
if self.__add_at["start"]:
add_at_start = self.__add_at["start"]
if len(neg) > 0:
ipp = 0
if neg.startswith(stn_sep):
ipp = len(stn_sep) # adjust for existing end separator
add_at_start.append(neg[ipp:])
neg = stn_sep.join(add_at_start)
if self.add_at["end"]:
add_at_end = self.add_at["end"]
if self.__add_at["end"]:
add_at_end = self.__add_at["end"]
if len(neg) > 0:
ipl = len(neg)
if neg.endswith(stn_sep):
ipl -= len(stn_sep) # adjust for existing start separator
ipl -= len(stn_sep) # adjust for existing start separator
add_at_end.insert(0, neg[:ipl])
neg = stn_sep.join(add_at_end)
# self.add_at = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []}
# self.insertion_at = [None for _ in range(10)]
self.result = pos + self.NEGATIVE_SEP + neg
self.__result = pos + self.NEGATIVE_SEP + neg
def start(self, tree):
self.result = ""
self.__result = ""
t1 = time.monotonic_ns()
self.__visit(tree.children)
self.__process_negtags()
+4 -12
View File
@@ -1,5 +1,6 @@
import fnmatch
import os
from pathlib import Path
from typing import Optional
import logging
import yaml
@@ -82,19 +83,10 @@ class PPPWildcards:
for fullpath in list(self.__wildcard_files.keys()):
if fullpath != self.LOCALINPUT_FILENAME:
path = os.path.dirname(fullpath)
if not os.path.exists(fullpath):
if not os.path.exists(fullpath) or not any(
Path(path).is_relative_to(folder) for folder in self.__wildcards_folders
):
self.__remove_wildcards_from_path(fullpath)
else:
a = False
for folder in self.__wildcards_folders:
try:
if os.path.commonpath([folder, path]) == folder:
a = True
break
except ValueError:
pass
if not a:
self.__remove_wildcards_from_path(fullpath)
elif wildcards_input is None:
self.__remove_wildcards_from_path(fullpath)
if wildcards_folders is not None or wildcards_input is not None:
+73 -27
View File
@@ -10,11 +10,11 @@ import numpy as np
sys.path.append(str(Path(__file__).parent)) # base path for the extension
from modules import scripts, shared, script_callbacks # pylint: disable=import-error
from modules.processing import StableDiffusionProcessing # pylint: disable=import-error
from modules.shared import opts # pylint: disable=import-error
from modules.paths import models_path # pylint: disable=import-error
import gradio as gr # pylint: disable=import-error
from modules import scripts, shared, script_callbacks # type: ignore
from modules.processing import StableDiffusionProcessing # type: ignore
from modules.shared import opts # type: ignore
from modules.paths import models_path # type: ignore
import gradio as gr # type: ignore
from ppp import PromptPostProcessor
from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, SUPPORTED_APPS, SUPPORTED_APPS_NAMES, PPPStateOptions
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory, log
@@ -140,7 +140,22 @@ class PromptPostProcessorA1111Script(scripts.Script):
# show_label=True,
elem_id="ppp_incremental_seed",
)
return [force_equal_seeds, unlink_seed, seed, incremental_seed]
gr.HTML("<br>")
with gr.Row(equal_height=True):
combinatorial = gr.Checkbox(
label="Combinatorial mode",
info="Generate all prompt combinations and cycle through them to fill the batch.",
value=PromptPostProcessor.DEFAULT_DO_COMBINATORIAL,
elem_id="ppp_combinatorial",
)
combinatorial_limit = gr.Number(
label="Combinations limit (0 = no limit)",
value=PromptPostProcessor.DEFAULT_COMBINATORIAL_LIMIT,
precision=0,
min_width=120,
elem_id="ppp_combinatorial_limit",
)
return [force_equal_seeds, unlink_seed, seed, incremental_seed, combinatorial, combinatorial_limit]
def process(
self,
@@ -149,6 +164,8 @@ class PromptPostProcessorA1111Script(scripts.Script):
input_unlink_seed,
input_seed,
input_incremental_seed,
input_combinatorial,
input_combinatorial_limit,
): # pylint: disable=arguments-differ
"""
Processes the prompts and applies post-processing operations.
@@ -159,6 +176,8 @@ class PromptPostProcessorA1111Script(scripts.Script):
input_unlink_seed (bool): Flag indicating whether to unlink the seed.
input_seed (int): The seed value.
input_incremental_seed (bool): Flag indicating whether to use incremental seed.
input_combinatorial (bool): Flag indicating whether to use combinatorial mode.
input_combinatorial_limit (int): Maximum number of combinations (0 = no limit).
Returns:
None
@@ -176,6 +195,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
)
)
)
num_seeds = len(getattr(p, "all_seeds", []))
options = PPPStateOptions(
debug_level=DEBUG_LEVEL(getattr(opts, "ppp_gen_debug_level", PromptPostProcessor.DEFAULT_DEBUG_LEVEL)),
on_warning=ONWARNING_CHOICES(getattr(opts, "ppp_gen_onwarning", PromptPostProcessor.DEFAULT_ON_WARNING)),
@@ -220,6 +240,8 @@ class PromptPostProcessorA1111Script(scripts.Script):
cup_remove_extranetwork_tags=getattr(
opts, "ppp_rem_removeextranetworktags", PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS
),
do_combinatorial=input_combinatorial,
combinatorial_limit=max(num_seeds, int(input_combinatorial_limit)) if input_combinatorial else 0,
)
if self.ppp_logger is None:
lf = PromptPostProcessorLogFactory()
@@ -251,6 +273,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
"PPP unlink seed": input_unlink_seed,
"PPP prompt seed": input_seed,
"PPP incremental seed": input_incremental_seed,
"PPP combinatorial": input_combinatorial,
}
)
@@ -315,7 +338,6 @@ class PromptPostProcessorA1111Script(scripts.Script):
calculated_seeds: list[int] = []
if input_unlink_seed:
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "Using unlinked seed")
num_seeds = len(getattr(p, "all_seeds", []))
if input_incremental_seed:
first_seed = np.random.randint(0, 2**32, dtype=np.int64) if input_seed == -1 else input_seed
calculated_seeds = [first_seed + i for i in range(num_seeds)]
@@ -340,6 +362,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
# (prompt type, typeindex) -> (new positive prompt, new negative prompt)
prompts_list: dict[tuple[str, int], tuple[str, str]] = {}
extra_params = {}
# adds prompts
regular_type = "regular"
@@ -356,25 +379,49 @@ class PromptPostProcessorA1111Script(scripts.Script):
if hiresfix_exists:
prompts_list[(hiresfix_type, i)] = None
# processes prompts
for prompttype, typeindex in prompts_list.keys():
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, f"processing prompts ({prompttype}[{typeindex+1}])")
key = (
(hash_fullenv, calculated_seeds[typeindex], rpr[typeindex], rnr[typeindex])
if prompttype == regular_type
else (hash_fullenv, calculated_seeds[typeindex], rph[typeindex], rnh[typeindex])
)
cached = self.lru_cache.get(key)
if cached is None:
(hsh, seed, prompt, negative_prompt) = key
posp, negp, _ = ppp.process_prompt(prompt, negative_prompt, seed)
cached = (posp, negp)
self.lru_cache.put(key, cached)
# adds also the result so i2i doesn't process it unnecessarily
self.lru_cache.put((hsh, seed, posp, negp), cached)
else:
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache")
prompts_list[(prompttype, typeindex)] = cached
if input_combinatorial:
seed_for_comb = calculated_seeds[0] if calculated_seeds else 0
regular_copy = (rpr.copy() if rpr else None, rnr.copy() if rnr else None)
hiresfix_copy = (rph.copy() if rph else None, rnh.copy() if rnh else None)
regular_changes = False
hiresfix_changes = False
if regular_exists:
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "processing prompts combinatorially (regular)")
comb_results = ppp.process_prompt(rpr[0], rnr[0], seed_for_comb)
num_comb = len(comb_results)
for i in range(len(rpr)): # pylint: disable=consider-using-enumerate
posp, negp, _ = comb_results[i % num_comb]
prompts_list[(regular_type, i)] = (posp, negp)
extra_params["PPP combination"] = [1+(i % num_comb) for i in range(len(rpr))]
if hiresfix_exists:
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "processing prompts combinatorially (hiresfix)")
comb_results_hr = ppp.process_prompt(rph[0], rnh[0], seed_for_comb)
num_comb_hr = len(comb_results_hr)
for i in range(len(rph)): # pylint: disable=consider-using-enumerate
posp, negp, _ = comb_results_hr[i % num_comb_hr]
prompts_list[(hiresfix_type, i)] = (posp, negp)
extra_params["PPP HR combination"] = [1+(i % num_comb_hr) for i in range(len(rph))]
else:
# processes prompts
for prompttype, typeindex in prompts_list.keys():
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, f"processing prompts ({prompttype}[{typeindex+1}])")
key = (
(hash_fullenv, calculated_seeds[typeindex], rpr[typeindex], rnr[typeindex])
if prompttype == regular_type
else (hash_fullenv, calculated_seeds[typeindex], rph[typeindex], rnh[typeindex])
)
cached = self.lru_cache.get(key)
if cached is None:
(hsh, seed, prompt, negative_prompt) = key
results = ppp.process_prompt(prompt, negative_prompt, seed)
posp, negp, _ = results[0]
cached = (posp, negp)
self.lru_cache.put(key, cached)
# adds also the result so i2i doesn't process it unnecessarily
self.lru_cache.put((hsh, seed, posp, negp), cached)
else:
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache")
prompts_list[(prompttype, typeindex)] = cached
# with open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "last_prompts.txt"), "w", encoding="utf-8") as f:
# for (prompttype, typeindex), (posp, negp) in prompts_list.items():
@@ -404,7 +451,6 @@ class PromptPostProcessorA1111Script(scripts.Script):
rnh[typeindex] = negp
# initialize extra generation parameters
extra_params = {}
if add_prompts:
if regular_changes:
extra_params["PPP original prompts"] = regular_copy[0]
+120 -43
View File
@@ -6,10 +6,10 @@ import unittest
import datetime
from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, PPPStateOptions
from ppp_enmappings import PPPExtraNetworkMappings # pylint: disable=import-error
from ppp_wildcards import PPPWildcards # pylint: disable=import-error
from ppp import PromptPostProcessor # pylint: disable=import-error
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # pylint: disable=import-error
from ppp_enmappings import PPPExtraNetworkMappings # type: ignore
from ppp_wildcards import PPPWildcards # type: ignore
from ppp import PromptPostProcessor # type: ignore
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # type: ignore
class PromptPair(NamedTuple):
@@ -17,6 +17,12 @@ class PromptPair(NamedTuple):
negative_prompt: str = ""
class OutputTuple(NamedTuple):
prompt: str = ""
negative_prompt: str = ""
variables: dict[str, str] = None
class TestPromptPostProcessorBase(unittest.TestCase):
"""
A test case class for testing the PromptPostProcessor class.
@@ -106,29 +112,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
def interrupt(self):
self.interrupted = True
def process(
self,
input_prompts: PromptPair,
expected_output_prompts: Optional[PromptPair | list[PromptPair]] = None,
seed: int = 1,
ppp: Optional[str | PromptPostProcessor] = None,
interrupted: bool = False,
variables: dict[str, str] | None = None,
):
"""
Process the prompt and compare the results with the expected prompts.
Args:
input_prompts (PromptPair): The input prompts.
expected_output_prompts (PromptPair | list[PromptPair], optional): The expected prompts.
seed (int, optional): The seed value. Defaults to 1.
ppp (Optional[str | PromptPostProcessor], optional): The PromptPostProcessor instance or type. Defaults to None.
interrupted (bool, optional): The interrupted flag. Defaults to False.
variables (dict[str,str]|None, optional): Output variables to check. Defaults to None.
Returns:
None
"""
def init_obj(self, ppp: Optional[str | PromptPostProcessor] = None) -> PromptPostProcessor:
if isinstance(ppp, str):
if ppp == "nocup":
the_obj = PromptPostProcessor(
@@ -180,30 +164,123 @@ class TestPromptPostProcessorBase(unittest.TestCase):
self.wildcards_obj,
self.extranetwork_maps_obj,
)
return the_obj
def process(
self,
input_prompts: PromptPair,
expected_output: Optional[OutputTuple | list[OutputTuple]] = None,
seed: int = 1,
ppp: Optional[str | PromptPostProcessor] = None,
interrupted: bool = False,
):
"""
Process the prompt and compare the results with the expected prompts.
Args:
input_prompts (PromptPair): The input prompts.
expected_output (OutputTuple | list[OutputTuple], optional): The expected output. When a list is provided, the test will run once for each expected output, using the same input prompt, but seed will be incremented for each iteration.
seed (int, optional): The seed value. Defaults to 1.
ppp (Optional[str | PromptPostProcessor], optional): The PromptPostProcessor instance or type. Defaults to None.
interrupted (bool, optional): The interrupted flag. Defaults to False.
Returns:
None
"""
the_obj = self.init_obj(ppp)
out = (
[PromptPair("", "")]
if expected_output_prompts is None
else expected_output_prompts if isinstance(expected_output_prompts, list) else [expected_output_prompts]
[OutputTuple("", "", None)]
if expected_output is None
else expected_output if isinstance(expected_output, list) else [expected_output]
)
for eo in out:
result_prompt, result_negative_prompt, output_variables = the_obj.process_prompt(
result = the_obj.process_prompt(
input_prompts.prompt,
input_prompts.negative_prompt,
seed,
)
self.assertEqual(self.interrupted, interrupted, "Interrupted flag is incorrect")
if not self.interrupted:
if expected_output_prompts is not None:
self.assertEqual(result_prompt, eo.prompt, "Incorrect prompt")
self.assertEqual(result_negative_prompt, eo.negative_prompt, "Incorrect negative prompt")
if variables is not None:
for var_name, var_value in variables.items():
self.assertIn(
var_name, output_variables, f"Variable '{var_name}' not found in output variables"
)
self.assertEqual(
output_variables[var_name],
var_value,
f"Variable '{var_name}' has incorrect value",
result_prompt, result_negative_prompt, output_variables = result[0]
if expected_output is not None:
self.assertTrue(
result_prompt == eo.prompt and result_negative_prompt == eo.negative_prompt,
f"Incorrect result '{eo.prompt}' / '{eo.negative_prompt}', got '{result_prompt}' / '{result_negative_prompt}'",
)
if eo.variables:
unmatched_vars = {}
expected_values = {}
for var_name, var_value in eo.variables.items():
if var_name not in output_variables or output_variables[var_name] != var_value:
unmatched_vars[var_name] = (
output_variables[var_name] if var_name in output_variables else None
)
expected_values[var_name] = var_value
self.assertTrue(
not unmatched_vars,
f"Result '{eo.prompt}' / '{eo.negative_prompt}' found, but variables do not match: expected {expected_values}, got {unmatched_vars}",
)
seed += 1
def process_combinatorial(
self,
input_prompts: PromptPair,
expected_output: Optional[OutputTuple | list[OutputTuple]] = None,
seed: int = 1,
ppp: Optional[str | PromptPostProcessor] = None,
interrupted: bool = False,
):
"""
Process the prompt and compare the results with the expected prompts.
Args:
input_prompts (PromptPair): The input prompts.
expected_output (OutputTuple | list[OutputTuple], optional): The expected output.
seed (int, optional): The seed value. Defaults to 1.
ppp (Optional[str | PromptPostProcessor], optional): The PromptPostProcessor instance or type. Defaults to None.
interrupted (bool, optional): The interrupted flag. Defaults to False.
Returns:
None
"""
the_obj = self.init_obj(ppp)
out = (
[OutputTuple("", "", None)]
if expected_output is None
else expected_output if isinstance(expected_output, list) else [expected_output]
)
result = the_obj.process_prompt(
input_prompts.prompt,
input_prompts.negative_prompt,
seed,
combinatorial=True,
)
self.assertEqual(self.interrupted, interrupted, "Interrupted flag is incorrect")
if not self.interrupted:
if expected_output is not None:
self.assertEqual(
len(result), len(out), f"Incorrect number of combinations (expected {len(out)}, got {len (result)})"
)
for out_prompt, out_negative_prompt, out_variables in out:
found = None
for r_prompt, r_negative_prompt, r_variables in result:
if r_prompt == out_prompt and r_negative_prompt == out_negative_prompt:
found = OutputTuple(r_prompt, r_negative_prompt, r_variables)
break
self.assertTrue(
bool(found),
f"Combination '{out_prompt}' / '{out_negative_prompt}' not found in output",
)
if found and out_variables:
unmatched_vars = {}
expected_values = {}
for var_name, var_value in out_variables.items():
if var_name not in found.variables or found.variables[var_name] != var_value:
unmatched_vars[var_name] = (
found.variables[var_name] if var_name in found.variables else None
)
expected_values[var_name] = var_value
self.assertTrue(
not unmatched_vars,
f"Combination '{out_prompt}' / '{out_negative_prompt}' found, but variables do not match: expected {expected_values}, got {unmatched_vars}",
)
+27 -12
View File
@@ -1,7 +1,7 @@
from dataclasses import replace
from ppp import PromptPostProcessor # pylint: disable=import-error
from .base_tests import PromptPair, TestPromptPostProcessorBase
from ppp import PromptPostProcessor # type: ignore
from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -17,14 +17,14 @@ class TestChoices(TestPromptPostProcessorBase):
def test_ch_choices(self): # simple choices with weights
self.process(
PromptPair("the choices are: {3::choice1|2::choice2|choice3}", ""),
PromptPair("the choices are: choice2", ""),
OutputTuple("the choices are: choice2", ""),
ppp="nocup",
)
def test_ch_unsupportedsampler(self): # unsupported sampler
self.process(
PromptPair("the choices are: {@choice1|choice2|choice3}", ""),
PromptPair("", ""),
OutputTuple("", ""),
ppp="nocup",
interrupted=True,
)
@@ -35,28 +35,28 @@ class TestChoices(TestPromptPostProcessorBase):
"the choices are: {\n3::choice1 # this is option 1\n|2::choice2\n# this was option 2\n|choice3 # this is option 3\n}",
"",
),
PromptPair("the choices are: choice2", ""),
OutputTuple("the choices are: choice2", ""),
ppp="nocup",
)
def test_ch_choices_multiple(self): # choices with multiple selection
self.process(
PromptPair("the choices are: {~2$$, $$3::choice1|2:: choice2 |choice3}", ""),
PromptPair("the choices are: choice2 , choice3", ""),
OutputTuple("the choices are: choice2 , choice3", ""),
ppp="nocup",
)
def test_ch_choices_if_multiple(self): # choices with if and multiple selection
self.process(
PromptPair("the choices are: {2$$, $$3::choice1|2 if _is_sd1::choice2|choice3}", ""),
PromptPair("the choices are: choice1, choice3", ""),
OutputTuple("the choices are: choice1, choice3", ""),
ppp="nocup",
)
def test_ch_choices_set_if_multiple(self): # choices with if user variable and multiple selection
self.process(
PromptPair("${var=test}the choices are: {2$$, $$3::choice1|2 if not var eq 'test'::choice2|choice3}", ""),
PromptPair("the choices are: choice1, choice3", ""),
OutputTuple("the choices are: choice1, choice3", ""),
ppp="nocup",
)
@@ -66,21 +66,21 @@ class TestChoices(TestPromptPostProcessorBase):
"${var=test}the choices are: {2$$, $$3::choice1${var2=test2} {if var2 eq 'test2'::choice11|choice12}|2 if not var eq 'test'::choice2|choice3}",
"",
),
PromptPair("the choices are: choice1 choice11, choice3", ""),
OutputTuple("the choices are: choice1 choice11, choice3", ""),
ppp="nocup",
)
def test_ch_choicesinsidelora(self): # simple choices inside a lora
self.process(
PromptPair("<lora:test1:1><lora:test__other__name:1><lora:test2:{0.2|0.5|0.7|1}>", ""),
PromptPair("<lora:test1:1><lora:test__other__name:1><lora:test2:0.7>", ""),
OutputTuple("<lora:test1:1><lora:test__other__name:1><lora:test2:0.7>", ""),
ppp="nocup",
)
def test_ch_removelorawithchoices(self):
self.process(
PromptPair("<lora:test1:1><lora:test2:{0.2|0.5|0.7|1}>", ""),
PromptPair("", ""),
OutputTuple("", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
self.def_env_info,
@@ -98,6 +98,21 @@ class TestChoices(TestPromptPostProcessorBase):
def test_ch_cmd_includewildcard(self):
self.process(
PromptPair("{ch_one|ch_two|%0.5::include yaml/wildcard1}", ""),
PromptPair("ch_two", ""),
OutputTuple("ch_two", ""),
ppp="nocup",
)
# Combinatorial
def test_ch_combinatorial(self):
self.process_combinatorial(
PromptPair("{choice1|choice2|choice3}, ${v:{option1|option2}}", ""),
[
OutputTuple("choice1, option1", ""),
OutputTuple("choice1, option2", ""),
OutputTuple("choice2, option1", ""),
OutputTuple("choice2, option2", ""),
OutputTuple("choice3, option1", ""),
OutputTuple("choice3, option2", "", {"v": "option2"}),
],
)
+15 -15
View File
@@ -1,8 +1,8 @@
import logging
from dataclasses import replace
from ppp import PromptPostProcessor # pylint: disable=import-error
from .base_tests import PromptPair, TestPromptPostProcessorBase
from ppp import PromptPostProcessor # type: ignore
from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
if __name__ == "__main__":
@@ -19,7 +19,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_simple(self): # simple cleanup
self.process(
PromptPair(" this is a ((test ), , , (), , [] ( , test ,:2.0):1.5), (red:1.5) ", " normal quality "),
PromptPair("this is a ((test), (test,:2):1.5), (red:1.5)", "normal quality"),
OutputTuple("this is a ((test), (test,:2):1.5), (red:1.5)", "normal quality"),
)
def test_cl_complex(self): # complex cleanup
@@ -28,7 +28,7 @@ class TestCleanup(TestPromptPostProcessorBase):
" this is BREAKABLE a ((test)), ,AND AND(() [] <lora:test> ANDERSON (test:2.0):1.5) :o BREAK \n BREAK (red:1.5) ",
" [:hands, feet, :0.15]normal quality ",
),
PromptPair(
OutputTuple(
"this is BREAKABLE a (test:1.21) AND(<lora:test> ANDERSON (test:2):1.5) :o BREAK (red:1.5)",
"[:hands, feet, :0.15]normal quality",
),
@@ -37,7 +37,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_removenetworktags(self): # remove network tags
self.process(
PromptPair("this is a <lora:test:1> test__yaml/wildcard7__", ""),
PromptPair("this is a test", ""),
OutputTuple("this is a test", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
self.def_env_info,
@@ -55,7 +55,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_dontremoveseparatorsoneol(self): # don't remove separators on eol
self.process(
PromptPair("this is a test,\nsecond line", ""),
PromptPair("this is a test,\nsecond line", ""),
OutputTuple("this is a test,\nsecond line", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
self.def_env_info,
@@ -79,7 +79,7 @@ class TestCleanup(TestPromptPostProcessorBase):
(d:0.9)""",
"",
),
PromptPair(
OutputTuple(
""" (l:1.1) (d:0.9),
(l:1.1)
(d:0.9)""",
@@ -115,7 +115,7 @@ class TestCleanup(TestPromptPostProcessorBase):
"this is (a test:0.9) of (attention (merging:1.2)) where ((this)) ((is joined:1.2)) and ([this too]:1.3)",
"",
),
PromptPair(
OutputTuple(
"this is [a test] of (attention (merging:1.2)) where (this:1.21) (is joined:1.32) and (this too:1.17)",
"",
),
@@ -127,7 +127,7 @@ class TestCleanup(TestPromptPostProcessorBase):
"this is (a test:0.9) of not (attention (merging:1.2)) where ((this)) ((is not joined:1.2)) and neither is ([this]:1.3)",
"",
),
PromptPair(
OutputTuple(
"this is (a test:0.9) of not (attention (merging:1.2)) where ((this)) ((is not joined:1.2)) and neither is ([this]:1.3)",
"",
),
@@ -140,7 +140,7 @@ class TestCleanup(TestPromptPostProcessorBase):
with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm:
self.process(
PromptPair("(unclosed paren", ""),
PromptPair("(unclosed paren", ""),
OutputTuple("(unclosed paren", ""),
)
self.assertTrue(
any("Unmatched" in msg for msg in cm.output),
@@ -151,7 +151,7 @@ class TestCleanup(TestPromptPostProcessorBase):
with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm:
self.process(
PromptPair("extra close paren)", ""),
PromptPair("extra close paren)", ""),
OutputTuple("extra close paren)", ""),
)
self.assertTrue(
any("Unmatched" in msg for msg in cm.output),
@@ -162,7 +162,7 @@ class TestCleanup(TestPromptPostProcessorBase):
with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm:
self.process(
PromptPair("(mismatched]", ""),
PromptPair("(mismatched]", ""),
OutputTuple("(mismatched]", ""),
)
self.assertTrue(
any("Mismatched" in msg or "Unmatched" in msg for msg in cm.output),
@@ -173,7 +173,7 @@ class TestCleanup(TestPromptPostProcessorBase):
with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm:
self.process(
PromptPair("unclosed [bracket", ""),
PromptPair("unclosed [bracket", ""),
OutputTuple("unclosed [bracket", ""),
)
self.assertTrue(
any("Unmatched" in msg for msg in cm.output),
@@ -184,7 +184,7 @@ class TestCleanup(TestPromptPostProcessorBase):
with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm:
self.process(
PromptPair("[(unmatched [bracket))", ""),
PromptPair("[(unmatched [bracket))", ""),
OutputTuple("[(unmatched [bracket))", ""),
)
self.assertTrue(
any("Unmatched" in msg for msg in cm.output),
@@ -195,5 +195,5 @@ class TestCleanup(TestPromptPostProcessorBase):
with self.assertNoLogs("PromptPostProcessor", level=logging.WARNING):
self.process(
PromptPair(r"text with \(escaped unmatched\]", ""),
PromptPair(r"text with \(escaped unmatched\]", ""),
OutputTuple(r"text with \(escaped unmatched\]", ""),
)
+22 -22
View File
@@ -1,5 +1,5 @@
from ppp import PromptPostProcessor # pylint: disable=import-error
from .base_tests import PromptPair, TestPromptPostProcessorBase
from ppp import PromptPostProcessor # type: ignore
from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -18,7 +18,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1] (test2) (test3:1.5) [(test4)]",
"",
),
PromptPair("(test1:0.9) (test2) (test3:1.5) (test4:0.99)", ""),
OutputTuple("(test1:0.9) (test2) (test3:1.5) (test4:0.99)", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -39,7 +39,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1] (test2) (test3:1.5)",
"",
),
PromptPair("test1 test2 test3", ""),
OutputTuple("test1 test2 test3", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -60,7 +60,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1] (test2) (test3:1.5)",
"",
),
PromptPair("", ""),
OutputTuple("", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -81,7 +81,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1] (test2) (test3:1.5)",
"",
),
PromptPair("", ""),
OutputTuple("", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -103,7 +103,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1:test2:0.5]",
"",
),
PromptPair("test1", ""),
OutputTuple("test1", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -124,7 +124,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1:test2:0.5]",
"",
),
PromptPair("test2", ""),
OutputTuple("test2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -145,7 +145,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1::0.5] [:test2:0.5] [test3:test4:0.5]",
"",
),
PromptPair("test1 test3", ""),
OutputTuple("test1 test3", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -166,7 +166,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1:test2:0.5]",
"",
),
PromptPair("", ""),
OutputTuple("", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -187,7 +187,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1:test2:0.5]",
"",
),
PromptPair("", ""),
OutputTuple("", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -209,7 +209,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1|test2|test3]",
"",
),
PromptPair("test1", ""),
OutputTuple("test1", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -230,7 +230,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1|test2|test3]",
"",
),
PromptPair("", ""),
OutputTuple("", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -251,7 +251,7 @@ class TestHosts(TestPromptPostProcessorBase):
"[test1|test2|test3]",
"",
),
PromptPair("", ""),
OutputTuple("", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -273,7 +273,7 @@ class TestHosts(TestPromptPostProcessorBase):
"test1 AND test2:2",
"",
),
PromptPair("test1\ntest2", ""),
OutputTuple("test1\ntest2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -294,7 +294,7 @@ class TestHosts(TestPromptPostProcessorBase):
"test1 AND test2:2",
"",
),
PromptPair("test1, test2", ""),
OutputTuple("test1, test2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -315,7 +315,7 @@ class TestHosts(TestPromptPostProcessorBase):
"test1 AND test2:2",
"",
),
PromptPair("test1 test2", ""),
OutputTuple("test1 test2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -336,7 +336,7 @@ class TestHosts(TestPromptPostProcessorBase):
"test1 AND test2:2",
"",
),
PromptPair("", ""),
OutputTuple("", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -358,7 +358,7 @@ class TestHosts(TestPromptPostProcessorBase):
"test1 BREAK test2",
"",
),
PromptPair("test1\ntest2", ""),
OutputTuple("test1\ntest2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -379,7 +379,7 @@ class TestHosts(TestPromptPostProcessorBase):
"test1 BREAK test2",
"",
),
PromptPair("test1, test2", ""),
OutputTuple("test1, test2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -400,7 +400,7 @@ class TestHosts(TestPromptPostProcessorBase):
"test1 BREAK test2",
"",
),
PromptPair("test1 test2", ""),
OutputTuple("test1 test2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -421,7 +421,7 @@ class TestHosts(TestPromptPostProcessorBase):
"test1 BREAK test2",
"",
),
PromptPair("", ""),
OutputTuple("", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
+12 -11
View File
@@ -1,4 +1,4 @@
from .base_tests import PromptPair, TestPromptPostProcessorBase
from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -17,7 +17,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
"flowers<ppp:stn>red<ppp:/stn>, <ppp:stn s>green<ppp:/stn>, <ppp:stn e>blue<ppp:/stn><ppp:stn p0>yellow<ppp:/stn>, <ppp:stn p1>purple<ppp:/stn><ppp:stn p2>black<ppp:/stn>",
"<ppp:stn i0/>normal quality<ppp:stn i1>, worse quality<ppp:stn i2/>",
),
PromptPair("flowers", "red, green, yellow, normal quality, purple, worse quality, black, blue"),
OutputTuple("flowers", "red, green, yellow, normal quality, purple, worse quality, black, blue"),
)
def test_stn_complex(self): # complex negtags
@@ -26,7 +26,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
"<ppp:stn>red<ppp:/stn> ((<ppp:stn s>pink<ppp:/stn>)), flowers <ppp:stn e>purple<ppp:/stn>, <ppp:stn p0>mauve<ppp:/stn><ppp:stn e>blue<ppp:/stn>, <ppp:stn p0>yellow<ppp:/stn> <ppp:stn p1>green<ppp:/stn>",
"normal quality, <ppp:stn i0/>, bad quality<ppp:stn i1/>, worse quality",
),
PromptPair(
OutputTuple(
"flowers",
"red, (pink:1.21), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue",
),
@@ -38,7 +38,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
"<ppp:stn>red<ppp:/stn> ((<ppp:stn s>pink<ppp:/stn>)), flowers <ppp:stn e>purple<ppp:/stn>, <ppp:stn p0>mauve<ppp:/stn><ppp:stn e>blue<ppp:/stn>, <ppp:stn p0>yellow<ppp:/stn> <ppp:stn p1>green<ppp:/stn>",
"normal quality, <ppp:stn i0/>, bad quality<ppp:stn i1/>, worse quality",
),
PromptPair(
OutputTuple(
" (()), flowers , , ",
"red, ((pink)), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue",
),
@@ -51,7 +51,8 @@ class TestSendToNegative(TestPromptPostProcessorBase):
"[<ppp:stn>neg1<ppp:/stn>] this is a ((test<ppp:stn e>neg2<ppp:/stn>) (test:2.0): 1.5 ) (red<ppp:stn>[square]<ppp:/stn>:1.5)",
"normal quality",
),
PromptPair(
OutputTuple(
"this is a ((test) (test:2):1.5) (red:1.5)", "[neg1], ([square]:1.5), normal quality, (neg2:1.65)"
),
)
@@ -62,7 +63,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
"this is a (([complex<ppp:stn>neg1<ppp:/stn>|simple<ppp:stn>neg2<ppp:/stn>|regular<ppp:stn>neg3<ppp:/stn>] test)(test:2.0):1.5)",
"normal quality",
),
PromptPair(
OutputTuple(
"this is a (([complex|simple|regular] test)(test:2):1.5)",
"([neg1||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality",
),
@@ -74,7 +75,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
"this is a (([complex<ppp:stn>neg1<ppp:/stn>[one|two<ppp:stn>neg12<ppp:/stn>||three|four(<ppp:stn>neg14<ppp:/stn>)]|simple<ppp:stn>neg2<ppp:/stn>|regular<ppp:stn>neg3<ppp:/stn>] test)(test:2.0):1.5)",
"normal quality",
),
PromptPair(
OutputTuple(
"this is a (([complex[one|two||three|four]|simple|regular] test)(test:2):1.5)",
"([neg1||]:1.65), ([[|neg12|||]||]:1.65), ([[||||(neg14)]||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality",
),
@@ -83,7 +84,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
def test_stn_inside_scheduling(self): # negtag inside scheduling
self.process(
PromptPair("this is [abc<ppp:stn>neg1<ppp:/stn>:def<ppp:stn e>neg2<ppp:/stn>: 5 ]", "normal quality"),
[PromptPair("this is [abc:def:5]", "[neg1::5], normal quality, [neg2:5]")],
OutputTuple("this is [abc:def:5]", "[neg1::5], normal quality, [neg2:5]"),
)
def test_stn_complex_features(self): # complex negtags with AND, BREAK and other features
@@ -92,7 +93,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
"[<ppp:stn>neg5<ppp:/stn>] this \\(is\\): a (([complex|simple<ppp:stn>neg6<ppp:/stn>|regular] test<ppp:stn>neg1<ppp:/stn>)(test:2.0):1.5) \nBREAK, BREAK with [abc<ppp:stn>neg4<ppp:/stn>:def<ppp:stn p0>neg2(neg3:1.6)<ppp:/stn>:5]:0.5 AND loratrigger <lora:xxx:1> AND AND hypernettrigger <hypernet:yyy>:0.3",
"normal quality, <ppp:stn i0/>",
),
PromptPair(
OutputTuple(
"this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger <lora:xxx:1> AND hypernettrigger <hypernet:yyy>:0.3",
"[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]",
),
@@ -104,7 +105,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
"[<ppp:stn>neg5<ppp:/stn>] this \\(is\\): a (([complex|simple<ppp:stn>neg6<ppp:/stn>|regular] test<ppp:stn>neg1<ppp:/stn>)(test:2.0):1.5) \nBREAK, BREAK with [abc<ppp:stn>neg4<ppp:/stn>:def<ppp:stn p0>neg2(neg3:1.6)<ppp:/stn>:5]:0.5 AND loratrigger <lora:xxx:1> AND AND hypernettrigger <hypernet:yyy>:0.3",
"normal quality, <ppp:stn i0/>",
),
PromptPair(
OutputTuple(
"this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger <lora:xxx:1> AND hypernettrigger <hypernet:yyy>:0.3",
"[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]",
),
@@ -116,7 +117,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
"[pos1<ppp:stn>neg1<ppp:/stn>[pos11|pos12<ppp:stn>neg12<ppp:/stn>||pos14|pos15<ppp:stn>neg15<ppp:/stn>]|pos2<ppp:stn>neg2<ppp:/stn>|pos3<ppp:stn>neg3<ppp:/stn>]",
"",
),
PromptPair(
OutputTuple(
"[pos1[pos11|pos12||pos14|pos15]|pos2|pos3]",
"[neg1||], [[|neg12|||]||], [[||||neg15]||], [|neg2|], [||neg3]",
# "[neg1[|neg12|||neg15]|neg2|neg3]", # expected output if the constructs were unified
+102 -118
View File
@@ -1,8 +1,8 @@
from dataclasses import replace
from ppp import PromptPostProcessor # pylint: disable=import-error
from ppp_classes import ONWARNING_CHOICES # pylint: disable=import-error
from .base_tests import PromptPair, TestPromptPostProcessorBase
from ppp import PromptPostProcessor # type: ignore
from ppp_classes import ONWARNING_CHOICES # type: ignore
from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -21,8 +21,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1=}<ppp:set v2><ppp:/set>${v3:}",
"",
),
PromptPair("", ""),
variables={"v1": "", "v2": "", "v3": ""},
OutputTuple("", "",{"v1": "", "v2": "", "v3": ""}),
)
# Echoed variables tests
@@ -34,8 +33,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"",
),
# v3 is echoed withs two defaults, the output prompt has both but the variable value is the last default
PromptPair("test3test4", ""),
variables={"v1": "test1", "v2": "test2", "v3": "test4"},
OutputTuple("test3test4", "", {"v1": "test1", "v2": "test2", "v3": "test4"}),
)
def test_unknown_echoed_variable(self):
@@ -44,8 +42,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1}",
"",
),
PromptPair("", ""),
variables={"v1": ""},
OutputTuple("", "", {"v1": ""}),
ppp=PromptPostProcessor(
self.ppp_logger,
self.def_env_info,
@@ -68,7 +65,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${f=filename}${w=0.5}<lora:${f}:${w}>",
"",
),
PromptPair("<lora:filename:0.5>", ""),
OutputTuple("<lora:filename:0.5>", ""),
)
# Variable nesting tests
@@ -79,8 +76,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1=test ${v2:OK}}${v1}",
"",
),
PromptPair("test OK", ""),
variables={"v1": "test OK", "v2": "OK"},
OutputTuple("test OK", "", {"v1": "test OK", "v2": "OK"}),
)
def test_var_nested_2(self): # variable set nested in variable default
@@ -89,8 +85,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1:test ${v2=OK}${v2}}",
"",
),
PromptPair("test OK", ""),
variables={"v1": "test OK", "v2": "OK"},
OutputTuple("test OK", "", {"v1": "test OK", "v2": "OK"}),
)
def test_var_nested_3(self): # variable default nested in variable default
@@ -99,8 +94,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1:test ${v2:OK}}",
"",
),
PromptPair("test OK", ""),
variables={"v1": "test OK", "v2": "OK"},
OutputTuple("test OK", "", {"v1": "test OK", "v2": "OK"}),
)
# Array variable tests
@@ -111,8 +105,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1[]=val1}${v1[]+=val2}${v1[]+=val3}${v1[1]:defval},${v1[]:defval2},${v1[&'.']:defval3}",
"",
),
PromptPair("val2,val1, val2, val3,val1.val2.val3", ""),
variables={"v1[]": "val1, val2, val3", "v1[1]": "val2", "v1[&'.']": "val1.val2.val3"},
OutputTuple("val2,val1, val2, val3,val1.val2.val3", "", {"v1[]": "val1, val2, val3", "v1[1]": "val2", "v1[&'.']": "val1.val2.val3"}),
)
def test_array_variable_2(self): # override of array variable value, test of default value when array variable is empty, test of default value when array variable is not set
@@ -121,8 +114,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1[]=val1}${v1[]=val2}${v1[]:defval},${v2[]:defval2},${v2[1]:defval3},${v3[]=}${v3[]:defval4}",
"",
),
PromptPair("val2,defval2,defval3", ""),
variables={"v1[]": "val2", "v2[]": "defval2", "v2[1]": "defval3", "v3[]": ""},
OutputTuple("val2,defval2,defval3", "", {"v1[]": "val2", "v2[]": "defval2", "v2[1]": "defval3", "v3[]": ""}),
)
def test_array_variable_3(self): # access array index by variable, set array variable to expanded array variable and add expanded array
@@ -131,8 +123,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1[]=val1}${v1[]+=val2}${v2=1}${v1[v2]:defval1}${v3[]=${v1[]}}${v3[]+=${v1[]}}, ${v3[&'.']}",
"",
),
PromptPair("val2, val1, val2.val1, val2", ""),
variables={"v1[]": "val1, val2", "v2": "1", "v1[v2]": "val2", "v3[]": "val1, val2, val1, val2", "v3[&'.']": "val1, val2.val1, val2"},
OutputTuple("val2, val1, val2.val1, val2", "", {"v1[]": "val1, val2", "v2": "1", "v1[v2]": "val2", "v3[]": "val1, val2, val1, val2", "v3[&'.']": "val1, val2.val1, val2"}),
)
def test_array_variable_4(self): # test list in array
@@ -141,8 +132,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1[]=val1}${v1[]+=val2}${v1[]+=val3}<ppp:if ('val1','val2') in v1[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
variables={"v1[]": "val1, val2, val3"},
OutputTuple("OK", "", {"v1[]": "val1, val2, val3"}),
)
def test_array_variable_5(self): # test empty array
@@ -151,8 +141,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1[]=}<ppp:if v1[]>OK<ppp:else>not OK<ppp:/if>,<ppp:if not v1[0]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK,OK", ""),
variables={"v1[]": ""},
OutputTuple("OK,OK", "", {"v1[]": ""}),
)
def test_array_variable_6(self): # array variable set and addition with expanded values from array variables
@@ -161,8 +150,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1[]=val1}${v1[]+=val2}${v2[]=val3}${v3[]=*v1[]}${v3[]+=*v2[]}",
"",
),
PromptPair("", ""),
variables={"v1[]": "val1, val2", "v2[]": "val3", "v3[]": "val1, val2, val3"},
OutputTuple("", "", {"v1[]": "val1, val2", "v2[]": "val3", "v3[]": "val1, val2, val3"}),
)
def test_array_variable_7(self): # array variable set and addition with expanded values from wildcards
@@ -171,8 +159,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1[]=*__yaml/wildcard1__}${v1[]+=*__yaml/wildcard2__}${v1[2]:defval}",
"",
),
PromptPair("choice3", ""),
variables={"v1[]": "choice2, choice1, choice3, choice1"},
OutputTuple("choice3", "", {"v1[]": "choice2, choice1, choice3, choice1"}),
)
def test_array_variable_8(self): # array variable set and addition with expanded values from lists
@@ -181,8 +168,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1[]=*()}${v1[]+=*('one','two')}${v2=three}${v1[]+=*(v2,'four')}${v1[2]:defval}",
"",
),
PromptPair("three", ""),
variables={"v1[]": "one, two, three, four", "v2": "three"},
OutputTuple("three", "", {"v1[]": "one, two, three, four", "v2": "three"}),
)
def test_array_variable_9(self): # array variable length
@@ -191,8 +177,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v1[]=val1}${v1[]+=val2}${v1[]+=val3}${v1[#]:defval}, <ppp:if v1[#] eq 3>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("3, OK", ""),
variables={"v1[]": "val1, val2, val3", "v1[#]": "3"},
OutputTuple("3, OK", "", {"v1[]": "val1, val2, val3", "v1[#]": "3"}),
)
def test_array_variable_10(self): # array variable set with expanded values from wildcards in command format
@@ -201,8 +186,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set v1[]>*__yaml/wildcard1__<ppp:/set><ppp:set v1[] add>*__yaml/wildcard2__<ppp:/set>${v1[2]:defval}",
"",
),
PromptPair("choice3", ""),
variables={"v1[]": "choice2, choice1, choice3, choice1"},
OutputTuple("choice3", "", {"v1[]": "choice2, choice1, choice3, choice1"}),
)
# Operator tests
@@ -215,7 +199,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=hello}${r2=hello}<ppp:if r1 eq r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_RnoteqR(self): # test for the not before the operator
@@ -224,7 +208,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=hello}${r2=bye}<ppp:if r1 not eq r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_RneR(self):
@@ -233,7 +217,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=hello}${r2=bye}<ppp:if r1 ne r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_RltR(self):
@@ -242,7 +226,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=1}${r2=2}<ppp:if r1 lt r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_RgtR(self):
@@ -251,7 +235,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=2}${r2=1}<ppp:if r1 gt r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_RleR(self):
@@ -260,7 +244,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=1}${r2=1}<ppp:if r1 le r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_RgeR(self):
@@ -269,7 +253,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=1}${r2=1}<ppp:if r1 ge r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_RinR(self):
@@ -278,7 +262,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=hello}${r2=hello world}<ppp:if r1 in r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_RcontainsR(self):
@@ -287,7 +271,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=hello world}${r2=hello}<ppp:if r1 contains r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
## A vs A
@@ -298,7 +282,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*('hello','world')}${a2[]=*('hello','world')}<ppp:if a1[] eq a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_AneA_1(self):
@@ -307,7 +291,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*('hello')}${a2[]=*('bye')}<ppp:if a1[] ne a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_AneA_2(self):
@@ -316,7 +300,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*('hello','world')}${a2[]=*('hello')}<ppp:if a1[] ne a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_AneA_3(self):
@@ -325,7 +309,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*('hello','world')}${a2[]=*('world','hello')}<ppp:if a1[] ne a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_AltA_1(self):
@@ -334,7 +318,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*(1,2,3)}${a2[]=*(2,3,4)}<ppp:if a1[] lt a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_AltA_2(self):
@@ -343,7 +327,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*(1,2)}${a2[]=*(2,3,4)}<ppp:if a1[] lt a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
)
def test_operator_AgtA(self):
@@ -352,7 +336,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*(2,3,4)}${a2[]=*(1,2,3)}<ppp:if a1[] gt a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_AleA(self):
@@ -361,7 +345,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*(1,2)}${a2[]=*(1,3)}<ppp:if a1[] le a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_AgeA(self):
@@ -370,7 +354,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*(1,3)}${a2[]=*(1,2)}<ppp:if a1[] ge a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_AinA(self):
@@ -379,7 +363,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*('hello')}${a2[]=*('hello', 'world')}<ppp:if a1[] in a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_AcontainsA(self):
@@ -388,7 +372,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*('hello','world')}${a2[]=*('hello')}<ppp:if a1[] contains a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
## A vs R
@@ -399,7 +383,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*('hello','world')}${r2=hello}<ppp:if a1[] eq r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp="nostrict",
)
@@ -409,7 +393,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*('hello')}${r2=bye}<ppp:if a1[] ne r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
ppp="nostrict",
)
@@ -419,7 +403,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*(1,2,3)}${r2=2}<ppp:if a1[] lt r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp="nostrict",
)
@@ -429,7 +413,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*(2,3,4)}${r2=2}<ppp:if a1[] gt r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp="nostrict",
)
@@ -439,7 +423,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*(1,2)}${r2=2}<ppp:if a1[] le r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
ppp="nostrict",
)
@@ -449,7 +433,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*(1,3)}${r2=2}<ppp:if a1[] ge r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp="nostrict",
)
@@ -459,7 +443,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*('hello', 'world')}${r2=hello world)}<ppp:if a1[] in r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_AcontainsR(self):
@@ -468,7 +452,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*('hello','world')}${r2=hello}<ppp:if a1[] contains r2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
## R vs A
@@ -479,7 +463,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=hello}${a2[]=*('hello','world')}<ppp:if r1 eq a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp="nostrict",
)
@@ -489,7 +473,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=bye}${a2[]=*('hello')}<ppp:if r1 ne a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
ppp="nostrict",
)
@@ -499,7 +483,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=2}${a2[]=*(1,2,3)}<ppp:if r1 lt a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp="nostrict",
)
@@ -509,7 +493,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=2}${a2[]=*(2,3,4)}<ppp:if r1 gt a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp="nostrict",
)
@@ -519,7 +503,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=2}${a2[]=*(1,2)}<ppp:if r1 le a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp="nostrict",
)
@@ -529,7 +513,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=2}${a2[]=*(1,3)}<ppp:if r1 ge a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp="nostrict",
)
@@ -539,7 +523,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=hello}${a2[]=*('hello', 'world')}<ppp:if r1 in a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_RcontainsA(self):
@@ -548,7 +532,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=hello world}${a2[]=*('hello','world')}<ppp:if r1 contains a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
## R vs V
@@ -559,7 +543,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=hello}<ppp:if r1 eq 'hello'>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_ReqV_str_fail(self):
@@ -568,7 +552,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=hello}<ppp:if r1 eq 42>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
interrupted=True,
)
@@ -578,7 +562,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=42}<ppp:if r1 eq 42>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_ReqV_num_fail(self):
@@ -587,7 +571,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=42}<ppp:if r1 eq '42'>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
interrupted=True,
)
@@ -597,7 +581,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=true}<ppp:if r1 eq true>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_operator_ReqV_bool_fail(self):
@@ -606,7 +590,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${r1=true}<ppp:if r1 eq 'true'>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
interrupted=True,
)
@@ -618,7 +602,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a1[]=*('hello','world')}<ppp:if a1[] in ('hello','world')>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
def test_listoperand_LinA(self):
@@ -627,7 +611,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a2[]=*('hello','world')}<ppp:if ('hello','world') in a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
# Indexed operands
@@ -638,7 +622,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a2[]=*('hello','world')}<ppp:if a2[0] in a2[]>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
# Float values
@@ -649,7 +633,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${a=1.5}<ppp:if a gt 1>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("OK", ""),
OutputTuple("OK", ""),
)
@@ -661,7 +645,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:if undefined_var gt 0>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
self.def_env_info,
@@ -682,7 +666,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:if undefined_var gt 0>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("", ""),
OutputTuple("", ""),
interrupted=True,
)
@@ -692,7 +676,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set myvar>abc<ppp:/set><ppp:if myvar gt 0>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
self.def_env_info,
@@ -713,7 +697,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set myvar>abc<ppp:/set><ppp:if myvar gt 0>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("", ""),
OutputTuple("", ""),
interrupted=True,
)
@@ -723,7 +707,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set myvar><ppp:/set><ppp:if myvar gt 0>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("not OK", ""),
OutputTuple("not OK", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
self.def_env_info,
@@ -746,7 +730,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"[<ppp:stn>neg5<ppp:/stn>] this \\(is\\): a (([complex|simple<ppp:stn>neg6<ppp:/stn>|regular] test<ppp:stn>neg1<ppp:/stn>)(test:2.0):1.5) \nBREAK, BREAK with [abc<ppp:stn>neg4<ppp:/stn>:def<ppp:stn p0>neg2(neg3:1.6)<ppp:/stn>:5]:0.5 AND loratrigger <lora:xxx:1> AND AND hypernettrigger <hypernet:yyy>:0.3",
"normal quality, <ppp:stn i0/>",
),
PromptPair(
OutputTuple(
"this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger <lora:xxx:1> AND hypernettrigger <hypernet:yyy>:0.3",
"[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]",
),
@@ -758,7 +742,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK, BREAK <ppp:if _is_sd1>with [abc<ppp:stn>neg4<ppp:/stn>:def:5]<ppp:/if>:0.5 AND <ppp:if _is_sd1>loratrigger <lora:xxx:1><ppp:elif _is_sdxl>hypernettrigger <hypernet:yyy><ppp:else>nothing<ppp:/if>:0.3",
"normal quality",
),
PromptPair(
OutputTuple(
"this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK :0.5 AND hypernettrigger <hypernet:yyy>:0.3",
"normal quality",
),
@@ -770,7 +754,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"this is <ppp:if _sd eq 'sd1'>SD1<ppp:else><ppp:if _is_pony>PONY<ppp:else>SD2<ppp:/if><ppp:/if><ppp:if _is_sdxl_no_pony>NOPONY<ppp:/if><ppp:if _is_pure_sdxl>NOPONY<ppp:/if>",
"",
),
PromptPair("this is PONY", ""),
OutputTuple("this is PONY", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -788,19 +772,19 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_cmd_set_if(self): # set and if commands
self.process(
PromptPair("<ppp:set v>value<ppp:/set>this test is <ppp:if v>OK<ppp:else>not OK<ppp:/if>", ""),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_empty(self): # set to empty
self.process(
PromptPair("<ppp:set v><ppp:/set>${v2=}this test is <ppp:if v or v2>not OK<ppp:else>OK<ppp:/if>", ""),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_eval_if(self): # set and if commands
self.process(
PromptPair("<ppp:set v evaluate>value<ppp:/set>this test is <ppp:if v>OK<ppp:else>not OK<ppp:/if>", ""),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_if_echo_nested(self): # nested set, if and echo commands
@@ -809,7 +793,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set v1>1<ppp:/set><ppp:if v1 gt 0><ppp:set v2>OK<ppp:/set><ppp:/if><ppp:if v2 eq 'OK'><ppp:echo v2/><ppp:else>not OK<ppp:/if> <ppp:echo v2>NOK<ppp:/echo> <ppp:echo v3>OK<ppp:/echo>",
"",
),
PromptPair("OK OK OK", ""),
OutputTuple("OK OK OK", ""),
)
def test_cmd_set_if_complex_conditions_1(self): # complex conditions (or)
@@ -818,7 +802,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set v1>true<ppp:/set><ppp:set v2>false<ppp:/set>this test is <ppp:if v1 or v2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_if_complex_conditions_2(self): # complex conditions (and)
@@ -827,13 +811,13 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set v1>true<ppp:/set><ppp:set v2>true<ppp:/set>this test is <ppp:if v1 and v2>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_if_complex_conditions_3(self): # complex conditions (not)
self.process(
PromptPair("<ppp:set v1>false<ppp:/set>this test is <ppp:if not v1>OK<ppp:else>not OK<ppp:/if>", ""),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_if_complex_conditions_4(self): # complex conditions (not, precedence)
@@ -842,7 +826,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set v1>true<ppp:/set><ppp:set v2>false<ppp:/set>this test is <ppp:if not (v1 and v2)>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_if_complex_conditions_5(self): # complex conditions (not, precedence, comparison)
@@ -851,7 +835,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set v1>1<ppp:/set><ppp:set v2>false<ppp:/set>this test is <ppp:if not(v1 eq 1 and v2)>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_if_complex_conditions_6(self): # complex conditions
@@ -860,7 +844,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set v1>1<ppp:/set><ppp:set v2>2<ppp:/set><ppp:set v3>3<ppp:/set>this test is <ppp:if v1 eq 1 and v2 eq 2 and v3 eq 3>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_if_complex_conditions_7(self): # complex conditions
@@ -869,7 +853,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set v1>1<ppp:/set><ppp:set v2>2<ppp:/set><ppp:set v3>3<ppp:/set>this test is <ppp:if v1 eq 1 and v2 not eq 2 or v3 eq 3>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_if2(self): # set and more complex if commands
@@ -878,7 +862,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"First: <ppp:set v>value1<ppp:/set>this test is <ppp:if v in ('value1','value2')>OK<ppp:elif v in ('value3')>OK2<ppp:else>not OK<ppp:/if>\nSecond: <ppp:set v2>value3<ppp:/set>this test is <ppp:if not v2 in ('value1','value2')>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("First: this test is OK\nSecond: this test is OK", ""),
OutputTuple("First: this test is OK\nSecond: this test is OK", ""),
)
def test_cmd_set_add_if(self): # set, add and if commands
@@ -887,7 +871,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set v>value<ppp:/set><ppp:set v add>2<ppp:/set>this test is <ppp:if v eq 'value2'>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_add_DP_if(self): # set, add (DP format) and if commands
@@ -896,7 +880,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v=value}${v+=2}this test is <ppp:if v eq 'value2'>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_immediateeval(self): # set (DP format) with mixed evaluation
@@ -905,7 +889,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${var=!__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var2:default}, ${var3=__yaml/wildcard1__}${var3}, ${var3}",
"",
),
PromptPair("the choices are: choice2, choice2, default, choice3, choice1", ""),
OutputTuple("the choices are: choice2, choice2, default, choice3, choice1", ""),
ppp="nocup",
)
@@ -915,7 +899,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${var=__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var+=, __yaml/wildcard2__}${var}, ${var}, ${var+=!, __yaml/wildcard3__}${var}, ${var}",
"",
),
PromptPair(
OutputTuple(
"the choices are: choice2, choice3, choice1, choice1- choice2 -choice3, choice2, choice2 -choice1-choice3, choice2, choice3-choice1- choice2 , choice1, choice2 , choice2, choice3-choice1- choice2 , choice1, choice2 ",
"",
),
@@ -928,7 +912,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set v ifundefined>value<ppp:/set>this test is <ppp:if v eq 'value'>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_ifundefined_if_2(self): # set, ifundefined and if commands
@@ -937,7 +921,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:set v>value<ppp:/set><ppp:set v ifundefined>value2<ppp:/set>this test is <ppp:if v eq 'value'>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_ifundefined_DP_if(self): # set, ifundefined (DP format) and if commands
@@ -946,7 +930,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v?=value}this test is <ppp:if v eq 'value'>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_set_ifundefined_DP_if_2(self): # set, ifundefined (DP format) and if commands
@@ -955,7 +939,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${v=!value}${v?=!value2}this test is <ppp:if v eq 'value'>OK<ppp:else>not OK<ppp:/if>",
"",
),
PromptPair("this test is OK", ""),
OutputTuple("this test is OK", ""),
)
def test_cmd_echo_sysvar(self):
@@ -964,7 +948,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"${_model:defval}",
"",
),
PromptPair("sdxl", ""),
OutputTuple("sdxl", ""),
)
def test_cmd_ext(self): # ext
@@ -973,7 +957,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:ext lora lora1name if not _is_pony>trigger1<ppp:/ext><ppp:ext lora 'lora2 name' -0.8 if not _is_pony>trigger2<ppp:/ext><ppp:ext lora lora3__name '0.5:0.8' if not _is_pony><ppp:ext lora lora4name>trigger4<ppp:/ext><ppp:ext lora \"lora5 (name)\" 1/>trigger5",
"",
),
PromptPair(
OutputTuple(
"<lora:lora1name:1>trigger1,<lora:lora2 name:-0.8>trigger2,<lora:lora3__name:0.5:0.8><lora:lora4name:1>trigger4,<lora:lora5 (name):1>trigger5",
"",
),
@@ -985,7 +969,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:ext $lora lora1/><ppp:ext $lora lora1>",
"",
),
PromptPair("triggergeneric1, triggergeneric2, two, triggergeneric1, triggergeneric2, two", ""),
OutputTuple("triggergeneric1, triggergeneric2, two, triggergeneric1, triggergeneric2, two", ""),
)
def test_cmd_ext_map1(self): # ext mapping, no lora
@@ -994,7 +978,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:ext $lora lora1>inlinetrigger<ppp:/ext>",
"",
),
PromptPair("inlinetrigger, triggergeneric1, triggergeneric2, two", ""),
OutputTuple("inlinetrigger, triggergeneric1, triggergeneric2, two", ""),
)
def test_cmd_ext_map2(self): # ext mapping, lora with weight
@@ -1003,7 +987,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:ext $lora lora1>inlinetrigger<ppp:/ext>",
"",
),
PromptPair("<lora:lorapony:0.8>inlinetrigger, triggerpony1, triggerpony2", ""),
OutputTuple("<lora:lorapony:0.8>inlinetrigger, triggerpony1, triggerpony2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -1024,7 +1008,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:ext $lora lora1 0.5>inlinetrigger<ppp:/ext>",
"",
),
PromptPair("<lora:lorapony:0.4>inlinetrigger, triggerpony1, triggerpony2", ""),
OutputTuple("<lora:lorapony:0.4>inlinetrigger, triggerpony1, triggerpony2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -1045,7 +1029,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:ext $lora lora1 '0.6:0.8'>inlinetrigger<ppp:/ext>",
"",
),
PromptPair("<lora:lorapony:0.6:0.8>inlinetrigger, triggerpony1, triggerpony2", ""),
OutputTuple("<lora:lorapony:0.6:0.8>inlinetrigger, triggerpony1, triggerpony2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
@@ -1066,7 +1050,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"<ppp:ext $lora lora1>inlinetrigger<ppp:/ext>",
"",
),
PromptPair("<lora:loraillustrious:0.9:0.8>inlinetrigger, triggerillustrious1, triggerillustrious2", ""),
OutputTuple("<lora:loraillustrious:0.9:0.8>inlinetrigger, triggerillustrious1, triggerillustrious2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
+3 -3
View File
@@ -1,8 +1,8 @@
from dataclasses import replace
from ppp import PromptPostProcessor
from ppp_classes import ONWARNING_CHOICES # pylint: disable=import-error
from .base_tests import PromptPair, TestPromptPostProcessorBase
from ppp_classes import ONWARNING_CHOICES # type: ignore
from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -21,7 +21,7 @@ class TestModelVariants(TestPromptPostProcessorBase):
"<ppp:if _is_test1>test1<ppp:/if><ppp:if _is_test2>test2<ppp:/if><ppp:if _is_test3>test3<ppp:/if><ppp:if _is_test4>test4<ppp:/if>",
"",
),
PromptPair("test1test2", ""),
OutputTuple("test1test2", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
{
+179 -49
View File
@@ -2,7 +2,7 @@ from dataclasses import replace
from ppp import PromptPostProcessor
from ppp_classes import IFWILDCARDS_CHOICES
from .base_tests import PromptPair, TestPromptPostProcessorBase
from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -18,7 +18,7 @@ class TestWildcards(TestPromptPostProcessorBase):
def test_wc_ignore(self): # wildcards with ignore option
self.process(
PromptPair("__bad_wildcard__", "{option1|option2}"),
PromptPair("__bad_wildcard__", "{option1|option2}"),
OutputTuple("__bad_wildcard__", "{option1|option2}"),
ppp=PromptPostProcessor(
self.ppp_logger,
self.def_env_info,
@@ -40,7 +40,7 @@ class TestWildcards(TestPromptPostProcessorBase):
"[<ppp:stn>neg5<ppp:/stn>] this is: __bad_wildcard__ a (([complex|simple<ppp:stn>neg6<ppp:/stn>|regular] test<ppp:stn>neg1<ppp:/stn>)(test:2.0):1.5) \nBREAK, BREAK with [abc<ppp:stn>neg4<ppp:/stn>:def<ppp:stn p0>neg2(neg3:1.6)<ppp:/stn>:5] <lora:xxx:1>",
"normal quality, <ppp:stn i0/> {option1|option2}",
),
PromptPair(
OutputTuple(
"this is: a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]<lora:xxx:1>",
"[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]",
),
@@ -62,7 +62,7 @@ class TestWildcards(TestPromptPostProcessorBase):
def test_wc_warn(self): # wildcards with warn option
self.process(
PromptPair("__bad_wildcard__", "{option1|option2}"),
PromptPair(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", "{option1|option2}"),
OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", "{option1|option2}"),
ppp=PromptPostProcessor(
self.ppp_logger,
self.def_env_info,
@@ -81,7 +81,7 @@ class TestWildcards(TestPromptPostProcessorBase):
def test_wc_stop(self): # wildcards with stop option
self.process(
PromptPair("__bad_wildcard__", "{option1|option2}"),
PromptPair(
OutputTuple(
PromptPostProcessor.WILDCARD_STOP.format("__bad_wildcard__") + "__bad_wildcard__",
"{option1|option2}",
),
@@ -104,7 +104,7 @@ class TestWildcards(TestPromptPostProcessorBase):
def test_wcinvar_warn(self): # wildcards in var with warn option
self.process(
PromptPair("${v=__bad_wildcard__}${v}", ""),
PromptPair(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", ""),
OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
self.def_env_info,
@@ -123,154 +123,154 @@ class TestWildcards(TestPromptPostProcessorBase):
def test_wc_invalid_name(self):
self.process(
PromptPair("the choices are: ___invalid__", ""),
PromptPair("the choices are: ___invalid__", ""),
OutputTuple("the choices are: ___invalid__", ""),
ppp="nocup",
)
def test_wc_wildcard1a_text(self): # simple text wildcard
self.process(
PromptPair("the choices are: __text/wildcard1__", ""),
PromptPair("the choices are: choice2", ""),
OutputTuple("the choices are: choice2", ""),
ppp="nocup",
)
def test_wc_wildcard1a_json(self): # simple json wildcard
self.process(
PromptPair("the choices are: __json/wildcard1__", ""),
PromptPair("the choices are: choice2", ""),
OutputTuple("the choices are: choice2", ""),
ppp="nocup",
)
def test_wc_wildcard1a_yaml(self): # simple yaml wildcard
self.process(
PromptPair("the choices are: __yaml/wildcard1__", ""),
PromptPair("the choices are: choice2", ""),
OutputTuple("the choices are: choice2", ""),
ppp="nocup",
)
def test_wc_wildcard1b_text(self): # simple text wildcard with multiple choices
self.process(
PromptPair("the choices are: __2-$$text/wildcard1__", ""),
PromptPair("the choices are: choice3, choice1", ""),
OutputTuple("the choices are: choice3, choice1", ""),
ppp="nocup",
)
def test_wc_wildcard1b_json(self): # simple json wildcard with multiple choices
self.process(
PromptPair("the choices are: __2-$$json/wildcard1__", ""),
PromptPair("the choices are: choice3, choice1", ""),
OutputTuple("the choices are: choice3, choice1", ""),
ppp="nocup",
)
def test_wc_wildcard1b_yaml(self): # simple yaml wildcard with multiple choices
self.process(
PromptPair("the choices are: __2-$$yaml/wildcard1__", ""),
PromptPair("the choices are: choice3, choice1", ""),
OutputTuple("the choices are: choice3, choice1", ""),
ppp="nocup",
)
def test_wc_wildcard2_text(self): # simple text wildcard with default options
self.process(
PromptPair("the choices are: __text/wildcard2__", ""),
PromptPair("the choices are: choice3-choice1", ""),
OutputTuple("the choices are: choice3-choice1", ""),
ppp="nocup",
)
def test_wc_wildcard2_json(self): # simple json wildcard with default options
self.process(
PromptPair("the choices are: __json/wildcard2__", ""),
PromptPair("the choices are: choice3-choice1", ""),
OutputTuple("the choices are: choice3-choice1", ""),
ppp="nocup",
)
def test_wc_wildcard2_yaml(self): # simple yaml wildcard with default options
self.process(
PromptPair("the choices are: __yaml/wildcard2__", ""),
PromptPair("the choices are: choice3-choice1", ""),
OutputTuple("the choices are: choice3-choice1", ""),
ppp="nocup",
)
def test_wc_test2_yaml(self): # simple yaml wildcard
self.process(
PromptPair("the choice is: __testwc/test2__", ""),
PromptPair("the choice is: 2", ""),
OutputTuple("the choice is: 2", ""),
ppp="nocup",
)
def test_wc_test3_yaml(self): # simple yaml wildcard
self.process(
PromptPair("the choice is: __testwc/test3__", ""),
PromptPair("the choice is: one choice", ""),
OutputTuple("the choice is: one choice", ""),
ppp="nocup",
)
def test_wc_wildcard_filter_index(self): # wildcard with positional index filter
self.process(
PromptPair("the choice is: __yaml/wildcard2'2'__", ""),
PromptPair("the choice is: choice3-choice3", ""),
OutputTuple("the choice is: choice3-choice3", ""),
ppp="nocup",
)
def test_wc_wildcard_filter_index_range(self): # wildcard with positional index range filter
self.process(
PromptPair("the choice is: __yaml/wildcard2'2-3'__", ""),
PromptPair("the choice is: choice3-choice3", ""),
OutputTuple("the choice is: choice3-choice3", ""),
ppp="nocup",
)
def test_wc_wildcard_filter_label(self): # wildcard with label filter
self.process(
PromptPair("the choice is: __yaml/wildcard2'label1'__", ""),
PromptPair("the choice is: choice3-choice1", ""),
OutputTuple("the choice is: choice3-choice1", ""),
ppp="nocup",
)
def test_wc_wildcard_filter_label2(self): # wildcard with label filter in multiple choices
self.process(
PromptPair("the choice is: __yaml/wildcard2'label2'__", ""),
PromptPair("the choice is: choice1-choice1", ""),
OutputTuple("the choice is: choice1-choice1", ""),
ppp="nocup",
)
def test_wc_wildcard_filter_label3(self): # wildcard with multiple label filter
self.process(
PromptPair("the choice is: __yaml/wildcard2'label1,label2'__", ""),
PromptPair("the choice is: choice3-choice1", ""),
OutputTuple("the choice is: choice3-choice1", ""),
ppp="nocup",
)
def test_wc_wildcard_filter_indexlabel(self): # wildcard with mixed index and label filter
self.process(
PromptPair("the choice is: __yaml/wildcard2'2,label2'__", ""),
PromptPair("the choice is: choice3-choice1", ""),
OutputTuple("the choice is: choice3-choice1", ""),
ppp="nocup",
)
def test_wc_wildcard_filter_compound(self): # wildcard with compound filter
self.process(
PromptPair("the choice is: __yaml/wildcard2'label1+label3'__", ""),
PromptPair("the choice is: choice3-choice1", ""),
OutputTuple("the choice is: choice3-choice1", ""),
ppp="nocup",
)
def test_wc_wildcard_filter_compound2(self): # wildcard with inherited compound filter
self.process(
PromptPair("the choice is: __yaml/wildcard2bis'#label1+label3'__", ""),
PromptPair("the choice is: choice3bis", ""),
OutputTuple("the choice is: choice3bis", ""),
ppp="nocup",
)
def test_wc_wildcard_filter_compound3(self): # wildcard with doubly inherited compound filter
self.process(
PromptPair("the choice is: __yaml/wildcard2bisbis'#label1+label3'__", ""),
PromptPair("the choice is: choice1bisbis", ""),
OutputTuple("the choice is: choice1bisbis", ""),
ppp="nocup",
)
def test_wc_wildcard_filter_compound4(self): # wildcard with doubly inherited compound filter with variable
self.process(
PromptPair("${v=label1}the choice is: __yaml/wildcard2bisbis'#${v}+label3'__", ""),
PromptPair("the choice is: choice1bisbis", ""),
OutputTuple("the choice is: choice1bisbis", ""),
ppp="nocup",
)
@@ -280,7 +280,7 @@ class TestWildcards(TestPromptPostProcessorBase):
"<ppp:setwcdeffilter 'yaml/wildcard2' 'label1+label3' />the choice is: __yaml/wildcard2__, <ppp:setwcdeffilter 'yaml/wildcard2' />__yaml/wildcard2__",
"",
),
PromptPair("the choice is: choice3-choice1, choice3-choice1- choice2 ", ""),
OutputTuple("the choice is: choice3-choice1, choice3-choice1- choice2 ", ""),
ppp="nocup",
)
@@ -290,49 +290,49 @@ class TestWildcards(TestPromptPostProcessorBase):
"${v=label1}<ppp:setwcdeffilter 'yaml/wildcard2' '${v}+label3' />the choice is: __yaml/wildcard2__, <ppp:setwcdeffilter 'yaml/wildcard2' />__yaml/wildcard2__",
"",
),
PromptPair("the choice is: choice3-choice1, choice3-choice1- choice2 ", ""),
OutputTuple("the choice is: choice3-choice1, choice3-choice1- choice2 ", ""),
ppp="nocup",
)
def test_wc_nested_wildcard_text(self): # nested text wildcard with repeating multiple choices
self.process(
PromptPair("the choices are: __r3$$-$$text/wildcard3__", ""),
PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""),
OutputTuple("the choices are: choice3,choice1- choice2 ,choice3", ""),
ppp="nocup",
)
def test_wc_nested_wildcard_json(self): # nested json wildcard with repeating multiple choices
self.process(
PromptPair("the choices are: __r3$$-$$json/wildcard3__", ""),
PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""),
OutputTuple("the choices are: choice3,choice1- choice2 ,choice3", ""),
ppp="nocup",
)
def test_wc_nested_wildcard_yaml(self): # nested yaml wildcard with repeating multiple choices
self.process(
PromptPair("the choices are: __r3$$-$$yaml/wildcard3__", ""),
PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""),
OutputTuple("the choices are: choice3,choice1- choice2 ,choice3", ""),
ppp="nocup",
)
def test_wc_wildcard_optional(self): # empty wildcard with no error
self.process(
PromptPair("the choices are: __yaml/empty_wildcard__", ""),
PromptPair("the choices are: ", ""),
OutputTuple("the choices are: ", ""),
ppp="nocup",
)
def test_wc_wildcard4_yaml(self): # simple yaml wildcard with one option
self.process(
PromptPair("the choices are: __yaml/wildcard4__", ""),
PromptPair("the choices are: inline text", ""),
OutputTuple("the choices are: inline text", ""),
ppp="nocup",
)
def test_wc_wildcard6_yaml(self): # simple yaml wildcard with object formatted choices
self.process(
PromptPair("the choices are: __yaml/wildcard6__", ""),
PromptPair("the choices are: choice2", ""),
OutputTuple("the choices are: choice2", ""),
ppp="nocup",
)
@@ -340,9 +340,9 @@ class TestWildcards(TestPromptPostProcessorBase):
self.process(
PromptPair("the choices are: {__~2$$yaml/wildcard2__|choice0}", ""),
[
PromptPair("the choices are: choice0", ""),
PromptPair("the choices are: choice1, choice3", ""),
PromptPair("the choices are: choice1, choice3", ""),
OutputTuple("the choices are: choice0", ""),
OutputTuple("the choices are: choice1, choice3", ""),
OutputTuple("the choices are: choice1, choice3", ""),
],
ppp="nocup",
)
@@ -350,7 +350,7 @@ class TestWildcards(TestPromptPostProcessorBase):
def test_wc_unsupportedsampler(self): # unsupported sampler
self.process(
PromptPair("the choices are: __@yaml/wildcard2__", ""),
PromptPair("", ""),
OutputTuple("", ""),
ppp="nocup",
interrupted=True,
)
@@ -358,42 +358,42 @@ class TestWildcards(TestPromptPostProcessorBase):
def test_wc_wildcard_globbing(self): # wildcard with globbing
self.process(
PromptPair("the choices are: __yaml/wildcard[12]__, __yaml/wildcard?__", ""),
PromptPair("the choices are: choice3-choice2, <lora:test2:1>- choice2 -choice3", ""),
OutputTuple("the choices are: choice3-choice2, <lora:test2:1>- choice2 -choice3", ""),
ppp="nocup",
)
def test_wc_wildcardwithvar(self): # wildcard with inline variable
self.process(
PromptPair("the choices are: __yaml/wildcard5(var=test)__, __yaml/wildcard5__", ""),
PromptPair("the choices are: inline test, inline default", ""),
OutputTuple("the choices are: inline test, inline default", ""),
ppp="nocup",
)
def test_wc_wildcardPS_yaml(self): # yaml wildcard with object formatted choices and options and prefix and suffix
self.process(
PromptPair("the choices are: __yaml/wildcardPS__", ""),
PromptPair("the choices are: prefix-choice2/choice3-suffix", ""),
OutputTuple("the choices are: prefix-choice2/choice3-suffix", ""),
ppp="nocup",
)
def test_wc_anonymouswildcard_yaml(self): # yaml anonymous wildcard
self.process(
PromptPair("the choices are: __yaml/anonwildcards__", ""),
PromptPair("the choices are: six", ""),
OutputTuple("the choices are: six", ""),
ppp="nocup",
)
def test_wc_wildcard_input(self): # simple yaml wildcard input
self.process(
PromptPair("the choices are: __yaml_input/wildcardI__", ""),
PromptPair("the choices are: choice2", ""),
OutputTuple("the choices are: choice2", ""),
ppp="nocup",
)
def test_wc_circular(self): # wildcard circular reference
self.process(
PromptPair("the choices are: __yaml/circular1__", ""),
PromptPair("", ""),
OutputTuple("", ""),
ppp="nocup",
interrupted=True,
)
@@ -401,14 +401,14 @@ class TestWildcards(TestPromptPostProcessorBase):
def test_wc_including(self): # wildcard including another wildcard
self.process(
PromptPair("the choices are: __yaml/including__", ""),
PromptPair("the choices are: choice4", ""),
OutputTuple("the choices are: choice4", ""),
ppp="nocup",
)
def test_wc_circular_including(self): # wildcard including another wildcard in a circular reference
self.process(
PromptPair("the choices are: __yaml/including1__", ""),
PromptPair("", ""),
OutputTuple("", ""),
ppp="nocup",
interrupted=True,
)
@@ -419,6 +419,136 @@ class TestWildcards(TestPromptPostProcessorBase):
"the choices are: ${x={1|2|3}}${w=yaml/wildcard${x}}__yaml/wildcard${x}__ __${w}__ __<ppp:echo w/>__",
"",
),
PromptPair("the choices are: choice1-choice3-choice1 choice3- choice2 - choice2 choice3", ""),
OutputTuple("the choices are: choice1-choice3-choice1 choice3- choice2 - choice2 choice3", ""),
ppp="nocup",
)
# Combinatorial
def test_wc_combinatorial_1(self): # combinatorial wildcard with variable
self.process_combinatorial(
PromptPair("the choices are: __2$$yaml/wildcard2__, ${v:{option1|option2}}", ""),
[ # 12 combinations
OutputTuple("the choices are: choice1, choice2, option1", "", {"v": "option1"}),
OutputTuple("the choices are: choice1, choice2, option2", "", {"v": "option2"}),
OutputTuple("the choices are: choice1, choice3, option1", "", {"v": "option1"}),
OutputTuple("the choices are: choice1, choice3, option2", "", {"v": "option2"}),
OutputTuple("the choices are: choice2, choice1, option1", "", {"v": "option1"}),
OutputTuple("the choices are: choice2, choice1, option2", "", {"v": "option2"}),
OutputTuple("the choices are: choice2, choice3, option1", "", {"v": "option1"}),
OutputTuple("the choices are: choice2, choice3, option2", "", {"v": "option2"}),
OutputTuple("the choices are: choice3, choice1, option1", "", {"v": "option1"}),
OutputTuple("the choices are: choice3, choice1, option2", "", {"v": "option2"}),
OutputTuple("the choices are: choice3, choice2, option1", "", {"v": "option1"}),
OutputTuple("the choices are: choice3, choice2, option2", "", {"v": "option2"}),
],
)
def test_wc_combinatorial_2(self): # combinatorial wildcard
self.process_combinatorial(
PromptPair("__yaml/wildcard2__", ""),
[ # 36 combinations
# groups of 3
## same choice repeated 3 times
OutputTuple("choice1-choice1-choice1", ""),
OutputTuple(" choice2 - choice2 - choice2 ", ""),
OutputTuple("choice3-choice3-choice3", ""),
## one choice repeated 2 times in all positions
OutputTuple("choice1-choice1- choice2 ", ""),
OutputTuple("choice1-choice1-choice3", ""),
OutputTuple(" choice2 - choice2 -choice1", ""),
OutputTuple(" choice2 - choice2 -choice3", ""),
OutputTuple("choice3-choice3-choice1", ""),
OutputTuple("choice3-choice3- choice2 ", ""),
OutputTuple(" choice2 -choice1-choice1", ""),
OutputTuple("choice3-choice1-choice1", ""),
OutputTuple("choice1- choice2 - choice2 ", ""),
OutputTuple("choice3- choice2 - choice2 ", ""),
OutputTuple("choice1-choice3-choice3", ""),
OutputTuple(" choice2 -choice3-choice3", ""),
OutputTuple("choice1- choice2 -choice1", ""),
OutputTuple("choice1-choice3-choice1", ""),
OutputTuple(" choice2 -choice1- choice2 ", ""),
OutputTuple(" choice2 -choice3- choice2 ", ""),
OutputTuple("choice3-choice1-choice3", ""),
OutputTuple("choice3- choice2 -choice3", ""),
## choices 1, 2, 3 in all positions
OutputTuple("choice1- choice2 -choice3", ""),
OutputTuple("choice1-choice3- choice2 ", ""),
OutputTuple(" choice2 -choice1-choice3", ""),
OutputTuple(" choice2 -choice3-choice1", ""),
OutputTuple("choice3-choice1- choice2 ", ""),
OutputTuple("choice3- choice2 -choice1", ""),
# groups of 2
## same choice repeated 2 times
OutputTuple("choice1-choice1", ""),
OutputTuple(" choice2 - choice2 ", ""),
OutputTuple("choice3-choice3", ""),
## choices 1 and 2 in all positions
OutputTuple("choice1- choice2 ", ""),
OutputTuple(" choice2 -choice1", ""),
## choices 2 and 3 in all positions
OutputTuple(" choice2 -choice3", ""),
OutputTuple("choice3- choice2 ", ""),
## choices 1 and 3 in all positions
OutputTuple("choice1-choice3", ""),
OutputTuple("choice3-choice1", ""),
],
ppp="nocup",
)
def test_wc_combinatorial_3(self): # combinatorial wildcard (keep choice order)
self.process_combinatorial(
PromptPair("__2-3$$-$$yaml/wildcard2__", ""),
[ # 4 combinations
# groups of 3
## choices 1, 2, 3
OutputTuple("choice1- choice2 -choice3", ""),
# groups of 2
## choices 1 and 2
OutputTuple("choice1- choice2 ", ""),
## choices 2 and 3
OutputTuple(" choice2 -choice3", ""),
## choices 1 and 3
OutputTuple("choice1-choice3", ""),
],
ppp=PromptPostProcessor(
self.ppp_logger,
self.def_env_info,
replace(
self.defopts,
keep_choices_order=True,
cup_do_cleanup=False,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
),
)
def test_wc_combinatorial_4(self): # combinatorial wildcard (don't keep choice order)
self.process_combinatorial(
PromptPair("__2-3$$-$$yaml/wildcard2__", ""),
[ # 12 combinations
# groups of 3
## choices 1, 2, 3 in all positions
OutputTuple("choice1- choice2 -choice3", ""),
OutputTuple("choice1-choice3- choice2 ", ""),
OutputTuple(" choice2 -choice1-choice3", ""),
OutputTuple(" choice2 -choice3-choice1", ""),
OutputTuple("choice3-choice1- choice2 ", ""),
OutputTuple("choice3- choice2 -choice1", ""),
# groups of 2
## choices 1 and 2 in all positions
OutputTuple("choice1- choice2 ", ""),
OutputTuple(" choice2 -choice1", ""),
## choices 2 and 3 in all positions
OutputTuple(" choice2 -choice3", ""),
OutputTuple("choice3- choice2 ", ""),
## choices 1 and 3 in all positions
OutputTuple("choice1-choice3", ""),
OutputTuple("choice3-choice1", ""),
],
ppp="nocup",
)