diff --git a/.github/instructions/tests.instructions.md b/.github/instructions/tests.instructions.md
index 89a011b..ac48c69 100644
--- a/.github/instructions/tests.instructions.md
+++ b/.github/instructions/tests.instructions.md
@@ -31,33 +31,59 @@ Examples: `test_cl_simple`, `test_ch_choices`, `test_wc_ignore`
## Running a Test Case via `self.process()`
-Use the `process()` helper from the base class — never instantiate `PromptPostProcessor` directly in test methods.
+Use the `process()` helper from the base class — never instantiate `PromptPostProcessor` directly in test methods unless there is a need to pass specific options not covered by the default setup or the "nocup" or "nostrict" options.
```python
def test_cl_simple(self):
"""simple cleanup"""
self.process(
- "input prompt", # positive prompt
- "", # negative prompt
- PromptPair("expected output", ""), # expected result
+ InputTuple(
+ "input prompt", # positive prompt
+ ""), # negative prompt
+ OutputTuple(
+ "expected output", # expected positive prompt
+ "", # expected negative prompt
+ {} # expected variables (optional)
+ ),
)
+
+def test_cl_combinatorial(self):
+ """simple cleanup"""
+ self.process(
+ InputTuple(
+ "input prompt", # positive prompt
+ ""), # negative prompt
+ [
+ OutputTuple(
+ "expected output", # expected positive prompt
+ "", # expected negative prompt
+ {} # expected variables (optional)
+ ),
+ OutputTuple(
+ "expected output", # expected positive prompt
+ "", # expected negative prompt
+ {} # expected variables (optional)
+ ),
+ ],
+ combinatorial=True,
+ )
+
```
### `process()` Signature (key parameters)
| Parameter | Type | Notes |
|-----------|------|-------|
-| `input_prompt` | `str` | Positive prompt input |
-| `input_negative_prompt` | `str` | Negative prompt input |
-| `expected_output` | `PromptPair \| list[PromptPair]` | Single or multiple valid outputs |
+| `input_prompts` | `InputTuple` | Prompts input |
+| `expected_output` | `OutputTuple \| list[OutputTuple]` | Single or multiple valid outputs |
| `seed` | `int` | Optional, defaults to fixed seed |
-| `ppp` | `PromptPostProcessor \| str \| None` | Pass `"nocup"` to skip creation |
+| `ppp` | `PromptPostProcessor \| str \| None` | Supported values `"nocup"`, `"nostrict"` or a specific instance |
| `interrupted` | `bool` | Expected interrupt flag |
-| `output_variables` | `dict[str, str] \| None` | Variables to validate after processing |
+| `combinatorial` | `bool` | Whether to run a combinatorial generation. If a specific ppp instance is used then it is ignored |
## Assertions
-Use `assertEqual` with a descriptive message string:
+Use `assertEqual` or similar methods with a descriptive message string:
```python
self.assertEqual(result, expected, "Descriptive failure message")
@@ -73,8 +99,24 @@ Override `self.defopts` or `self.def_env_info` to pass non-default options — d
```python
def test_cl_custom(self):
"""cleanup with custom separator"""
- opts = {**self.defopts, "ppp_stn_separator": " | "}
- self.process("a, , b", "", PromptPair("a | b", ""), ppp_opts=opts)
+ self.process(
+ InputTuple("a, , b", ""),
+ OutputTuple("a | b", ""),
+ ppp=PromptPostProcessor(
+ self.ppp_logger,
+ self.def_env_info,
+ replace(
+ self.defopts,
+ keep_choices_order=True,
+ cup_do_cleanup=False,
+ do_combinatorial=True,
+ ),
+ self.grammar_content,
+ self.interrupt,
+ self.wildcards_obj,
+ self.extranetwork_maps_obj,
+ ),
+ )
```
## Entry Point
diff --git a/.gitignore b/.gitignore
index c17d6eb..0651593 100644
--- a/.gitignore
+++ b/.gitignore
@@ -6,8 +6,7 @@ venv
!.vscode/settings.json
!.vscode/launch.json
+logs
tests/tests_local.py
tests/local_wildcards
tests/logs
-
-scripts/last_prompts.txt
diff --git a/README.md b/README.md
index f76bd7f..3c90f10 100644
--- a/README.md
+++ b/README.md
@@ -78,7 +78,7 @@ 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.
+To develop, I suggest doing so with the extension isolated from the UI (you can use a symlink to test it in the UI), and with its own virtual environment (venv or .venv), so the tests work and can be debugged properly.
## License
diff --git a/ppp.py b/ppp.py
index e632d6e..f825e3f 100644
--- a/ppp.py
+++ b/ppp.py
@@ -84,6 +84,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS = defopt["cup_remove_extranetwork_tags"]
DEFAULT_STRICT_OPERATORS = defopt["strict_operators"]
DEFAULT_DO_COMBINATORIAL = defopt["do_combinatorial"]
+ DEFAULT_COMBINATORIAL_SHUFFLE = defopt["combinatorial_shuffle"]
DEFAULT_COMBINATORIAL_LIMIT = defopt["combinatorial_limit"]
WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK '
WILDCARD_STOP = "INVALID WILDCARD! {0}\nBREAK "
@@ -882,7 +883,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
v.update(variables)
return prompt, negative_prompt, v
- def __processprompts(self, rng, prompt, negative_prompt) -> list[tuple[str, str, dict[str, str | None]]]:
+ def __processprompts(
+ self, rng: np.random.Generator, prompt: str, negative_prompt: str
+ ) -> list[tuple[str, str, dict[str, str | None]]]:
"""
Process the prompt and negative prompt.
@@ -914,6 +917,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
try:
results = processor.start_visit(parsed)
except PPPInterrupt as e:
+ results = []
self.log(logging.ERROR, e.message)
if e.pos_prefix:
prompt = e.pos_prefix + prompt
@@ -931,6 +935,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
final_results.append(self.__postprocess_result(r))
if self.state.options.do_combinatorial:
self.log(logging.INFO, f"Total combinations: {len(final_results)}")
+ if self.state.options.combinatorial_shuffle:
+ rng.shuffle(final_results)
+ self.log(logging.INFO, "Combinations shuffled")
return final_results
def process_prompt(
@@ -974,7 +981,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, {}]
+ 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, {}]
+ return [(original_prompt, original_negative_prompt, {})]
diff --git a/ppp_cache.py b/ppp_cache.py
index 44ec45f..2e15cdf 100644
--- a/ppp_cache.py
+++ b/ppp_cache.py
@@ -7,7 +7,7 @@ from ppp_logging import DEBUG_LEVEL
class PPPLRUCache:
- ProcessInput = Tuple[int, int, str, str] # (seed, wildcards_hash, positive_prompt, negative_prompt)
+ ProcessInput = Tuple[int, int, str, str] # (env_hash, seed, positive_prompt, negative_prompt)
ProcessResult = Tuple[str, str] # (positive_prompt, negative_prompt)
def __init__(self, capacity: int, logger: Logger = None, debug_level: DEBUG_LEVEL = DEBUG_LEVEL.none):
diff --git a/ppp_classes.py b/ppp_classes.py
index d8905e0..b4ed3db 100644
--- a/ppp_classes.py
+++ b/ppp_classes.py
@@ -79,7 +79,7 @@ class ModelDetectConfig(BaseModel):
@model_validator(mode="after")
def check_class_or_property(self) -> "ModelDetectConfig":
if self.class_ is None and self.property is None:
- raise ValueError("either 'class' or 'property' must be specified")
+ raise ValueError("Either 'class' or 'property' must be specified")
return self
@@ -101,17 +101,17 @@ class FindInFilenamePattern(BaseModel):
flag_value = 0
for flag in v:
if not isinstance(flag, str) or not hasattr(re, flag):
- raise ValueError(f"invalid regex flag '{flag}'")
+ raise ValueError(f"Invalid regex flag '{flag}'")
flag_value |= getattr(re, flag)
return flag_value
- raise ValueError(f"expected int or list of flag-name strings, got {type(v).__name__}")
+ raise ValueError(f"Expected int or list of flag-name strings, got {type(v).__name__}")
@model_validator(mode="after")
def validate_regex(self) -> "FindInFilenamePattern":
try:
re.compile(self.regex, self.flags)
except re.error as exc:
- raise ValueError(f"invalid regex pattern '{self.regex}': {exc}") from exc
+ raise ValueError(f"Invalid regex pattern '{self.regex}': {exc}") from exc
return self
@@ -136,9 +136,9 @@ class VariantConfig(BaseModel):
elif isinstance(item, dict):
normalized.append(item)
else:
- raise ValueError(f"expected str or dict in 'find_in_filename' list, got {type(item).__name__}")
+ raise ValueError(f"Expected str or dict in 'find_in_filename' list, got {type(item).__name__}")
return normalized
- raise ValueError(f"expected str, dict, or list for 'find_in_filename', got {type(v).__name__}")
+ raise ValueError(f"Expected str, dict, or list for 'find_in_filename', got {type(v).__name__}")
# ------------------- Model configuration -------------------
@@ -153,7 +153,7 @@ class ModelConfig(BaseModel):
@model_validator(mode="after")
def check_detect_or_variants(self) -> "ModelConfig":
if self.detect is None and self.variants is None:
- raise ValueError("at least one of 'detect' or 'variants' must be specified")
+ raise ValueError("At least one of 'detect' or 'variants' must be specified")
return self
@@ -169,7 +169,7 @@ class PPPConfig(BaseModel):
@model_validator(mode="after")
def check_hosts_or_models(self) -> "PPPConfig":
if self.hosts is None and self.models is None:
- raise ValueError("at least one of 'hosts' or 'models' must be specified")
+ raise ValueError("At least one of 'hosts' or 'models' must be specified")
return self
# ------------------- State object -------------------
@@ -202,6 +202,7 @@ class PPPStateOptions:
cup_remove_extranetwork_tags: bool = False
strict_operators: bool = True
do_combinatorial: bool = False
+ combinatorial_shuffle: bool = False
combinatorial_limit: int = 100 # 0 = no limit
def __post_init__(self):
diff --git a/ppp_comfyui.py b/ppp_comfyui.py
index 1fb7571..d3aa4d9 100644
--- a/ppp_comfyui.py
+++ b/ppp_comfyui.py
@@ -189,6 +189,15 @@ class PromptPostProcessorComfyUINode:
"label_off": "No",
},
),
+ "combinatorial_shuffle": (
+ "BOOLEAN",
+ {
+ "default": PromptPostProcessor.DEFAULT_COMBINATORIAL_SHUFFLE,
+ "tooltip": "Shuffle the combinatorial results",
+ "label_on": "Yes",
+ "label_off": "No",
+ },
+ ),
"combinatorial_limit": (
"INT",
{
@@ -277,6 +286,7 @@ class PromptPostProcessorComfyUINode:
do_cleanup,
cleanup_variables,
do_combinatorial,
+ combinatorial_shuffle,
combinatorial_limit,
wc_options=None,
stn_options=None,
@@ -371,6 +381,7 @@ class PromptPostProcessorComfyUINode:
else PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS
),
do_combinatorial=do_combinatorial,
+ combinatorial_shuffle=combinatorial_shuffle,
combinatorial_limit=combinatorial_limit,
)
self.wildcards_obj.refresh_wildcards(
@@ -393,12 +404,19 @@ class PromptPostProcessorComfyUINode:
self.extranetwork_mappings_obj,
)
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,
- # )
+
+ # with open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "logs", "last_prompts_comfyui.txt"), "w", encoding="utf-8") as f:
+ # f.write(f"Seed: {seed if seed is not None else 1}\n")
+ # f.write(f"In Positive: {pos_prompt}\n")
+ # f.write(f"In Negative: {neg_prompt}\n")
+ # f.write("\n")
+ # for i, (posp, negp, var) in enumerate(results):
+ # f.write(f"Index: {i}\n")
+ # f.write(f"Out Positive: {posp}\n")
+ # f.write(f"Out Negative: {negp}\n")
+ # f.write(f"Out Variables: {var}\n")
+ # f.write("\n")
+
return tuple(zip(*results)) # unzip the list of tuples into tuple of lists
def interrupt(self):
diff --git a/ppp_tree.py b/ppp_tree.py
index aa731db..0768333 100644
--- a/ppp_tree.py
+++ b/ppp_tree.py
@@ -124,8 +124,12 @@ class TreeProcessor(lark.visitors.Interpreter):
)
return tuple(self.__comb_trace)
+ limit_reached = False
+
def _dfs(forced_path: tuple[int, ...]):
+ nonlocal limit_reached
if 0 < limit <= len(results):
+ limit_reached = True
return
trace = _run(forced_path)
# For each decision that was reached but not forced, spawn branches for all
@@ -133,19 +137,21 @@ class TreeProcessor(lark.visitors.Interpreter):
# 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):
+ limit_reached = True
return
num_options = trace[i]
for opt in range(1, num_options):
if 0 < limit <= len(results):
+ limit_reached = True
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):
+ if limit_reached:
self.log(
- logging.WARNING, f"Combinatorial limit of {limit} reached; some combinations may have been skipped."
+ logging.WARNING, f"Combinatorial limit of {limit} reached; some combinations have been skipped."
)
return results
@@ -1354,25 +1360,26 @@ class TreeProcessor(lark.visitors.Interpreter):
enmapping = self.state.extranetwork_mappings_obj.extranetwork_mappings.get(extnet_id, None)
if enmapping:
for v in enmapping.variants:
- if str(v.condition):
+ cond = str(v.condition) if v.condition is not None else None
+ if cond:
try:
cnd = parse_prompt(
self.state,
"condition",
- str(v.condition),
+ cond,
self.state.parsers["condition"],
True,
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
- f"Error parsing condition '{escape_single_quotes(str(v.condition))}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}",
+ f"Error parsing condition '{escape_single_quotes(cond)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}",
e,
)
cnd = None
else:
cnd = "True"
if cnd is not None and (cnd == "True" or self.__eval_condition(cnd)):
- if str(v.condition):
+ if cond:
found_mappings.append(v)
else:
else_mapping = v
diff --git a/ppp_variables.py b/ppp_variables.py
index 5c3e784..b5c23fc 100644
--- a/ppp_variables.py
+++ b/ppp_variables.py
@@ -32,14 +32,14 @@ class VariableRepository:
def set_system(self, name: str, value: Any) -> None:
"""Set a system variable."""
if not self.name_is_system(name):
- raise ValueError(f"invalid system variable name '{name}': must start with an underscore")
+ raise ValueError(f"Invalid system variable name '{name}': must start with an underscore")
self._system[name] = value
def update_system(self, mapping: dict[str, Any]) -> None:
"""Bulk-update system variables from *mapping*."""
for name in mapping:
if not self.name_is_system(name):
- raise ValueError(f"invalid system variable name '{name}': must start with an underscore")
+ raise ValueError(f"Invalid system variable name '{name}': must start with an underscore")
self._system.update(mapping)
def clear_system(self) -> None:
@@ -59,7 +59,7 @@ class VariableRepository:
def set_user(self, name: str, value: Any) -> None:
"""Set a user variable."""
if self.name_is_system(name):
- raise ValueError(f"invalid user variable name '{name}': must not start with an underscore")
+ raise ValueError(f"Invalid user variable name '{name}': must not start with an underscore")
self._user[name] = value
def delete_user(self, name: str) -> None:
diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py
index 68d1017..ae0dffe 100644
--- a/scripts/ppp_script.py
+++ b/scripts/ppp_script.py
@@ -103,8 +103,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
elem_id="ppp_force_equal_seeds",
)
gr.HTML("
")
- gr.Markdown(
- """
+ gr.Markdown("""
Unlink the seed to use the specified one for the prompts instead of the image seed.
* A seed of -1 and "Incremental seed" checked will use a random seed for the first prompt and consecutive values for the rest. This is the same as when you use -1 for the image seed.
@@ -113,8 +112,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
* Any other seed value and "Incremental seed" unchecked will use the specified seed for all the prompts.
Seeds are only used for the wildcards and choice constructs.
- """
- )
+ """)
gr.HTML("
")
with gr.Row(equal_height=True):
unlink_seed = gr.Checkbox(
@@ -148,6 +146,12 @@ class PromptPostProcessorA1111Script(scripts.Script):
value=PromptPostProcessor.DEFAULT_DO_COMBINATORIAL,
elem_id="ppp_combinatorial",
)
+ combinatorial_shuffle = gr.Checkbox(
+ label="Shuffle combinations",
+ info="Shuffle the combinatorial results.",
+ value=PromptPostProcessor.DEFAULT_COMBINATORIAL_SHUFFLE,
+ elem_id="ppp_combinatorial_shuffle",
+ )
combinatorial_limit = gr.Number(
label="Combinations limit (0 = no limit)",
value=PromptPostProcessor.DEFAULT_COMBINATORIAL_LIMIT,
@@ -155,7 +159,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
min_width=120,
elem_id="ppp_combinatorial_limit",
)
- return [force_equal_seeds, unlink_seed, seed, incremental_seed, combinatorial, combinatorial_limit]
+ return [force_equal_seeds, unlink_seed, seed, incremental_seed, combinatorial, combinatorial_shuffle, combinatorial_limit]
def process(
self,
@@ -165,6 +169,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
input_seed,
input_incremental_seed,
input_combinatorial,
+ input_combinatorial_shuffle,
input_combinatorial_limit,
): # pylint: disable=arguments-differ
"""
@@ -177,6 +182,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
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_shuffle (bool): Flag indicating whether to shuffle the combinatorial results.
input_combinatorial_limit (int): Maximum number of combinations (0 = no limit).
Returns:
@@ -241,6 +247,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
opts, "ppp_rem_removeextranetworktags", PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS
),
do_combinatorial=input_combinatorial,
+ combinatorial_shuffle=input_combinatorial_shuffle,
combinatorial_limit=max(num_seeds, int(input_combinatorial_limit)) if input_combinatorial else 0,
)
if self.ppp_logger is None:
@@ -392,19 +399,40 @@ class PromptPostProcessorA1111Script(scripts.Script):
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))]
+ 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))]
+ hiresfix_equal = regular_exists and rph == rpr and rnh == rnr
+ if hiresfix_equal:
+ log(
+ self.ppp_logger,
+ self.ppp_debug_level,
+ logging.INFO,
+ "hiresfix prompts are the same as regular prompts, skipping combinatorial processing for hiresfix",
+ )
+ for i in range(len(rph)): # pylint: disable=consider-using-enumerate
+ prompts_list[(hiresfix_type, i)] = prompts_list.get((regular_type, i))
+ else:
+ 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}])")
+ 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
@@ -412,7 +440,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
)
cached = self.lru_cache.get(key)
if cached is None:
- (hsh, seed, prompt, negative_prompt) = key
+ hsh, seed, prompt, negative_prompt = key
results = ppp.process_prompt(prompt, negative_prompt, seed)
posp, negp, _ = results[0]
cached = (posp, negp)
@@ -423,14 +451,14 @@ class PromptPostProcessorA1111Script(scripts.Script):
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:
+ # with open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "..", "logs", f"last_prompts_{app.value}.txt"), "w", encoding="utf-8") as f:
# for (prompttype, typeindex), (posp, negp) in prompts_list.items():
- # f.write(f"Key: {prompttype} {typeindex}\n")
+ # f.write(f"Key: {prompttype}[{typeindex}]\n")
# f.write(f"Seed: {calculated_seeds[typeindex]}\n")
- # f.write(f"Old Positive: {rpr[typeindex] if prompttype == regular_type else rph[typeindex]}\n")
- # f.write(f"Old Negative: {rnr[typeindex] if prompttype == regular_type else rnh[typeindex]}\n")
- # f.write(f"New Positive: {posp}\n")
- # f.write(f"New Negative: {negp}\n")
+ # f.write(f"In Positive: {rpr[typeindex] if prompttype == regular_type else rph[typeindex]}\n")
+ # f.write(f"In Negative: {rnr[typeindex] if prompttype == regular_type else rnh[typeindex]}\n")
+ # f.write(f"Out Positive: {posp}\n")
+ # f.write(f"Out Negative: {negp}\n")
# f.write("\n")
# updates the prompts
diff --git a/tests/base_tests.py b/tests/base_tests.py
index cd9e93b..2b40e6d 100644
--- a/tests/base_tests.py
+++ b/tests/base_tests.py
@@ -1,7 +1,7 @@
from dataclasses import replace
import os
import logging
-from typing import NamedTuple, Optional
+from typing import Any, NamedTuple, Optional
import unittest
import datetime
@@ -12,7 +12,7 @@ from ppp import PromptPostProcessor # type: ignore
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # type: ignore
-class PromptPair(NamedTuple):
+class InputTuple(NamedTuple):
prompt: str = ""
negative_prompt: str = ""
@@ -20,7 +20,7 @@ class PromptPair(NamedTuple):
class OutputTuple(NamedTuple):
prompt: str = ""
negative_prompt: str = ""
- variables: dict[str, str] = None
+ variables: dict[str, Any] = None
class TestPromptPostProcessorBase(unittest.TestCase):
@@ -177,7 +177,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
def process(
self,
- input_prompts: PromptPair,
+ input_prompts: InputTuple,
expected_output: Optional[OutputTuple | list[OutputTuple]] = None,
seed: int = 1,
ppp: Optional[str | PromptPostProcessor] = None,
@@ -188,7 +188,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
Process the prompt and compare the results with the expected prompts.
Args:
- input_prompts (PromptPair): The input prompts.
+ input_prompts (InputTuple): 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.
@@ -214,7 +214,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
)
if self.interrupted != interrupted:
errors.append(f"Interrupted flag is incorrect: expected {interrupted}, got {self.interrupted}")
- elif expected_output is not None:
+ elif not self.interrupted and expected_output is not None:
if len(result) != len(out):
errors.append(f"Incorrect number of combinations (expected {len(out)}, got {len(result)})")
for out_prompt, out_negative_prompt, out_variables in out:
@@ -253,8 +253,8 @@ class TestPromptPostProcessorBase(unittest.TestCase):
)
if self.interrupted != interrupted:
errors.append(f"Interrupted flag is incorrect: expected {interrupted}, got {self.interrupted}")
- elif expected_output is not None:
- result_prompt, result_negative_prompt, output_variables = result[0]
+ elif not self.interrupted and expected_output is not None:
+ result_prompt, result_negative_prompt, output_variables = (result[0] if result else (None, None, None))
if result_prompt != eo.prompt or result_negative_prompt != eo.negative_prompt:
errors.append(
f"Incorrect result '{eo.prompt}' / '{eo.negative_prompt}', got '{result_prompt}' / '{result_negative_prompt}'"
diff --git a/tests/tests_choices.py b/tests/tests_choices.py
index f1fb35c..b2ed022 100644
--- a/tests/tests_choices.py
+++ b/tests/tests_choices.py
@@ -1,7 +1,7 @@
from dataclasses import replace
from ppp import PromptPostProcessor # type: ignore
-from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
+from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -16,14 +16,14 @@ class TestChoices(TestPromptPostProcessorBase):
def test_ch_choices(self): # simple choices with weights
self.process(
- PromptPair("the choices are: {3::choice1|2::choice2|choice3}", ""),
+ InputTuple("the choices are: {3::choice1|2::choice2|choice3}", ""),
OutputTuple("the choices are: choice2", ""),
ppp="nocup",
)
def test_ch_unsupportedsampler(self): # unsupported sampler
self.process(
- PromptPair("the choices are: {@choice1|choice2|choice3}", ""),
+ InputTuple("the choices are: {@choice1|choice2|choice3}", ""),
OutputTuple("", ""),
ppp="nocup",
interrupted=True,
@@ -31,7 +31,7 @@ class TestChoices(TestPromptPostProcessorBase):
def test_ch_choices_withcomments(self): # choices with comments and multiline
self.process(
- PromptPair(
+ InputTuple(
"the choices are: {\n3::choice1 # this is option 1\n|2::choice2\n# this was option 2\n|choice3 # this is option 3\n}",
"",
),
@@ -41,28 +41,28 @@ class TestChoices(TestPromptPostProcessorBase):
def test_ch_choices_multiple(self): # choices with multiple selection
self.process(
- PromptPair("the choices are: {~2$$, $$3::choice1|2:: choice2 |choice3}", ""),
+ InputTuple("the choices are: {~2$$, $$3::choice1|2:: 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}", ""),
+ InputTuple("the choices are: {2$$, $$3::choice1|2 if _is_sd1::choice2|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}", ""),
+ InputTuple("${var=test}the choices are: {2$$, $$3::choice1|2 if not var eq 'test'::choice2|choice3}", ""),
OutputTuple("the choices are: choice1, choice3", ""),
ppp="nocup",
)
def test_ch_choices_set_if_nested(self): # nested choices with if user variable and multiple selection
self.process(
- PromptPair(
+ InputTuple(
"${var=test}the choices are: {2$$, $$3::choice1${var2=test2} {if var2 eq 'test2'::choice11|choice12}|2 if not var eq 'test'::choice2|choice3}",
"",
),
@@ -72,14 +72,14 @@ class TestChoices(TestPromptPostProcessorBase):
def test_ch_choicesinsidelora(self): # simple choices inside a lora
self.process(
- PromptPair("", ""),
+ InputTuple("", ""),
OutputTuple("", ""),
ppp="nocup",
)
def test_ch_removelorawithchoices(self):
self.process(
- PromptPair("", ""),
+ InputTuple("", ""),
OutputTuple("", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
@@ -97,7 +97,7 @@ class TestChoices(TestPromptPostProcessorBase):
def test_ch_cmd_includewildcard(self):
self.process(
- PromptPair("{ch_one|ch_two|%0.5::include yaml/wildcard1}", ""),
+ InputTuple("{ch_one|ch_two|%0.5::include yaml/wildcard1}", ""),
OutputTuple("ch_two", ""),
ppp="nocup",
)
@@ -106,7 +106,7 @@ class TestChoices(TestPromptPostProcessorBase):
def test_ch_combinatorial(self):
self.process(
- PromptPair("{choice1|choice2|choice3}, ${v:{option1|option2}}", ""),
+ InputTuple("{choice1|choice2|choice3}, ${v:{option1|option2}}", ""),
[
OutputTuple("choice1, option1", ""),
OutputTuple("choice1, option2", ""),
diff --git a/tests/tests_cleanup.py b/tests/tests_cleanup.py
index a05a7d0..b74e11c 100644
--- a/tests/tests_cleanup.py
+++ b/tests/tests_cleanup.py
@@ -2,7 +2,7 @@ import logging
from dataclasses import replace
from ppp import PromptPostProcessor # type: ignore
-from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
+from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
if __name__ == "__main__":
@@ -18,13 +18,13 @@ 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 "),
+ InputTuple(" this is a ((test ), , , (), , [] ( , test ,:2.0):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
self.process(
- PromptPair(
+ InputTuple(
" this is BREAKABLE a ((test)), ,AND AND(() [] ANDERSON (test:2.0):1.5) :o BREAK \n BREAK (red:1.5) ",
" [:hands, feet, :0.15]normal quality ",
),
@@ -36,7 +36,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_removenetworktags(self): # remove network tags
self.process(
- PromptPair("this is a test__yaml/wildcard7__", ""),
+ InputTuple("this is a test__yaml/wildcard7__", ""),
OutputTuple("this is a test", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
@@ -54,7 +54,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_dontremoveseparatorsoneol(self): # don't remove separators on eol
self.process(
- PromptPair("this is a test,\nsecond line", ""),
+ InputTuple("this is a test,\nsecond line", ""),
OutputTuple("this is a test,\nsecond line", ""),
ppp=PromptPostProcessor(
self.ppp_logger,
@@ -73,7 +73,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_separatorswitheol(self): # don't remove eols with the separators
self.process(
- PromptPair(
+ InputTuple(
"""{ (d:0.9) ,, (l:1.1) | (l:1.1) (d:0.9),,, }
(l:1.1)
(d:0.9)""",
@@ -111,7 +111,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_mergeattention(self): # merge attention
self.process(
- PromptPair(
+ InputTuple(
"this is (a test:0.9) of (attention (merging:1.2)) where ((this)) ((is joined:1.2)) and ([this too]:1.3)",
"",
),
@@ -123,7 +123,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_not_mergeattention(self): # not merge attention
self.process(
- PromptPair(
+ InputTuple(
"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)",
"",
),
@@ -139,7 +139,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_warn_unmatched_open_paren(self): # unmatched open parenthesis triggers warning
with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm:
self.process(
- PromptPair("(unclosed paren", ""),
+ InputTuple("(unclosed paren", ""),
OutputTuple("(unclosed paren", ""),
)
self.assertTrue(
@@ -150,7 +150,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_warn_unmatched_close_paren(self): # unmatched close parenthesis triggers warning
with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm:
self.process(
- PromptPair("extra close paren)", ""),
+ InputTuple("extra close paren)", ""),
OutputTuple("extra close paren)", ""),
)
self.assertTrue(
@@ -161,7 +161,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_warn_mismatched_brackets(self): # mismatched bracket types trigger warning
with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm:
self.process(
- PromptPair("(mismatched]", ""),
+ InputTuple("(mismatched]", ""),
OutputTuple("(mismatched]", ""),
)
self.assertTrue(
@@ -172,7 +172,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_warn_unmatched_open_bracket(self): # unmatched open bracket triggers warning
with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm:
self.process(
- PromptPair("unclosed [bracket", ""),
+ InputTuple("unclosed [bracket", ""),
OutputTuple("unclosed [bracket", ""),
)
self.assertTrue(
@@ -183,7 +183,7 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_warn_unmatched_complex(self): # unmatched complex case triggers warning
with self.assertLogs("PromptPostProcessor", level=logging.WARNING) as cm:
self.process(
- PromptPair("[(unmatched [bracket))", ""),
+ InputTuple("[(unmatched [bracket))", ""),
OutputTuple("[(unmatched [bracket))", ""),
)
self.assertTrue(
@@ -194,6 +194,6 @@ class TestCleanup(TestPromptPostProcessorBase):
def test_cl_warn_escaped_unmatched_no_false_warning(self): # escaped unmatched paren/bracket does not trigger warning
with self.assertNoLogs("PromptPostProcessor", level=logging.WARNING):
self.process(
- PromptPair(r"text with \(escaped unmatched\]", ""),
+ InputTuple(r"text with \(escaped unmatched\]", ""),
OutputTuple(r"text with \(escaped unmatched\]", ""),
)
diff --git a/tests/tests_host.py b/tests/tests_host.py
index b58f6be..033931d 100644
--- a/tests/tests_host.py
+++ b/tests/tests_host.py
@@ -1,5 +1,5 @@
from ppp import PromptPostProcessor # type: ignore
-from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
+from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -14,7 +14,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_attention_parentheses(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1] (test2) (test3:1.5) [(test4)]",
"",
),
@@ -35,7 +35,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_attention_disable(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1] (test2) (test3:1.5)",
"",
),
@@ -56,7 +56,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_attention_remove(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1] (test2) (test3:1.5)",
"",
),
@@ -77,7 +77,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_attention_error(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1] (test2) (test3:1.5)",
"",
),
@@ -99,7 +99,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_scheduling_before(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1:test2:0.5]",
"",
),
@@ -120,7 +120,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_scheduling_after(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1:test2:0.5]",
"",
),
@@ -141,7 +141,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_scheduling_first(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1::0.5] [:test2:0.5] [test3:test4:0.5]",
"",
),
@@ -162,7 +162,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_scheduling_remove(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1:test2:0.5]",
"",
),
@@ -183,7 +183,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_scheduling_error(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1:test2:0.5]",
"",
),
@@ -205,7 +205,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_alternation_first(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1|test2|test3]",
"",
),
@@ -226,7 +226,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_alternation_remove(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1|test2|test3]",
"",
),
@@ -247,7 +247,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_alternation_error(self):
self.process(
- PromptPair(
+ InputTuple(
"[test1|test2|test3]",
"",
),
@@ -269,7 +269,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_and_eol(self):
self.process(
- PromptPair(
+ InputTuple(
"test1 AND test2:2",
"",
),
@@ -290,7 +290,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_and_comma(self):
self.process(
- PromptPair(
+ InputTuple(
"test1 AND test2:2",
"",
),
@@ -311,7 +311,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_and_remove(self):
self.process(
- PromptPair(
+ InputTuple(
"test1 AND test2:2",
"",
),
@@ -332,7 +332,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_and_error(self):
self.process(
- PromptPair(
+ InputTuple(
"test1 AND test2:2",
"",
),
@@ -354,7 +354,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_break_eol(self):
self.process(
- PromptPair(
+ InputTuple(
"test1 BREAK test2",
"",
),
@@ -375,7 +375,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_break_comma(self):
self.process(
- PromptPair(
+ InputTuple(
"test1 BREAK test2",
"",
),
@@ -396,7 +396,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_break_remove(self):
self.process(
- PromptPair(
+ InputTuple(
"test1 BREAK test2",
"",
),
@@ -417,7 +417,7 @@ class TestHosts(TestPromptPostProcessorBase):
def test_host_break_error(self):
self.process(
- PromptPair(
+ InputTuple(
"test1 BREAK test2",
"",
),
diff --git a/tests/tests_performance.py b/tests/tests_performance.py
index 668b643..a5be804 100644
--- a/tests/tests_performance.py
+++ b/tests/tests_performance.py
@@ -1,4 +1,4 @@
-from .base_tests import PromptPair, TestPromptPostProcessorBase
+from .base_tests import InputTuple, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -18,7 +18,7 @@ class TestPerformance(TestPromptPostProcessorBase):
["(this:1.2) is a [test] using a [simple|low complexity] prompt with "] * 15
)
self.process(
- PromptPair(large_prompt, ""),
+ InputTuple(large_prompt, ""),
ppp="nocup",
)
@@ -30,7 +30,7 @@ class TestPerformance(TestPromptPostProcessorBase):
["(this:1.2) is a [test] using a [simple|low complexity] prompt with "] * 15
)
self.process(
- PromptPair(large_prompt, ""),
+ InputTuple(large_prompt, ""),
ppp="nocup",
)
@@ -39,7 +39,7 @@ class TestPerformance(TestPromptPostProcessorBase):
): # performance test with a large prompt with new constructs (full parser)
large_prompt = ", ".join(["__yaml/wildcard1__, (__yaml/wildcard2__), __yaml/wildcard3__, {one|two|three}"] * 15)
self.process(
- PromptPair(large_prompt, ""),
+ InputTuple(large_prompt, ""),
ppp="nocup",
)
@@ -49,27 +49,27 @@ class TestPerformance(TestPromptPostProcessorBase):
def test_parser_performance_simple_attention(self): # performance test with only attention
large_prompt = ", ".join(["(one:1.2) two (three) four [five] six"] * 20)
self.process(
- PromptPair(large_prompt, ""),
+ InputTuple(large_prompt, ""),
ppp="nocup",
)
def test_parser_performance_simple_schedules(self): # performance test with only schedules
large_prompt = ", ".join(["[one:1:0.5] two [three:0.8] four [five:5:0.2] six"] * 20)
self.process(
- PromptPair(large_prompt, ""),
+ InputTuple(large_prompt, ""),
ppp="nocup",
)
def test_parser_performance_simple_alternation(self): # performance test with only alternation
large_prompt = ", ".join(["[one|1] two [three|3] four [five|5] six"] * 20)
self.process(
- PromptPair(large_prompt, ""),
+ InputTuple(large_prompt, ""),
ppp="nocup",
)
def test_parser_performance_simple_extranetwork(self): # performance test with only extra networks
large_prompt = ", ".join([" two four six"] * 20)
self.process(
- PromptPair(large_prompt, ""),
+ InputTuple(large_prompt, ""),
ppp="nocup",
)
diff --git a/tests/tests_stn.py b/tests/tests_stn.py
index 5cc2e40..50fe1a5 100644
--- a/tests/tests_stn.py
+++ b/tests/tests_stn.py
@@ -1,4 +1,4 @@
-from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
+from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -13,7 +13,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
def test_stn_simple(self): # negtags with different parameters and separations
self.process(
- PromptPair(
+ InputTuple(
"flowersred, green, blueyellow, purpleblack",
"normal quality, worse quality",
),
@@ -22,7 +22,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
def test_stn_complex(self): # complex negtags
self.process(
- PromptPair(
+ InputTuple(
"red ((pink)), flowers purple, mauveblue, yellow green",
"normal quality, , bad quality, worse quality",
),
@@ -34,7 +34,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
def test_stn_complex_nocleanup(self): # complex negtags with no cleanup
self.process(
- PromptPair(
+ InputTuple(
"red ((pink)), flowers purple, mauveblue, yellow green",
"normal quality, , bad quality, worse quality",
),
@@ -47,7 +47,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
def test_stn_inside_attention(self): # negtag inside attention
self.process(
- PromptPair(
+ InputTuple(
"[neg1] this is a ((testneg2) (test:2.0): 1.5 ) (red[square]:1.5)",
"normal quality",
),
@@ -59,7 +59,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
def test_stn_inside_alternation(self): # negtag inside alternation
self.process(
- PromptPair(
+ InputTuple(
"this is a (([complexneg1|simpleneg2|regularneg3] test)(test:2.0):1.5)",
"normal quality",
),
@@ -71,7 +71,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
def test_stn_inside_alternation_recursive(self): # negtag inside alternation (recursive alternation)
self.process(
- PromptPair(
+ InputTuple(
"this is a (([complexneg1[one|twoneg12||three|four(neg14)]|simpleneg2|regularneg3] test)(test:2.0):1.5)",
"normal quality",
),
@@ -83,13 +83,13 @@ class TestSendToNegative(TestPromptPostProcessorBase):
def test_stn_inside_scheduling(self): # negtag inside scheduling
self.process(
- PromptPair("this is [abcneg1:defneg2: 5 ]", "normal quality"),
+ InputTuple("this is [abcneg1:defneg2: 5 ]", "normal quality"),
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
self.process(
- PromptPair(
+ InputTuple(
"[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loratrigger AND AND hypernettrigger :0.3",
"normal quality, ",
),
@@ -101,7 +101,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
def test_stn_complex_features_newformat(self): # complex negtags with AND, BREAK and other features (new format)
self.process(
- PromptPair(
+ InputTuple(
"[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loratrigger AND AND hypernettrigger :0.3",
"normal quality, ",
),
@@ -113,7 +113,7 @@ class TestSendToNegative(TestPromptPostProcessorBase):
def test_stn_inside_alternation_recursive_2(self): # negtag inside alternation (recursive alternation)
self.process(
- PromptPair(
+ InputTuple(
"[pos1neg1[pos11|pos12neg12||pos14|pos15neg15]|pos2neg2|pos3neg3]",
"",
),
diff --git a/tests/tests_varcomms.py b/tests/tests_varcomms.py
index d85bb3e..d5139f4 100644
--- a/tests/tests_varcomms.py
+++ b/tests/tests_varcomms.py
@@ -2,7 +2,7 @@ from dataclasses import replace
from ppp import PromptPostProcessor # type: ignore
from ppp_classes import ONWARNING_CHOICES # type: ignore
-from .base_tests import OutputTuple, PromptPair, TestPromptPostProcessorBase
+from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -17,7 +17,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_empty_variable(self):
self.process(
- PromptPair(
+ InputTuple(
"${v1=}${v3:}",
"",
),
@@ -28,7 +28,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_echoed_variable(self):
self.process(
- PromptPair(
+ InputTuple(
"${v1=test1}test2${v3:test3}${v3:test4}",
"",
),
@@ -38,7 +38,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_unknown_echoed_variable(self):
self.process(
- PromptPair(
+ InputTuple(
"${v1}",
"",
),
@@ -61,7 +61,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_variable_in_extranetwork(self):
self.process(
- PromptPair(
+ InputTuple(
"${f=filename}${w=0.5}",
"",
),
@@ -72,7 +72,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_var_nested_1(self): # variable default nested in variable set
self.process(
- PromptPair(
+ InputTuple(
"${v1=test ${v2:OK}}${v1}",
"",
),
@@ -81,7 +81,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_var_nested_2(self): # variable set nested in variable default
self.process(
- PromptPair(
+ InputTuple(
"${v1:test ${v2=OK}${v2}}",
"",
),
@@ -90,7 +90,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_var_nested_3(self): # variable default nested in variable default
self.process(
- PromptPair(
+ InputTuple(
"${v1:test ${v2:OK}}",
"",
),
@@ -101,7 +101,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_array_variable_1(self): # array variable set with += and test of index value and full array with and without default separator
self.process(
- PromptPair(
+ InputTuple(
"${v1[]=val1}${v1[]+=val2}${v1[]+=val3}${v1[1]:defval},${v1[]:defval2},${v1[&'.']:defval3}",
"",
),
@@ -110,7 +110,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
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
self.process(
- PromptPair(
+ InputTuple(
"${v1[]=val1}${v1[]=val2}${v1[]:defval},${v2[]:defval2},${v2[1]:defval3},${v3[]=}${v3[]:defval4}",
"",
),
@@ -119,7 +119,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_array_variable_3(self): # access array index by variable, set array variable to expanded array variable and add expanded array
self.process(
- PromptPair(
+ InputTuple(
"${v1[]=val1}${v1[]+=val2}${v2=1}${v1[v2]:defval1}${v3[]=${v1[]}}${v3[]+=${v1[]}}, ${v3[&'.']}",
"",
),
@@ -128,7 +128,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_array_variable_4(self): # test list in array
self.process(
- PromptPair(
+ InputTuple(
"${v1[]=val1}${v1[]+=val2}${v1[]+=val3}OKnot OK",
"",
),
@@ -137,7 +137,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_array_variable_5(self): # test empty array
self.process(
- PromptPair(
+ InputTuple(
"${v1[]=}OKnot OK,OKnot OK",
"",
),
@@ -146,7 +146,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_array_variable_6(self): # array variable set and addition with expanded values from array variables
self.process(
- PromptPair(
+ InputTuple(
"${v1[]=val1}${v1[]+=val2}${v2[]=val3}${v3[]=*v1[]}${v3[]+=*v2[]}",
"",
),
@@ -155,7 +155,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_array_variable_7(self): # array variable set and addition with expanded values from wildcards
self.process(
- PromptPair(
+ InputTuple(
"${v1[]=*__yaml/wildcard1__}${v1[]+=*__yaml/wildcard2__}${v1[2]:defval}",
"",
),
@@ -164,7 +164,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_array_variable_8(self): # array variable set and addition with expanded values from lists
self.process(
- PromptPair(
+ InputTuple(
"${v1[]=*()}${v1[]+=*('one','two')}${v2=three}${v1[]+=*(v2,'four')}${v1[2]:defval}",
"",
),
@@ -173,7 +173,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_array_variable_9(self): # array variable length
self.process(
- PromptPair(
+ InputTuple(
"${v1[]=val1}${v1[]+=val2}${v1[]+=val3}${v1[#]:defval}, OKnot OK",
"",
),
@@ -182,7 +182,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_array_variable_10(self): # array variable set with expanded values from wildcards in command format
self.process(
- PromptPair(
+ InputTuple(
"*__yaml/wildcard1__*__yaml/wildcard2__${v1[2]:defval}",
"",
),
@@ -195,7 +195,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_ReqR(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=hello}${r2=hello}OKnot OK",
"",
),
@@ -204,7 +204,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RnoteqR(self): # test for the not before the operator
self.process(
- PromptPair(
+ InputTuple(
"${r1=hello}${r2=bye}OKnot OK",
"",
),
@@ -213,7 +213,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RneR(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=hello}${r2=bye}OKnot OK",
"",
),
@@ -222,7 +222,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RltR(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=1}${r2=2}OKnot OK",
"",
),
@@ -231,7 +231,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RgtR(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=2}${r2=1}OKnot OK",
"",
),
@@ -240,7 +240,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RleR(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=1}${r2=1}OKnot OK",
"",
),
@@ -249,7 +249,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RgeR(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=1}${r2=1}OKnot OK",
"",
),
@@ -258,7 +258,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RinR(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=hello}${r2=hello world}OKnot OK",
"",
),
@@ -267,7 +267,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RcontainsR(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=hello world}${r2=hello}OKnot OK",
"",
),
@@ -278,7 +278,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AeqA(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*('hello','world')}${a2[]=*('hello','world')}OKnot OK",
"",
),
@@ -287,7 +287,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AneA_1(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*('hello')}${a2[]=*('bye')}OKnot OK",
"",
),
@@ -296,7 +296,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AneA_2(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*('hello','world')}${a2[]=*('hello')}OKnot OK",
"",
),
@@ -305,7 +305,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AneA_3(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*('hello','world')}${a2[]=*('world','hello')}OKnot OK",
"",
),
@@ -314,7 +314,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AltA_1(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*(1,2,3)}${a2[]=*(2,3,4)}OKnot OK",
"",
),
@@ -323,7 +323,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AltA_2(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*(1,2)}${a2[]=*(2,3,4)}OKnot OK",
"",
),
@@ -332,7 +332,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AgtA(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*(2,3,4)}${a2[]=*(1,2,3)}OKnot OK",
"",
),
@@ -341,7 +341,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AleA(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*(1,2)}${a2[]=*(1,3)}OKnot OK",
"",
),
@@ -350,7 +350,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AgeA(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*(1,3)}${a2[]=*(1,2)}OKnot OK",
"",
),
@@ -359,7 +359,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AinA(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*('hello')}${a2[]=*('hello', 'world')}OKnot OK",
"",
),
@@ -368,7 +368,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AcontainsA(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*('hello','world')}${a2[]=*('hello')}OKnot OK",
"",
),
@@ -379,7 +379,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AeqR(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*('hello','world')}${r2=hello}OKnot OK",
"",
),
@@ -389,7 +389,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AneR(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*('hello')}${r2=bye}OKnot OK",
"",
),
@@ -399,7 +399,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AltR(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*(1,2,3)}${r2=2}OKnot OK",
"",
),
@@ -409,7 +409,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AgtR(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*(2,3,4)}${r2=2}OKnot OK",
"",
),
@@ -419,7 +419,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AleR(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*(1,2)}${r2=2}OKnot OK",
"",
),
@@ -429,7 +429,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AgeR(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*(1,3)}${r2=2}OKnot OK",
"",
),
@@ -439,7 +439,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AinR(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*('hello', 'world')}${r2=hello world)}OKnot OK",
"",
),
@@ -448,7 +448,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_AcontainsR(self):
self.process(
- PromptPair(
+ InputTuple(
"${a1[]=*('hello','world')}${r2=hello}OKnot OK",
"",
),
@@ -459,7 +459,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_ReqA(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=hello}${a2[]=*('hello','world')}OKnot OK",
"",
),
@@ -469,7 +469,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RneA(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=bye}${a2[]=*('hello')}OKnot OK",
"",
),
@@ -479,7 +479,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RltA(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=2}${a2[]=*(1,2,3)}OKnot OK",
"",
),
@@ -489,7 +489,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RgtA(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=2}${a2[]=*(2,3,4)}OKnot OK",
"",
),
@@ -499,7 +499,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RleA(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=2}${a2[]=*(1,2)}OKnot OK",
"",
),
@@ -509,7 +509,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RgeA(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=2}${a2[]=*(1,3)}OKnot OK",
"",
),
@@ -519,7 +519,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RinA(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=hello}${a2[]=*('hello', 'world')}OKnot OK",
"",
),
@@ -528,7 +528,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
def test_operator_RcontainsA(self):
self.process(
- PromptPair(
+ InputTuple(
"${r1=hello world}${a2[]=*('hello','world')}OKnot OK