Files
acorderob-sd-webui-prompt-p…/ppp_tree.py
T

2065 lines
97 KiB
Python

from collections import namedtuple
import logging
import math
import re
import textwrap
import time
from typing import Optional
import lark
import numpy as np
from ppp_classes import IFWILDCARDS_CHOICES, PPPState
from ppp_enmappings import PPPENMappingVariant
from ppp_logging import DEBUG_LEVEL, log
from ppp_utils import escape_single_quotes
from ppp_common import parse_prompt, warn_or_stop
from ppp_wildcards import PPPWildcard
class TreeProcessor(lark.visitors.Interpreter):
"""
A class for interpreting and processing a tree generated by the prompt parser.
Args:
state (PPPState): The state object containing the current processing state.
rng (numpy.random.Generator): The random number generator.
Attributes:
add_at (dict): The dictionary to store the content to be added at different positions of the negative prompt.
insertion_at (list): The list of insertion points in the negative prompt.
detectedWildcards (list): The list of detected invalid wildcards or choices.
result (str): The final processed prompt.
"""
def __init__(self, state: PPPState, rng: np.random.Generator):
super().__init__()
self.state = state
self.__debug_level = state.options.debug_level
self.__rng = rng
self.AccumulatedShell = namedtuple("AccumulatedShell", ["type", "data"])
self.NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"])
self.__shell: list[self.AccumulatedShell] = [] # type: ignore
self.__negtags: list[self.NegTag] = [] # type: ignore
self.__already_processed: list[str] = []
self.__is_negative = False
self.__wildcard_filters = {}
self.__seen_wildcards: list[str] = []
self.add_at: dict = {"start": [], "insertion_point": [[] for x in range(10)], "end": []}
self.insertion_at: list[tuple[int, int]] = [None for x in range(10)]
self.detectedWildcards: list[str] = []
self.result = ""
def log(self, kind, message: str, min_level: DEBUG_LEVEL | None = None):
log(self.state.logger, self.state.options.debug_level, kind, message, min_level)
def warn_or_stop(self, message: str, e: Exception = None):
warn_or_stop(self.state, self.__is_negative, message, e)
def start_visit(self, prompt_description: str, parsed_prompt: lark.Tree, is_negative: bool = False) -> str:
"""
Start the visit process.
Args:
prompt_description (str): The description of the prompt.
parsed_prompt (Tree): The parsed prompt.
is_negative (bool): Whether the prompt is negative or not.
Returns:
str: The processed prompt.
"""
t1 = time.monotonic_ns()
self.__is_negative = is_negative
self.log(logging.INFO, f"Processing {prompt_description}...")
self.visit(parsed_prompt)
t2 = time.monotonic_ns()
self.log(logging.INFO, f"Process {prompt_description} time: {(t2 - t1) / 1_000_000_000:.3f} seconds")
return self.result
def __visit(
self,
node: lark.Tree | lark.Token | list[lark.Tree | lark.Token] | None,
restore_state: bool = False,
discard_content: bool = False,
) -> str:
"""
Visit a node in the tree and process it or accumulate its value if it is a Token.
Args:
node (Tree|Token|list): The node or list of nodes to visit.
restore_state (bool): Whether to restore the state after visiting the node.
discard_content (bool): Whether to discard the content of the node.
Returns:
str: The result of the visit.
"""
backup_result = self.result
# self.log(logging.DEBUG, f"Visiting node {node}.")
if restore_state:
# self.log(logging.DEBUG, "Backing up state before visiting.")
backup_shell = self.__shell.copy()
backup_negtags = self.__negtags.copy()
backup_already_processed = self.__already_processed.copy()
backup_add_at = self.add_at.copy()
backup_insertion_at = self.insertion_at.copy()
backup_detectedwildcards = self.detectedWildcards.copy()
backup_user_variables = self.state.user_variables.copy()
backup_echoed_variables = self.state.echoed_variables.copy()
if node is not None:
if isinstance(node, list):
for child in node:
self.__visit(child)
elif isinstance(node, lark.Tree):
self.visit(node)
elif isinstance(node, lark.Token):
self.result += node
len_backup = len(backup_result)
# if self.result[:len_backup] == backup_result: # this is only necessary if we call parse_prompt with a parser from "start", because it resets the result
added_result = self.result[len_backup:]
# else:
# added_result = self.result
if discard_content or restore_state:
self.result = backup_result
if restore_state:
# self.log(logging.DEBUG, "Restoring state after visiting.")
self.__shell = backup_shell
self.__negtags = backup_negtags
self.__already_processed = backup_already_processed
self.add_at = backup_add_at
self.insertion_at = backup_insertion_at
self.detectedWildcards = backup_detectedwildcards
self.state.user_variables.clear()
self.state.user_variables.update(backup_user_variables)
self.state.echoed_variables.clear()
self.state.echoed_variables.update(backup_echoed_variables)
return added_result
def __get_original_node_content(self, node: lark.Tree | lark.Token, default=None) -> str:
"""
Get the original content of a node.
Args:
node (Tree|Token): The node to get the content from.
default: The default value to return if the content is not found.
Returns:
str: The original content of the node.
"""
return node.meta.content if hasattr(node, "meta") and node.meta is not None and not node.meta.empty else default
def __parse_array_specifier(self, specifier: str | None) -> tuple[Optional[int], Optional[str], Optional[bool]]:
"""
Convert an index/separator/count string to a specific type.
Args:
specifier (str|None): The specifier string.
Returns:
tuple: A tuple containing the index, separator and count boolean.
"""
if specifier is None:
return None, None, None
if specifier == "#": # special value to indicate length of the array variable
return None, None, True
if specifier.startswith("&"): # special value to indicate a separator
return None, specifier[2:-1], False
if not specifier.isdecimal():
# bare identifier: resolve as variable
specifier = self.get_final_user_variable(specifier)
if specifier.isdecimal():
return int(specifier), None, False
# invalid specifier
return None, None, False
def __get_user_variable_value(
self, name: str, specifier: str | None = None, evaluate=True, visit=False
) -> str | list[str] | None:
"""
Get the value of a user variable.
Args:
name (str): The name of the user variable.
specifier (str|None): The specifier for an array variable.
evaluate (bool): Whether to evaluate the variable.
visit (bool): Whether to also visit the variable (add to result).
Returns:
str|list[str]|None: The value of the user variable.
"""
def visit_value(v):
visited = False
if isinstance(v, lark.Tree):
if evaluate:
v = self.__visit(v, restore_state=not visit, discard_content=not visit)
visited = visit
else:
v = self.__get_original_node_content(v, "")
elif isinstance(v, lark.Token):
v = str(v)
if visit and not visited:
self.result += v
return v
v = self.state.user_variables.get(name, None)
if v is None:
return None
is_array = name[-2:] == "[]"
if is_array:
if isinstance(v, list):
idx, sep, cnt = self.__parse_array_specifier(specifier)
if cnt:
v = len(v)
if visit:
self.result += str(v)
elif idx is not None:
if 0 <= idx < len(v):
v = visit_value(v[idx])
else:
v = None # invalid index
else:
if sep is None:
sep = self.state.options.choice_separator
v2 = []
for i, item in enumerate(v):
v2.append(visit_value(item))
if visit and i < len(v) - 1:
self.result += sep
v = v2
else:
v = None # error
elif specifier is not None:
v = None # error
else:
v = visit_value(v)
return v
def __separate_arrayref(self, name_specifier: str):
"""
Separate the name and specifier part of a variable reference.
Args:
name_specifier (str): The variable reference string.
Returns:
tuple: A tuple containing the name and specifier part.
"""
name, isarray, specifier = re.match(r"^([^\[]+)(\[([^[]*)\])?$", name_specifier).groups()
if isarray:
name += "[]"
if specifier == "":
specifier = None
return name, specifier
def get_final_user_variable(self, name_specifier: str) -> str:
"""
Get the final value of a user variable, resolving any references if needed.
Args:
name_specifier (str): The variable reference string.
Returns:
str: The final value of the user variable.
"""
name, specifier = self.__separate_arrayref(name_specifier)
v = self.__get_user_variable_value(name, specifier, True, False)
if isinstance(v, list):
_, sep, _ = self.__parse_array_specifier(specifier)
if sep is None:
sep = self.state.options.choice_separator
v = sep.join(str(item) for item in v)
return str(v)
def __set_user_variable_value(self, name: str, value: str | lark.Tree | list):
"""
Set the value of a user variable.
Args:
name (str): The name of the user variable.
value (str|lark.Tree|list): The value to be set.
"""
self.state.user_variables[name] = value
def __remove_user_variable(self, name: str):
"""
Remove a user variable.
Args:
name (str): The name of the user variable.
"""
if name in self.state.user_variables:
del self.state.user_variables[name]
def __debug_end(self, construct: str, start_result: str, duration: int, info=None):
"""
Log the end of a construct processing.
Args:
construct (str): The name of the construct.
start_result (str): The initial result.
duration (int): The duration of the processing in ns.
info: Additional information to log.
"""
if self.__debug_level == DEBUG_LEVEL.full:
info = f"({info}) " if info is not None and info != "" else ""
output = self.result[len(start_result) :]
if output != "":
output = f" >> '{escape_single_quotes(output)}'"
self.log(logging.DEBUG, f"TreeProcessor.{construct} {info}({duration / 1_000_000_000:.3f} seconds){output}")
def __adjust_strnum(self, s: str) -> str | int | float:
"""
Adjust a string that may represent a number to its appropriate type.
If it is a number, it is converted to an integer or float.
Args:
s (str): The string to adjust.
Returns:
str | int | float: The adjusted string, integer, or float.
"""
try:
return int(s)
except ValueError:
pass
if bool(re.match(r"^[+-]?\d+\.\d+$", s)):
return float(s)
return s.lower()
def __warn_mixedtype(self, desc: str, operand1, operand2, operation):
"""
Warn the user if mixed type values are used in a comparison.
Args:
desc (str): Description of the comparison.
operand1: The first operand.
operand2: The second operand.
operation: The operation to perform if the types are compatible.
"""
if operand1 is None or operand2 is None:
self.warn_or_stop(f"Undefined value used in comparison: '{escape_single_quotes(desc)}'")
return False
if (
isinstance(operand1, (str, int, float, bool))
and isinstance(operand2, (str, int, float, bool))
and operand1.__class__ != operand2.__class__
):
self.warn_or_stop(f"Mixed type values used in comparison: '{escape_single_quotes(desc)}'")
return False
return operation(operand1, operand2)
def __resolve_operand(self, c: str) -> str | bool | int | float:
"""
Resolve an operand value.
Args:
c (str): The operand value to resolve.
Returns:
str | bool | int | float: The resolved operand value (in lowercase for strings).
"""
if c.startswith('"') and c.endswith('"') or c.startswith("'") and c.endswith("'"):
return c[1:-1].lower() if self.state.options.strict_operators else self.__adjust_strnum(c[1:-1])
try:
return int(c)
except ValueError:
pass
if bool(re.match(r"^[+-]?\d+\.\d+$", c)):
return float(c)
if c.lower() in ("false", ""):
return False
if c.lower() == "true":
return True
# Bare identifier - resolve as variable reference
if c.startswith("_"):
vartype = "system"
val = self.state.system_variables.get(c, None)
else:
vartype = "user"
varname, varspecifier = self.__separate_arrayref(c)
val = self.__get_user_variable_value(varname, varspecifier)
if val is None:
val = ""
self.warn_or_stop(f"Unknown {vartype} variable '{escape_single_quotes(c)}'")
if isinstance(val, str):
try:
return int(val)
except ValueError:
pass
if bool(re.match(r"^[+-]?\d+\.\d+$", val)):
return float(val)
if val in ("false", ""):
return False
if val == "true":
return True
val = val.lower()
return val
def __wmt(self, cond_desc, a, b, op):
return self.__warn_mixedtype(cond_desc, a, b, op)
def __pairwise_all(self, cond_desc, op1v, op2v, op):
return all(self.__wmt(cond_desc, a, b, op) for a, b in zip(op1v, op2v))
def __alltoone_all(self, cond_desc, op1v, op2v, op):
if isinstance(op1v, list):
return all(self.__wmt(cond_desc, a, op2v, op) for a in op1v)
else:
return all(self.__wmt(cond_desc, op1v, b, op) for b in op2v)
def __alltoone_any(self, cond_desc, op1v, op2v, op):
if isinstance(op1v, list):
return any(self.__wmt(cond_desc, a, op2v, op) for a in op1v)
else:
return any(self.__wmt(cond_desc, op1v, b, op) for b in op2v)
def __eval_basiccondition(
self,
cond_desc: str,
operand1: str | list[str],
operator: str,
operand2: str | list[str],
) -> bool:
"""
Evaluate a condition based on the given operands and operator.
Args:
cond_desc (str): The description of the condition (for logging).
operand1 (str | list[str]): The first operand.
operator (str): The operator.
operand2 (str | list[str]): The second operand.
Returns:
bool: The result of the condition evaluation.
"""
if isinstance(operand1, list):
operand1_value = list(self.__resolve_operand(c) for c in operand1)
else:
operand1_value = self.__resolve_operand(operand1)
operand1_isarray = isinstance(operand1_value, list)
if isinstance(operand2, list):
operand2_value = list(self.__resolve_operand(c) for c in operand2)
else:
operand2_value = self.__resolve_operand(operand2)
operand2_isarray = isinstance(operand2_value, list)
if operator == "truthy":
result = bool(operand1_value)
else:
if not operand1_isarray and not operand2_isarray:
operations = {
"eq": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x == y),
"ne": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x != y),
"gt": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x > y),
"lt": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x < y),
"ge": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x >= y),
"le": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x <= y),
"in": lambda: self.__wmt(cond_desc, str(operand1_value), str(operand2_value), lambda x, y: x in y),
"any_in": None, # does not make sense
"contains": lambda: self.__wmt(
cond_desc, str(operand1_value), str(operand2_value), lambda x, y: y in x
),
"contains_any": None, # does not make sense
}
elif operand1_isarray and operand2_isarray:
operations = {
"eq": lambda: len(operand1_value) == len(operand2_value)
and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x == y),
"ne": lambda: len(operand1_value) != len(operand2_value)
or self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x != y),
"gt": lambda: len(operand1_value) == len(operand2_value)
and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x > y),
"lt": lambda: len(operand1_value) == len(operand2_value)
and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x < y),
"ge": lambda: len(operand1_value) == len(operand2_value)
and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x >= y),
"le": lambda: len(operand1_value) == len(operand2_value)
and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x <= y),
"in": lambda: self.__alltoone_all(cond_desc, operand1_value, operand2_value, lambda x, y: x in y),
# all(
# self.__wmt(cond_desc, a, operand2_value, lambda x, y: x in y) for a in operand1_value
# ),
"any_in": lambda: self.__alltoone_any(
cond_desc, operand1_value, operand2_value, lambda x, y: x in y
),
# any(
# self.__wmt(cond_desc, a, operand2_value, lambda x, y: x in y) for a in operand1_value
# ),
"contains": lambda: self.__alltoone_all(
cond_desc, operand2_value, operand1_value, lambda x, y: x in y
),
# all(
# self.__wmt(cond_desc, a, operand1_value, lambda x, y: x in y) for a in operand2_value
# ),
"contains_any": lambda: self.__alltoone_any(
cond_desc, operand2_value, operand1_value, lambda x, y: x in y
),
# any(
# self.__wmt(cond_desc, a, operand1_value, lambda x, y: x in y) for a in operand2_value
# ),
}
elif operand1_isarray and not operand2_isarray:
if self.state.options.strict_operators:
operations = {
"eq": None, # does not make sense
"ne": None, # does not make sense
"gt": None, # does not make sense
"lt": None, # does not make sense
"ge": None, # does not make sense
"le": None, # does not make sense
}
else:
operations = {
"eq": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x == y
),
"ne": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x != y
),
"gt": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x > y
),
"lt": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x < y
),
"ge": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x >= y
),
"le": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x <= y
),
}
operations.update(
{
"in": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: str(x) in str(y)
),
"any_in": lambda: self.__alltoone_any(
cond_desc, operand1_value, operand2_value, lambda x, y: str(x) in str(y)
),
"contains": lambda: self.__wmt(cond_desc, operand2_value, operand1_value, lambda x, y: x in y),
"contains_any": None, # does not make sense
}
)
elif not operand1_isarray and operand2_isarray:
if self.state.options.strict_operators:
operations = {
"eq": None, # does not make sense
"ne": None, # does not make sense
"gt": None, # does not make sense
"lt": None, # does not make sense
"ge": None, # does not make sense
"le": None, # does not make sense
}
else:
operations = {
"eq": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x == y
),
"ne": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x != y
),
"gt": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x > y
),
"lt": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x < y
),
"ge": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x >= y
),
"le": lambda: self.__alltoone_all(
cond_desc, operand1_value, operand2_value, lambda x, y: x <= y
),
}
operations.update(
{
"in": lambda: self.__wmt(cond_desc, operand1_value, operand2_value, lambda x, y: x in y),
"any_in": None, # does not make sense
"contains": lambda: self.__alltoone_all(
cond_desc, operand2_value, operand1_value, lambda x, y: str(x) in str(y)
),
"contains_any": lambda: self.__alltoone_any(
cond_desc, operand2_value, operand1_value, lambda x, y: str(x) in str(y)
),
}
)
else:
operations = {}
operation = operations.get(operator, None)
if operation is None:
self.warn_or_stop(
f"Unsupported operator '{escape_single_quotes(operator)}' in condition '{escape_single_quotes(cond_desc)}'"
)
return False
result = operation()
return result
def __separate_vardescriptor(self, vardescriptor: lark.Tree) -> tuple[str, str | None]:
"""
Separate the name and index/sep part of a variable descriptor.
Args:
vardescriptor (lark.Tree): The variable descriptor tree.
Returns:
tuple[str, str | None]: A tuple containing the name and index/sep part of the variable descriptor.
"""
vardescriptor_name = str(vardescriptor.children[0])
vardescriptor_specifier = None
if vardescriptor.children[1] is not None:
vardescriptor_name += "[]"
if vardescriptor.children[2] is not None:
if isinstance(vardescriptor.children[2], lark.Token):
vardescriptor_specifier = str(vardescriptor.children[2])
else:
vardescriptor_specifier = "".join(vardescriptor.children[2].children)
return vardescriptor_name, vardescriptor_specifier
def __get_complex_element(self, value_node: lark.Tree | lark.Token) -> str:
"""
Return in string form the value of a complex value node, which can be either a simple value or a variable descriptor. Does not evaluate variables.
Args:
value_node (lark.Tree | lark.Token): The complex value node to be evaluated.
Returns:
str: The value.
"""
if isinstance(value_node, lark.Tree):
# it's a vardescriptor_get
vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor(value_node)
return (
vardescriptor_name[0:-2] + f"[{vardescriptor_specifier}]"
if vardescriptor_specifier is not None
else vardescriptor_name
)
# it's a SIMPLEVALUE
return str(value_node)
def __get_cond_operand(self, value_node: lark.Tree) -> str | list[str]:
"""
Returns a simple value or a list of simple values or variables. Does not evaluate them.
Args:
value_node (lark.Tree): The value tree to be evaluated.
Returns:
str | list[str]: The result.
"""
if isinstance(value_node, lark.Tree) and value_node.data == "listvalue":
return list(self.__get_complex_element(v) for v in value_node.children)
return self.__get_complex_element(value_node)
def __eval_condition(self, condition: lark.Tree) -> bool:
"""
Evaluate an if condition based on the given condition tree.
Args:
condition (lark.Tree): The condition tree to be evaluated.
Returns:
bool: The result of the if condition evaluation.
"""
# self.log(logging.DEBUG, f"__eval_condition {condition.data}")
if condition.data == "operation_and":
cond_result = True
for c in condition.children:
cond_result = cond_result and self.__eval_condition(c)
if not cond_result:
break
elif condition.data == "operation_or":
cond_result = False
for c in condition.children:
cond_result = cond_result or self.__eval_condition(c)
if cond_result:
break
elif condition.data == "operation_not":
cond_result = not self.__eval_condition(condition.children[0])
else: # truthy_operand / comparison
# we get the name of the variable
cond_operand1 = self.__get_cond_operand(condition.children[0])
poscomp = 1
invert = False
if poscomp >= len(condition.children):
# no condition, just a variable
cond_operation = "truthy"
cond_operand2 = "true"
cond_desc = (
cond_operand1
if isinstance(cond_operand1, str)
else "(" + ", ".join(str(c) for c in cond_operand1) + ")"
)
else:
# we get the comparison (with possible not) and the value
cond_operation = str(condition.children[poscomp])
if cond_operation == "not":
invert = not invert
poscomp += 1
cond_operation = str(condition.children[poscomp])
poscomp += 1
cond_value_node = condition.children[poscomp]
cond_operand2 = self.__get_cond_operand(cond_value_node)
cond_desc = f"{cond_operand1} {cond_operation} {cond_operand2 if isinstance(cond_operand2, str) else '(' + ', '.join(str(c) for c in cond_operand2) + ')'}"
cond_result = self.__eval_basiccondition(cond_desc, cond_operand1, cond_operation, cond_operand2)
if invert:
cond_result = not cond_result
return cond_result
def promptcomp(self, tree: lark.Tree):
"""
Process a prompt composition construct in the tree.
"""
start_result = self.result
t1 = time.monotonic_ns()
self.__visit(tree.children[0])
and_processing = self.state.host_config.and_
if len(tree.children) > 1:
and_replacements = {
"eol": ("replaced with EOL", "\n"),
"comma": ("replaced with COMMA", ", "),
"remove": ("removed", " "),
}
if tree.children[1] is not None:
self.result += f":{tree.children[1]}"
for i in range(2, len(tree.children), 3):
if and_processing in and_replacements.keys():
self.result = (
self.result.rstrip()
+ and_replacements[and_processing][1]
+ self.__visit(tree.children[i + 1], False, True).lstrip()
)
self.log(logging.DEBUG, f"AND construct {and_replacements[and_processing][0]}")
elif and_processing == "error":
self.warn_or_stop("AND constructs are not allowed!")
else: # and_processing == "ok":
if self.state.options.cup_ands:
self.result = re.sub(r"[, ]+$", "\n" if self.state.options.cup_ands_eol else " ", self.result)
if self.result[-1:].isalnum(): # add space if needed
self.result += " "
self.result += "AND"
added_result = self.__visit(tree.children[i + 1], False, True)
if self.state.options.cup_ands:
added_result = re.sub(r"^[, ]+", " ", added_result)
if added_result[0:1].isalnum(): # add space if needed
added_result = " " + added_result
self.result += added_result
if tree.children[i + 2] is not None:
self.result += f":{tree.children[i+2]}"
t2 = time.monotonic_ns()
self.__debug_end("promptcomp", start_result, t2 - t1)
def scheduled(self, tree: lark.Tree):
"""
Process a scheduling construct in the tree and add it to the accumulated shell.
"""
start_result = self.result
t1 = time.monotonic_ns()
before = tree.children[0]
after = tree.children[-2]
pos_str = tree.children[-1]
pos = float(pos_str)
if pos >= 1:
pos = int(pos)
scheduling_processing = self.state.host_config.scheduling
if scheduling_processing == "before":
self.log(logging.DEBUG, "Scheduling construct removed, taking before option")
if before is not None:
self.__visit(before)
elif scheduling_processing == "after":
self.log(logging.DEBUG, "Scheduling construct removed, taking after option")
if after is not None:
self.__visit(after)
elif scheduling_processing == "first":
self.log(logging.DEBUG, "Scheduling construct removed, taking first option")
if before is not None:
self.__visit(before)
elif after is not None:
self.__visit(after)
elif scheduling_processing == "remove":
self.log(logging.DEBUG, "Scheduling construct removed")
elif scheduling_processing == "error":
self.warn_or_stop("Scheduling constructs are not allowed!")
else: # scheduling_processing == "ok"
# self.__shell.append(self.AccumulatedShell("sc", pos))
self.result += "["
if before is not None:
self.log(logging.DEBUG, f"Shell scheduled before with position {pos}")
self.__shell.append(self.AccumulatedShell("scb", pos))
self.__visit(before)
self.__shell.pop()
self.log(logging.DEBUG, f"Shell scheduled after with position {pos}")
self.__shell.append(self.AccumulatedShell("sca", pos))
self.result += ":"
self.__visit(after)
self.__shell.pop()
if self.state.options.cup_empty_constructs and re.fullmatch(
re.escape(start_result) + r"\[:\s*", self.result
):
self.result = start_result
else:
self.result += f":{pos_str}]"
# self.__shell.pop()
t2 = time.monotonic_ns()
self.__debug_end("scheduled", start_result, t2 - t1, pos_str)
def alternate(self, tree: lark.Tree):
"""
Process an alternation construct in the tree and add it to the accumulated shell.
"""
start_result = self.result
t1 = time.monotonic_ns()
alternation_processing = self.state.host_config.alternation
if alternation_processing == "first":
self.log(logging.DEBUG, "Alternation construct removed, taking first option")
self.__visit(tree.children[0])
elif alternation_processing == "remove":
self.log(logging.DEBUG, "Alternation construct removed")
elif alternation_processing == "error":
self.warn_or_stop("Alternation constructs are not allowed!")
else: # alternation_processing == "ok"
# self.__shell.append(self.AccumulatedShell("al", len(tree.children)))
self.result += "["
for i, opt in enumerate(tree.children):
self.log(logging.DEBUG, f"Shell alternate option {i+1}")
self.__shell.append(self.AccumulatedShell("alo", {"pos": i + 1, "len": len(tree.children)}))
if i > 0:
self.result += "|"
self.__visit(opt)
self.__shell.pop()
self.result += "]"
if self.state.options.cup_empty_constructs and re.fullmatch(
re.escape(start_result) + r"\[\s*\]", self.result
):
self.result = start_result
# self.__shell.pop()
t2 = time.monotonic_ns()
self.__debug_end("alternate", start_result, t2 - t1)
def attention(self, tree: lark.Tree):
"""
Process a attention change construct in the tree and add it to the accumulated shell.
"""
start_result = self.result
t1 = time.monotonic_ns()
# weight_kind: -1: remove, 0=none, 1=decrease, 2=increase, 3=specific
if len(tree.children) == 2:
weight_str = tree.children[-1]
if weight_str is not None:
weight_kind = 3 # specific weight
weight = float(weight_str)
else:
weight_kind = 2 # increase attention
weight = 1.1
weight_str = "1.1"
else:
weight_kind = 1 # decrease attention
weight = 0.9
weight_str = "0.9"
self.log(logging.DEBUG, f"Shell attention with weight {weight}")
current_tree = tree.children[0]
if self.state.options.cup_merge_attention:
while isinstance(current_tree, lark.Tree) and current_tree.data == "attention":
# we merge the weights
if len(current_tree.children) == 2:
inner_weight = current_tree.children[-1]
if inner_weight is not None:
inner_weight = float(inner_weight)
else:
inner_weight = 1.1
else:
inner_weight = 0.9
weight *= inner_weight
current_tree = current_tree.children[0]
weight = math.floor(weight * 100) / 100 # we round to 2 decimals
weight_str = f"{weight:.2f}".rstrip("0").rstrip(".")
if weight_str == "0.9":
weight_kind = 1
elif weight_str == "1.1":
weight_kind = 2
else:
weight_kind = 3
attention_processing = self.state.host_config.attention
if attention_processing == "parentheses":
if weight_kind == 1:
weight_kind = 3
weight_str = "0.9"
self.log(logging.DEBUG, "Converted to parentheses format")
elif attention_processing == "disable":
weight_kind = 0
self.log(logging.DEBUG, "Attention construct disabled")
elif attention_processing == "remove":
weight_kind = -1
self.log(logging.DEBUG, "Attention construct removed")
elif attention_processing == "error":
self.warn_or_stop("Attention constructs are not allowed!")
# else: attention_processing == "ok":
if weight_kind == -1:
# we just ignore the attention construct
pass
elif weight_kind == 0:
# we just visit the content without adding any attention
self.__visit(current_tree)
else:
self.__shell.append(self.AccumulatedShell("at", (weight_kind, weight_str)))
if weight_kind == 1:
starttag = "["
self.result += starttag
self.__visit(current_tree)
endtag = "]"
elif weight_kind == 2:
starttag = "("
self.result += starttag
self.__visit(current_tree)
endtag = ")"
else: # weight_kind == 3
starttag = "("
self.result += starttag
self.__visit(current_tree)
endtag = f":{weight_str})"
if self.state.options.cup_empty_constructs and re.fullmatch(
re.escape(start_result + starttag) + r"\s*", self.result
):
self.result = start_result
else:
self.result += endtag
self.__shell.pop()
t2 = time.monotonic_ns()
self.__debug_end("attention", start_result, t2 - t1, weight_str)
def commandstn(self, tree: lark.Tree):
"""
Process a send to negative command in the tree and add it to the list of negative tags.
"""
start_result = self.result
info = None
t1 = time.monotonic_ns()
if not self.__is_negative:
negtagparameters = tree.children[0]
if negtagparameters is not None:
parameters = str(negtagparameters)
else:
parameters = ""
content = self.__visit(tree.children[1::], False, True)
self.__negtags.append(
self.NegTag(len(self.result), len(self.result), content, parameters, self.__shell.copy())
)
info = f"with {escape_single_quotes(parameters) or 'no parameters'} : {escape_single_quotes(content)}"
else:
self.warn_or_stop("Ignored negative command in negative prompt")
self.__visit(tree.children[1::])
t2 = time.monotonic_ns()
self.__debug_end("commandstn", start_result, t2 - t1, info)
def commandstni(self, tree: lark.Tree):
"""
Process a send to negative insertion point command in the tree and add it to the list of negative tags.
"""
start_result = self.result
info = None
t1 = time.monotonic_ns()
if self.__is_negative:
negtagparameters = tree.children[0]
if negtagparameters is not None:
parameters = str(negtagparameters)
else:
parameters = ""
self.__negtags.append(self.NegTag(len(self.result), len(self.result), "", parameters, self.__shell.copy()))
info = f"with {parameters or 'no parameters'}"
else:
self.warn_or_stop("Ignored negative insertion point command in positive prompt")
t2 = time.monotonic_ns()
self.__debug_end("commandstni", start_result, t2 - t1, info)
def __varset(
self,
command: str,
variable_name: str,
variable_specifier: str | None,
modifiers: lark.Tree | None,
content: lark.Tree | None,
):
"""
Process a generic set command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
if variable_name.startswith("_"):
self.warn_or_stop(
f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! System variables cannot be set."
)
return
info = variable_name
is_array = variable_name[-2:] == "[]"
if variable_specifier is not None:
if not is_array:
self.warn_or_stop(
f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! Only array variables can be indexed."
)
return
info = f"{variable_name[0:-2]}[{variable_specifier}]"
value_description = self.__get_original_node_content(content, None)
value = content
raw_oldvalue = self.state.user_variables.get(variable_name, None)
newvalue = None
some_error = False
if variable_specifier is not None:
if (
is_array
and raw_oldvalue is not None
and isinstance(raw_oldvalue, list)
and not 0 <= int(variable_specifier) < len(raw_oldvalue)
):
self.warn_or_stop(
f"Invalid index {variable_specifier} for variable '{escape_single_quotes(variable_name)}'! Index out of bounds."
)
some_error = True
if not is_array:
self.warn_or_stop(
f"Invalid index for '{escape_single_quotes(variable_name)}'! Only array variables can be indexed."
)
some_error = True
if raw_oldvalue is None or not isinstance(raw_oldvalue, list):
self.warn_or_stop(
f"Invalid index for '{escape_single_quotes(variable_name)}'! Cannot set index of an undefined or non-array variable."
)
some_error = True
adding = False
if not some_error:
modifiers_str: list[str] = [str(m) for m in modifiers.children] if modifiers is not None else []
if any(item in modifiers_str for item in ["+", "add"]):
adding = True
info += f" += '{escape_single_quotes(value_description or '')}'"
if raw_oldvalue is None:
newvalue = value
self.warn_or_stop(f"Unknown variable {escape_single_quotes(variable_name)}")
elif is_array:
if not isinstance(raw_oldvalue, list):
self.warn_or_stop(
f"Invalid variable value for '{escape_single_quotes(variable_name)}'! Cannot add to a non-array value."
)
else:
newvalue = value
elif isinstance(raw_oldvalue, str):
newvalue = lark.Tree(
lark.Token("RULE", "varvalue"),
[lark.Token("plain", raw_oldvalue), value],
# Meta should be {"content": raw_oldvalue + value},
)
else: # lark.Tree
newvalue = lark.Tree(
lark.Token("RULE", "varvalue"),
[raw_oldvalue, value],
# Meta should be {"content": raw_oldvalue.meta.content + value.meta.content},
)
elif any(item in modifiers_str for item in ["?", "ifundefined"]):
info += f" ?= '{escape_single_quotes(value_description or '')}'"
if raw_oldvalue is None:
newvalue = value
else:
info += " (not set)"
else:
newvalue = value
if newvalue is not None:
is_starred = isinstance(newvalue, lark.Tree) and newvalue.data == "starredvalue"
access_full_array = is_array and variable_specifier is None
if any(item in modifiers_str for item in ["!", "evaluate"]):
newvalue = self.__visit(newvalue, False, True)
info += " =! "
else:
info += " = "
if is_starred:
if access_full_array and isinstance(newvalue.children[0], lark.Tree):
if newvalue.children[0].data == "vardescriptor_get":
vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor(newvalue.children[0])
if vardescriptor_specifier is not None:
newvalue = None
else:
newvalue = self.__get_user_variable_value(vardescriptor_name, None)
elif newvalue.children[0].data == "listvalue":
newvalue = list(
self.__resolve_operand(c) for c in self.__get_cond_operand(newvalue.children[0])
)
elif newvalue.children[0].data == "wildcard":
backup_result = self.result
newvalue = self.__process_wildcard(newvalue.children[0])
self.result = backup_result
else:
newvalue = None
else:
newvalue = None
if newvalue is None:
self.warn_or_stop(
f"Invalid use of starred value for '{escape_single_quotes(variable_name)}'! Starred values can only be assigned or added to unindexed array variables."
)
newvalue = ""
if is_array:
if variable_specifier is not None:
# Accessing an existing index, we need to update the array
newarray = raw_oldvalue.copy()
newarray[int(variable_specifier)] = newvalue
newvalue = newarray
else:
# Accessing the whole array
if isinstance(newvalue, list):
if raw_oldvalue is None or not adding:
pass
else:
newvalue = raw_oldvalue + newvalue
else:
if raw_oldvalue is None or not adding:
newvalue = [newvalue]
else:
newvalue = raw_oldvalue + [newvalue]
self.__set_user_variable_value(variable_name, newvalue)
currentvalue = self.__get_user_variable_value(variable_name, variable_specifier, False)
if currentvalue is None:
info += "error"
elif isinstance(currentvalue, list):
info += "[" + ", ".join(f"'{escape_single_quotes(str(v))}'" for v in currentvalue) + "]"
else:
info += f"'{escape_single_quotes(currentvalue)}'"
t2 = time.monotonic_ns()
self.__debug_end(command, start_result, t2 - t1, info)
def variableset(self, tree: lark.Tree):
"""
Process a DP set variable command in the tree and add it to the dictionary of variables.
"""
modifiers = tree.children[1] or lark.Tree(lark.Token("RULE", "variablesetmodifiers"), [])
immediate = tree.children[2]
if immediate is not None:
modifiers.children = modifiers.children.copy()
modifiers.children.append(immediate)
vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor(tree.children[0])
self.__varset("variableset", vardescriptor_name, vardescriptor_specifier, modifiers, tree.children[3])
def commandset(self, tree: lark.Tree):
"""
Process a set command in the tree and add it to the dictionary of variables.
"""
vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor(tree.children[0])
self.__varset("commandset", vardescriptor_name, vardescriptor_specifier, tree.children[1], tree.children[2])
def __varecho(
self,
command: str,
variable_name: str,
variable_specifier: str | None,
default: lark.Tree | None,
):
"""
Process a generic echo command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
default_value = None
# if default is not None:
# default_value = self.__visit(default, True) # for log
is_array = variable_name[-2:] == "[]"
vname = f"{variable_name[0:-2]}[{variable_specifier}]" if variable_specifier is not None else variable_name
if variable_name.startswith("_"):
is_systemvar = True
value = self.state.system_variables.get(variable_name, None)
if value is not None:
self.result += value
else:
is_systemvar = False
value = self.__get_user_variable_value(variable_name, variable_specifier, True, True)
if value is None:
if default is not None:
self.log(logging.DEBUG, f"Variable '{escape_single_quotes(vname)}' not found, using default value")
value = self.__visit(default, False, True)
self.result += value
default_value = value
else:
self.warn_or_stop(f"Unknown variable {escape_single_quotes(vname)}")
default_value = ""
value = ""
if not is_systemvar:
self.state.echoed_variables[vname] = value
t2 = time.monotonic_ns()
info = variable_name
if is_array and variable_specifier is not None:
info = info[0:-2] + f"[{variable_specifier}]"
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):
"""
Process a DP use variable command in the tree.
"""
vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor(tree.children[0])
self.__varecho(
"variableuse",
vardescriptor_name,
vardescriptor_specifier,
tree.children[1] if len(tree.children) > 1 else None,
)
def commandecho(self, tree: lark.Tree):
"""
Process an echo command in the tree.
"""
vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor(tree.children[0])
self.__varecho(
"commandecho",
vardescriptor_name,
vardescriptor_specifier,
tree.children[1] if len(tree.children) > 1 else None,
)
def commandif(self, tree: lark.Tree):
"""
Process an if command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
for i, n in enumerate(tree.children):
content = n.children[-1]
if len(n.children) == 2: # its not an else
# has a condition
condition = n.children[0]
c = self.__get_original_node_content(condition, f"condition {i}")
if self.__eval_condition(condition):
self.__visit(content)
t2 = time.monotonic_ns()
self.__debug_end("commandif", start_result, t2 - t1, c)
return
else: # its an else
self.__visit(content)
t2 = time.monotonic_ns()
self.__debug_end("commandif", start_result, t2 - t1, "else")
return
def commandext(self, tree: lark.Tree):
"""
Process an extranetwork command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
extnet = "(ignored)"
if not self.state.options.cup_remove_extranetwork_tags:
extnet_type: str = (tree.children[0].children[0] or "") + str(tree.children[0].children[1])
is_mapping = extnet_type.startswith("$")
if is_mapping:
extnet_type = extnet_type[1:]
extnet_id: str = str(tree.children[1])
if extnet_id.startswith("'") or extnet_id.startswith('"'):
extnet_id = extnet_id[1:-1]
extnet_id = re.sub(r"\\(.)", r"\1", extnet_id) # so we can escape some special characters
parameters: str = ""
parameters_defaulted = False
if tree.children[2]:
parameters = str(tree.children[2])
elif extnet_type in ("lora", "hypernet"):
parameters = "1"
parameters_defaulted = True
if parameters.startswith("'") or parameters.startswith('"'):
parameters = parameters[1:-1]
parameters_is_number = bool(re.match(r"^[-+]?\d*\.?\d+$", parameters or ""))
condition = tree.children[3]
if not condition or self.__eval_condition(condition):
extnet_id = f"{extnet_type}:{extnet_id}"
triggers = tree.children[4] if len(tree.children) > 4 else None
extra_triggers = None
compiled_extra_triggers = None
if is_mapping:
found = self.state.extranetwork_mappings_obj.cached_mappings.get(extnet_id, None)
# we assume the conditions do not change inside the prompt
found_in_cache = found is not None
if found is None:
found_mappings: list[PPPENMappingVariant] = []
else_mapping = None
if self.state.extranetwork_mappings_obj:
enmapping = self.state.extranetwork_mappings_obj.extranetwork_mappings.get(extnet_id, None)
if enmapping:
for v in enmapping.variants:
if v.condition:
try:
cnd = parse_prompt(
self.state,
"condition",
v.condition,
self.state.parsers["condition"],
True,
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
f"Error parsing condition '{escape_single_quotes(v.condition)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}",
e,
)
cnd = None
else:
cnd = "True"
if cnd is not None and (cnd == "True" or self.__eval_condition(cnd)):
if v.condition:
found_mappings.append(v)
else:
else_mapping = v
if found_mappings:
found = found_mappings[
self.__rng.choice(
len(found_mappings),
p=[v.weight or 1 for v in found_mappings],
)
]
else:
found = else_mapping
self.state.extranetwork_mappings_obj.cached_mappings[extnet_id] = found
if found:
if found.name:
if not found_in_cache:
self.log(
logging.INFO,
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
if not f_parameters and extnet_type in ("lora", "hypernet"):
f_parameters = "1"
found_parameters_is_number = True
else:
found_parameters_is_number = f_parameters and bool(
re.match(r"^[-+]?\d*\.?\d+$", str(f_parameters) or "")
)
if found_parameters_is_number and parameters_is_number:
parameters = f"{float(f_parameters) * float(parameters):.2f}".rstrip("0").rstrip(".")
elif f_parameters is not None and parameters_defaulted:
parameters = f_parameters
elif found.triggers:
if not found_in_cache:
self.log(
logging.INFO,
f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to just triggers",
)
extnet_id = None
else:
if not found_in_cache:
self.log(
logging.INFO, f"Mapping extranetwork '{escape_single_quotes(extnet_id)}' to nothing"
)
extnet_id = None
if found.triggers:
extra_triggers = ", ".join(found.triggers)
try:
compiled_extra_triggers = parse_prompt(
self.state, "triggers", extra_triggers, self.state.parsers["content"], True
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
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 '{escape_single_quotes(extnet_id)}' not found!")
if extnet_id:
extnet = f"<{extnet_id}:{parameters}>"
self.result += extnet
elif triggers or compiled_extra_triggers:
extnet = "(only triggers)"
if triggers or compiled_extra_triggers:
if extnet_id:
if not self.state.options.cup_extranetwork_tags:
self.result += " "
else:
self.result += ", "
if triggers:
self.result += self.__visit(triggers, True, True)
if compiled_extra_triggers:
if triggers:
self.result += ", "
self.result += self.__visit(compiled_extra_triggers, True, True)
if triggers or compiled_extra_triggers:
self.result += ", "
t2 = time.monotonic_ns()
self.__debug_end("commandext", start_result, t2 - t1, extnet)
def commandsetwcdeffilter(self, tree: lark.Tree):
"""
Process a setwcdeffilter (Set Wildcard Default Filter) command in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
wildcard_key: str = self.__visit(tree.children[0].children[1], False, True)
selected_wildcards = [x.key for x in self.state.wildcards_obj.get_wildcards(wildcard_key)]
if not selected_wildcards:
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:
self.log(logging.DEBUG, f"Removed default filter for wildcard '{escape_single_quotes(wc)}'")
self.state.wildcards_obj.set_wildcard_default_filter(wc, None)
else:
filter_specifier = self.__extract_filter_specifiers(filter_object)
for wc in selected_wildcards:
self.log(logging.DEBUG, f"Set default filter for wildcard '{escape_single_quotes(wc)}'")
self.state.wildcards_obj.set_wildcard_default_filter(wc, filter_specifier)
t2 = time.monotonic_ns()
self.__debug_end("commandsetwcdeffilter", start_result, t2 - t1)
def extranetworktag(self, tree: lark.Tree):
"""
Process an extra network construct in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
if not self.state.options.cup_remove_extranetwork_tags:
self.result += f"<{tree.children[0]}"
self.__visit(tree.children[1])
self.result += ">"
t2 = time.monotonic_ns()
self.__debug_end("extranetworktag", start_result, t2 - t1)
def __get_choices_internal_get(
self,
choice_values: list[dict],
filter_specifier: Optional[list[list[str]]] = None,
wildcard_key: str = None,
) -> list[dict]:
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):
passes = False
for or_ in filter_specifier:
tmp_pass = True
for and_ in or_:
if and_.isdecimal():
if int(and_) != i:
tmp_pass = False
break
elif re.match(r"^\d+-\d+$", and_):
try:
start, end = map(int, and_.split("-", 1))
if not start <= i <= end:
tmp_pass = False
break
except ValueError:
tmp_pass = False
break
elif and_.lower() not in c.get("labels", []):
tmp_pass = False
break
if tmp_pass:
passes = True
break
if passes:
filtered_choice_values.append(c)
if not filtered_choice_values:
self.warn_or_stop(
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()
expanded_choice_values = []
for i, c in enumerate(filtered_choice_values):
if c.get("command", False):
content_text = self.__visit(c.get("content", ""), False, True).strip()
(cmd, cmd_args) = content_text.split()
if cmd == "include":
wcs = self.state.wildcards_obj.get_wildcards(cmd_args)
if not wcs:
self.warn_or_stop(
f"Included wildcard '{escape_single_quotes(cmd_args)}' not found 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 '{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)
self.log(logging.DEBUG, f"Seen wildcard '{escape_single_quotes(wc.key)}'")
self.log(logging.DEBUG, f"Including choices from wildcard '{escape_single_quotes(wc.key)}'")
(_, choice_values) = self.__check_wildcard_initialization(wc)
if choice_values is not None:
ch_values = self.__get_choices_internal_get(choice_values, None, wc.key)
for cv in ch_values:
expanded_choice_values.append(
{
**cv,
"weight": float(cv.get("weight", 1.0) * c_weight), # we adjust the weight
}
)
else:
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
def __get_choices_select(
self,
options: dict | None,
choice_values: list[dict],
filter_specifier: Optional[list[list[str]]] = None,
wildcard_key: str = None,
) -> tuple[str, list[str], str, str]:
"""
Select choices based on the options.
Args:
options (dict): The object representing the options construct.
choice_values (list[dict]): A list of choice objects.
filter_specifier (list[list[str]]): The filter specifier.
wildcard_key (str): The wildcard key if it is a wildcard.
Returns:
tuple: A tuple containing the prefix, selected choices, separator and suffix
"""
seen_wildcards_len = len(self.__seen_wildcards)
if options is None:
options = {}
sampler: str = options.get("sampler", "~")
repeating: bool = options.get("repeating", False)
optional: bool = options.get("optional", False)
if "count" in options:
from_value = options["count"]
to_value = from_value
else:
from_value: int = options.get("from", 1)
to_value: int = options.get("to", 1)
separator: str = options.get("separator", self.state.options.choice_separator)
msg_where = f"wildcard '{escape_single_quotes(wildcard_key)}'" if wildcard_key else "choices"
if sampler != "~":
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] = []
weights = []
included_choices = 0
excluded_choices = 0
excluded_weights_sum = 0
for i, c in enumerate(expanded_choice_values):
c["choice_index"] = i # we index them to later sort the results
weight = float(c.get("weight", 1.0))
condition = c.get("if", None)
if weight > 0 and (condition is None or self.__eval_condition(condition)):
available_choices.append(c)
weights.append(weight)
included_choices += 1
else:
weights.append(-1)
excluded_choices += 1
excluded_weights_sum += weight
if excluded_choices > 0: # we need to redistribute the excluded weights
weights = [weight + excluded_weights_sum / included_choices for weight in weights if weight >= 0]
weights = np.array(weights)
weights /= weights.sum() # normalize weights
if available_choices:
if from_value < 0:
from_value = 1
elif from_value > len(available_choices):
from_value = len(available_choices)
if to_value < 1:
to_value = 1
elif (to_value > len(available_choices) and not repeating) or from_value > to_value:
to_value = len(available_choices)
num_choices = (
self.__rng.integers(from_value, to_value, endpoint=True) if from_value < to_value else from_value
)
else:
num_choices = 0
if not optional and from_value > 0:
self.warn_or_stop(f"Not enough choices found for {msg_where}!")
if num_choices < 2:
repeating = False
self.log(
logging.DEBUG,
f"Selecting {'optional ' if optional else ''}{'repeating ' if repeating else ''}{num_choices} choice"
+ ("s" if num_choices != 1 else "")
+ (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else ""),
)
if num_choices > 0:
selected_choices: list[dict] = (
list(self.__rng.choice(available_choices, size=num_choices, p=weights, replace=repeating))
if available_choices
else []
)
if self.state.options.keep_choices_order:
selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"])
selected_choices_text = []
prefix: str = (
self.__visit(options.get("prefix", None), False, True)
if options.get("prefix", None) is not None
else ""
)
if prefix != "" and re.match(r"\w", prefix[-1]):
prefix += " "
for i, c in enumerate(selected_choices):
t1 = time.monotonic_ns()
choice_content_obj = c.get("content", c.get("text", None))
if isinstance(choice_content_obj, str):
choice_content = choice_content_obj
else:
choice_content = self.__visit(choice_content_obj, False, True)
t2 = time.monotonic_ns()
self.log(
logging.DEBUG,
f"Adding choice {i+1} ({(t2-t1) / 1_000_000_000:.3f} seconds):\n"
+ textwrap.indent(re.sub(r"\n$", "", choice_content), " "),
)
selected_choices_text.append(choice_content)
suffix: str = (
self.__visit(options.get("suffix", None), False, True)
if options.get("suffix", None) is not None
else ""
)
if suffix != "" and re.match(r"\w", suffix[0]):
suffix = " " + suffix
# remove comments
results = [re.sub(r"\s*#[^\n]*(?:\n|$)", "", r, flags=re.DOTALL) for r in selected_choices_text]
else:
prefix = ""
suffix = ""
results = []
self.log(
logging.DEBUG,
"Unseen wildcards: "
+ ", ".join([f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]),
)
self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len]
return (prefix, results, separator, suffix)
def __convert_choices_options(self, options: Optional[lark.Tree], is_wcdef: bool = False) -> dict:
"""
Convert the choices options to a dictionary.
Args:
options (Tree): The choices options tree.
Returns:
dict: The converted choices options.
"""
if options is None:
return None
options_dict = {}
if len(options.children) == 1:
if options.children[0] is not None:
options_dict["sampler"] = str(options.children[0])
else:
if options.children[0] is not None:
options_dict["sampler"] = str(options.children[0].children[0])
if options.children[1] is not None:
flags = str(options.children[1].children[0])
options_dict["repeating"] = "r" in flags
options_dict["optional"] = "o" in flags
irange = 2
idesc = 3
isep = 4
if options.children[irange] is not None:
if len(options.children[irange].children) == 1:
if options.children[irange].children[0] is not None:
options_dict["count"] = int(options.children[irange].children[0])
else:
options_dict["from"] = (
int(options.children[irange].children[0])
if options.children[irange].children[0] is not None
else 1
)
options_dict["to"] = (
int(options.children[irange].children[1])
if options.children[irange].children[1] is not None
else 1
)
if is_wcdef:
if options.children[idesc] is not None:
options_dict["description"] = str(options.children[idesc].children[0])[1:-1]
else:
isep -= 1 # only wildcard definition options have a description
if options.children[isep] is not None:
options_dict["separator"] = self.__visit(options.children[isep], False, True)
if not options_dict:
options_dict = None
return options_dict
def __convert_choice(self, choice: lark.Tree) -> dict:
"""
Convert the choice to a dictionary.
Args:
choice (Tree): The choice tree.
Returns:
dict: The converted choice.
"""
choice_dict = {}
choice_dict["command"] = choice.children[0] is not None
c_label_obj = choice.children[1]
choice_dict["labels"] = [str(x).lower() for x in c_label_obj.children[1:-1]] if c_label_obj is not None else []
choice_dict["weight"] = float(choice.children[2].children[0]) if choice.children[2] is not None else 1.0
choice_dict["if"] = choice.children[3].children[0] if choice.children[3] is not None else None
choice_dict["content"] = choice.children[-1]
return choice_dict
def __check_wildcard_initialization(self, wildcard: PPPWildcard) -> tuple[dict | None, list[dict] | None]:
"""
Initializes a wildcard if it hasn't been yet.
Args:
wildcard (PPPWildcard): The wildcard to check.
Returns:
tuple: A tuple containing the options and choice values of the wildcard.
"""
choice_values = wildcard.choices
options = wildcard.options
if choice_values is None:
t1 = time.monotonic_ns()
choice_values = []
options, n = self.get_wildcard_options(wildcard)
# we process the choices
for cv in wildcard.unprocessed_choices[n:]:
if isinstance(cv, dict):
if self.state.wildcards_obj.is_dict_choice_options(cv):
condition = cv.get("if", None)
if condition is not None and isinstance(condition, str):
try:
cv["if"] = parse_prompt(
self.state, "condition", condition, self.state.parsers["condition"], True
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
f"Error parsing condition '{escape_single_quotes(condition)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
e,
)
cv["if"] = None
content = cv.get("content", cv.get("text", None))
cv["content"] = content
if "text" in cv:
del cv["text"]
if content is not None and isinstance(content, str):
try:
cv["content"] = parse_prompt(
self.state, "choicevalue", content, self.state.parsers["choicevalue"], True
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
f"Error parsing choice content '{escape_single_quotes(content)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
e,
)
cv["content"] = None
if cv["content"] is not None:
self.log(logging.DEBUG, f"Processed choice {cv}")
choice_values.append(cv)
else:
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 '{escape_single_quotes(wildcard.key)}'!")
else:
try:
choice_values.append(
self.__convert_choice(
parse_prompt(self.state, "choice", cv, self.state.parsers["choice"], True)
)
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
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()
self.log(
logging.DEBUG,
f"Processed choices for wildcard '{escape_single_quotes(wildcard.key)}' ({(t2-t1) / 1_000_000_000:.3f} seconds)",
)
return (self.__clean_wildcard_options(options), choice_values)
def get_wildcard_options(self, wildcard: PPPWildcard) -> tuple[dict | None, int]:
options = wildcard.options
n = 0
# we check the first choice to see if it is actually options
if isinstance(wildcard.unprocessed_choices[0], dict):
if self.state.wildcards_obj.is_dict_wcdef_options(wildcard.unprocessed_choices[0]):
options = wildcard.unprocessed_choices[0]
prefix = options.get("prefix", None)
if prefix is not None and isinstance(prefix, str):
try:
options["prefix"] = parse_prompt(
self.state, "choicevalue", prefix, self.state.parsers["choicevalue"], True
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
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)
if suffix is not None and isinstance(suffix, str):
try:
options["suffix"] = parse_prompt(
self.state, "choicevalue", suffix, self.state.parsers["choicevalue"], True
)
except lark.exceptions.UnexpectedInput as e:
self.warn_or_stop(
f"Error parsing choice suffix '{escape_single_quotes(suffix)}' in wildcard '{escape_single_quotes(wildcard.key)}'! : {e.__class__.__name__}",
e,
)
n = 1
else:
if wildcard.unprocessed_choices[0].endswith("$$"):
try:
options = self.__convert_choices_options(
parse_prompt(
self.state,
"as wildcard options",
wildcard.unprocessed_choices[0][:-2].strip(),
self.state.parsers["wcdefoptions"],
True,
),
True,
)
n = 1
except lark.exceptions.UnexpectedInput:
options = None
if options is None:
self.log(logging.DEBUG, "Does not have options")
wildcard.options = options
return options, n
def __clean_wildcard_options(self, options: dict) -> dict:
if options is not None:
options.pop("description", None) # description is only for wildcard definitions, not for usage
if not options:
options = None
return options
def __process_wildcard(self, tree: lark.Tree) -> list:
"""
Process a wildcard in the tree.
Returns:
list: A list containing the selected elements.
"""
t1 = time.monotonic_ns()
chosen_choices = []
start_result = self.result
seen_wildcards_len = len(self.__seen_wildcards)
applied_options = self.__clean_wildcard_options(self.__convert_choices_options(tree.children[0], False))
wildcard_key: str = self.__visit(tree.children[1], False, True)
wc = self.__get_original_node_content(tree, f"?__{wildcard_key}__")
if self.state.options.process_wildcards:
self.log(logging.DEBUG, f"Processing wildcard: {wildcard_key}")
selected_wildcards = self.state.wildcards_obj.get_wildcards(wildcard_key)
if not selected_wildcards:
self.detectedWildcards.append(wc)
self.result += wc
t2 = time.monotonic_ns()
self.__debug_end("wildcard", start_result, t2 - t1, wc)
return []
filter_specifier: list[list[str]] = None
filter_object = tree.children[2]
if filter_object is not None:
if ( # it's an inherited filter from another wildcard
isinstance(filter_object.children[1], lark.Token)
and filter_object.children[1] is not None
and "^" in str(filter_object.children[1])
):
filter_wildcard_key = self.__visit(filter_object.children[2], False, True)
filter_specifier = self.__wildcard_filters.get(filter_wildcard_key, None)
self.log(logging.DEBUG, "Filtering choices with inherited filter")
else:
filter_specifier = self.__extract_filter_specifiers(filter_object.children[2])
self.log(logging.DEBUG, "Filtering choices")
self.__wildcard_filters[wildcard_key] = filter_specifier
if filter_object.children[1] is not None and "#" in str(
filter_object.children[1]
): # means do not use the filter in this wildcard
self.log(logging.DEBUG, "Ignoring filter")
filter_specifier = None
else:
filter_specifier = self.state.wildcards_obj.get_wildcard_default_filter(wildcard_key)
if filter_specifier is not None:
self.__wildcard_filters[wildcard_key] = filter_specifier
self.log(logging.DEBUG, "Applying default filter")
if (
len(selected_wildcards) > 1
and filter_specifier is not None
and any(y[0].isdecimal() for x in filter_specifier for y in x)
):
self.log(
logging.WARNING,
f"Using a globbing wildcard '{escape_single_quotes(wildcard_key)}' with positional index filters is not recommended!",
)
var_object = tree.children[3]
variablename = None
variablebackup = None
if var_object is not None:
vardescriptor_name, vardescriptor_specifier = self.__separate_vardescriptor(var_object.children[0])
variablename = vardescriptor_name
# variablevalue = self.__visit(var_object.children[1], False, True)
variablebackup = self.state.user_variables.get(variablename, None)
# self.__remove_user_variable(variablename)
# self.__set_user_variable_value(variablename, variablevalue)
self.__varset("wildcard", variablename, vardescriptor_specifier, None, var_object.children[1])
choice_values_all = []
for wildcard in selected_wildcards:
if wildcard is None:
self.detectedWildcards.append(wc)
self.result += wc
t2 = time.monotonic_ns()
self.__debug_end("wildcard", start_result, t2 - t1, wc)
return []
if wildcard.key in self.__seen_wildcards:
self.warn_or_stop(
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)
self.log(logging.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:
self.log(
logging.DEBUG, f"Options for wildcard '{escape_single_quotes(wildcard.key)}' are ignored!"
)
choice_values_all += choice_values
prefix, chosen_choices, separator, suffix = self.__get_choices_select(
applied_options, choice_values_all, filter_specifier, wildcard_key
)
if chosen_choices:
self.result += prefix + separator.join(chosen_choices) + suffix
if wildcard_key in self.__wildcard_filters:
del self.__wildcard_filters[wildcard_key]
if variablename is not None:
self.__remove_user_variable(variablename)
if variablebackup is not None:
self.state.user_variables[variablename] = variablebackup
elif self.state.options.if_wildcards != IFWILDCARDS_CHOICES.remove:
self.detectedWildcards.append(wc)
self.result += wc
if self.__debug_level == DEBUG_LEVEL.full:
list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]
self.log(logging.DEBUG, f"Unseen wildcards: {', '.join(list_unseen)}")
self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len]
t2 = time.monotonic_ns()
self.__debug_end("wildcard", start_result, t2 - t1, f"'{escape_single_quotes(wc)}'")
return chosen_choices
def wildcard(self, tree: lark.Tree):
"""
Process a wildcard construct in the tree.
"""
self.__process_wildcard(tree)
def __extract_filter_specifiers(self, filters: lark.Tree) -> list[list[str]]:
filter_specifier = []
for or_ in filters.children:
for and_ in or_.children:
label = and_.children[0]
if isinstance(label, lark.Token):
# it's a literal, we can use it directly
filter_specifier.append([str(label)])
else:
# it's a variable, we need to evaluate it
v = self.__visit(label, False, True)
filter_specifier.append([v])
return filter_specifier
def choices(self, tree: lark.Tree):
"""
Process a choices construct in the tree.
"""
t1 = time.monotonic_ns()
start_result = self.result
options = self.__convert_choices_options(tree.children[0], False)
choice_values = [self.__convert_choice(c) for c in tree.children[1::]]
ch = self.__get_original_node_content(tree, "?{...}")
if self.state.options.process_wildcards:
self.log(logging.DEBUG, "Processing choices:")
prefix, chosen_choices, separator, suffix = self.__get_choices_select(options, choice_values)
if chosen_choices:
self.result += prefix + separator.join(chosen_choices) + suffix
elif self.state.options.if_wildcards != IFWILDCARDS_CHOICES.remove:
self.detectedWildcards.append(ch)
self.result += ch
t2 = time.monotonic_ns()
self.__debug_end("choices", start_result, t2 - t1, f"'{escape_single_quotes(ch)}'")
def __default__(self, tree):
t1 = time.monotonic_ns()
start_result = self.result
self.__visit(tree.children)
t2 = time.monotonic_ns()
self.__debug_end(tree.data.value, start_result, t2 - t1)
def start(self, tree):
self.result = ""
t1 = time.monotonic_ns()
self.__visit(tree.children)
attention_processing = self.state.host_config.attention
# process the found negative tags
for negtag in self.__negtags:
if self.state.options.cup_merge_attention:
# join consecutive attention elements
for i in range(len(negtag.shell) - 1, 0, -1):
if negtag.shell[i].type == "at" and negtag.shell[i - 1].type == "at":
new_weight = ( # we limit the new weight to two decimals
math.floor(100 * float(negtag.shell[i - 1].data[1]) * float(negtag.shell[i].data[1])) / 100
)
new_weight_str = f"{new_weight:.2f}".rstrip("0").rstrip(".")
if new_weight_str == "0.9" and attention_processing != "parentheses":
new_kind = 1
elif new_weight_str == "1.1":
new_kind = 2
else:
new_kind = 3
negtag.shell[i - 1] = self.AccumulatedShell(
"at",
(new_kind, new_weight_str),
)
negtag.shell.pop(i)
start = ""
end = ""
for s in negtag.shell:
match s.type:
case "at":
if s.data[0] == 1:
start += "["
end = "]" + end
elif s.data[0] == 2:
start += "("
end = ")" + end
else: # 3
start += "("
end = f":{s.data[1]})" + end
# case "sc":
case "scb":
start += "["
end = f"::{s.data}]" + end
case "sca":
start += "["
end = f":{s.data}]" + end
# case "al":
case "alo":
start += "[" + ("|" * int(s.data["pos"] - 1))
end = ("|" * int(s.data["len"] - s.data["pos"])) + "]" + end
content = start + negtag.content + end
position = negtag.parameters or "s"
if position.startswith("i"):
n = int(position[1])
self.insertion_at[n] = [negtag.start, negtag.end]
elif len(content) > 0:
if content not in self.__already_processed:
if self.state.options.stn_ignore_repeats:
self.__already_processed.append(content)
self.log(logging.DEBUG, f"Adding content at position {position}: {content}")
if position == "e":
self.add_at["end"].append(content)
elif position.startswith("p"):
n = int(position[1])
self.add_at["insertion_point"][n].append(content)
else: # position == "s" or invalid
self.add_at["start"].append(content)
else:
self.log(logging.WARNING, f"Ignoring repeated content: {content}")
t2 = time.monotonic_ns()
self.__debug_end("start", "", t2 - t1)