* Escape quotes on log messages
* Separated tests in multiple files. * Added pydantic validation for configuration. * Improved backup/restore on visit of nodes. * Echoing of variables with default no longer visits the default when not used (and no longer creates variables that were not actually used).
This commit is contained in:
Vendored
-8
@@ -9,14 +9,6 @@
|
||||
"type": "debugpy",
|
||||
"request": "attach",
|
||||
"processId": "${command:pickProcess}"
|
||||
},
|
||||
{
|
||||
"name": "Tests",
|
||||
"type": "debugpy",
|
||||
"request": "launch",
|
||||
"program": "tests/tests.py",
|
||||
"console": "integratedTerminal",
|
||||
"justMyCode": true
|
||||
}
|
||||
]
|
||||
}
|
||||
+9
-9
@@ -60,15 +60,15 @@ The only command available is `include wildcard`, which will include the choices
|
||||
|
||||
These are examples of formats you can use to insert a choice construct:
|
||||
|
||||
| Construct | Result |
|
||||
| --------- | ------ |
|
||||
| `{choice1\|5::choice2\|3::choice3}` | select 1 choice, two of them have weights |
|
||||
| `{3$$choice1\|5 if _is_sd1::choice2\|choice3}` | select 3 choices, one has a weight and a condition |
|
||||
| `{2-3$$2::choice1\|choice2\|choice3}` | select 2 to 3 choices, one of them has a weight |
|
||||
| `{r2-3$$choice1\|choice2\|choice3}` | select 2 to 3 choices allowing repetition |
|
||||
| `{2-3$$ / $$choice1\|choice2\|choice3}` | select 2 to 3 choices with separator " / " |
|
||||
| `{o$$if _is_sd1::choice1\|if _is_sd2::choice2}`| select 1 choice, both have conditions, if none matches it is allowed because we indicate that it is optional |
|
||||
| `{choice1\|choice2\|%0.5::path/wildcard}` | select 1 choice from the two specified and the ones inside the path/wildcard wildcard, which will be weighted with half their weights |
|
||||
| Construct | Result |
|
||||
| --------- | ------ |
|
||||
| `{choice1\|5::choice2\|3::choice3}` | select 1 choice, two of them have weights |
|
||||
| `{3$$choice1\|5 if _is_sd1::choice2\|choice3}` | select 3 choices, one has a weight and a condition |
|
||||
| `{2-3$$2::choice1\|choice2\|choice3}` | select 2 to 3 choices, one of them has a weight |
|
||||
| `{r2-3$$choice1\|choice2\|choice3}` | select 2 to 3 choices allowing repetition |
|
||||
| `{2-3$$ / $$choice1\|choice2\|choice3}` | select 2 to 3 choices with separator " / " |
|
||||
| `{o$$if _is_sd1::choice1\|if _is_sd2::choice2}` | select 1 choice, both have conditions, if none matches it is allowed because we indicate that it is optional |
|
||||
| `{choice1\|choice2\|%0.5::include path/wildcard}` | select 1 choice from the two specified and the ones inside the path/wildcard wildcard, which will be weighted with half their weights |
|
||||
|
||||
Notes:
|
||||
|
||||
|
||||
@@ -12,10 +12,11 @@ import lark
|
||||
import numpy as np
|
||||
import yaml
|
||||
|
||||
from ppp_hosts import SUPPORTED_APPS # pylint: disable=import-error
|
||||
from ppp_logging import DEBUG_LEVEL # pylint: disable=import-error
|
||||
from ppp_wildcards import PPPWildcard, PPPWildcards # pylint: disable=import-error
|
||||
from ppp_enmappings import PPPENMappingVariant, PPPExtraNetworkMappings # pylint: disable=import-error
|
||||
from ppp_classes import SUPPORTED_APPS
|
||||
from ppp_logging import DEBUG_LEVEL
|
||||
from ppp_utils import escape_single_quotes
|
||||
from ppp_wildcards import PPPWildcard, PPPWildcards
|
||||
from ppp_enmappings import PPPENMappingVariant, PPPExtraNetworkMappings
|
||||
|
||||
|
||||
class PPPInterrupt(Exception):
|
||||
@@ -132,14 +133,15 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
self.config: dict[str, Any] = yaml.safe_load(f)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
self.config = {}
|
||||
raise PPPInterrupt(f"Failed to load default configuration from '{default_config_file}'.") from exc
|
||||
raise PPPInterrupt(
|
||||
f"Failed to load default configuration from '{escape_single_quotes(default_config_file)}'."
|
||||
) from exc
|
||||
validate_def_cfg = self.__validate_normalize_configuration(self.config, "default configuration file")
|
||||
if validate_def_cfg != 0:
|
||||
errmsg = "Default configuration file has errors. Please restore the default configuration file and, per instructions, use a copy to adapt it."
|
||||
if validate_def_cfg == 2:
|
||||
raise PPPInterrupt(errmsg)
|
||||
else:
|
||||
self.logger.warning(errmsg)
|
||||
self.logger.warning(errmsg)
|
||||
|
||||
user_config_file = self.env_info.get("ppp_config", "")
|
||||
user_config: dict[str, Any] = {}
|
||||
@@ -162,8 +164,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
if user_config_file and os.path.exists(user_config_file):
|
||||
with open(user_config_file, "r", encoding="utf-8") as f:
|
||||
user_config = yaml.safe_load(f)
|
||||
self.__validate_normalize_configuration(user_config, "user configuration")
|
||||
self.__merge_configuration(user_config)
|
||||
self.__validate_normalize_configuration(user_config, "user configuration")
|
||||
if user_config:
|
||||
self.__merge_configuration(user_config)
|
||||
|
||||
self.models_config: dict[str, dict[str, Any] | None] = self.config.get("models") or {}
|
||||
self.known_models: list[str] = list(self.models_config.keys())
|
||||
@@ -179,7 +182,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
self.host_config: dict[str, Any] = (self.config.get("hosts") or {}).get(self.env_info.get("app", ""))
|
||||
if self.host_config is None:
|
||||
raise PPPInterrupt(
|
||||
f"No host configuration found for app '{self.env_info.get('app', '')}'. Please check your configuration."
|
||||
f"No host configuration found for app '{escape_single_quotes(self.env_info.get('app', ''))}'. Please check your configuration."
|
||||
)
|
||||
|
||||
# Update env_info with model detection
|
||||
@@ -208,7 +211,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
self.variants_definitions[v] = (m, vo["find_in_filename"])
|
||||
else:
|
||||
self.logger.warning(
|
||||
f"Variant name '{v}' in model '{m}' conflicts with a known model name. Discarding variant."
|
||||
f"Variant name '{escape_single_quotes(v)}' in model '{escape_single_quotes(m)}' conflicts with a known model name. Discarding variant."
|
||||
)
|
||||
|
||||
if self.debug_level != DEBUG_LEVEL.none:
|
||||
@@ -442,27 +445,27 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
return {"regex": find_in_filename, "flags": re.IGNORECASE}
|
||||
except re.error:
|
||||
self.logger.warning(
|
||||
f"{where.title()}: Invalid regex pattern for variant '{variant_key}' in model '{model_key}'. Discarding variant."
|
||||
f"{where.title()}: Invalid regex pattern for variant '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant."
|
||||
)
|
||||
elif isinstance(find_in_filename, dict):
|
||||
regex = find_in_filename.get("regex", "")
|
||||
flags = find_in_filename.get("flags", [])
|
||||
if not isinstance(regex, str) or not isinstance(flags, list) or not all(isinstance(f, str) for f in flags):
|
||||
self.logger.warning(
|
||||
f"{where.title()}: Invalid format for 'find_in_filename' for variant '{variant_key}' in model '{model_key}'. Discarding variant."
|
||||
f"{where.title()}: Invalid format for 'find_in_filename' for variant '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant."
|
||||
)
|
||||
else:
|
||||
fl = self.__re_flags_from_list(flags)
|
||||
if fl == 0 and len(flags):
|
||||
self.logger.warning(
|
||||
f"{where.title()}: Invalid regex flags for variant '{variant_key}' in model '{model_key}'. Discarding variant."
|
||||
f"{where.title()}: Invalid regex flags for variant '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant."
|
||||
)
|
||||
try:
|
||||
re.compile(regex, fl)
|
||||
return {"regex": regex, "flags": fl}
|
||||
except re.error:
|
||||
self.logger.warning(
|
||||
f"{where.title()}: Invalid regex pattern for variant '{variant_key}' in model '{model_key}'. Discarding variant."
|
||||
f"{where.title()}: Invalid regex pattern for variant '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant."
|
||||
)
|
||||
return None
|
||||
|
||||
@@ -493,14 +496,18 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
defcfg_hosts: dict[str, Any] = cfg.get("hosts", {})
|
||||
for host_key, host_value in dict(defcfg_hosts).items():
|
||||
if host_key not in SUPPORTED_APPS._value2member_map_: # pylint: disable=protected-access
|
||||
self.logger.warning(f"{where.capitalize()}: Unsupported host '{host_key}'. Discarding host.")
|
||||
self.logger.warning(
|
||||
f"{where.capitalize()}: Unsupported host '{escape_single_quotes(host_key)}'. Discarding host."
|
||||
)
|
||||
defcfg_hosts.pop(host_key, None)
|
||||
result = 1
|
||||
elif host_value is not None and (
|
||||
not isinstance(host_value, dict)
|
||||
or not all(k in ["attention", "scheduling", "alternation", "and", "break"] for k in host_value)
|
||||
):
|
||||
self.logger.warning(f"{where.capitalize()}: Invalid format for host '{host_key}'. Discarding host.")
|
||||
self.logger.warning(
|
||||
f"{where.capitalize()}: Invalid format for host '{escape_single_quotes(host_key)}'. Discarding host."
|
||||
)
|
||||
defcfg_hosts.pop(host_key, None)
|
||||
result = 1
|
||||
defcfg_models: dict[str, Any] = cfg.get("models", {})
|
||||
@@ -510,7 +517,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
or model_value.get("detect") is None
|
||||
or not isinstance(model_value["detect"], dict)
|
||||
):
|
||||
self.logger.warning(f"{where.capitalize()}: Invalid format for model '{model_key}'. Discarding model.")
|
||||
self.logger.warning(
|
||||
f"{where.capitalize()}: Invalid format for model '{escape_single_quotes(model_key)}'. Discarding model."
|
||||
)
|
||||
defcfg_models.pop(model_key, None)
|
||||
result = 1
|
||||
else:
|
||||
@@ -518,14 +527,14 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
for host_key, host_value in dict(defcfg_m_detect).items():
|
||||
if host_key not in SUPPORTED_APPS._value2member_map_: # pylint: disable=protected-access
|
||||
self.logger.warning(
|
||||
f"{where.capitalize()}: Unsupported host '{host_key}' in 'detect' for model '{model_key}'. Discarding host."
|
||||
f"{where.capitalize()}: Unsupported host '{escape_single_quotes(host_key)}' in 'detect' for model '{escape_single_quotes(model_key)}'. Discarding host."
|
||||
)
|
||||
defcfg_m_detect.pop(host_key, None)
|
||||
result = 1
|
||||
elif host_value is not None:
|
||||
if not isinstance(host_value, dict):
|
||||
self.logger.warning(
|
||||
f"{where.capitalize()}: Invalid format for host '{host_key}' in 'detect' for model '{model_key}'. Discarding host."
|
||||
f"{where.capitalize()}: Invalid format for host '{escape_single_quotes(host_key)}' in 'detect' for model '{escape_single_quotes(model_key)}'. Discarding host."
|
||||
)
|
||||
defcfg_m_detect.pop(host_key, None)
|
||||
result = 1
|
||||
@@ -534,27 +543,27 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
isinstance(c, str) for c in host_value["class"]
|
||||
):
|
||||
self.logger.warning(
|
||||
f"{where.capitalize()}: Invalid format for 'class' in host '{host_key}' in 'detect' for model '{model_key}'. Discarding host."
|
||||
f"{where.capitalize()}: Invalid format for 'class' in host '{escape_single_quotes(host_key)}' in 'detect' for model '{escape_single_quotes(model_key)}'. Discarding host."
|
||||
)
|
||||
defcfg_m_detect.pop(host_key, None)
|
||||
result = 1
|
||||
elif "property" in host_value:
|
||||
if not isinstance(host_value["property"], str):
|
||||
self.logger.warning(
|
||||
f"{where.capitalize()}: Invalid format for 'property' in host '{host_key}' in 'detect' for model '{model_key}'. Discarding host."
|
||||
f"{where.capitalize()}: Invalid format for 'property' in host '{escape_single_quotes(host_key)}' in 'detect' for model '{escape_single_quotes(model_key)}'. Discarding host."
|
||||
)
|
||||
defcfg_m_detect.pop(host_key, None)
|
||||
result = 1
|
||||
else:
|
||||
self.logger.warning(
|
||||
f"{where.capitalize()}: Neither 'class' nor 'property' specified for host '{host_key}' in 'detect' for model '{model_key}'. Discarding host."
|
||||
f"{where.capitalize()}: Neither 'class' nor 'property' specified for host '{escape_single_quotes(host_key)}' in 'detect' for model '{escape_single_quotes(model_key)}'. Discarding host."
|
||||
)
|
||||
defcfg_m_detect.pop(host_key, None)
|
||||
result = 1
|
||||
if "variants" in model_value:
|
||||
if not isinstance(model_value["variants"], dict):
|
||||
self.logger.warning(
|
||||
f"{where.capitalize()}: Invalid format for 'variants' in model '{model_key}'. Discarding model."
|
||||
f"{where.capitalize()}: Invalid format for 'variants' in model '{escape_single_quotes(model_key)}'. Discarding model."
|
||||
)
|
||||
defcfg_models.pop(model_key, None)
|
||||
result = 1
|
||||
@@ -563,7 +572,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
for variant_key, variant_value in dict(defcfg_m_variants).items():
|
||||
if not isinstance(variant_key, str) or not variant_key.isidentifier():
|
||||
self.logger.warning(
|
||||
f"{where.capitalize()}: Invalid variant name '{variant_key}' in model '{model_key}'. Discarding variant."
|
||||
f"{where.capitalize()}: Invalid variant name '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant."
|
||||
)
|
||||
defcfg_m_variants.pop(variant_key, None)
|
||||
result = 1
|
||||
@@ -571,7 +580,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
variant_value.get("find_in_filename"), (str, dict, list)
|
||||
):
|
||||
self.logger.warning(
|
||||
f"{where.capitalize()}: Invalid format for variant '{variant_key}' in model '{model_key}'. Discarding variant."
|
||||
f"{where.capitalize()}: Invalid format for variant '{escape_single_quotes(variant_key)}' in model '{escape_single_quotes(model_key)}'. Discarding variant."
|
||||
)
|
||||
defcfg_m_variants.pop(variant_key, None)
|
||||
result = 1
|
||||
@@ -1193,7 +1202,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
parsed_prompt = None
|
||||
try:
|
||||
if self.debug_level == DEBUG_LEVEL.full:
|
||||
self.logger.debug(self.format_output(f"Parsing {prompt_description}: '{prompt}'"))
|
||||
self.logger.debug(self.format_output(f"Parsing {prompt_description}: '{escape_single_quotes(prompt)}'"))
|
||||
parsed_prompt = parser.parse(prompt)
|
||||
# we store the contents so we can use them later even if the meta position is not valid anymore
|
||||
if isinstance(parsed_prompt, lark.Tree):
|
||||
@@ -1206,7 +1215,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
except lark.exceptions.UnexpectedInput:
|
||||
if raise_parsing_error:
|
||||
raise
|
||||
self.logger.exception(self.format_output(f"Parsing failed on prompt!: {prompt}"))
|
||||
self.logger.exception(self.format_output(f"Parsing failed on prompt!: {escape_single_quotes(prompt)}"))
|
||||
t2 = time.monotonic_ns()
|
||||
if self.debug_level == DEBUG_LEVEL.full:
|
||||
self.logger.debug(f"Parse {prompt_description} time: {(t2 - t1) / 1_000_000_000:.3f} seconds")
|
||||
@@ -1306,13 +1315,19 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
str: The result of the visit.
|
||||
"""
|
||||
backup_result = self.result
|
||||
# if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
# self.__ppp.logger.debug(f"Visiting node {node}.")
|
||||
if restore_state:
|
||||
# if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
# self.__ppp.logger.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_user_variables = {k: v.copy() for k, v in self.__ppp.user_variables.items()}
|
||||
backup_echoed_variables = self.__ppp.echoed_variables.copy()
|
||||
if node is not None:
|
||||
if isinstance(node, list):
|
||||
for child in node:
|
||||
@@ -1329,12 +1344,16 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
if discard_content or restore_state:
|
||||
self.result = backup_result
|
||||
if restore_state:
|
||||
# if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
# self.__ppp.logger.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.__ppp.user_variables = backup_user_variables
|
||||
self.__ppp.echoed_variables = backup_echoed_variables
|
||||
return added_result
|
||||
|
||||
def __get_original_node_content(self, node: lark.Tree | lark.Token, default=None) -> str:
|
||||
@@ -1416,7 +1435,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
info = f"({info}) " if info is not None and info != "" else ""
|
||||
output = self.result[len(start_result) :]
|
||||
if output != "":
|
||||
output = f" >> '{output}'"
|
||||
output = f" >> '{escape_single_quotes(output)}'"
|
||||
self.__ppp.logger.debug(
|
||||
self.__ppp.format_output(
|
||||
f"TreeProcessor.{construct} {info}({duration / 1_000_000_000:.3f} seconds){output}"
|
||||
@@ -1512,7 +1531,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
try:
|
||||
var_value_adjusted = int(var_value)
|
||||
except (ValueError, TypeError):
|
||||
self.warn_or_stop(f"Cannot convert variable value '{var_value}' to integer for comparison")
|
||||
self.warn_or_stop(
|
||||
f"Cannot convert variable value '{escape_single_quotes(var_value)}' to integer for comparison"
|
||||
)
|
||||
return False
|
||||
result = comp_ops[cond_comp](var_value_adjusted, c)
|
||||
if result:
|
||||
@@ -1820,7 +1841,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
self.__negtags.append(
|
||||
self.NegTag(len(self.result), len(self.result), content, parameters, self.__shell.copy())
|
||||
)
|
||||
info = f"with {parameters or 'no parameters'} : {content}"
|
||||
info = f"with {escape_single_quotes(parameters) or 'no parameters'} : {escape_single_quotes(content)}"
|
||||
else:
|
||||
self.warn_or_stop("Ignored negative command in negative prompt")
|
||||
self.__visit(tree.children[1::])
|
||||
@@ -1862,18 +1883,20 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
t1 = time.monotonic_ns()
|
||||
start_result = self.result
|
||||
if variable.startswith("_"):
|
||||
self.warn_or_stop(f"Invalid variable name '{variable}' detected! System variables cannot be set.")
|
||||
self.warn_or_stop(
|
||||
f"Invalid variable name '{escape_single_quotes(variable)}' detected! System variables cannot be set."
|
||||
)
|
||||
return
|
||||
info = variable
|
||||
value_description = self.__get_original_node_content(content, None)
|
||||
value = content
|
||||
modifiers_str: list[str] = [m.value for m in modifiers.children] if modifiers is not None else []
|
||||
if any(item in modifiers_str for item in ["+", "add"]):
|
||||
info += f" += '{value_description}'"
|
||||
info += f" += '{escape_single_quotes(value_description or '')}'"
|
||||
raw_oldvalue = self.__ppp.user_variables.get(variable, None)
|
||||
if raw_oldvalue is None:
|
||||
newvalue = value
|
||||
self.warn_or_stop(f"Unknown variable {variable}")
|
||||
self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable)}")
|
||||
elif isinstance(raw_oldvalue, str):
|
||||
newvalue = lark.Tree(
|
||||
lark.Token("RULE", "varvalue"),
|
||||
@@ -1887,7 +1910,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
# Meta should be {"content": raw_oldvalue.meta.content + value.meta.content},
|
||||
)
|
||||
elif any(item in modifiers_str for item in ["?", "ifundefined"]):
|
||||
info += f" ?= '{value_description}'"
|
||||
info += f" ?= '{escape_single_quotes(value_description or '')}'"
|
||||
raw_oldvalue = self.__ppp.user_variables.get(variable, None)
|
||||
if raw_oldvalue is None:
|
||||
newvalue = value
|
||||
@@ -1907,7 +1930,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
if currentvalue is None:
|
||||
info += "not evaluated yet"
|
||||
else:
|
||||
info += f"'{currentvalue}'"
|
||||
info += f"'{escape_single_quotes(currentvalue)}'"
|
||||
t2 = time.monotonic_ns()
|
||||
self.__debug_end(command, start_result, t2 - t1, info)
|
||||
|
||||
@@ -1934,22 +1957,28 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
"""
|
||||
t1 = time.monotonic_ns()
|
||||
start_result = self.result
|
||||
if default is not None:
|
||||
default_value = self.__visit(default, True) # for log
|
||||
default_value = None
|
||||
# if default is not None:
|
||||
# default_value = self.__visit(default, True) # for log
|
||||
value = self.__get_user_variable_value(variable, True, True)
|
||||
if value is None:
|
||||
if default is not None:
|
||||
if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
self.__ppp.logger.debug(
|
||||
f"Variable '{escape_single_quotes(variable)}' not found, using default value"
|
||||
)
|
||||
v = self.__visit(default, False, True)
|
||||
default_value = v
|
||||
self.__ppp.echoed_variables[variable] = v
|
||||
self.result += v
|
||||
else:
|
||||
self.warn_or_stop(f"Unknown variable {variable}")
|
||||
self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable)}")
|
||||
else:
|
||||
self.__ppp.echoed_variables[variable] = value
|
||||
t2 = time.monotonic_ns()
|
||||
info = variable
|
||||
if default is not None:
|
||||
info += f" with default '{default_value}'"
|
||||
if default_value is not None:
|
||||
info += f" with default '{escape_single_quotes(default_value)}'"
|
||||
self.__debug_end(command, start_result, t2 - t1, info)
|
||||
|
||||
def variableuse(self, tree: lark.Tree):
|
||||
@@ -2039,7 +2068,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
)
|
||||
except lark.exceptions.UnexpectedInput as e:
|
||||
self.warn_or_stop(
|
||||
f"Error parsing condition '{v.condition}' in extranetwork mapping '{extnet_id}'! : {e.__class__.__name__}",
|
||||
f"Error parsing condition '{escape_single_quotes(v.condition)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}",
|
||||
e,
|
||||
)
|
||||
cnd = None
|
||||
@@ -2064,7 +2093,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
if found.name:
|
||||
if not found_in_cache and self.__ppp.debug_level != DEBUG_LEVEL.none:
|
||||
self.__ppp.logger.info(
|
||||
f"Mapping extranetwork '{extnet_id}' to '{extnet_type}:{found.name}'"
|
||||
f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to '{escape_single_quotes(extnet_type)}:{escape_single_quotes(found.name)}'"
|
||||
)
|
||||
extnet_id = f"{extnet_type}:{found.name}"
|
||||
f_parameters = found.parameters
|
||||
@@ -2083,11 +2112,15 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
parameters = f_parameters
|
||||
elif found.triggers:
|
||||
if not found_in_cache and self.__ppp.debug_level != DEBUG_LEVEL.none:
|
||||
self.__ppp.logger.info(f"Mapping extranetwork '{extnet_id}' to just triggers")
|
||||
self.__ppp.logger.info(
|
||||
f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to just triggers"
|
||||
)
|
||||
extnet_id = None
|
||||
else:
|
||||
if not found_in_cache and self.__ppp.debug_level != DEBUG_LEVEL.none:
|
||||
self.__ppp.logger.info(f"Mapping extranetwork '{extnet_id}' to nothing")
|
||||
self.__ppp.logger.info(
|
||||
f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to nothing"
|
||||
)
|
||||
extnet_id = None
|
||||
if found.triggers:
|
||||
extra_triggers = ", ".join(found.triggers)
|
||||
@@ -2097,12 +2130,12 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
)
|
||||
except lark.exceptions.UnexpectedInput as e:
|
||||
self.warn_or_stop(
|
||||
f"Error parsing triggers '{extra_triggers}' in extranetwork mapping '{extnet_id}'! : {e.__class__.__name__}",
|
||||
f"Error parsing triggers '{escape_single_quotes(extra_triggers)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}",
|
||||
e,
|
||||
)
|
||||
compiled_extra_triggers = None
|
||||
else:
|
||||
self.warn_or_stop(f"Extranetwork mapping '{extnet_id}' not found!")
|
||||
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
|
||||
@@ -2134,13 +2167,15 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
wildcard_key: str = self.__visit(tree.children[0].children[1], False, True)
|
||||
selected_wildcards = [x.key for x in self.__ppp.wildcard_obj.get_wildcards(wildcard_key)]
|
||||
if not selected_wildcards:
|
||||
self.warn_or_stop(f"Wildcard '{wildcard_key}' not found for default filter setting!")
|
||||
self.warn_or_stop(
|
||||
f"Wildcard '{escape_single_quotes(wildcard_key)}' not found for default filter setting!"
|
||||
)
|
||||
else:
|
||||
filter_object = tree.children[1].children[1] if tree.children[1] is not None else None
|
||||
if filter_object is None:
|
||||
for wc in selected_wildcards:
|
||||
if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
self.__ppp.logger.debug(f"Removed default filter for wildcard '{wc}'")
|
||||
self.__ppp.logger.debug(f"Removed default filter for wildcard '{escape_single_quotes(wc)}'")
|
||||
self.__ppp.wildcard_obj.set_wildcard_default_filter(wc, None)
|
||||
else:
|
||||
filter_specifier: list[list[str]] = [
|
||||
@@ -2148,7 +2183,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
]
|
||||
for wc in selected_wildcards:
|
||||
if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
self.__ppp.logger.debug(f"Set default filter for wildcard '{wc}'")
|
||||
self.__ppp.logger.debug(f"Set default filter for wildcard '{escape_single_quotes(wc)}'")
|
||||
self.__ppp.wildcard_obj.set_wildcard_default_filter(wc, filter_specifier)
|
||||
t2 = time.monotonic_ns()
|
||||
self.__debug_end("commandsetwcdeffilter", start_result, t2 - t1)
|
||||
@@ -2172,7 +2207,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
filter_specifier: Optional[list[list[str]]] = None,
|
||||
wildcard_key: str = None,
|
||||
) -> list[dict]:
|
||||
msg_where = f"wildcard '{wildcard_key}'" if wildcard_key else "choices"
|
||||
msg_where = f"wildcard '{escape_single_quotes(wildcard_key)}'" if wildcard_key else "choices"
|
||||
if filter_specifier is not None:
|
||||
filtered_choice_values = []
|
||||
for i, c in enumerate(choice_values):
|
||||
@@ -2194,7 +2229,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
filtered_choice_values.append(c)
|
||||
if not filtered_choice_values:
|
||||
self.warn_or_stop(
|
||||
f"Wildcard filter specifier '{','.join(['+'.join(y for y in x) for x in filter_specifier])}' found no matches in choices for wildcard '{wildcard_key}'!"
|
||||
f"Wildcard filter specifier '{escape_single_quotes(','.join(['+'.join(y for y in x) for x in filter_specifier]))}' found no matches in choices for wildcard '{escape_single_quotes(wildcard_key)}'!"
|
||||
)
|
||||
else:
|
||||
filtered_choice_values = choice_values.copy()
|
||||
@@ -2206,18 +2241,22 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
if cmd == "include":
|
||||
wcs = self.__ppp.wildcard_obj.get_wildcards(cmd_args)
|
||||
if not wcs:
|
||||
self.warn_or_stop(f"Not found included wildcard '{cmd_args}' at {msg_where}!")
|
||||
self.warn_or_stop(
|
||||
f"Not found included wildcard '{escape_single_quotes(cmd_args)}' at {msg_where}!"
|
||||
)
|
||||
c_weight = float(c.get("weight", 1.0))
|
||||
for wc in wcs:
|
||||
if wc.key in self.__seen_wildcards:
|
||||
self.warn_or_stop(
|
||||
f"Circular reference detected including wildcard '{wc.key}' at {msg_where} (chain starts at '{self.__seen_wildcards[0]}')!"
|
||||
f"Circular reference detected including wildcard '{escape_single_quotes(wc.key)}' at {msg_where} (chain starts at '{escape_single_quotes(self.__seen_wildcards[0])}')!"
|
||||
)
|
||||
continue
|
||||
self.__seen_wildcards.append(wc.key)
|
||||
if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
self.__ppp.logger.debug(f"Seen wildcard '{wc.key}'")
|
||||
self.__ppp.logger.debug(f"Including choices from wildcard '{wc.key}'")
|
||||
self.__ppp.logger.debug(f"Seen wildcard '{escape_single_quotes(wc.key)}'")
|
||||
self.__ppp.logger.debug(
|
||||
f"Including choices from wildcard '{escape_single_quotes(wc.key)}'"
|
||||
)
|
||||
(_, 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)
|
||||
@@ -2229,7 +2268,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
}
|
||||
)
|
||||
else:
|
||||
self.warn_or_stop(f"Unsupported choice command '{cmd}' at {msg_where}!")
|
||||
self.warn_or_stop(f"Unsupported choice command '{escape_single_quotes(cmd)}' at {msg_where}!")
|
||||
else:
|
||||
expanded_choice_values.append(c)
|
||||
return expanded_choice_values
|
||||
@@ -2266,9 +2305,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
from_value: int = options.get("from", 1)
|
||||
to_value: int = options.get("to", 1)
|
||||
separator: str = options.get("separator", self.__ppp.wil_choice_separator)
|
||||
msg_where = f"wildcard '{wildcard_key}'" if wildcard_key else "choices"
|
||||
msg_where = f"wildcard '{escape_single_quotes(wildcard_key)}'" if wildcard_key else "choices"
|
||||
if sampler != "~":
|
||||
self.warn_or_stop(f"Unsupported sampler '{sampler}' at {msg_where} options!")
|
||||
self.warn_or_stop(f"Unsupported sampler '{escape_single_quotes(sampler)}' at {msg_where} options!")
|
||||
sampler = "~"
|
||||
expanded_choice_values = self.__get_choices_internal_get(choice_values, filter_specifier, wildcard_key)
|
||||
available_choices: list[dict] = []
|
||||
@@ -2317,7 +2356,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
self.__ppp.format_output(
|
||||
f"Selecting {'optional ' if optional else ''}{'repeating ' if repeating else ''}{num_choices} choice"
|
||||
+ ("s" if num_choices != 1 else "")
|
||||
+ (f" and separating with '{separator}'" if num_choices > 1 else "")
|
||||
+ (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else "")
|
||||
)
|
||||
)
|
||||
if num_choices > 0:
|
||||
@@ -2364,7 +2403,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
suffix = ""
|
||||
results = []
|
||||
if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
list_unseen = [f"'{x}'" for x in self.__seen_wildcards[seen_wildcards_len:]]
|
||||
list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]
|
||||
self.__ppp.logger.debug(f"Unseen wildcards: {', '.join(list_unseen)}")
|
||||
self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len]
|
||||
return (prefix, results, separator, suffix)
|
||||
@@ -2473,7 +2512,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
)
|
||||
except lark.exceptions.UnexpectedInput as e:
|
||||
self.warn_or_stop(
|
||||
f"Error parsing choice prefix '{prefix}' in wildcard '{wildcard.key}'! : {e.__class__.__name__}",
|
||||
f"Error parsing choice prefix '{escape_single_quotes(prefix)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
|
||||
e,
|
||||
)
|
||||
suffix = options.get("suffix", None)
|
||||
@@ -2484,7 +2523,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
)
|
||||
except lark.exceptions.UnexpectedInput as e:
|
||||
self.warn_or_stop(
|
||||
f"Error parsing choice suffix '{suffix}' in wildcard '{wildcard.key}'! : {e.__class__.__name__}",
|
||||
f"Error parsing choice suffix '{escape_single_quotes(suffix)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
|
||||
e,
|
||||
)
|
||||
n = 1
|
||||
@@ -2517,7 +2556,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
)
|
||||
except lark.exceptions.UnexpectedInput as e:
|
||||
self.warn_or_stop(
|
||||
f"Error parsing condition '{condition}' in wildcard '{wildcard.key}'! : {e.__class__.__name__}",
|
||||
f"Error parsing condition '{escape_single_quotes(condition)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
|
||||
e,
|
||||
)
|
||||
cv["if"] = None
|
||||
@@ -2532,7 +2571,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
)
|
||||
except lark.exceptions.UnexpectedInput as e:
|
||||
self.warn_or_stop(
|
||||
f"Error parsing choice content '{content}' in wildcard '{wildcard.key}'! : {e.__class__.__name__}",
|
||||
f"Error parsing choice content '{escape_single_quotes(content)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
|
||||
e,
|
||||
)
|
||||
cv["content"] = None
|
||||
@@ -2541,9 +2580,13 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
self.__ppp.logger.debug(f"Processed choice {cv}")
|
||||
choice_values.append(cv)
|
||||
else:
|
||||
self.warn_or_stop(f"Invalid choice {cv} in wildcard '{wildcard.key}'!")
|
||||
self.warn_or_stop(
|
||||
f"Invalid choice {cv} in wildcard '{escape_single_quotes(wildcard.key)}'!"
|
||||
)
|
||||
else:
|
||||
self.warn_or_stop(f"Invalid choice {cv} in wildcard '{wildcard.key}'!")
|
||||
self.warn_or_stop(
|
||||
f"Invalid choice {cv} in wildcard '{escape_single_quotes(wildcard.key)}'!"
|
||||
)
|
||||
else:
|
||||
try:
|
||||
choice_values.append(
|
||||
@@ -2553,13 +2596,14 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
)
|
||||
except lark.exceptions.UnexpectedInput as e:
|
||||
self.warn_or_stop(
|
||||
f"Error parsing choice '{cv}' in wildcard '{wildcard.key}'! : {e.__class__.__name__}", e
|
||||
f"Error parsing choice '{escape_single_quotes(cv)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
|
||||
e,
|
||||
)
|
||||
wildcard.choices = choice_values
|
||||
t2 = time.monotonic_ns()
|
||||
if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
self.__ppp.logger.debug(
|
||||
f"Processed choices for wildcard '{wildcard.key}' ({(t2-t1) / 1_000_000_000:.3f} seconds)"
|
||||
f"Processed choices for wildcard '{escape_single_quotes(wildcard.key)}' ({(t2-t1) / 1_000_000_000:.3f} seconds)"
|
||||
)
|
||||
return (options, choice_values)
|
||||
|
||||
@@ -2618,7 +2662,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
and any(x.isdecimal() for x in filter_specifier)
|
||||
):
|
||||
self.__ppp.logger.warning(
|
||||
f"Using a globbing wildcard '{wildcard_key}' with positional index filters is not recommended!"
|
||||
f"Using a globbing wildcard '{escape_single_quotes(wildcard_key)}' with positional index filters is not recommended!"
|
||||
)
|
||||
var_object = tree.children[3]
|
||||
variablename = None
|
||||
@@ -2639,19 +2683,21 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
return
|
||||
if wildcard.key in self.__seen_wildcards:
|
||||
self.warn_or_stop(
|
||||
f"Circular reference detected with wildcard '{self.__seen_wildcards[-1]}' (chain starts at '{self.__seen_wildcards[0]}')!"
|
||||
f"Circular reference detected with wildcard '{escape_single_quotes(self.__seen_wildcards[-1])}' (chain starts at '{escape_single_quotes(self.__seen_wildcards[0])}')!"
|
||||
)
|
||||
continue
|
||||
self.__seen_wildcards.append(wildcard.key)
|
||||
if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
self.__ppp.logger.debug(f"Seen wildcard '{wildcard.key}'")
|
||||
self.__ppp.logger.debug(f"Seen wildcard '{escape_single_quotes(wildcard.key)}'")
|
||||
(options, choice_values) = self.__check_wildcard_initialization(wildcard)
|
||||
if options is not None:
|
||||
if applied_options is None:
|
||||
applied_options = options
|
||||
else:
|
||||
if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
self.__ppp.logger.debug(f"Options for wildcard '{wildcard.key}' are ignored!")
|
||||
self.__ppp.logger.debug(
|
||||
f"Options for wildcard '{escape_single_quotes(wildcard.key)}' are ignored!"
|
||||
)
|
||||
choice_values_all += choice_values
|
||||
self.result += self.__get_choices(applied_options, choice_values_all, filter_specifier, wildcard_key)
|
||||
if wildcard_key in self.__wildcard_filters:
|
||||
@@ -2664,11 +2710,11 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
self.detectedWildcards.append(wc)
|
||||
self.result += wc
|
||||
if self.__ppp.debug_level == DEBUG_LEVEL.full:
|
||||
list_unseen = [f"'{x}'" for x in self.__seen_wildcards[seen_wildcards_len:]]
|
||||
list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]
|
||||
self.__ppp.logger.debug(f"Unseen wildcards: {', '.join(list_unseen)}")
|
||||
self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len]
|
||||
t2 = time.monotonic_ns()
|
||||
self.__debug_end("wildcard", start_result, t2 - t1, f"'{wc}'")
|
||||
self.__debug_end("wildcard", start_result, t2 - t1, f"'{escape_single_quotes(wc)}'")
|
||||
|
||||
def choices(self, tree: lark.Tree):
|
||||
"""
|
||||
@@ -2687,7 +2733,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
self.detectedWildcards.append(ch)
|
||||
self.result += ch
|
||||
t2 = time.monotonic_ns()
|
||||
self.__debug_end("choices", start_result, t2 - t1, f"'{ch}'")
|
||||
self.__debug_end("choices", start_result, t2 - t1, f"'{escape_single_quotes(ch)}'")
|
||||
|
||||
def __default__(self, tree):
|
||||
t1 = time.monotonic_ns()
|
||||
|
||||
+1
-1
@@ -2,7 +2,7 @@ from collections import OrderedDict
|
||||
from logging import Logger
|
||||
from typing import Tuple
|
||||
|
||||
from ppp_logging import DEBUG_LEVEL # pylint: disable=import-error
|
||||
from ppp_logging import DEBUG_LEVEL
|
||||
|
||||
|
||||
class PPPLRUCache:
|
||||
|
||||
+151
@@ -0,0 +1,151 @@
|
||||
"""Pydantic models for the PPP configuration file structure (ppp_config.yaml)."""
|
||||
|
||||
import re
|
||||
from enum import Enum
|
||||
from typing import Literal, Optional
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
||||
class SUPPORTED_APPS(Enum):
|
||||
comfyui = "comfyui"
|
||||
a1111 = "a1111"
|
||||
forge = "forge"
|
||||
reforge = "reforge"
|
||||
sdnext = "sdnext"
|
||||
tests = "tests" # for testing purposes only, not a real app
|
||||
|
||||
SUPPORTED_APPS_NAMES = {
|
||||
SUPPORTED_APPS.comfyui: "ComfyUI",
|
||||
SUPPORTED_APPS.sdnext: "SD.Next",
|
||||
SUPPORTED_APPS.forge: "Forge",
|
||||
SUPPORTED_APPS.reforge: "reForge",
|
||||
SUPPORTED_APPS.a1111: "A1111 (or compatible)",
|
||||
SUPPORTED_APPS.tests: "Tests",
|
||||
}
|
||||
|
||||
# ------------------- Host configuration -------------------
|
||||
|
||||
AttentionOption = Literal["ok", "parentheses", "disable", "remove", "error"]
|
||||
SchedulingOption = Literal["ok", "before", "after", "first", "remove", "error"]
|
||||
AlternationOption = Literal["ok", "first", "remove", "error"]
|
||||
AndOption = Literal["ok", "eol", "comma", "remove", "error"]
|
||||
BreakOption = Literal["ok", "eol", "comma", "remove", "error"]
|
||||
|
||||
|
||||
class HostConfig(BaseModel):
|
||||
"""Configuration for a specific host application."""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True, extra="forbid")
|
||||
|
||||
attention: AttentionOption = "ok"
|
||||
scheduling: SchedulingOption = "ok"
|
||||
alternation: AlternationOption = "ok"
|
||||
and_: AndOption = Field("ok", alias="and")
|
||||
break_: BreakOption = Field("ok", alias="break")
|
||||
|
||||
|
||||
# ------------------- Model detection -------------------
|
||||
|
||||
|
||||
class ModelDetectConfig(BaseModel):
|
||||
"""Detection configuration for a specific host when loading a model."""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
class_: Optional[list[str]] = Field(None, alias="class")
|
||||
property: Optional[str] = None
|
||||
|
||||
@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")
|
||||
return self
|
||||
|
||||
|
||||
# ------------------- Variant find_in_filename -------------------
|
||||
|
||||
|
||||
class FindInFilenamePattern(BaseModel):
|
||||
"""A regex pattern with optional flags used to identify a model variant in the filename."""
|
||||
|
||||
regex: str
|
||||
flags: int = 0
|
||||
|
||||
@field_validator("flags", mode="before")
|
||||
@classmethod
|
||||
def parse_flags(cls, v: object) -> int:
|
||||
if isinstance(v, int):
|
||||
return v
|
||||
if isinstance(v, list):
|
||||
flag_value = 0
|
||||
for flag in v:
|
||||
if not isinstance(flag, str) or not hasattr(re, 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__}")
|
||||
|
||||
@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
|
||||
return self
|
||||
|
||||
|
||||
class VariantConfig(BaseModel):
|
||||
"""Configuration for a specific model variant."""
|
||||
|
||||
find_in_filename: list[FindInFilenamePattern]
|
||||
|
||||
@field_validator("find_in_filename", mode="before")
|
||||
@classmethod
|
||||
def normalize_find_in_filename(cls, v: object) -> list:
|
||||
"""Normalize str / dict / list input to always be a list of FindInFilenamePattern-compatible dicts."""
|
||||
if isinstance(v, str):
|
||||
return [{"regex": v, "flags": re.IGNORECASE}]
|
||||
if isinstance(v, dict):
|
||||
return [v]
|
||||
if isinstance(v, list):
|
||||
normalized = []
|
||||
for item in v:
|
||||
if isinstance(item, str):
|
||||
normalized.append({"regex": item, "flags": re.IGNORECASE})
|
||||
elif isinstance(item, dict):
|
||||
normalized.append(item)
|
||||
else:
|
||||
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__}")
|
||||
|
||||
|
||||
# ------------------- Model configuration -------------------
|
||||
|
||||
|
||||
class ModelConfig(BaseModel):
|
||||
"""Configuration for a supported base model."""
|
||||
|
||||
detect: Optional[dict[str, Optional[ModelDetectConfig]]] = None
|
||||
variants: Optional[dict[str, VariantConfig]] = None
|
||||
|
||||
@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")
|
||||
return self
|
||||
|
||||
|
||||
# ------------------- Top-level configuration -------------------
|
||||
|
||||
|
||||
class PPPConfig(BaseModel):
|
||||
"""Top-level PPP configuration structure matching ppp_config.yaml."""
|
||||
|
||||
hosts: Optional[dict[str, Optional[HostConfig]]] = None
|
||||
models: Optional[dict[str, Optional[ModelConfig | None]]] = None
|
||||
|
||||
@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")
|
||||
return self
|
||||
+9
-9
@@ -1,14 +1,14 @@
|
||||
import os
|
||||
|
||||
# pylint: disable=import-error
|
||||
import folder_paths # type: ignore
|
||||
import nodes # type: ignore
|
||||
import folder_paths # pylint: disable=import-error # type: ignore
|
||||
import nodes # pylint: disable=import-error # type: ignore
|
||||
|
||||
from .ppp import PromptPostProcessor
|
||||
from .ppp_hosts import SUPPORTED_APPS
|
||||
from .ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory
|
||||
from .ppp_wildcards import PPPWildcards
|
||||
from .ppp_enmappings import PPPExtraNetworkMappings
|
||||
from ppp import PromptPostProcessor
|
||||
from ppp_classes import SUPPORTED_APPS
|
||||
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory
|
||||
from ppp_utils import escape_single_quotes
|
||||
from ppp_wildcards import PPPWildcards
|
||||
from ppp_enmappings import PPPExtraNetworkMappings
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit("This script must be run from ComfyUI")
|
||||
@@ -166,7 +166,7 @@ class PromptPostProcessorComfyUINode:
|
||||
for input_name, input_type in input_types.items():
|
||||
t = expected[input_name]
|
||||
if input_type != t:
|
||||
return f"Invalid type for input '{input_name}': {input_type} (expected {t})"
|
||||
return f"Invalid type for input '{escape_single_quotes(input_name)}': {input_type} (expected {t})"
|
||||
return True
|
||||
|
||||
RETURN_TYPES = (
|
||||
|
||||
+18
-10
@@ -3,8 +3,8 @@ from typing import Optional
|
||||
import logging
|
||||
import yaml
|
||||
|
||||
from ppp_logging import DEBUG_LEVEL # pylint: disable=import-error
|
||||
from ppp_utils import deep_freeze # pylint: disable=import-error
|
||||
from ppp_logging import DEBUG_LEVEL
|
||||
from ppp_utils import deep_freeze, escape_single_quotes
|
||||
|
||||
|
||||
class PPPENMappingVariant:
|
||||
@@ -209,25 +209,27 @@ class PPPExtraNetworkMappings:
|
||||
full_path (str): The path to the file that contains it.
|
||||
"""
|
||||
if not isinstance(content, dict):
|
||||
self.__logger.warning(f"Invalid extra network mapping in file '{full_path}'!")
|
||||
self.__logger.warning(f"Invalid extra network mapping in file '{escape_single_quotes(full_path)}'!")
|
||||
return
|
||||
for kind, maps in content.items():
|
||||
if not isinstance(maps, dict):
|
||||
self.__logger.warning(f"Invalid extra network mapping definition for '{kind}:*' in file '{full_path}'!")
|
||||
self.__logger.warning(
|
||||
f"Invalid extra network mapping definition for '{escape_single_quotes(kind)}:*' in file '{escape_single_quotes(full_path)}'!"
|
||||
)
|
||||
else:
|
||||
for name, variants in maps.items():
|
||||
key = f"{kind}:{name}"
|
||||
if not isinstance(variants, list):
|
||||
self.__logger.warning(
|
||||
f"Invalid extra network mapping definition for '{key}' in file '{full_path}'!"
|
||||
f"Invalid extra network mapping definition for '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}'!"
|
||||
)
|
||||
elif self.extranetwork_mappings.get(key, None) is not None:
|
||||
self.__logger.warning(
|
||||
f"Duplicate extra network mapping '{key}' in file '{full_path}' and '{self.extranetwork_mappings[key].file}'!"
|
||||
f"Duplicate extra network mapping '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.extranetwork_mappings[key].file)}'!"
|
||||
)
|
||||
elif not isinstance(variants, list) or not all(isinstance(v, dict) for v in variants):
|
||||
self.__logger.warning(
|
||||
f"Invalid extra network mapping definition for '{key}' in file '{full_path}'!"
|
||||
f"Invalid extra network mapping definition for '{escape_single_quotes(key)}' in file '{escape_single_quotes(full_path)}'!"
|
||||
)
|
||||
else:
|
||||
self.extranetwork_mappings[key] = PPPENMapping(full_path, kind, name, variants)
|
||||
@@ -245,12 +247,16 @@ class PPPExtraNetworkMappings:
|
||||
with open(full_path, "r", encoding="utf-8") as file:
|
||||
content = yaml.safe_load(file)
|
||||
except: # pylint: disable=bare-except
|
||||
self.__logger.warning(f"Could not read file '{full_path}' with utf-8 encoding, trying windows-1252...")
|
||||
self.__logger.warning(
|
||||
f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252..."
|
||||
)
|
||||
with open(full_path, "r", encoding="windows-1252") as file:
|
||||
content = yaml.safe_load(file)
|
||||
self.__add_extranetwork_mapping(content, full_path)
|
||||
except Exception as e: # pylint: disable=broad-except
|
||||
self.__logger.error(f"Error reading extra network mappings from file '{full_path}': {e}")
|
||||
self.__logger.error(
|
||||
f"Error reading extra network mappings from file '{escape_single_quotes(full_path)}': {e}"
|
||||
)
|
||||
|
||||
def __get_extranetwork_mappings_in_directory(self, directory: str):
|
||||
"""
|
||||
@@ -260,7 +266,9 @@ class PPPExtraNetworkMappings:
|
||||
directory (str): The path to the directory.
|
||||
"""
|
||||
if not os.path.exists(directory):
|
||||
self.__logger.warning(f"Extra network mappings directory '{directory}' does not exist!")
|
||||
self.__logger.warning(
|
||||
f"Extra network mappings directory '{escape_single_quotes(directory)}' does not exist!"
|
||||
)
|
||||
return
|
||||
for filename in os.listdir(directory):
|
||||
full_path = os.path.abspath(os.path.join(directory, filename))
|
||||
|
||||
@@ -1,19 +0,0 @@
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class SUPPORTED_APPS(Enum):
|
||||
comfyui = "comfyui"
|
||||
a1111 = "a1111"
|
||||
forge = "forge"
|
||||
reforge = "reforge"
|
||||
sdnext = "sdnext"
|
||||
tests = "tests" # for testing purposes only, not a real app
|
||||
|
||||
SUPPORTED_APPS_NAMES = {
|
||||
SUPPORTED_APPS.comfyui: "ComfyUI",
|
||||
SUPPORTED_APPS.sdnext: "SD.Next",
|
||||
SUPPORTED_APPS.forge: "Forge",
|
||||
SUPPORTED_APPS.reforge: "reForge",
|
||||
SUPPORTED_APPS.a1111: "A1111 (or compatible)",
|
||||
SUPPORTED_APPS.tests: "Tests",
|
||||
}
|
||||
+1
-1
@@ -2,7 +2,7 @@ from enum import Enum
|
||||
import logging
|
||||
import sys
|
||||
import copy
|
||||
from ppp_hosts import SUPPORTED_APPS # pylint: disable=import-error
|
||||
from ppp_classes import SUPPORTED_APPS
|
||||
|
||||
|
||||
class DEBUG_LEVEL(Enum):
|
||||
|
||||
@@ -15,3 +15,27 @@ def deep_freeze(obj):
|
||||
if isinstance(obj, set):
|
||||
return tuple(deep_freeze(i) for i in sorted(obj))
|
||||
return obj
|
||||
|
||||
def escape_single_quotes(s: str):
|
||||
"""
|
||||
Escape single quotes in a string.
|
||||
|
||||
Args:
|
||||
s (str): The string to escape.
|
||||
|
||||
Returns:
|
||||
str: The escaped string.
|
||||
"""
|
||||
return s.replace("'", "\\'")
|
||||
|
||||
def escape_double_quotes(s: str):
|
||||
"""
|
||||
Escape double quotes in a string.
|
||||
|
||||
Args:
|
||||
s (str): The string to escape.
|
||||
|
||||
Returns:
|
||||
str: The escaped string.
|
||||
"""
|
||||
return s.replace('"', '\\"')
|
||||
|
||||
+28
-16
@@ -4,8 +4,8 @@ from typing import Optional
|
||||
import logging
|
||||
import yaml
|
||||
|
||||
from ppp_logging import DEBUG_LEVEL # pylint: disable=import-error
|
||||
from ppp_utils import deep_freeze # pylint: disable=import-error
|
||||
from ppp_logging import DEBUG_LEVEL
|
||||
from ppp_utils import deep_freeze, escape_single_quotes
|
||||
|
||||
|
||||
class PPPWildcard:
|
||||
@@ -208,7 +208,7 @@ class PPPWildcards:
|
||||
self.__get_wildcards_in_structured_file(full_path, base)
|
||||
self.__wildcard_files[full_path] = last_modified
|
||||
except Exception as e: # pylint: disable=broad-except
|
||||
self.__logger.error(f"Error reading wildcard file '{full_path}': {e}")
|
||||
self.__logger.error(f"Error reading wildcard file '{escape_single_quotes(full_path)}': {e}")
|
||||
|
||||
def __get_wildcards_in_input(self, wildcards_input: str):
|
||||
"""
|
||||
@@ -284,7 +284,9 @@ class PPPWildcards:
|
||||
if isinstance(obj, (int, float, bool)):
|
||||
return [str(obj)]
|
||||
if not isinstance(obj, list) or len(obj) == 0:
|
||||
self.__logger.warning(f"Invalid format in wildcard '{'/'.join(key_parts)}' in file '{full_path}'!")
|
||||
self.__logger.warning(
|
||||
f"Invalid format in wildcard '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(full_path)}'!"
|
||||
)
|
||||
return None
|
||||
choices = []
|
||||
for i, c in enumerate(obj):
|
||||
@@ -297,7 +299,7 @@ class PPPWildcards:
|
||||
choices.append(self.__process_dict_choice(c, full_path, key_parts, i))
|
||||
else:
|
||||
self.__logger.warning(
|
||||
f"Invalid choice {i+1} in wildcard '{'/'.join(key_parts)}' in file '{full_path}'!"
|
||||
f"Invalid choice {i+1} in wildcard '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(full_path)}'!"
|
||||
)
|
||||
return choices
|
||||
|
||||
@@ -328,7 +330,9 @@ class PPPWildcards:
|
||||
# we assume it is an anonymous wildcard with options
|
||||
firstkey = list(c.keys())[0]
|
||||
return self.__create_anonymous_wildcard(full_path, key_parts, i, c[firstkey], firstkey)
|
||||
self.__logger.warning(f"Invalid choice {i+1} in wildcard '{'/'.join(key_parts)}' in file '{full_path}'!")
|
||||
self.__logger.warning(
|
||||
f"Invalid choice {i+1} in wildcard '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(full_path)}'!"
|
||||
)
|
||||
return None
|
||||
|
||||
def __create_anonymous_wildcard(self, full_path, key_parts, i, content, options=None):
|
||||
@@ -371,16 +375,18 @@ class PPPWildcards:
|
||||
fullkey = "/".join(tmp_key_parts)
|
||||
if self.wildcards.get(fullkey, None) is not None:
|
||||
self.__logger.warning(
|
||||
f"Duplicate wildcard '{fullkey}' in file '{full_path}' and '{self.wildcards[fullkey].file}'!"
|
||||
f"Duplicate wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.wildcards[fullkey].file)}'!"
|
||||
)
|
||||
else:
|
||||
obj = self.__get_nested(content, key)
|
||||
choices = self.__get_choices(obj, full_path, tmp_key_parts)
|
||||
if choices is None:
|
||||
self.__logger.warning(f"Invalid wildcard '{fullkey}' in file '{full_path}'!")
|
||||
self.__logger.warning(
|
||||
f"Invalid wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'!"
|
||||
)
|
||||
elif fullkey.startswith("_"):
|
||||
self.__logger.warning(
|
||||
f"Invalid wildcard name '{fullkey}' in file '{full_path}'! (cannot start with underscore)"
|
||||
f"Invalid wildcard name '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'! (cannot start with underscore)"
|
||||
)
|
||||
else:
|
||||
self.wildcards[fullkey] = PPPWildcard(full_path, fullkey, choices)
|
||||
@@ -390,20 +396,22 @@ class PPPWildcards:
|
||||
elif isinstance(content, (int, float, bool)):
|
||||
content = [str(content)]
|
||||
if not isinstance(content, list):
|
||||
self.__logger.warning(f"Invalid wildcard in file '{full_path}'!")
|
||||
self.__logger.warning(f"Invalid wildcard in file '{escape_single_quotes(full_path)}'!")
|
||||
return
|
||||
fullkey = "/".join(key_parts)
|
||||
if self.wildcards.get(fullkey, None) is not None:
|
||||
self.__logger.warning(
|
||||
f"Duplicate wildcard '{fullkey}' in file '{full_path}' and '{self.wildcards[fullkey].file}'!"
|
||||
f"Duplicate wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}' and '{escape_single_quotes(self.wildcards[fullkey].file)}'!"
|
||||
)
|
||||
else:
|
||||
choices = self.__get_choices(content, full_path, key_parts)
|
||||
if choices is None:
|
||||
self.__logger.warning(f"Invalid wildcard '{fullkey}' in file '{full_path}'!")
|
||||
self.__logger.warning(
|
||||
f"Invalid wildcard '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'!"
|
||||
)
|
||||
elif fullkey.startswith("_"):
|
||||
self.__logger.warning(
|
||||
f"Invalid wildcard name '{fullkey}' in file '{full_path}'! (cannot start with underscore)"
|
||||
f"Invalid wildcard name '{escape_single_quotes(fullkey)}' in file '{escape_single_quotes(full_path)}'! (cannot start with underscore)"
|
||||
)
|
||||
else:
|
||||
self.wildcards[fullkey] = PPPWildcard(full_path, fullkey, choices)
|
||||
@@ -422,7 +430,9 @@ class PPPWildcards:
|
||||
with open(full_path, "r", encoding="utf-8") as file:
|
||||
content = yaml.safe_load(file)
|
||||
except: # pylint: disable=bare-except
|
||||
self.__logger.warning(f"Could not read file '{full_path}' with utf-8 encoding, trying windows-1252...")
|
||||
self.__logger.warning(
|
||||
f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252..."
|
||||
)
|
||||
with open(full_path, "r", encoding="windows-1252") as file:
|
||||
content = yaml.safe_load(file)
|
||||
self.__add_wildcard(content, full_path, external_key_parts)
|
||||
@@ -441,7 +451,9 @@ class PPPWildcards:
|
||||
with open(full_path, "r", encoding="utf-8") as file:
|
||||
text_content = map(lambda x: x.strip("\n\r"), file.readlines())
|
||||
except: # pylint: disable=bare-except
|
||||
self.__logger.warning(f"Could not read file '{full_path}' with utf-8 encoding, trying windows-1252...")
|
||||
self.__logger.warning(
|
||||
f"Could not read file '{escape_single_quotes(full_path)}' with utf-8 encoding, trying windows-1252..."
|
||||
)
|
||||
with open(full_path, "r", encoding="windows-1252") as file:
|
||||
text_content = map(lambda x: x.strip("\n\r"), file.readlines())
|
||||
text_content = list(filter(lambda x: x.strip() != "" and not x.strip().startswith("#"), text_content))
|
||||
@@ -457,7 +469,7 @@ class PPPWildcards:
|
||||
directory (str): The path to the directory.
|
||||
"""
|
||||
if not os.path.exists(directory):
|
||||
self.__logger.warning(f"Wildcard directory '{directory}' does not exist!")
|
||||
self.__logger.warning(f"Wildcard directory '{escape_single_quotes(directory)}' does not exist!")
|
||||
return
|
||||
for filename in os.listdir(directory):
|
||||
full_path = os.path.abspath(os.path.join(directory, filename))
|
||||
|
||||
@@ -14,12 +14,12 @@ from modules.processing import StableDiffusionProcessing # pylint: disable=impo
|
||||
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 ppp import PromptPostProcessor # pylint: disable=import-error
|
||||
from ppp_hosts import SUPPORTED_APPS, SUPPORTED_APPS_NAMES # pylint: disable=import-error
|
||||
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # pylint: disable=import-error
|
||||
from ppp_cache import PPPLRUCache # pylint: disable=import-error
|
||||
from ppp_wildcards import PPPWildcards # pylint: disable=import-error
|
||||
from ppp_enmappings import PPPExtraNetworkMappings # pylint: disable=import-error
|
||||
from ppp import PromptPostProcessor
|
||||
from ppp_classes import SUPPORTED_APPS, SUPPORTED_APPS_NAMES
|
||||
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory
|
||||
from ppp_cache import PPPLRUCache
|
||||
from ppp_wildcards import PPPWildcards
|
||||
from ppp_enmappings import PPPExtraNetworkMappings
|
||||
|
||||
|
||||
class PromptPostProcessorA1111Script(scripts.Script):
|
||||
@@ -294,7 +294,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_mappings_obj,
|
||||
)
|
||||
hash_options = ppp.options_hash()
|
||||
hash_options = ppp.options_hash()
|
||||
hash_envinfo = ppp.envinfo_hash()
|
||||
prompts_list = []
|
||||
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
import os
|
||||
import logging
|
||||
from typing import NamedTuple, Optional
|
||||
import unittest
|
||||
import datetime
|
||||
|
||||
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
|
||||
|
||||
|
||||
class PromptPair(NamedTuple):
|
||||
prompt: str = ""
|
||||
negative_prompt: str = ""
|
||||
|
||||
|
||||
class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
"""
|
||||
A test case class for testing the PromptPostProcessor class.
|
||||
"""
|
||||
|
||||
def setUp(self, enable_file_logging=False):
|
||||
"""
|
||||
Set up the test case by initializing the necessary objects and configurations.
|
||||
|
||||
Args:
|
||||
enable_file_logging (bool): Whether to enable logging to a file. Defaults to True.
|
||||
"""
|
||||
self.enable_file_logging = enable_file_logging
|
||||
test_name = self.id().split(".")[-1] # Extract the test method name
|
||||
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
|
||||
if self.enable_file_logging:
|
||||
log_filename = f"tests/logs/{test_name}_{timestamp}.log"
|
||||
else:
|
||||
log_filename = None # Disable file logging
|
||||
|
||||
self.lf = PromptPostProcessorLogFactory(None, log_filename)
|
||||
self.ppp_logger = self.lf.log
|
||||
self.ppp_logger.setLevel(logging.DEBUG)
|
||||
self.grammar_content = None
|
||||
self.interrupted = False
|
||||
self.defopts = {
|
||||
"debug_level": DEBUG_LEVEL.full.value,
|
||||
"on_warning": PromptPostProcessor.ONWARNING_CHOICES.stop.value,
|
||||
"process_wildcards": True,
|
||||
"if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value,
|
||||
"choice_separator": ", ",
|
||||
"keep_choices_order": False,
|
||||
"stn_separator": ", ",
|
||||
"stn_ignore_repeats": True,
|
||||
"do_cleanup": True,
|
||||
"cleanup_variables": True,
|
||||
"cleanup_empty_constructs": True,
|
||||
"cleanup_extra_separators": True,
|
||||
"cleanup_extra_separators2": True,
|
||||
"cleanup_extra_separators_include_eol": False,
|
||||
"cleanup_extra_spaces": True,
|
||||
"cleanup_breaks": True,
|
||||
"cleanup_breaks_eol": False,
|
||||
"cleanup_ands": True,
|
||||
"cleanup_ands_eol": False,
|
||||
"cleanup_extranetwork_tags": True,
|
||||
"cleanup_merge_attention": True,
|
||||
"remove_extranetwork_tags": False,
|
||||
}
|
||||
self.def_env_info = {
|
||||
"app": "tests",
|
||||
"ppp_config": None,
|
||||
"model_class": "SDXL",
|
||||
"property_base": {"is_sdxl": True},
|
||||
"models_path": "./webui/models",
|
||||
"model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors",
|
||||
}
|
||||
self.interrupted = False
|
||||
self.wildcards_obj = PPPWildcards(self.lf.log)
|
||||
self.extranetwork_maps_obj = PPPExtraNetworkMappings(self.lf.log)
|
||||
self.wildcards_obj.refresh_wildcards(
|
||||
DEBUG_LEVEL.full,
|
||||
[
|
||||
os.path.abspath(os.path.join(os.path.dirname(__file__), "wildcards")),
|
||||
os.path.abspath(os.path.join(os.path.dirname(__file__), "wildcards2")),
|
||||
],
|
||||
"""
|
||||
yaml_input:
|
||||
wildcardI:
|
||||
- choice1
|
||||
- choice2
|
||||
- choice3
|
||||
""",
|
||||
)
|
||||
self.extranetwork_maps_obj.refresh_extranetwork_mappings(
|
||||
DEBUG_LEVEL.full,
|
||||
[os.path.abspath(os.path.join(os.path.dirname(__file__), "enmappings"))],
|
||||
"""
|
||||
""",
|
||||
)
|
||||
grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "../grammar.lark")
|
||||
with open(grammar_filename, "r", encoding="utf-8") as file:
|
||||
self.grammar_content = file.read()
|
||||
|
||||
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
|
||||
"""
|
||||
if isinstance(ppp, str):
|
||||
if ppp == "nocup":
|
||||
the_obj = PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{
|
||||
**self.defopts,
|
||||
"do_cleanup": False,
|
||||
"cleanup_variables": False,
|
||||
"cleanup_empty_constructs": False,
|
||||
"cleanup_extra_separators": False,
|
||||
"cleanup_extra_separators2": False,
|
||||
"cleanup_extra_separators_include_eol": False,
|
||||
"cleanup_extra_spaces": False,
|
||||
"cleanup_breaks": False,
|
||||
"cleanup_breaks_eol": False,
|
||||
"cleanup_ands": False,
|
||||
"cleanup_ands_eol": False,
|
||||
"cleanup_extranetwork_tags": False,
|
||||
"cleanup_merge_attention": False,
|
||||
},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
)
|
||||
# elif ppp == "comfyui":
|
||||
# the_obj = PromptPostProcessor(
|
||||
# self.ppp_logger,
|
||||
# self.interrupt,
|
||||
# {
|
||||
# **self.def_env_info,
|
||||
# "app": "comfyui",
|
||||
# "model_class": "SDXL",
|
||||
# },
|
||||
# self.defopts,
|
||||
# self.grammar_content,
|
||||
# self.wildcards_obj,
|
||||
# self.extranetwork_maps_obj,
|
||||
# )
|
||||
else:
|
||||
the_obj = ppp
|
||||
if not the_obj:
|
||||
the_obj = PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
)
|
||||
out = (
|
||||
[PromptPair("", "")]
|
||||
if expected_output_prompts is None
|
||||
else expected_output_prompts if isinstance(expected_output_prompts, list) else [expected_output_prompts]
|
||||
)
|
||||
for eo in out:
|
||||
result_prompt, result_negative_prompt, output_variables = 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",
|
||||
)
|
||||
seed += 1
|
||||
-1961
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,101 @@
|
||||
import unittest
|
||||
|
||||
from ppp import PromptPostProcessor # pylint: disable=import-error
|
||||
from .base_tests import PromptPair, TestPromptPostProcessorBase
|
||||
|
||||
|
||||
class TestChoices(TestPromptPostProcessorBase):
|
||||
|
||||
def setUp(self): # pylint: disable=arguments-differ
|
||||
super().setUp(enable_file_logging=False)
|
||||
|
||||
# Choices tests
|
||||
|
||||
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", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_ch_unsupportedsampler(self): # unsupported sampler
|
||||
self.process(
|
||||
PromptPair("the choices are: {@choice1|choice2|choice3}", ""),
|
||||
PromptPair("", ""),
|
||||
ppp="nocup",
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
def test_ch_choices_withcomments(self): # choices with comments and multiline
|
||||
self.process(
|
||||
PromptPair(
|
||||
"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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_ch_choices_set_if_nested(self): # nested choices with if user variable and multiple selection
|
||||
self.process(
|
||||
PromptPair(
|
||||
"${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", ""),
|
||||
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>", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_ch_removelorawithchoices(self):
|
||||
self.process(
|
||||
PromptPair("<lora:test1:1><lora:test2:{0.2|0.5|0.7|1}>", ""),
|
||||
PromptPair("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{**self.defopts, "remove_extranetwork_tags": True},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_ch_cmd_includewildcard(self):
|
||||
self.process(
|
||||
PromptPair("{ch_one|ch_two|%0.5::include yaml/wildcard1}", ""),
|
||||
PromptPair("ch_two", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,132 @@
|
||||
import unittest
|
||||
|
||||
from ppp import PromptPostProcessor # pylint: disable=import-error
|
||||
|
||||
from .base_tests import PromptPair, TestPromptPostProcessorBase
|
||||
|
||||
|
||||
class TestCleanup(TestPromptPostProcessorBase):
|
||||
|
||||
def setUp(self): # pylint: disable=arguments-differ
|
||||
super().setUp(enable_file_logging=False)
|
||||
|
||||
# Cleanup tests
|
||||
|
||||
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"),
|
||||
)
|
||||
|
||||
def test_cl_complex(self): # complex cleanup
|
||||
self.process(
|
||||
PromptPair(
|
||||
" 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(
|
||||
"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",
|
||||
),
|
||||
)
|
||||
|
||||
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", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{**self.defopts, "remove_extranetwork_tags": True},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
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", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{
|
||||
**self.defopts,
|
||||
"cleanup_extra_separators2": False,
|
||||
"cleanup_extra_separators_include_eol": False,
|
||||
},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_cl_separatorswitheol(self): # don't remove eols with the separators
|
||||
self.process(
|
||||
PromptPair(
|
||||
"""{ (d:0.9) ,, (l:1.1) | (l:1.1) (d:0.9),,, }
|
||||
(l:1.1)
|
||||
(d:0.9)""",
|
||||
"",
|
||||
),
|
||||
PromptPair(
|
||||
""" (l:1.1) (d:0.9),
|
||||
(l:1.1)
|
||||
(d:0.9)""",
|
||||
"",
|
||||
),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{
|
||||
**self.defopts,
|
||||
"cleanup_empty_constructs": False,
|
||||
"cleanup_extra_separators": True,
|
||||
"cleanup_extra_separators2": False,
|
||||
"cleanup_extra_separators_include_eol": False,
|
||||
"cleanup_extra_spaces": False,
|
||||
"cleanup_breaks": False,
|
||||
"cleanup_breaks_eol": False,
|
||||
"cleanup_ands": False,
|
||||
"cleanup_ands_eol": False,
|
||||
"cleanup_extranetwork_tags": False,
|
||||
"cleanup_merge_attention": False,
|
||||
},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_cl_mergeattention(self): # merge attention
|
||||
self.process(
|
||||
PromptPair(
|
||||
"this is (a test:0.9) of (attention (merging:1.2)) where ((this)) ((is joined:1.2)) and ([this too]:1.3)",
|
||||
"",
|
||||
),
|
||||
PromptPair(
|
||||
"this is [a test] of (attention (merging:1.2)) where (this:1.21) (is joined:1.32) and (this too:1.17)",
|
||||
"",
|
||||
),
|
||||
)
|
||||
|
||||
def test_cl_not_mergeattention(self): # not merge attention
|
||||
self.process(
|
||||
PromptPair(
|
||||
"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(
|
||||
"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)",
|
||||
"",
|
||||
),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,348 @@
|
||||
import unittest
|
||||
|
||||
from ppp import PromptPostProcessor # pylint: disable=import-error
|
||||
from .base_tests import PromptPair, TestPromptPostProcessorBase
|
||||
|
||||
|
||||
class TestCommands(TestPromptPostProcessorBase):
|
||||
|
||||
def setUp(self): # pylint: disable=arguments-differ
|
||||
super().setUp(enable_file_logging=False)
|
||||
|
||||
# Command tests
|
||||
|
||||
def test_cmd_stn_complex_features(self): # complex stn command with AND, BREAK and other features
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[<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(
|
||||
"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]",
|
||||
),
|
||||
)
|
||||
|
||||
def test_cmd_if_complex_features(self): # complex if command
|
||||
self.process(
|
||||
PromptPair(
|
||||
"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(
|
||||
"this \\(is\\): a (([complex|simple|regular] test)(test:2):1.5)\nBREAK :0.5 AND hypernettrigger <hypernet:yyy>:0.3",
|
||||
"normal quality",
|
||||
),
|
||||
)
|
||||
|
||||
def test_cmd_if_nested(self): # nested if command
|
||||
self.process(
|
||||
PromptPair(
|
||||
"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", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
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", ""),
|
||||
)
|
||||
|
||||
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", ""),
|
||||
)
|
||||
|
||||
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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_if_echo_nested(self): # nested set, if and echo commands
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_if_complex_conditions_1(self): # complex conditions (or)
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_if_complex_conditions_2(self): # complex conditions (and)
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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", ""),
|
||||
)
|
||||
|
||||
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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_if_complex_conditions_4(self): # complex conditions (not, precedence)
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_if_complex_conditions_5(self): # complex conditions (not, precedence, comparison)
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_if_complex_conditions_6(self): # complex conditions
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_if_complex_conditions_7(self): # complex conditions
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_if2(self): # set and more complex if commands
|
||||
self.process(
|
||||
PromptPair(
|
||||
"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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_add_if(self): # set, add and if commands
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_add_DP_if(self): # set, add (DP format) and if commands
|
||||
self.process(
|
||||
PromptPair(
|
||||
"${v=value}${v+=2}this test is <ppp:if v eq 'value2'>OK<ppp:else>not OK<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("this test is OK", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_immediateeval(self): # set (DP format) with mixed evaluation
|
||||
self.process(
|
||||
PromptPair(
|
||||
"${var=!__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var2:default}, ${var3=__yaml/wildcard1__}${var3}, ${var3}",
|
||||
"",
|
||||
),
|
||||
PromptPair("the choices are: choice2, choice2, default, choice3, choice1", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_cmd_set_mixeval(self): # set and add (DP format) with mixed evaluation
|
||||
self.process(
|
||||
PromptPair(
|
||||
"${var=__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var+=, __yaml/wildcard2__}${var}, ${var}, ${var+=!, __yaml/wildcard3__}${var}, ${var}",
|
||||
"",
|
||||
),
|
||||
PromptPair(
|
||||
"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 ",
|
||||
"",
|
||||
),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_cmd_set_ifundefined_if(self): # set, ifundefined and if commands
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_ifundefined_if_2(self): # set, ifundefined and if commands
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_ifundefined_DP_if(self): # set, ifundefined (DP format) and if commands
|
||||
self.process(
|
||||
PromptPair(
|
||||
"${v?=value}this test is <ppp:if v eq 'value'>OK<ppp:else>not OK<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("this test is OK", ""),
|
||||
)
|
||||
|
||||
def test_cmd_set_ifundefined_DP_if_2(self): # set, ifundefined (DP format) and if commands
|
||||
self.process(
|
||||
PromptPair(
|
||||
"${v=!value}${v?=!value2}this test is <ppp:if v eq 'value'>OK<ppp:else>not OK<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("this test is OK", ""),
|
||||
)
|
||||
|
||||
def test_cmd_ext(self): # ext
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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(
|
||||
"<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",
|
||||
"",
|
||||
),
|
||||
)
|
||||
|
||||
def test_cmd_ext_map_notrigger(self): # ext mapping, no trigger
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:ext $lora lora1/><ppp:ext $lora lora1>",
|
||||
"",
|
||||
),
|
||||
PromptPair("triggergeneric1, triggergeneric2, two, triggergeneric1, triggergeneric2, two", ""),
|
||||
)
|
||||
|
||||
def test_cmd_ext_map1(self): # ext mapping, no lora
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:ext $lora lora1>inlinetrigger<ppp:/ext>",
|
||||
"",
|
||||
),
|
||||
PromptPair("inlinetrigger, triggergeneric1, triggergeneric2, two", ""),
|
||||
)
|
||||
|
||||
def test_cmd_ext_map2(self): # ext mapping, lora with weight
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:ext $lora lora1>inlinetrigger<ppp:/ext>",
|
||||
"",
|
||||
),
|
||||
PromptPair("<lora:lorapony:0.8>inlinetrigger, triggerpony1, triggerpony2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_cmd_ext_map3(self): # ext mapping, lora with weight adjusted
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:ext $lora lora1 0.5>inlinetrigger<ppp:/ext>",
|
||||
"",
|
||||
),
|
||||
PromptPair("<lora:lorapony:0.4>inlinetrigger, triggerpony1, triggerpony2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_cmd_ext_map4(self): # ext mapping, lora with parameters
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:ext $lora lora1 '0.6:0.8'>inlinetrigger<ppp:/ext>",
|
||||
"",
|
||||
),
|
||||
PromptPair("<lora:lorapony:0.6:0.8>inlinetrigger, triggerpony1, triggerpony2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_cmd_ext_map5(self): # ext mapping, lora with no parameters
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:ext $lora lora1>inlinetrigger<ppp:/ext>",
|
||||
"",
|
||||
),
|
||||
PromptPair("<lora:loraillustrious:0.9:0.8>inlinetrigger, triggerillustrious1, triggerillustrious2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/ilxlmodel.safetensors",
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,442 @@
|
||||
import unittest
|
||||
|
||||
from ppp import PromptPostProcessor # pylint: disable=import-error
|
||||
|
||||
from .base_tests import PromptPair, TestPromptPostProcessorBase
|
||||
|
||||
|
||||
class TestHosts(TestPromptPostProcessorBase):
|
||||
|
||||
def setUp(self): # pylint: disable=arguments-differ
|
||||
super().setUp(enable_file_logging=False)
|
||||
|
||||
# Hosts tests
|
||||
|
||||
def test_host_attention_parentheses(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1] (test2) (test3:1.5) [(test4)]",
|
||||
"",
|
||||
),
|
||||
PromptPair("(test1:0.9) (test2) (test3:1.5) (test4:0.99)", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"attention": "parentheses"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_attention_disable(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1] (test2) (test3:1.5)",
|
||||
"",
|
||||
),
|
||||
PromptPair("test1 test2 test3", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"attention": "disable"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_attention_remove(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1] (test2) (test3:1.5)",
|
||||
"",
|
||||
),
|
||||
PromptPair("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"attention": "remove"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_attention_error(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1] (test2) (test3:1.5)",
|
||||
"",
|
||||
),
|
||||
PromptPair("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"attention": "error"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
def test_host_scheduling_before(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1:test2:0.5]",
|
||||
"",
|
||||
),
|
||||
PromptPair("test1", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"scheduling": "before"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_scheduling_after(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1:test2:0.5]",
|
||||
"",
|
||||
),
|
||||
PromptPair("test2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"scheduling": "after"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_scheduling_first(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1::0.5] [:test2:0.5] [test3:test4:0.5]",
|
||||
"",
|
||||
),
|
||||
PromptPair("test1 test3", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"scheduling": "first"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_scheduling_remove(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1:test2:0.5]",
|
||||
"",
|
||||
),
|
||||
PromptPair("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"scheduling": "remove"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_scheduling_error(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1:test2:0.5]",
|
||||
"",
|
||||
),
|
||||
PromptPair("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"scheduling": "error"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
def test_host_alternation_first(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1|test2|test3]",
|
||||
"",
|
||||
),
|
||||
PromptPair("test1", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"alternation": "first"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_alternation_remove(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1|test2|test3]",
|
||||
"",
|
||||
),
|
||||
PromptPair("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"alternation": "remove"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_alternation_error(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[test1|test2|test3]",
|
||||
"",
|
||||
),
|
||||
PromptPair("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"alternation": "error"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
def test_host_and_eol(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"test1 AND test2:2",
|
||||
"",
|
||||
),
|
||||
PromptPair("test1\ntest2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"and": "eol"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_and_comma(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"test1 AND test2:2",
|
||||
"",
|
||||
),
|
||||
PromptPair("test1, test2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"and": "comma"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_and_remove(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"test1 AND test2:2",
|
||||
"",
|
||||
),
|
||||
PromptPair("test1 test2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"and": "remove"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_and_error(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"test1 AND test2:2",
|
||||
"",
|
||||
),
|
||||
PromptPair("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"and": "error"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
def test_host_break_eol(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"test1 BREAK test2",
|
||||
"",
|
||||
),
|
||||
PromptPair("test1\ntest2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"break": "eol"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_break_comma(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"test1 BREAK test2",
|
||||
"",
|
||||
),
|
||||
PromptPair("test1, test2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"break": "comma"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_break_remove(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"test1 BREAK test2",
|
||||
"",
|
||||
),
|
||||
PromptPair("test1 test2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"break": "remove"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_host_break_error(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"test1 BREAK test2",
|
||||
"",
|
||||
),
|
||||
PromptPair("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"break": "error"}}},
|
||||
},
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,78 @@
|
||||
import unittest
|
||||
|
||||
from .base_tests import PromptPair, TestPromptPostProcessorBase
|
||||
|
||||
|
||||
class TestPerformance(TestPromptPostProcessorBase):
|
||||
|
||||
def setUp(self): # pylint: disable=arguments-differ
|
||||
super().setUp(enable_file_logging=False)
|
||||
|
||||
# Performance tests
|
||||
|
||||
def test_parser_performance_simple_simpleparser(
|
||||
self,
|
||||
): # performance test with a large prompt without new constructs
|
||||
large_prompt = ", ".join(
|
||||
["(this:1.2) is a [test] using a [simple|low complexity] prompt with <lora:test:1>"] * 15
|
||||
)
|
||||
self.process(
|
||||
PromptPair(large_prompt, ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_parser_performance_simple_fullparser(
|
||||
self,
|
||||
): # performance test with a large prompt without new constructs but using full parser
|
||||
# we trick it to use the full parser by including some characters
|
||||
large_prompt = "{__${x:}}" + ", ".join(
|
||||
["(this:1.2) is a [test] using a [simple|low complexity] prompt with <lora:test:1>"] * 15
|
||||
)
|
||||
self.process(
|
||||
PromptPair(large_prompt, ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_parser_performance_complex_fullparser(
|
||||
self,
|
||||
): # 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, ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
# the following tests are performance tests with only one kind of the old constructs
|
||||
# same number of constructs and approximately the same full length
|
||||
|
||||
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, ""),
|
||||
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, ""),
|
||||
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, ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_parser_performance_simple_extranetwork(self): # performance test with only extra networks
|
||||
large_prompt = ", ".join(["<lora:one:1> two <lora:three:1> four <lora:five:1> six"] * 20)
|
||||
self.process(
|
||||
PromptPair(large_prompt, ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,127 @@
|
||||
import unittest
|
||||
|
||||
from .base_tests import PromptPair, TestPromptPostProcessorBase
|
||||
|
||||
|
||||
class TestSendToNegative(TestPromptPostProcessorBase):
|
||||
|
||||
def setUp(self): # pylint: disable=arguments-differ
|
||||
super().setUp(enable_file_logging=False)
|
||||
|
||||
# Send To Negative tests
|
||||
|
||||
def test_stn_simple(self): # negtags with different parameters and separations
|
||||
self.process(
|
||||
PromptPair(
|
||||
"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"),
|
||||
)
|
||||
|
||||
def test_stn_complex(self): # complex negtags
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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(
|
||||
"flowers",
|
||||
"red, (pink:1.21), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue",
|
||||
),
|
||||
)
|
||||
|
||||
def test_stn_complex_nocleanup(self): # complex negtags with no cleanup
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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(
|
||||
" (()), flowers , , ",
|
||||
"red, ((pink)), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue",
|
||||
),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_stn_inside_attention(self): # negtag inside attention
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[<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(
|
||||
"this is a ((test) (test:2):1.5) (red:1.5)", "[neg1], ([square]:1.5), normal quality, (neg2:1.65)"
|
||||
),
|
||||
)
|
||||
|
||||
def test_stn_inside_alternation(self): # negtag inside alternation
|
||||
self.process(
|
||||
PromptPair(
|
||||
"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(
|
||||
"this is a (([complex|simple|regular] test)(test:2):1.5)",
|
||||
"([neg1||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality",
|
||||
),
|
||||
)
|
||||
|
||||
def test_stn_inside_alternation_recursive(self): # negtag inside alternation (recursive alternation)
|
||||
self.process(
|
||||
PromptPair(
|
||||
"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(
|
||||
"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",
|
||||
),
|
||||
)
|
||||
|
||||
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]")],
|
||||
)
|
||||
|
||||
def test_stn_complex_features(self): # complex negtags with AND, BREAK and other features
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[<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(
|
||||
"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]",
|
||||
),
|
||||
)
|
||||
|
||||
def test_stn_complex_features_newformat(self): # complex negtags with AND, BREAK and other features (new format)
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[<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(
|
||||
"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]",
|
||||
),
|
||||
)
|
||||
|
||||
def test_stn_inside_alternation_recursive_2(self): # negtag inside alternation (recursive alternation)
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[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(
|
||||
"[pos1[pos11|pos12||pos14|pos15]|pos2|pos3]",
|
||||
"[neg1||], [[|neg12|||]||], [[||||neg15]||], [|neg2|], [||neg3]",
|
||||
# "[neg1[|neg12|||neg15]|neg2|neg3]", # expected output if the constructs were unified
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,160 @@
|
||||
import unittest
|
||||
|
||||
from ppp import PromptPostProcessor # pylint: disable=import-error
|
||||
from .base_tests import PromptPair, TestPromptPostProcessorBase
|
||||
|
||||
|
||||
class TestVariables(TestPromptPostProcessorBase):
|
||||
|
||||
def setUp(self): # pylint: disable=arguments-differ
|
||||
super().setUp(enable_file_logging=False)
|
||||
|
||||
# Variable nesting tests
|
||||
|
||||
def test_var_nested_1(self): # variable default nested in variable set
|
||||
self.process(
|
||||
PromptPair(
|
||||
"${v1=test ${v2:OK}}${v1}",
|
||||
"",
|
||||
),
|
||||
PromptPair("test OK", ""),
|
||||
variables={"v1": "test OK", "v2": "OK"},
|
||||
)
|
||||
|
||||
def test_var_nested_2(self): # variable set nested in variable default
|
||||
self.process(
|
||||
PromptPair(
|
||||
"${v1:test ${v2=OK}${v2}}",
|
||||
"",
|
||||
),
|
||||
PromptPair("test OK", ""),
|
||||
variables={"v1": "test OK", "v2": "OK"},
|
||||
)
|
||||
|
||||
def test_var_nested_3(self): # variable default nested in variable default
|
||||
self.process(
|
||||
PromptPair(
|
||||
"${v1:test ${v2:OK}}",
|
||||
"",
|
||||
),
|
||||
PromptPair("test OK", ""),
|
||||
variables={"v1": "test OK", "v2": "OK"},
|
||||
)
|
||||
|
||||
# Variable-vs-variable comparison tests
|
||||
|
||||
def test_cmd_if_var_vs_var_eq(self): # var eq var: both set to same value, if-branch taken
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:set v1>hello<ppp:/set><ppp:set v2>hello<ppp:/set><ppp:if v1 eq v2>YES<ppp:else>NO<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("YES", ""),
|
||||
)
|
||||
|
||||
def test_cmd_if_var_vs_var_ne(self): # var ne var: different values, ne condition true
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:set v1>apple<ppp:/set><ppp:set v2>orange<ppp:/set><ppp:if v1 ne v2>YES<ppp:else>NO<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("YES", ""),
|
||||
)
|
||||
|
||||
def test_cmd_if_var_vs_var_contains(self): # var contains var: var1 contains var2's value
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:set v1>hello world<ppp:/set><ppp:set v2>hello<ppp:/set><ppp:if v1 contains v2>YES<ppp:else>NO<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("YES", ""),
|
||||
)
|
||||
|
||||
def test_cmd_if_var_vs_var_not_contains(self): # var not contains var: var1 does not contain var2's value
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:set v1>hello world<ppp:/set><ppp:set v2>goodbye<ppp:/set><ppp:if v1 not contains v2>YES<ppp:else>NO<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("YES", ""),
|
||||
)
|
||||
|
||||
# NaN/undefined variable integer comparison tests
|
||||
|
||||
def test_cmd_if_undefined_var_int_compare_warn(self): # undefined var integer compare with on_warning=warn
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:if undefined_var gt 0>YES<ppp:else>NO<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("NO", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_cmd_if_undefined_var_int_compare_stop(self): # undefined var integer compare with on_warning=stop
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:if undefined_var gt 0>YES<ppp:else>NO<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("", ""),
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
def test_cmd_if_nonnumeric_var_int_compare_warn(self): # non-numeric var integer compare with on_warning=warn
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:set myvar>abc<ppp:/set><ppp:if myvar gt 0>YES<ppp:else>NO<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("NO", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_cmd_if_nonnumeric_var_int_compare_stop(self): # non-numeric var integer compare with on_warning=stop
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:set myvar>abc<ppp:/set><ppp:if myvar gt 0>YES<ppp:else>NO<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("", ""),
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
def test_cmd_if_empty_var_int_compare(self): # empty string var integer compare with on_warning=warn
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:set myvar><ppp:/set><ppp:if myvar gt 0>YES<ppp:else>NO<ppp:/if>",
|
||||
"",
|
||||
),
|
||||
PromptPair("NO", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{**self.defopts, "on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,75 @@
|
||||
import unittest
|
||||
|
||||
from ppp import PromptPostProcessor # pylint: disable=import-error
|
||||
from .base_tests import PromptPair, TestPromptPostProcessorBase
|
||||
|
||||
|
||||
class TestModelVariants(TestPromptPostProcessorBase):
|
||||
|
||||
def setUp(self): # pylint: disable=arguments-differ
|
||||
super().setUp(enable_file_logging=False)
|
||||
|
||||
# Model variants tests
|
||||
|
||||
def test_variants(self):
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<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", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors",
|
||||
"ppp_config": {
|
||||
"models": {
|
||||
"sd1": {
|
||||
"detect": {"tests": {"class": ["SD15", "SD15_instructpix2pix"]}},
|
||||
"variants": {
|
||||
"test3": {"find_in_filename": "testmodel"},
|
||||
"sdxl": {"find_in_filename": "testmodel"},
|
||||
},
|
||||
},
|
||||
"sdxl": {
|
||||
"detect": {
|
||||
"tests": {
|
||||
"class": [
|
||||
"SDXL",
|
||||
"SDXLRefiner",
|
||||
"SDXL_instructpix2pix",
|
||||
"Segmind_Vega",
|
||||
"KOALA_700M",
|
||||
"KOALA_1B",
|
||||
]
|
||||
}
|
||||
},
|
||||
"variants": {
|
||||
"test1": {"find_in_filename": "testmodel"},
|
||||
"test2": {"find_in_filename": "testmodel"},
|
||||
},
|
||||
},
|
||||
"something": {
|
||||
"detect": {"tests": {"class": ["something"]}},
|
||||
"variants": {
|
||||
"test4": {"find_in_filename": "testmodel"},
|
||||
},
|
||||
},
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
**self.defopts,
|
||||
"on_warning": PromptPostProcessor.ONWARNING_CHOICES.warn.value,
|
||||
},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,400 @@
|
||||
import unittest
|
||||
|
||||
from ppp import PromptPostProcessor
|
||||
from .base_tests import PromptPair, TestPromptPostProcessorBase
|
||||
|
||||
|
||||
class TestWildcards(TestPromptPostProcessorBase):
|
||||
|
||||
def setUp(self): # pylint: disable=arguments-differ
|
||||
super().setUp(enable_file_logging=False)
|
||||
|
||||
# Wildcards tests
|
||||
|
||||
def test_wc_ignore(self): # wildcards with ignore option
|
||||
self.process(
|
||||
PromptPair("__bad_wildcard__", "{option1|option2}"),
|
||||
PromptPair("__bad_wildcard__", "{option1|option2}"),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{
|
||||
**self.defopts,
|
||||
"process_wildcards": False,
|
||||
"if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value,
|
||||
},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_wc_remove(self): # wildcards with remove option
|
||||
self.process(
|
||||
PromptPair(
|
||||
"[<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(
|
||||
"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]",
|
||||
),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{
|
||||
**self.defopts,
|
||||
"process_wildcards": False,
|
||||
"if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.remove.value,
|
||||
},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_wc_warn(self): # wildcards with warn option
|
||||
self.process(
|
||||
PromptPair("__bad_wildcard__", "{option1|option2}"),
|
||||
PromptPair(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", "{option1|option2}"),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{
|
||||
**self.defopts,
|
||||
"process_wildcards": False,
|
||||
"if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.warn.value,
|
||||
},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_wc_stop(self): # wildcards with stop option
|
||||
self.process(
|
||||
PromptPair("__bad_wildcard__", "{option1|option2}"),
|
||||
PromptPair(
|
||||
PromptPostProcessor.WILDCARD_STOP.format("__bad_wildcard__") + "__bad_wildcard__",
|
||||
"{option1|option2}",
|
||||
),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{
|
||||
**self.defopts,
|
||||
"process_wildcards": False,
|
||||
"if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.stop.value,
|
||||
},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
def test_wcinvar_warn(self): # wildcards in var with warn option
|
||||
self.process(
|
||||
PromptPair("${v=__bad_wildcard__}${v}", ""),
|
||||
PromptPair(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.interrupt,
|
||||
self.def_env_info,
|
||||
{
|
||||
**self.defopts,
|
||||
"process_wildcards": False,
|
||||
"if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.warn.value,
|
||||
},
|
||||
self.grammar_content,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
),
|
||||
)
|
||||
|
||||
def test_wc_invalid_name(self):
|
||||
self.process(
|
||||
PromptPair("the choices are: ___invalid__", ""),
|
||||
PromptPair("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", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_wc_wildcard1a_json(self): # simple json wildcard
|
||||
self.process(
|
||||
PromptPair("the choices are: __json/wildcard1__", ""),
|
||||
PromptPair("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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_wc_test2_yaml(self): # simple yaml wildcard
|
||||
self.process(
|
||||
PromptPair("the choice is: __testwc/test2__", ""),
|
||||
PromptPair("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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
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-choice3", ""),
|
||||
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", ""),
|
||||
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: choice3bisbis", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_wc_wildcard_default_filter(self): # wildcard with default filter
|
||||
self.process(
|
||||
PromptPair(
|
||||
"<ppp:setwcdeffilter 'yaml/wildcard2' 'label1+label3' />the choice is: __yaml/wildcard2__, <ppp:setwcdeffilter 'yaml/wildcard2' />__yaml/wildcard2__",
|
||||
"",
|
||||
),
|
||||
PromptPair("the choice is: choice3-choice3, 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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
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: ", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_wc_choice_wildcard_mix(self): # choices with wildcard mix
|
||||
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", ""),
|
||||
],
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_wc_unsupportedsampler(self): # unsupported sampler
|
||||
self.process(
|
||||
PromptPair("the choices are: __@yaml/wildcard2__", ""),
|
||||
PromptPair("", ""),
|
||||
ppp="nocup",
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
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", ""),
|
||||
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", ""),
|
||||
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", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_wc_anonymouswildcard_yaml(self): # yaml anonymous wildcard
|
||||
self.process(
|
||||
PromptPair("the choices are: __yaml/anonwildcards__", ""),
|
||||
PromptPair("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", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_wc_circular(self): # wildcard circular reference
|
||||
self.process(
|
||||
PromptPair("the choices are: __yaml/circular1__", ""),
|
||||
PromptPair("", ""),
|
||||
ppp="nocup",
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
def test_wc_including(self): # wildcard including another wildcard
|
||||
self.process(
|
||||
PromptPair("the choices are: __yaml/including__", ""),
|
||||
PromptPair("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("", ""),
|
||||
ppp="nocup",
|
||||
interrupted=True,
|
||||
)
|
||||
|
||||
def test_wc_dynamicwildcard(self): # wildcard built from variables
|
||||
self.process(
|
||||
PromptPair(
|
||||
"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", ""),
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user