diff --git a/docs/CONFIG.md b/docs/CONFIG.md index afc1752..d52b4bd 100644 --- a/docs/CONFIG.md +++ b/docs/CONFIG.md @@ -25,6 +25,7 @@ Inputs: * **neg_prompt**: Connect here the negative prompt text, or fill it as a widget. * **debug_level**: What to write to the console. * **on_warnings**: Warn on the console or stop the generation. +* **strict_mode**: Sets the strict mode in comparison operations. * **process_wildcards**: Activates the wildcard processing. * **do_cleanup**: Activates the cleanup processing. * **cleanup_variables**: Do a cleanup of the output variables (depends on do_cleanup). @@ -130,7 +131,8 @@ Options for extranetworks mapping, in case you want to change them from the defa ### General settings * **Debug level**: What to write to the console. Note: in *SD.Next* debug messages only show if you launch it with the `--debug` argument. -* **What to do on invalid content warnings?**: Warn on the console or stop the generation. This also affects the use of unknown variables, and integer comparisons with undefined or non-numeric variables: in *warn* mode the comparison evaluates to false, in *stop* mode the generation is stopped with an error. +* **What to do on invalid content warnings**: Warn on the console or stop the generation. This also affects the use of unknown variables, and integer comparisons with undefined or non-numeric variables: in *warn* mode the comparison evaluates to false, in *stop* mode the generation is stopped with an error. +* **Use strict operators**: Sets strict operations in comparisons. * **Apply in img2img**: Check if you want to do the processing in img2img processes. * **Add original prompts to metadata**: Adds original prompts to the metadata if they have changed. * **Extranetwork Mappings folders**: You can enter multiple folders separated by commas. diff --git a/docs/SYNTAX.md b/docs/SYNTAX.md index 74f154a..3aaf923 100644 --- a/docs/SYNTAX.md +++ b/docs/SYNTAX.md @@ -236,27 +236,38 @@ An echoed variable that uses the default because it doesn't have a value will be ## Array variables -There is support for array variables. They use brackets `[]` to differenciate from regular variables. Empty brackets mean the whole array (used when initializing or when echoing the whole array) and a number means an indexed value. These variables can only be set with the *Dynamic Prompts* format. +There is support for array variables. They use brackets `[]` to differenciate from regular variables. -Some examples of how to use them: +They can be initialized in several ways: | Construct | Meaning | | --------- | ------- | | `${var[]=value}` | initialize and set the first value | -| `${var[]+=}` | add an empty element to the array | -| `${var[]+=value}` | add a value to the array | | `${var[]=*()}` | initialize an empty array | | `${var[]=*var2[]}` | initialize an array from another array | | `${var[]=*__wildcard__}` | initialize an array from a wildcard | | `${var[]=*('value',var2)}` | initialize an array from a list of values (strings or variables) | +| `${var[]+=}` | add an empty element to the array | +| `${var[]+=value}` | add a value to the array | | `${var[]+=*var2[]}` | add elements from another array | | `${var[]+=*__wildcard__}` | add elements from a wildcard | + +The star operator `*` is always with inmediate evaluation, and can only be used with the *Dynamic Prompts* format. + +And can be accesed/echoed with: + +* Empty brackets mean the whole array (used when initializing or when echoing the whole array) +* An integer inside the brackets means an indexed value +* A hash inside the brackets is used to get the length of the array +* A string inside the brackets (with quotes) is used to get the full array joined with a separator. + +| Construct | Meaning | +| --------- | ------- | | `${var[]}` | echo all elements with a default separator | | `${var[' / ']}` | echo all elements with a specific separator | | `${var[n]}` | echo an element from the array | | `${var[n]:default}` | echo an element with a default | - -Initialization with the star operator `*` is always with inmediate evaluation. +| `${var[#]}` | echo the length of the array | ## If command @@ -283,20 +294,24 @@ Variable values `true` and `false` are considered a boolean, and an all digits v The operation can be preceded by `not` for readability, instead of using it in the front. -The supported operations are: `eq`, `ne`, `gt`, `lt`, `ge`, `le`, `in` and `contains`. This list shows what they do depending on the kind of operand (R = regular variable, A = array variable). +The supported operations are: `eq`, `ne`, `gt`, `lt`, `ge`, `le`, `in`, `any_in`, `contains` and `contains_any`. The `any_in` and `contains_any` variants exist because `any` and `contains` consider all elements. -| Operation | R1 op R2 | A1 op A2 | A1 op R2 | R1 op A2 | -| --------- | -------- | -------- | -------- | -------- | -| `eq` | OK | OK (pairwise) | Always False | Always False | -| `ne` | OK | OK (pairwise) | Always True | Always True | -| `gt` | OK | OK (pairwise) | Error | Error | -| `lt` | OK | OK (pairwise) | Error | Error | -| `ge` | OK | OK (pairwise) | Error | Error | -| `le` | OK | OK (pairwise) | Error | Error | -| `in` | OK (substring, R1 in R2) | OK (all A1 in A2) | OK (any substring in A1 in R2) | OK (R1 in A2) | -| `contains` | OK (substring, R2 in R1) | OK (all A2 in A1) | OK (R2 in A1) | OK (all substrings A2 in R1) | +This list shows what they do depending on the kind of operand (R = regular variable, A = array variable). -When a comparison tries to compare undefined variables or the values have different types (f.e. an integer and a string), the behavior depends on the `on_warning` setting: in `warn` mode the comparison evaluates to false, and in `stop` mode an error is raised. +| Operation | R1 op R2 | A1 op A2 | A1 op R2 | R1 op A2 | +| --------- | -------- | -------- | -------- | -------- | +| `eq` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise | +| `ne` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise | +| `gt` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise | +| `lt` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise | +| `ge` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise | +| `le` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise | +| `in` | OK (substring, R1 in R2) | OK (all A1 in A2) | OK (all substring in A1 in R2) | OK (R1 in A2) | +| `any_in` | Error | OK (any A1 in A2) | OK (any substring in A1 in R2) | Error | +| `contains` | OK (substring, R2 in R1) | OK (all A2 in A1) | OK (R2 in A1) | OK (all substrings A2 in R1) | +| `contains_any` | Error | OK (any A2 in A1) | Error | OK (any substrings A2 in R1) | + +When a comparison tries to compare undefined variables or the values have different types (f.e. an integer and a string), the behavior depends on the `on_warning` setting: in `warn` mode the comparison evaluates to false, and in `stop` mode an error is raised. In non strict mode a numeric string literal (with no leading zeros) will be considered an integer. The variable can be one set with the `set` command (user variables) or you can use system variables like these (names starting with an underscore are reserved for system variables): diff --git a/grammar.lark b/grammar.lark index 9c65af6..c117d6c 100644 --- a/grammar.lark +++ b/grammar.lark @@ -94,7 +94,7 @@ extranetworktag: "<" /(?!ppp:)\w+:/ encontent ">" ?encontent.3: content_en //#if ALLOW_COMMVARS or ALLOW_CHOICES or ALLOW_WILDCARDS - vardescriptor_get.8: VARNAME [ /\[/ [ SIGNED_NUMBER | STRING | IDENTIFIER ] /\]/ ] + vardescriptor_get.8: VARNAME [ /\[/ [ SIGNED_NUMBER | IDENTIFIER | STRING | /#/ ] /\]/ ] vardescriptor_set.9: VARNAME [ /\[/ [ SIGNED_NUMBER | IDENTIFIER ] /\]/ ] // conditions @@ -108,7 +108,7 @@ extranetworktag: "<" /(?!ppp:)\w+:/ encontent ">" operation_not: "not" ( ( _WHITESPACE ungrouped_condition ) | ( _WHITESPACE? grouped_condition ) ) ?complexvalue: vardescriptor_get | SIMPLEVALUE truthy_operand: vardescriptor_get - comparison: ( complexvalue | listvalue ) _WHITESPACE ( /not/ _WHITESPACE )? /eq|ne|lt|gt|le|ge|contains|in/ _WHITESPACE ( complexvalue | listvalue ) + comparison: ( complexvalue | listvalue ) _WHITESPACE ( /not/ _WHITESPACE )? /eq|ne|lt|gt|le|ge|contains|in|any_in|contains_any/ _WHITESPACE ( complexvalue | listvalue ) listvalue.9: "(" ( _WHITESPACE? complexvalue ( _WHITESPACE? "," _WHITESPACE? complexvalue )* )? _WHITESPACE? ")" //#endif @@ -125,7 +125,7 @@ extranetworktag: "<" /(?!ppp:)\w+:/ encontent ">" ifvalue.3: content // command: set - commandset: "" commandsetcontent "" + commandset: "" ( starredvalue | commandsetcontent ) "" commandsetmodifiers: (_WHITESPACE /evaluate|ifundefined|add/ )+ ?commandsetcontent.3: content diff --git a/ppp.py b/ppp.py index 0fde8fe..f6461d7 100644 --- a/ppp.py +++ b/ppp.py @@ -81,6 +81,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in DEFAULT_CUP_EXTRANETWORK_TAGS = defopt["cup_extranetwork_tags"] DEFAULT_CUP_MERGE_ATTENTION = defopt["cup_merge_attention"] DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS = defopt["cup_remove_extranetwork_tags"] + DEFAULT_STRICT_OPERATORS = defopt["strict_operators"] WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK ' WILDCARD_STOP = "INVALID WILDCARD! {0}\nBREAK " UNPROCESSED_STOP = "UNPROCESSED CONSTRUCTS!\nBREAK " diff --git a/ppp_classes.py b/ppp_classes.py index c555968..29cf90d 100644 --- a/ppp_classes.py +++ b/ppp_classes.py @@ -199,6 +199,7 @@ class PPPStateOptions: cup_extranetwork_tags: bool = False cup_merge_attention: bool = True cup_remove_extranetwork_tags: bool = False + strict_operators: bool = True @dataclass(frozen=True) diff --git a/ppp_comfyui.py b/ppp_comfyui.py index e4c82ab..29edaef 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -144,6 +144,15 @@ class PromptPostProcessorComfyUINode: "tooltip": "How to handle invalid content warnings", }, ), + "strict_operators": ( + "BOOLEAN", + { + "default": PromptPostProcessor.DEFAULT_STRICT_OPERATORS, + "tooltip": "Use strict operators", + "label_on": "Yes", + "label_off": "No", + }, + ), "process_wildcards": ( "BOOLEAN", { @@ -250,6 +259,7 @@ class PromptPostProcessorComfyUINode: stn_options=None, cup_options=None, en_options=None, + strict_operators=None, ): modelclass = ( model.model.model_config.__class__.__name__ if model is not None and not isinstance(model, str) else model @@ -282,6 +292,7 @@ class PromptPostProcessorComfyUINode: options = PPPStateOptions( debug_level=DEBUG_LEVEL(debug_level), on_warning=ONWARNING_CHOICES(on_warnings) if on_warnings else PromptPostProcessor.DEFAULT_ON_WARNING, + strict_operators=strict_operators if strict_operators is not None else PromptPostProcessor.DEFAULT_STRICT_OPERATORS, process_wildcards=process_wildcards, if_wildcards=(wc_options["wc_if_wildcards"] if wc_options else IFWILDCARDS_CHOICES.stop.value), choice_separator=( diff --git a/ppp_tree.py b/ppp_tree.py index 1935ef2..eda1aff 100644 --- a/ppp_tree.py +++ b/ppp_tree.py @@ -146,27 +146,29 @@ class TreeProcessor(lark.visitors.Interpreter): """ return node.meta.content if hasattr(node, "meta") and node.meta is not None and not node.meta.empty else default - def __idxsep_to_idx_sep(self, idxsep: str | None) -> tuple[Optional[int], Optional[str]]: + def __parse_idxsep(self, idxsep: str | None) -> tuple[Optional[int], Optional[str], Optional[bool]]: """ - Convert an index or separator string to an index and a separator. + Convert an index/separator/count string to a specific type. Args: idxsep (str|None): The index or separator string. Returns: - tuple: A tuple containing the index and separator. + tuple: A tuple containing the index, separator and count boolean. """ if idxsep is None: - return None, None + return None, None, None is_quoted = idxsep.startswith(("'", '"')) and idxsep.endswith(("'", '"')) - if not idxsep.isdigit() and not is_quoted: + if idxsep == "#": # special value to indicate length of the array variable + return None, None, True + if not idxsep.isdecimal() and not is_quoted: # bare identifier: resolve as variable idxsep = self.get_final_user_variable(idxsep) is_quoted = False - if idxsep.isdigit(): - return int(idxsep), None + if idxsep.isdecimal(): + return int(idxsep), None, False # separator: strip surrounding quotes if present - return None, idxsep[1:-1] if is_quoted else idxsep + return None, idxsep[1:-1] if is_quoted else idxsep, False def __get_user_variable_value( self, name: str, idxsep: str | None = None, evaluate=True, visit=False @@ -192,6 +194,8 @@ class TreeProcessor(lark.visitors.Interpreter): 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 @@ -202,8 +206,12 @@ class TreeProcessor(lark.visitors.Interpreter): is_array = name[-2:] == "[]" if is_array: if isinstance(v, list): - idx, sep = self.__idxsep_to_idx_sep(idxsep) - if idx is not None: + idx, sep, cnt = self.__parse_idxsep(idxsep) + 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: @@ -255,11 +263,11 @@ class TreeProcessor(lark.visitors.Interpreter): name, idxsep = self.__separate_arrayref(nameidx) v = self.__get_user_variable_value(name, idxsep, True, False) if isinstance(v, list): - _, sep = self.__idxsep_to_idx_sep(idxsep) + _, sep, _ = self.__parse_idxsep(idxsep) if sep is None: sep = self.state.options.choice_separator v = sep.join(str(item) for item in v) - return v + return str(v) def __set_user_variable_value(self, name: str, value: str | lark.Tree | list): """ @@ -298,22 +306,26 @@ class TreeProcessor(lark.visitors.Interpreter): 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: + 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 digit and does not have leading zeros (unless it's just "0"), it is converted to an integer. + If it is a number, it is converted to an integer or float. Args: s (str): The string to adjust. Returns: - str | int: The adjusted string or integer. + str | int | float: The adjusted string, integer, or float. """ - if s.isdigit() and not (s.startswith("0") and len(s) > 1): + 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): + def __warn_mixedtype(self, desc: str, operand1, operand2, operation): """ Warn the user if mixed type values are used in a comparison. @@ -327,15 +339,15 @@ class TreeProcessor(lark.visitors.Interpreter): self.warn_or_stop(f"Undefined value used in comparison: '{escape_single_quotes(desc)}'") return False if ( - isinstance(operand1, (str, int, bool)) - and isinstance(operand2, (str, int, bool)) + 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: + def __resolve_operand(self, c: str) -> str | bool | int | float: """ Resolve an operand value. @@ -343,12 +355,16 @@ class TreeProcessor(lark.visitors.Interpreter): c (str): The operand value to resolve. Returns: - str | bool | int: The resolved operand value (in lowercase for strings). + 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() # self.__adjust_strnum(c[1:-1]) - if c.isdigit(): + 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": @@ -365,16 +381,37 @@ class TreeProcessor(lark.visitors.Interpreter): val = "" self.warn_or_stop(f"Unknown {vartype} variable '{escape_single_quotes(c)}'") if isinstance(val, str): - if val.isdigit(): - val = int(val) - else: - val = self.__adjust_strnum(val) - if val in ("false", ""): - return False - if val == "true": - return True + 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, @@ -408,64 +445,144 @@ class TreeProcessor(lark.visitors.Interpreter): if operator == "truthy": result = bool(operand1_value) else: - def pairwise_all(op): - return all(op(a, b) for a, b in zip(operand1_value, operand2_value)) - - def wmt(a, b, op): - return self.warn_mixedtype(cond_desc, a, b, op) - if not operand1_isarray and not operand2_isarray: operations = { - "eq": lambda: wmt(operand1_value, operand2_value, lambda x, y: x == y), - "ne": lambda: wmt(operand1_value, operand2_value, lambda x, y: x != y), - "gt": lambda: wmt(operand1_value, operand2_value, lambda x, y: x > y), - "lt": lambda: wmt(operand1_value, operand2_value, lambda x, y: x < y), - "ge": lambda: wmt(operand1_value, operand2_value, lambda x, y: x >= y), - "le": lambda: wmt(operand1_value, operand2_value, lambda x, y: x <= y), - "in": lambda: wmt(str(operand1_value), str(operand2_value), lambda x, y: x in y), - "contains": lambda: wmt(str(operand1_value), str(operand2_value), lambda x, y: y in x), + "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 pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x == y)), + and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x == y), "ne": lambda: len(operand1_value) != len(operand2_value) - or pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x != y)), + or self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x != y), "gt": lambda: len(operand1_value) == len(operand2_value) - and pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x > y)), + and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x > y), "lt": lambda: len(operand1_value) == len(operand2_value) - and pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x < y)), + and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x < y), "ge": lambda: len(operand1_value) == len(operand2_value) - and pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x >= y)), + and self.__pairwise_all(cond_desc, operand1_value, operand2_value, lambda x, y: x >= y), "le": lambda: len(operand1_value) == len(operand2_value) - and pairwise_all(lambda a, b: wmt(a, b, lambda x, y: x <= y)), - "in": lambda: all(wmt(a, operand2_value, lambda x, y: x in y) for a in operand1_value), - "contains": lambda: all(wmt(a, operand1_value, lambda x, y: x in y) for a in 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: - operations = { - "eq": lambda: False, - "ne": lambda: True, - # "gt": lambda: False, - # "lt": lambda: False, - # "ge": lambda: False, - # "le": lambda: False, - "in": lambda: any(wmt(str(a), str(operand2_value), lambda x, y: x in y) for a in operand1_value), - "contains": lambda: wmt(operand1_value, operand2_value, lambda x, y: y in x), - } + 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: - operations = { - "eq": lambda: False, - "ne": lambda: True, - # "gt": lambda: False, - # "lt": lambda: False, - # "ge": lambda: False, - # "le": lambda: False, - "in": lambda: wmt(operand1_value, operand2_value, lambda x, y: x in y), - "contains": lambda: all( - wmt(str(operand1_value), str(a), lambda x, y: y in x) for a in operand2_value - ), - } + 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) @@ -565,7 +682,11 @@ class TreeProcessor(lark.visitors.Interpreter): # 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) + ")" + 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]) @@ -1034,19 +1155,25 @@ class TreeProcessor(lark.visitors.Interpreter): # default_value = self.__visit(default, True) # for log is_array = variable_name[-2:] == "[]" vname = f"{variable_name[0:-2]}[{variable_idxsep}]" if variable_idxsep is not None else variable_name - value = self.__get_user_variable_value(variable_name, variable_idxsep, True, True) + 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_idxsep, 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") - v = self.__visit(default, False, True) - self.result += v - default_value = v - self.state.echoed_variables[vname] = v + 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 = "" - self.state.echoed_variables[vname] = "" - else: + value = "" + if not is_systemvar: self.state.echoed_variables[vname] = value t2 = time.monotonic_ns() info = variable_name diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py index 5cf5b03..b1e6d20 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -179,6 +179,7 @@ class PromptPostProcessorA1111Script(scripts.Script): options = PPPStateOptions( debug_level=DEBUG_LEVEL(getattr(opts, "ppp_gen_debug_level", PromptPostProcessor.DEFAULT_DEBUG_LEVEL)), on_warning=ONWARNING_CHOICES(getattr(opts, "ppp_gen_onwarning", PromptPostProcessor.DEFAULT_ON_WARNING)), + strict_operators=getattr(opts, "ppp_gen_strict_operators", PromptPostProcessor.DEFAULT_STRICT_OPERATORS), process_wildcards=getattr(opts, "ppp_wil_processwildcards", PromptPostProcessor.DEFAULT_PROCESS_WILDCARDS), if_wildcards=IFWILDCARDS_CHOICES( getattr(opts, "ppp_wil_ifwildcards", PromptPostProcessor.DEFAULT_IF_WILDCARDS) @@ -300,17 +301,18 @@ class PromptPostProcessorA1111Script(scripts.Script): self.wildcards_obj, self.extranetwork_mappings_obj, ) - hash_options = ppp.options_hash() - hash_envinfo = ppp.envinfo_hash() - prompts_list = [] + hash_fullenv = hash( + (ppp.envinfo_hash(), ppp.options_hash(), self.wildcards_obj, self.extranetwork_mappings_obj) + ) if input_force_equal_seeds: log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "Forcing equal seeds") - seeds = getattr(p, "all_seeds", []) - subseeds = getattr(p, "all_subseeds", []) + seeds: list[int] = getattr(p, "all_seeds", []) + subseeds: list[int] = getattr(p, "all_subseeds", []) p.all_seeds = [seeds[0] for _ in seeds] p.all_subseeds = [subseeds[0] for _ in subseeds] + calculated_seeds: list[int] = [] if input_unlink_seed: log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "Using unlinked seed") num_seeds = len(getattr(p, "all_seeds", [])) @@ -322,9 +324,9 @@ class PromptPostProcessorA1111Script(scripts.Script): else: calculated_seeds = [input_seed for _ in range(num_seeds)] else: - seeds = getattr(p, "all_seeds", []) - subseeds = getattr(p, "all_subseeds", []) - subseed_strength = getattr(p, "subseed_strength", 0.0) + seeds: list[int] = getattr(p, "all_seeds", []) + subseeds: list[int] = getattr(p, "all_subseeds", []) + subseed_strength: float = getattr(p, "subseed_strength", 0.0) if subseed_strength > 0: calculated_seeds = [ int(subseed * subseed_strength + seed * (1 - subseed_strength)) @@ -336,93 +338,80 @@ class PromptPostProcessorA1111Script(scripts.Script): else: calculated_seeds = seeds - # initialize extra generation parameters - extra_params = {} + # (prompt type, typeindex) -> (new positive prompt, new negative prompt) + prompts_list: dict[tuple[str, int], tuple[str, str]] = {} - # adds regular prompts + # adds prompts + regular_type = "regular" rpr: list[str] = getattr(p, "all_prompts", None) rnr: list[str] = getattr(p, "all_negative_prompts", None) - if rpr is not None and rnr is not None: - prompts_list += [ - ("regular", seed, prompt, negative_prompt) - for seed, prompt, negative_prompt in zip(calculated_seeds, rpr, rnr) - if (seed, prompt, negative_prompt) not in prompts_list - ] - # make it compatible with A1111 hires fix + regular_exists = rpr is not None and rnr is not None + hiresfix_type = "hiresfix" rph: list[str] = getattr(p, "all_hr_prompts", None) rnh: list[str] = getattr(p, "all_hr_negative_prompts", None) - if rph is not None and rnh is not None and (rph != rpr or rnh != rnr): - prompts_list += [ - ("hiresfix", seed, prompt, negative_prompt) - for seed, prompt, negative_prompt in zip(calculated_seeds, rph, rnh) - if (seed, prompt, negative_prompt) not in prompts_list - ] + hiresfix_exists = rph is not None and rnh is not None + for i in range(len(calculated_seeds)): + if regular_exists: + prompts_list[(regular_type, i)] = None + if hiresfix_exists: + prompts_list[(hiresfix_type, i)] = None # processes prompts - for i, (prompttype, seed, prompt, negative_prompt) in enumerate(prompts_list): - log(self.ppp_logger, self.ppp_debug_level, logging.INFO, f"processing prompts[{i+1}] ({prompttype})") - if ( - self.lru_cache.get( - (hash_envinfo, hash_options, seed, hash(self.wildcards_obj), prompt, negative_prompt) - ) - is None - ): + for prompttype, typeindex in prompts_list.keys(): + log(self.ppp_logger, self.ppp_debug_level, logging.INFO, f"processing prompts ({prompttype}[{typeindex+1}])") + key = ( + (hash_fullenv, calculated_seeds[typeindex], rpr[typeindex], rnr[typeindex]) + if prompttype == regular_type + else (hash_fullenv, calculated_seeds[typeindex], rph[typeindex], rnh[typeindex]) + ) + cached = self.lru_cache.get(key) + if cached is None: + (hsh, seed, prompt, negative_prompt) = key posp, negp, _ = ppp.process_prompt(prompt, negative_prompt, seed) - self.lru_cache.put( - (hash_envinfo, hash_options, seed, hash(self.wildcards_obj), prompt, negative_prompt), (posp, negp) - ) + cached = (posp, negp) + self.lru_cache.put(key, cached) # adds also the result so i2i doesn't process it unnecessarily - self.lru_cache.put( - (hash_envinfo, hash_options, seed, hash(self.wildcards_obj), posp, negp), (posp, negp) - ) + self.lru_cache.put((hsh, seed, posp, negp), cached) else: log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache") + prompts_list[(prompttype, typeindex)] = cached + + # with open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "last_prompts.txt"), "w", encoding="utf-8") as f: + # for (prompttype, typeindex), (posp, negp) in prompts_list.items(): + # f.write(f"Key: {prompttype} {typeindex}\n") + # f.write(f"Seed: {calculated_seeds[typeindex]}\n") + # f.write(f"Old Positive: {rpr[typeindex] if prompttype == regular_type else rph[typeindex]}\n") + # f.write(f"Old Negative: {rnr[typeindex] if prompttype == regular_type else rnh[typeindex]}\n") + # f.write(f"New Positive: {posp}\n") + # f.write(f"New Negative: {negp}\n") + # f.write("\n") # updates the prompts - rpr_copy = None - rnr_copy = None - if rpr is not None and rnr is not None: - rpr_changes = False - rnr_changes = False - rpr_copy = rpr.copy() - rnr_copy = rnr.copy() - for i, (seed, prompt, negative_prompt) in enumerate(zip(calculated_seeds, rpr, rnr)): - found = self.lru_cache.get( - (hash_envinfo, hash_options, seed, hash(self.wildcards_obj), prompt, negative_prompt) - ) - if found is not None: - if rpr[i].strip() != found[0].strip(): - rpr_changes = True - if rnr[i].strip() != found[1].strip(): - rnr_changes = True - rpr[i] = found[0] - rnr[i] = found[1] - if add_prompts: - if rpr_changes: - extra_params["PPP original prompts"] = rpr_copy - if rnr_changes: - extra_params["PPP original negative prompts"] = rnr_copy - if rph is not None and rnh is not None: - rph_changes = False - rnh_changes = False - rph_copy = rph.copy() - rnh_copy = rnh.copy() - for i, (seed, prompt, negative_prompt) in enumerate(zip(calculated_seeds, rph, rnh)): - found = self.lru_cache.get( - (hash_envinfo, hash_options, seed, hash(self.wildcards_obj), prompt, negative_prompt) - ) - if found is not None: - if rph[i].strip() != found[0].strip() and (not rpr_copy or rph[i].strip() != rpr_copy[i].strip()): - rph_changes = True - if rnh[i].strip() != found[1].strip() and (not rnr_copy or rnh[i].strip() != rnr_copy[i].strip()): - rnh_changes = True - rph[i] = found[0] - rnh[i] = found[1] - if add_prompts: - if rph_changes: - extra_params["PPP original HR prompts"] = rph_copy - if rnh_changes: - extra_params["PPP original HR negative prompts"] = rnh_copy + regular_copy = (rpr.copy() if rpr else None, rnr.copy() if rnr else None) + hiresfix_copy = (rph.copy() if rph else None, rnh.copy() if rnh else None) + regular_changes = False + hiresfix_changes = False + for (prompttype, typeindex), (posp, negp) in prompts_list.items(): + if prompttype == regular_type: + if rpr[typeindex].strip() != posp.strip() or rnr[typeindex].strip() != negp.strip(): + regular_changes = True + rpr[typeindex] = posp + rnr[typeindex] = negp + elif prompttype == hiresfix_type: + if rph[typeindex].strip() != posp.strip() or rnh[typeindex].strip() != negp.strip(): + hiresfix_changes = True + rph[typeindex] = posp + rnh[typeindex] = negp + + # initialize extra generation parameters + extra_params = {} + if add_prompts: + if regular_changes: + extra_params["PPP original prompts"] = regular_copy[0] + extra_params["PPP original negative prompts"] = regular_copy[1] + if hiresfix_changes: + extra_params["PPP original HR prompts"] = hiresfix_copy[0] + extra_params["PPP original HR negative prompts"] = hiresfix_copy[1] # fill extra generation parameters only if not already present for k, v in extra_params.items(): @@ -508,7 +497,7 @@ def on_ui_settings(): key="ppp_gen_onwarning", info=shared.OptionInfo( default=ONWARNING_CHOICES.warn.value, - label="What to do on invalid content warnings?", + label="What to do on invalid content warnings", component=gr.Radio, component_args={ "choices": ( @@ -519,6 +508,14 @@ def on_ui_settings(): section=section, ), ) + shared.opts.add_option( + key="ppp_gen_strict_operators", + info=shared.OptionInfo( + default=PromptPostProcessor.DEFAULT_STRICT_OPERATORS, + label="Use strict operators", + section=section, + ), + ) shared.opts.add_option( key="ppp_gen_doi2i", info=shared.OptionInfo( diff --git a/tests/base_tests.py b/tests/base_tests.py index acee10d..a3bf851 100644 --- a/tests/base_tests.py +++ b/tests/base_tests.py @@ -46,6 +46,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): self.defopts = PPPStateOptions( debug_level=DEBUG_LEVEL.full, on_warning=ONWARNING_CHOICES.stop, + strict_operators=True, process_wildcards=True, if_wildcards=IFWILDCARDS_CHOICES.ignore, choice_separator=", ", @@ -154,6 +155,19 @@ class TestPromptPostProcessorBase(unittest.TestCase): self.wildcards_obj, self.extranetwork_maps_obj, ) + elif ppp == "nostrict": + the_obj = PromptPostProcessor( + self.ppp_logger, + self.def_env_info, + replace( + self.defopts, + strict_operators=False, + ), + self.grammar_content, + self.interrupt, + self.wildcards_obj, + self.extranetwork_maps_obj, + ) else: the_obj = ppp if not the_obj: diff --git a/tests/tests_varcomms.py b/tests/tests_varcomms.py index 5208b88..ebb195c 100644 --- a/tests/tests_varcomms.py +++ b/tests/tests_varcomms.py @@ -185,6 +185,17 @@ class TestVarCommands(TestPromptPostProcessorBase): variables={"v1[]": "one, two, three, four", "v2": "three"}, ) + def test_array_variable_9(self): # array variable length + self.process( + PromptPair( + "${v1[]=val1}${v1[]+=val2}${v1[]+=val3}${v1[#]:defval}, OKnot OK", + "", + ), + PromptPair("3, OK", ""), + variables={"v1[]": "val1, val2, val3", "v1[#]": "3"}, + ) + + # Operator tests ## R vs R @@ -219,7 +230,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RltR(self): self.process( PromptPair( - "${r1=01}${r2=2}OKnot OK", + "${r1=1}${r2=2}OKnot OK", "", ), PromptPair("OK", ""), @@ -228,7 +239,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RgtR(self): self.process( PromptPair( - "${r1=2}${r2=01}OKnot OK", + "${r1=2}${r2=1}OKnot OK", "", ), PromptPair("OK", ""), @@ -237,7 +248,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RleR(self): self.process( PromptPair( - "${r1=01}${r2=1}OKnot OK", + "${r1=1}${r2=1}OKnot OK", "", ), PromptPair("OK", ""), @@ -246,7 +257,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_RgeR(self): self.process( PromptPair( - "${r1=1}${r2=01}OKnot OK", + "${r1=1}${r2=1}OKnot OK", "", ), PromptPair("OK", ""), @@ -311,7 +322,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AltA_1(self): self.process( PromptPair( - "${a1[]=*(01,2,3)}${a2[]=*(2,3,4)}OKnot OK", + "${a1[]=*(1,2,3)}${a2[]=*(2,3,4)}OKnot OK", "", ), PromptPair("OK", ""), @@ -320,7 +331,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AltA_2(self): self.process( PromptPair( - "${a1[]=*(01,2)}${a2[]=*(2,3,4)}OKnot OK", + "${a1[]=*(1,2)}${a2[]=*(2,3,4)}OKnot OK", "", ), PromptPair("not OK", ""), @@ -329,7 +340,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AgtA(self): self.process( PromptPair( - "${a1[]=*(2,3,4)}${a2[]=*(01,2,3)}OKnot OK", + "${a1[]=*(2,3,4)}${a2[]=*(1,2,3)}OKnot OK", "", ), PromptPair("OK", ""), @@ -338,7 +349,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AleA(self): self.process( PromptPair( - "${a1[]=*(01,2)}${a2[]=*(1,3)}OKnot OK", + "${a1[]=*(1,2)}${a2[]=*(1,3)}OKnot OK", "", ), PromptPair("OK", ""), @@ -347,7 +358,7 @@ class TestVarCommands(TestPromptPostProcessorBase): def test_operator_AgeA(self): self.process( PromptPair( - "${a1[]=*(1,3)}${a2[]=*(01,2)}OKnot OK", + "${a1[]=*(1,3)}${a2[]=*(1,2)}OKnot OK", "", ), PromptPair("OK", ""), @@ -380,6 +391,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "", ), PromptPair("not OK", ""), + ppp="nostrict", ) def test_operator_AneR(self): @@ -389,16 +401,17 @@ class TestVarCommands(TestPromptPostProcessorBase): "", ), PromptPair("OK", ""), + ppp="nostrict", ) def test_operator_AltR(self): self.process( PromptPair( - "${a1[]=*(01,2,3)}${r2=2}OKnot OK", + "${a1[]=*(1,2,3)}${r2=2}OKnot OK", "", ), PromptPair("not OK", ""), - interrupted=True, + ppp="nostrict", ) def test_operator_AgtR(self): @@ -408,17 +421,17 @@ class TestVarCommands(TestPromptPostProcessorBase): "", ), PromptPair("not OK", ""), - interrupted=True, + ppp="nostrict", ) def test_operator_AleR(self): self.process( PromptPair( - "${a1[]=*(01,2)}${r2=2}OKnot OK", + "${a1[]=*(1,2)}${r2=2}OKnot OK", "", ), - PromptPair("not OK", ""), - interrupted=True, + PromptPair("OK", ""), + ppp="nostrict", ) def test_operator_AgeR(self): @@ -428,7 +441,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "", ), PromptPair("not OK", ""), - interrupted=True, + ppp="nostrict", ) def test_operator_AinR(self): @@ -458,6 +471,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "", ), PromptPair("not OK", ""), + ppp="nostrict", ) def test_operator_RneA(self): @@ -467,16 +481,17 @@ class TestVarCommands(TestPromptPostProcessorBase): "", ), PromptPair("OK", ""), + ppp="nostrict", ) def test_operator_RltA(self): self.process( PromptPair( - "${r1=2}${a2[]=*(01,2,3)}OKnot OK", + "${r1=2}${a2[]=*(1,2,3)}OKnot OK", "", ), PromptPair("not OK", ""), - interrupted=True, + ppp="nostrict", ) def test_operator_RgtA(self): @@ -486,17 +501,17 @@ class TestVarCommands(TestPromptPostProcessorBase): "", ), PromptPair("not OK", ""), - interrupted=True, + ppp="nostrict", ) def test_operator_RleA(self): self.process( PromptPair( - "${r1=2}${a2[]=*(01,2)}OKnot OK", + "${r1=2}${a2[]=*(1,2)}OKnot OK", "", ), PromptPair("not OK", ""), - interrupted=True, + ppp="nostrict", ) def test_operator_RgeA(self): @@ -506,7 +521,7 @@ class TestVarCommands(TestPromptPostProcessorBase): "", ), PromptPair("not OK", ""), - interrupted=True, + ppp="nostrict", ) def test_operator_RinA(self): @@ -922,6 +937,15 @@ class TestVarCommands(TestPromptPostProcessorBase): PromptPair("this test is OK", ""), ) + def test_cmd_echo_sysvar(self): + self.process( + PromptPair( + "${_model:defval}", + "", + ), + PromptPair("sdxl", ""), + ) + def test_cmd_ext(self): # ext self.process( PromptPair(