from collections import namedtuple from functools import reduce from itertools import combinations, combinations_with_replacement, permutations, product import logging import math import re import textwrap import time from typing import Any, Callable, Optional import lark import numpy as np from ppp_classes import IFWILDCARDS_CHOICES, SUPPORTED_APPS, 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. """ NEGATIVE_SEP = "\x1d" AccumulatedShell = namedtuple("AccumulatedShell", ["type", "data"]) NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"]) def __init__(self, state: PPPState, rng: np.random.Generator, on_model_info_update: Optional[Callable[[], None]] = None): super().__init__() self.state = state self.__on_model_info_update = on_model_info_update self.__debug_level = state.options.debug_level self.__rng = rng self.__shell: list[TreeProcessor.AccumulatedShell] = [] # type: ignore self.__negtags: list[TreeProcessor.NegTag] = [] # type: ignore self.__already_processed: list[str] = [] self.__is_negative = False self.__wildcard_filters = {} self.__seen_wildcards: list[str] = [] self.__add_at: dict[str, list] = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []} self.__insertion_at: list[tuple[int, int]] = [None for _ in range(10)] self.__detectedWildcards: list[tuple[str, bool]] = [] self.__result = "" self.__comb_forced_path: list[int] = [] self.__comb_trace: list[int] = [] self.__cycl_forced_path: list[int] = [] self.__cycl_trace: list[int] = [] def log(self, kind, message: str, min_level: DEBUG_LEVEL | None = None): log(self.state.logger, self.state.options.debug_level, kind, message, min_level) def warn_or_stop(self, message: str, e: Exception = None): warn_or_stop(self.state, self.__is_negative, message, e) def __reset_run_state(self): """Reset all per-run mutable state for a fresh combinatorial pass.""" self.__shell = [] self.__negtags = [] self.__already_processed = [] self.__is_negative = False self.__wildcard_filters = {} self.__seen_wildcards = [] self.__add_at = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []} self.__insertion_at = [None for _ in range(10)] self.__detectedWildcards = [] self.__result = "" if self.state.extranetwork_mappings_obj is not None: self.state.extranetwork_mappings_obj.cached_mappings.clear() def start_visit( self, parsed: lark.Tree, ) -> list[tuple[str, list[tuple[str, bool]], tuple[dict[str, Any], dict[str, str]]]]: """ Process the positive and negative prompts in a unified way using the same processor. STN insertions are applied to the negative result directly inside this processor. Args: parsed (Tree): The parsed unified prompt. Returns: list[tuple[str, list[tuple[str,bool]], tuple[dict[str, Any], dict[str, str]]]]: A list of (processed prompt, detected wildcards, variables snapshot) triples — one entry per combination in combinatorial mode, or a single entry otherwise. The variables snapshot is the return value of ``state.variables.backup_user_and_echoed()`` captured after processing. """ self.log(logging.INFO, "Processing prompt...") self.__detectedWildcards = [] self.__is_negative = False self.__result = "" if not self.state.options.do_combinatorial: self.__cycl_forced_path = list(self.state.cyclical_state.current_path) self.__cycl_trace = [] self.visit(parsed) self.__finalize_echoed_variables() if self.__cycl_trace: self.state.cyclical_state.last_trace = self.__cycl_trace[:] self.state.cyclical_state.advance() return [(self.__result, self.__detectedWildcards, self.state.variables.backup_user_and_echoed())] # Combinatorial mode: explore every possible path through choices and wildcards via DFS. # __comb_forced_path drives which option is selected at each decision point; # __comb_trace records how many options were available at each point so the DFS can # correctly enumerate unexplored branches after each run. initial_vars = self.state.variables.backup_user_and_echoed() results: list[tuple[str, list[tuple[str, bool]], tuple]] = [] limit = self.state.options.combinatorial_limit def _run(forced_path: tuple[int, ...]) -> tuple[int, ...]: self.log(logging.DEBUG, f"Running combinatorial path: {forced_path}") self.__comb_forced_path = list(forced_path) self.__comb_trace = [] self.__reset_run_state() self.state.variables.restore_user_and_echoed(initial_vars) self.visit(parsed) self.__finalize_echoed_variables() results.append( (self.__result, self.__detectedWildcards.copy(), self.state.variables.backup_user_and_echoed()) ) if len(results) == 1: first_run_estimate = reduce(lambda x, y: x * y, self.__comb_trace, 1) self.log(logging.INFO, f"Estimated combinations (lower bound): {first_run_estimate}") self.log(logging.INFO, f"Added combination {len(results)}") return tuple(self.__comb_trace) limit_reached = False def _dfs(forced_path: tuple[int, ...]): nonlocal limit_reached if 0 < limit <= len(results): limit_reached = True return trace = _run(forced_path) # For each decision that was reached but not forced, spawn branches for all # options beyond the default (index 0). # Iterate in reverse so later decisions vary fastest, producing lexicographic order. for i in range(len(trace) - 1, len(forced_path) - 1, -1): if 0 < limit <= len(results): limit_reached = True return num_options = trace[i] for opt in range(1, num_options): if 0 < limit <= len(results): limit_reached = True return # Pad with zeros for intermediate decisions so they keep the default. new_path = forced_path + (0,) * (i - len(forced_path)) + (opt,) _dfs(new_path) _dfs(()) if limit_reached: self.log(logging.WARNING, f"Combinatorial limit of {limit} reached; some combinations have been skipped.") return results def __finalize_echoed_variables(self): var_keys = self.state.variables.all_user_or_echoed_keys() for k in var_keys: ev = self.state.variables.get_echoed_value(k) if ev is None: ev = self.state.variables.get_user(k) if ev is None or ev.__class__ != str: # strict check to avoid problems with Tokens self.log(logging.DEBUG, f"Completing variable: {k}") ev = self.get_final_variable(k) self.state.variables.echo(k, ev) # ensure all variables are echoed so they are included in the snapshot def __visit( self, 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_vars = self.state.variables.backup_user_and_echoed() 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.variables.restore_user_and_echoed(backup_vars) 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_variable(specifier) if specifier.isdecimal(): return int(specifier), None, False # invalid specifier return None, None, False def __get_variable_value( self, name: str, specifier: str | None = None, evaluate=True, visit=False ) -> str | list[str] | None: """ Get the value of a variable. Args: name (str): The name of the 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 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.variables.get(name) 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_variable(self, name_specifier: str) -> str: """ Get the final value of a variable, resolving any references if needed. Args: name_specifier (str): The variable reference string. Returns: str: The final value of the variable. """ name, specifier = self.__separate_arrayref(name_specifier) v = self.__get_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 __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 compatible_types = [ (str, str), (int, int), (float, float), (bool, bool), (int, float), (float, int), (str, list), (int, list), (float, list), (bool, list), ] if not any(isinstance(operand1, t1) and isinstance(operand2, t2) for t1, t2 in compatible_types): self.warn_or_stop( f"Mixed type values ({type(operand1).__name__}, {type(operand2).__name__}) used in comparison: '{escape_single_quotes(desc)}'" ) 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 varname, varspecifier = self.__separate_arrayref(c) val = self.__get_variable_value(varname, varspecifier) if val is None: val = "" vartype = "system" if self.state.variables.name_is_system(c) else "user" 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 negative_sep(self, tree: lark.Tree): """ Process a negative prompt separator in the tree. """ start_result = self.__result t1 = time.monotonic_ns() x = tree.children[0] self.__result += x.value self.__is_negative = True t2 = time.monotonic_ns() self.__debug_end("negative_sep", start_result, t2 - t1) 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(TreeProcessor.AccumulatedShell("sc", pos)) self.__result += "[" if before is not None: self.log(logging.DEBUG, f"Shell scheduled before with position {pos}") self.__shell.append(TreeProcessor.AccumulatedShell("scb", pos)) self.__visit(before) self.__shell.pop() self.log(logging.DEBUG, f"Shell scheduled after with position {pos}") self.__shell.append(TreeProcessor.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(TreeProcessor.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(TreeProcessor.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) @staticmethod def _try_extract_attention(s: str) -> tuple[str, float] | None: """ If s is entirely a single attention wrapper - (inner), (inner:W), or [inner] - return (inner_content, weight). Otherwise return None. Used for post-visit merging when a wildcard or choices construct expands to a single attention that can be merged with an enclosing outer attention. """ if not s: return None open_char = s[0] if open_char == "(": close_char = ")" elif open_char == "[": close_char = "]" else: return None depth = 0 for i, c in enumerate(s): if c == open_char: depth += 1 elif c == close_char: depth -= 1 if depth == 0: if i != len(s) - 1: return None # wrapper closes before end of string -> multiple items break else: return None # never fully closed inner = s[1:-1] if open_char == "[": # Disambiguate from alternation [a|b] and scheduling [before:after:N]. # Both use [...] but are not attention constructs. paren_depth = 0 bracket_depth = 0 top_level_pipes = 0 top_level_colons = 0 last_colon_pos = -1 for i, c in enumerate(inner): if c == "(": paren_depth += 1 elif c == ")": paren_depth -= 1 elif c == "[": bracket_depth += 1 elif c == "]": bracket_depth -= 1 elif paren_depth == 0 and bracket_depth == 0: if c == "|": top_level_pipes += 1 elif c == ":": top_level_colons += 1 last_colon_pos = i if top_level_pipes > 0: return None # alternation construct if top_level_colons >= 2 and last_colon_pos >= 0: try: float(inner[last_colon_pos + 1 :]) return None # scheduling construct: [before:after:N] except ValueError: pass return (inner, 0.9) # Parenthesis form - scan backwards for a top-level :weight suffix depth = 0 for i in range(len(inner) - 1, -1, -1): c = inner[i] if c in ")]": depth += 1 elif c in "([": depth -= 1 elif c == ":" and depth == 0: try: w = float(inner[i + 1 :]) return (inner[:i], w) except ValueError: break return (inner, 1.1) 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: # we check while the children are attentions, in which case we merge the weights 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 self.log( logging.DEBUG, f"Merging nested attention with weight {inner_weight}, cumulative weight now {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(TreeProcessor.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})" # Post-visit merge: if the entire visited content is a single attention wrapper # (e.g. from a wildcard or choices expansion), merge weights here. # The static tree-walk above only covers direct attention children; this # handles the case where the inner attention came from an expanded wildcard. if self.state.options.cup_merge_attention: visited_content = self.__result[len(start_result) + len(starttag) :] merge = TreeProcessor._try_extract_attention(visited_content) if merge is not None: inner_content, inner_weight = merge weight = math.floor(weight * inner_weight * 100) / 100 self.log( logging.DEBUG, f"Merging nested attention with weight {inner_weight}, cumulative weight now {weight}", ) weight_str = f"{weight:.2f}".rstrip("0").rstrip(".") if weight_str == "1.1": weight_kind = 2 starttag = "(" endtag = ")" elif weight_str == "0.9" and attention_processing != "parentheses": weight_kind = 1 starttag = "[" endtag = "]" else: weight_kind = 3 starttag = "(" endtag = f":{weight_str})" self.__result = start_result + starttag + inner_content 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_nodes = tree.children[1::] attention_processing = self.state.host_config.attention peeled = False if ( self.state.options.cup_merge_attention and attention_processing in ("ok", "parentheses") and len(content_nodes) == 1 and isinstance(content_nodes[0], lark.Tree) and content_nodes[0].data == "attention" ): inner_tree = content_nodes[0] weight = 1.0 while isinstance(inner_tree, lark.Tree) and inner_tree.data == "attention": if len(inner_tree.children) == 2: w = inner_tree.children[-1] inner_weight = float(w) if w is not None else 1.1 else: inner_weight = 0.9 weight *= inner_weight inner_tree = inner_tree.children[0] weight = math.floor(weight * 100) / 100 weight_str = f"{weight:.2f}".rstrip("0").rstrip(".") if weight_str == "0.9" and attention_processing != "parentheses": weight_kind = 1 elif weight_str == "1.1": weight_kind = 2 else: weight_kind = 3 if attention_processing == "parentheses" and weight_kind == 1: weight_kind = 3 weight_str = "0.9" self.__shell.append(TreeProcessor.AccumulatedShell("at", (weight_kind, weight_str))) content = self.__visit(inner_tree, False, True) peeled = True else: content = self.__visit(content_nodes, False, True) self.__negtags.append( TreeProcessor.NegTag(len(self.__result), len(self.__result), content, parameters, self.__shell.copy()) ) if peeled: self.__shell.pop() 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( TreeProcessor.NegTag(len(self.__result), len(self.__result), "", parameters, self.__shell.copy()) ) info = f"with {parameters or 'no parameters'}" else: self.warn_or_stop("Ignored negative insertion point command in positive prompt") 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 settable_sysvars = {"_modelfullname": "model_filename", "_modelclass": "model_class"} if self.state.variables.name_is_system(variable_name): if (variable_name not in settable_sysvars and variable_name != "_modelinfo"): self.warn_or_stop( f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! System variables cannot be set." ) return app = self.state.env_info.get("app", "") if app not in (SUPPORTED_APPS.comfyui.value, SUPPORTED_APPS.tests.value): self.warn_or_stop( f"Setting '{escape_single_quotes(variable_name)}' is only supported in ComfyUI." ) return evaluated = self.__visit(content, restore_state=False, discard_content=True) if variable_name == "_modelinfo": # Format: @ at_pos = evaluated.find("@") if at_pos < 0: self.warn_or_stop( f"Invalid value for '_modelinfo': expected '@', got '{escape_single_quotes(evaluated)}'." ) return self.state.env_info["model_class"] = evaluated[:at_pos] self.state.env_info["model_filename"] = evaluated[at_pos + 1 :] else: self.state.env_info[settable_sysvars[variable_name]] = evaluated if self.__on_model_info_update is not None: self.__on_model_info_update() info = variable_name + " = " + f"'{escape_single_quotes(evaluated)}'" t2 = time.monotonic_ns() self.__debug_end(command, start_result, t2 - t1, info) 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.variables.get_user(variable_name) 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_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 = "" elif not isinstance(newvalue, list): self.warn_or_stop( f"Invalid starred value for '{escape_single_quotes(variable_name)}'! Starred values should be a list." ) 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.state.variables.set_user(variable_name, newvalue) currentvalue = self.__get_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 value = self.__get_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 self.state.variables.name_is_system(variable_name): self.state.variables.echo(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: cond = str(v.condition) if v.condition is not None else None if cond: try: cnd = parse_prompt( self.state, "condition", cond, self.state.parsers["condition"], True, ) except lark.exceptions.UnexpectedInput as e: self.warn_or_stop( f"Error parsing condition '{escape_single_quotes(cond)}' in extranetwork mapping '{escape_single_quotes(extnet_id)}'! : {e.__class__.__name__}", e, ) cnd = None else: cnd = "True" if cnd is not None and (cnd == "True" or self.__eval_condition(cnd)): if cond: found_mappings.append(v) else: else_mapping = v num_mappings = len(found_mappings) if num_mappings > 0: if self.state.options.do_combinatorial: decision_idx = len(self.__comb_trace) self.__comb_trace.append(num_mappings) chosen_idx = ( min(self.__comb_forced_path[decision_idx], num_mappings - 1) if decision_idx < len(self.__comb_forced_path) else 0 ) found = found_mappings[chosen_idx] elif num_mappings == 1: found = found_mappings[0] else: found = found_mappings[ self.__rng.choice( num_mappings, p=[v.weight or 1 for v in found_mappings], ) ] else: found = else_mapping if num_mappings < 2: 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[lark.Tree, list[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[lark.Tree,list[str]]: The resulting container and list of chosen choices. """ 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 not in ("~", "@"): 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) comb_chosen_selection: Optional[list[dict]] = None if self.state.options.do_combinatorial or sampler == "@": # Enumerate every distinct selection of choices, accounting for count range and repetition. all_selections: list[tuple] = [] # When keep_choices_order is False the output depends on the selection order, # so we must enumerate ordered sequences (permutations / product). # When keep_choices_order is True selections are sorted afterward, so all # orderings of the same items produce identical output and we only need # unordered iterators (combinations / combinations_with_replacement). for k in range(from_value, to_value + 1): if repeating: if self.state.options.keep_choices_order: all_selections.extend(combinations_with_replacement(available_choices, k)) else: all_selections.extend(product(available_choices, repeat=k)) else: if self.state.options.keep_choices_order: all_selections.extend(combinations(available_choices, k)) else: all_selections.extend(permutations(available_choices, k)) num_selections = len(all_selections) if self.state.options.do_combinatorial: decision_idx = len(self.__comb_trace) self.__comb_trace.append(num_selections) chosen_idx = ( min(self.__comb_forced_path[decision_idx], num_selections - 1) if decision_idx < len(self.__comb_forced_path) else 0 ) else: # sampler == "@" cycl_decision_idx = len(self.__cycl_trace) self.__cycl_trace.append(num_selections) chosen_idx = ( self.__cycl_forced_path[cycl_decision_idx] % num_selections if cycl_decision_idx < len(self.__cycl_forced_path) else 0 ) comb_chosen_selection = list(all_selections[chosen_idx]) num_choices = len(comb_chosen_selection) else: num_choices = ( self.__rng.integers(from_value, to_value, endpoint=True) if from_value < to_value else from_value ) else: num_choices = 0 if not optional and from_value > 0: 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: if comb_chosen_selection is not None: selected_choices: list[dict] = comb_chosen_selection else: selected_choices: list[dict] = ( list(self.__rng.choice(available_choices, size=num_choices, p=weights, replace=repeating)) if available_choices else [] ) if self.state.options.keep_choices_order: selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"]) selected_choices_text = [] 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) # remove comments results = [re.sub(r"\s*#[^\n]*(?:\n|$)", "", r, flags=re.DOTALL) for r in selected_choices_text] else: results = [] container = options.get("container", None) if container is None: separator = options.get("separator", self.state.options.choice_separator) container = lark.Tree( lark.Token("RULE", "content"), [ lark.Tree( lark.Token("RULE", "variableuse"), [ lark.Tree( lark.Token("RULE", "vardescriptor_get"), [ lark.Token("__identifier", "_choices"), lark.Token("__openbracket", "["), lark.Tree( lark.Token("RULE", "separator_descriptor"), [ lark.Token("__separatorflag", "&"), lark.Token("STRING", "'" + separator + "'"), ], ), lark.Token("__closebracket", "]"), ], ), None, ], ), ], ) 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 container, results def __apply_container(self, container: lark.Tree, choices: list[str]) -> str: # we save the choices variable in case there are nested choices old_choices = self.state.variables.get_system("_choices[]", None) self.state.variables.set_system("_choices[]", choices) joined_results = self.__visit(container, False, True) # we restore the old choices variable self.state.variables.set_system("_choices[]", old_choices) return joined_results 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] container = options.get("container", None) container_kind = "specified" if container is None: container_kind = "assembled" has_separator = "separator" in options has_prefix = "prefix" in options has_suffix = "suffix" in options if has_separator or has_prefix or has_suffix: separator = options.get("separator", self.state.options.choice_separator) container = "${_choices[&'" + separator + "']}" prefix = options.get("prefix", None) if prefix is not None and isinstance(prefix, str): if prefix != "" and re.match(r"\w", prefix[-1]): prefix += " " container = prefix + container suffix = options.get("suffix", None) if suffix is not None and isinstance(suffix, str): if suffix != "" and re.match(r"\w", suffix[0]): suffix = " " + suffix container += suffix options.pop("separator", None) options.pop("prefix", None) options.pop("suffix", None) options.pop("container", None) if container is not None and isinstance(container, str): try: options["container"] = parse_prompt( self.state, "choicevalue", container, self.state.parsers["choicevalue"], True ) except lark.exceptions.UnexpectedInput as e: self.warn_or_stop( f"Error parsing choice {container_kind} container '{escape_single_quotes(container)}' 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.__is_negative)) 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.variables.get_user(variablename) # self.state.variables.delete_user(variablename) # self.state.variables.set_user(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.__is_negative)) 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 container, chosen_choices = self.__get_choices_select( applied_options, choice_values_all, filter_specifier, wildcard_key ) if chosen_choices: self.__result += self.__apply_container(container, chosen_choices) if wildcard_key in self.__wildcard_filters: del self.__wildcard_filters[wildcard_key] if variablename is not None: self.state.variables.delete_user(variablename) if variablebackup is not None: self.state.variables.set_user(variablename, variablebackup) elif self.state.options.if_wildcards != IFWILDCARDS_CHOICES.remove: self.__detectedWildcards.append((wc, self.__is_negative)) self.__result += wc 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]]: if ( len(filters.children) == 1 and len(filters.children[0].children) == 1 and isinstance(filters.children[0].children[0], lark.Tree) ): # special case when the whole filter is in a variable # note that an individual label in a variable will also come through here f = self.__visit(filters.children[0].children[0], False, True) filters = parse_prompt( self.state, "filter specifier", str(f), self.state.parsers["wc_filter_or"], True, ) filter_specifier = [] for or_ in filters.children: and_group = [] for and_ in or_.children: label = and_.children[0] if isinstance(label, lark.Token): # it's a literal, we can use it directly and_group.append(str(label)) else: # it's a variable, we need to evaluate it v = self.__visit(label, False, True) # we remove commas and pluses to avoid confusion with the filter specifier syntax and_group.append(v.replace(",", "").replace("+", "")) filter_specifier.append(and_group) 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:") container, chosen_choices = self.__get_choices_select(options, choice_values) self.__result += self.__apply_container(container, chosen_choices) elif self.state.options.if_wildcards != IFWILDCARDS_CHOICES.remove: self.__detectedWildcards.append((ch, self.__is_negative)) self.__result += ch t2 = time.monotonic_ns() self.__debug_end("choices", start_result, t2 - t1, f"'{escape_single_quotes(ch)}'") def __default__(self, tree): t1 = time.monotonic_ns() start_result = self.__result self.__visit(tree.children) t2 = time.monotonic_ns() self.__debug_end(tree.data.value, start_result, t2 - t1) def __process_negtags(self): # process the found negative tags 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 self.state.host_config.attention != "parentheses": new_kind = 1 elif new_weight_str == "1.1": new_kind = 2 else: new_kind = 3 negtag.shell[i - 1] = TreeProcessor.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}") self.__negtags = [] def __apply_stn_insertions(self): """ Apply all accumulated STN content from add_at to self.result using the recorded insertion_at positions, then reset both so ppp.py does not re-apply them. """ pos, neg = self.__result.split(self.NEGATIVE_SEP, 1) neg_start = len(pos) + len(self.NEGATIVE_SEP) stn_sep = self.state.options.stn_separator self.log(logging.DEBUG, f"Applying STN additions to negative: {self.__add_at}") self.log(logging.DEBUG, f"Applying STN indexes: {self.__insertion_at}") ordered_range = sorted( range(10), key=lambda x: self.__insertion_at[x][0] if self.__insertion_at[x] is not None else float("-inf"), reverse=True, ) for n in ordered_range: if self.__insertion_at[n] is not None: insertion_point_n: list[str] = self.__add_at["insertion_point"][n] ipp = self.__insertion_at[n][0] - neg_start ipl = self.__insertion_at[n][1] - self.__insertion_at[n][0] if neg[ipp - len(stn_sep) : ipp] == stn_sep: ipp -= len(stn_sep) # adjust for existing start separator ipl += len(stn_sep) insertion_point_n.insert(0, neg[:ipp]) if neg[ipp + ipl : ipp + ipl + len(stn_sep)] == stn_sep: ipl += len(stn_sep) # adjust for existing end separator end_part = neg[ipp + ipl :] if len(end_part) > 0: insertion_point_n.append(end_part) neg = stn_sep.join(insertion_point_n) else: ipp = 0 if neg.startswith(stn_sep): ipp = len(stn_sep) self.__add_at["insertion_point"][n].append(neg[ipp:]) neg = stn_sep.join(self.__add_at["insertion_point"][n]) if self.__add_at["start"]: add_at_start = self.__add_at["start"] if len(neg) > 0: ipp = 0 if neg.startswith(stn_sep): ipp = len(stn_sep) # adjust for existing end separator add_at_start.append(neg[ipp:]) neg = stn_sep.join(add_at_start) if self.__add_at["end"]: add_at_end = self.__add_at["end"] if len(neg) > 0: ipl = len(neg) if neg.endswith(stn_sep): ipl -= len(stn_sep) # adjust for existing start separator add_at_end.insert(0, neg[:ipl]) neg = stn_sep.join(add_at_end) # self.add_at = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []} # self.insertion_at = [None for _ in range(10)] self.__result = pos + self.NEGATIVE_SEP + neg def start(self, tree): self.__result = "" t1 = time.monotonic_ns() self.__visit(tree.children) self.__process_negtags() if self.__is_negative: self.__apply_stn_insertions() t2 = time.monotonic_ns() self.__debug_end("start", "", t2 - t1)