* 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:
co-authored by
Copilot
parent
ad5cad30ff
commit
a31070b986
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
3.10.11
|
||||
Vendored
+7
@@ -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"
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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": []
|
||||
}
|
||||
}
|
||||
@@ -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, {}]
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||
{
|
||||
|
||||
@@ -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
@@ -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",
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user