From 54a20117ce0b78dcd500a3af9da8d5e788e979ec Mon Sep 17 00:00:00 2001 From: Antonio Cordero Balcazar Date: Mon, 1 Sep 2025 16:50:46 +0200 Subject: [PATCH] * New extranetworks command to group the triggers with the extranetwork. * Mapping of extranetworks through the new command. This allows you to use a "virtual" extranetwork that is converted to a real one depending on conditions, like the model kind/variant. * Fixed a parsing bug with the if command. * Some refactoring and additional checks. --- README.md | 1 + docs/CONFIG.md | 4 +- docs/SYNTAX.md | 78 +++++++++- grammar.lark | 42 ++++-- ppp.py | 237 +++++++++++++++++++++++++----- ppp_comfyui.py | 67 ++++++++- ppp_enmappings.py | 257 +++++++++++++++++++++++++++++++++ ppp_utils.py | 17 +++ ppp_wildcards.py | 20 +-- scripts/ppp_script.py | 30 +++- tests/enmappings/mappings.yaml | 11 ++ tests/tests.py | 139 +++++++++++++++++- 12 files changed, 820 insertions(+), 83 deletions(-) create mode 100644 ppp_enmappings.py create mode 100644 ppp_utils.py create mode 100644 tests/enmappings/mappings.yaml diff --git a/README.md b/README.md index 0ce737d..fbc321b 100644 --- a/README.md +++ b/README.md @@ -15,6 +15,7 @@ Currently this extension has these functions: * Set and modify local variables. * Filter content based on the loaded SD model or a variable. * Process wildcards. Compatible with Dynamic Prompts formats. Can also detect invalid wildcards and act as you choose. +* Map extranetworks depending on conditions (like the loaded model variant). * Clean up the prompt and negative prompt. Note: when used in an *A1111* compatible webui, the extension must be loaded after any other extension that modifies the prompt (like another wildcards extension). Usually extensions load by their folder name in alphanumeric order, so if the extensions are not loading in the correct order just rename this extension's folder so the ordering works out. When in doubt, just rename this extension's folder with a "z" in front (for example) so that it is the last one to load, or manually set such folder name when installing it. diff --git a/docs/CONFIG.md b/docs/CONFIG.md index 69078ba..1e47c11 100644 --- a/docs/CONFIG.md +++ b/docs/CONFIG.md @@ -10,6 +10,7 @@ * **pos_prompt**: Connect here the prompt text, or fill it as a widget. * **neg_prompt**: Connect here the negative prompt text, or fill it as a widget. * **wc_wildcards_input**: Wildcards definitions (in yaml or json format). Direct input added to the ones found in the wildcards folders. Allows wildcards to be included in the workflow. +* **en_mappings_input**: Extranetwork Mappings definitions (in yaml format). Direct input added to the ones found in the extranetwork mappings folders. Allows the mappings to be included in the workflow. Other common settings (see [below](#common-settings)) also appear as inputs or widgets. @@ -43,11 +44,12 @@ With this prompt: `__quality__, 1girl, ${head:__eyes__, __hair__, __expression__ The default value defines strings for *Pony* and *Illustrious* models. * **Apply in img2img**: check if you want to do the processing in img2img processes (*does not apply to the ComfyUI node*). * **Add original prompts to metadata**: adds original prompts to the metadata if they have changed (*does not apply to the ComfyUI node*). +* **Extranetwork Mappings folders**: you can enter multiple folders separated by commas. In *ComfyUI* you can leave it empty and add a `ppp_extranetworkmappings` entry in the **extra_model_paths.yaml** file. ### Wildcard settings * **Process wildcards**: you can choose to process wildcards and choices with this extension or use a different one. -* **Wildcards folders**: you can enter multiple folders separated by commas. In *ComfyUI* you can leave it empty and add a `wildcards` entry in the **extra_model_paths.yaml** file. +* **Wildcards folders**: you can enter multiple folders separated by commas. In *ComfyUI* you can leave it empty and add a `ppp_wildcards` or `wildcards` entry in the **extra_model_paths.yaml** file. * **What to do with remaining wildcards?**: select what do you want to do with any found wildcards/choices (when process wildcards is off or after the processing). * **Ignore**: do not try to detect wildcards. * **Remove**: detect wildcards and remove them. diff --git a/docs/SYNTAX.md b/docs/SYNTAX.md index 528e656..16b3f4e 100644 --- a/docs/SYNTAX.md +++ b/docs/SYNTAX.md @@ -26,14 +26,14 @@ The construct parameters can be written with the following options (all are opti * "**r**": means it allows repetition of the choices. * "**n**" or "**n-m**" or "**n-**" or "**-m**": number or range of choices to select. Allows zero as the start of a range. Default is 1. * "**$$sep**": separator when multiple choices are selected. Default is set in settings. -* "**$$**": end of the parameters. +* "**$$**": end of the parameters (not optional if any parameters). The choice options are as follows: * "**'identifiers'**": comma separated labels for the choice (optional, quotes can be single or double). Only makes sense inside a wildcard definition. Can be used when specifying the wildcard to select this specific choice. It's case insensitive. * "**n**": weight of the choice (optional, default 1). * "**if condition**": filters out the choice if the condition is false (optional; this is an extension to the *Dynamic Prompts* syntax). Same conditions as in the `if` command. -* "**::**": end of choice options +* "**::**": end of choice options (not optional if any options) Whitespace is allowed between parameters/options. @@ -221,6 +221,80 @@ The variable can be one set with the `set` or `add` commands (user variables) or Only one of the options will end up in the prompt, depending on the loaded model. +## ExtraNetwork command + +This command is a shortcut to add an extranetwork (usually a lora), and its triggers, with conditions. More legible and sometimes shorter than adding regular extranetworks inside if commands. + +The full format is: + +`[triggers]` +`` + +The `type` is the kind of extranetwork, like `lora` or `hypernet`. + +The `name` is the extranetwork identifier. If it is not a regular identifier (i.e. starts with a number or contains spaces or symbols) it should be inside quotes. + +The `parameters` is optional and its format depends on the extranetwork type. With loras or hypernets it is usually a single weight number, so if the type is one of those and there are no parameters it will default to `1`. If it is not a number it should go inside quotes. + +The `condition` uses the same format as in the `if` command, and it is also optional. + +The `triggers` are also optional, and can be any content. If there are no triggers the command ending can be omitted. + +If the condition passes (or if there is no condition) the extranetwork tag will be built and added to the result along with any triggers. + +### Examples + +(multiline to be easier to read) + +```text +test sd1x +test pony + +test sdxl +``` + +Will turn into one of these (or none) depending on the model: + +* `test sd1x` +* `test pony` +* `` +* `test sdxl` + +### Extranetworks mappings + +The extranetwork command supports specifying mappings of extranetworks, so a different lora can be used depending on the loaded model. + +If the type of extranetwork is prefixed with a `$` the command will look for a mapping. + +The mappings are configured in yaml files in any of the configured extranetwork mappings folders. The format is like this: + +```yaml +extnettype: + mappingname: + - condition: "" + name: "" + parameters: "" + triggers: [] + ... +``` + +Used like this: + +```text + +inline triggers +``` + +Each mapping can have any number of elements in its list of mappings. There are no mandatory properties for a mapping. The properties mean the following: + +* `condition`: the condition to check for this mapping to be used (usually it should be one of the `_is_*` variables). If the conditions of multiple mappings evaluate to True, one will be chosen randomly. If the condition is missing it is considered True, to be used in the last mapping to catch as an "else" condition, and will be used if no other mapping applies. +* `name`: name of the real extranetwork. If it is missing no extranetwork tag will be added. +* `parameters`: parameters for the real extranetwork. If it is missing it is assumed "1" for loras and hypernets. If both this parameter and the parameter in the ext command are numbers they are multiplied for the result. In other case the parameter of the ext command, if it exists, is used. +* `triggers`: list of trigger strings. If it is missing, only the inline triggers in the ext command will be added. +* `weight`: weight for this variant, in case multiple of them apply, to choose one. Default is 1. + +See the file in the tests folder as an example. + ## Sending content to the negative prompt The new format for this command is like this: diff --git a/grammar.lark b/grammar.lark index ce752b9..e3cf085 100644 --- a/grammar.lark +++ b/grammar.lark @@ -6,7 +6,7 @@ BOOLEAN: /true|false/i WILDCARD_NAME: /(?:(?!__|\$\$|[('"])\S)+/ INDEX: INT | IDENTIFIER IDENTIFIER: CNAME -SIMPLEVALUE: STRING | NUMBER | BOOLEAN +SIMPLEVALUE: STRING | SIGNED_NUMBER | BOOLEAN // plain text and weights ?plain: /((?!__|\bAND\b|\${)[^\\()\[\]:<>${]|\\.)+/s // exclude only the starting ones @@ -37,14 +37,14 @@ promptcomppart: content ?content_negtag.2: ( old_content | new_content_negtag | plain | specialchars_negtag )* ?content_alternate.2: ( old_content | new_content | plain_alternate | specialchars_alternate )* //#if ALLOW_WILDCARDS ALLOW_CHOICES ALLOW_COMMVARS - ?new_content.3: ( variableset | variableuse | commandstn | commandstni | commandset | commandecho | commandif | wildcard | choices )+ - ?new_content_negtag.3: ( variableset | variableuse | commandset | commandecho | commandif | wildcard | choices )+ + ?new_content.3: ( variableset | variableuse | commandstn | commandstni | commandset | commandecho | commandif | commandext | wildcard | choices )+ + ?new_content_negtag.3: ( variableset | variableuse | commandset | commandecho | commandif | commandext | wildcard | choices )+ //#elif ALLOW_WILDCARDS !ALLOW_CHOICES ALLOW_COMMVARS - ?new_content.3: ( variableset | variableuse | commandstn | commandstni | commandset | commandecho | commandif | wildcard )+ - ?new_content_negtag.3: ( variableset | variableuse | commandset | commandecho | commandif | wildcard )+ + ?new_content.3: ( variableset | variableuse | commandstn | commandstni | commandset | commandecho | commandif | commandext | wildcard )+ + ?new_content_negtag.3: ( variableset | variableuse | commandset | commandecho | commandif | commandext | wildcard )+ //#elif !ALLOW_WILDCARDS ALLOW_CHOICES ALLOW_COMMVARS - ?new_content.3: ( variableset | variableuse | commandstn | commandstni | commandset | commandecho | commandif | choices )+ - ?new_content_negtag.3: ( variableset | variableuse | commandset | commandecho | commandif | choices )+ + ?new_content.3: ( variableset | variableuse | commandstn | commandstni | commandset | commandecho | commandif | commandext | choices )+ + ?new_content_negtag.3: ( variableset | variableuse | commandset | commandecho | commandif | commandext | choices )+ //#elif ALLOW_WILDCARDS ALLOW_CHOICES !ALLOW_COMMVARS ?new_content.3: ( wildcard | choices )+ ?new_content_negtag.3: ( wildcard | choices )+ @@ -55,11 +55,11 @@ promptcomppart: content ?new_content.3: ( choices )+ ?new_content_negtag.3: ( choices )+ //#elif !ALLOW_WILDCARDS !ALLOW_CHOICES ALLOW_COMMVARS - ?new_content.3: ( variableset | variableuse | commandstn | commandstni | commandset | commandecho | commandif )+ - ?new_content_negtag.3: ( variableset | variableuse | commandset | commandecho | commandif )+ + ?new_content.3: ( variableset | variableuse | commandstn | commandstni | commandset | commandecho | commandif | commandext )+ + ?new_content_negtag.3: ( variableset | variableuse | commandset | commandecho | commandif | commandext )+ //#else - ?new_content.3: ( variableset | variableuse | commandstn | commandstni | commandset | commandecho | commandif | wildcard | choices )+ - ?new_content_negtag.3: ( variableset | variableuse | commandset | commandecho | commandif | wildcard | choices )+ + ?new_content.3: ( variableset | variableuse | commandstn | commandstni | commandset | commandecho | commandif | commandext | wildcard | choices )+ + ?new_content_negtag.3: ( variableset | variableuse | commandset | commandecho | commandif | wildcard | commandext | choices )+ //#endif //#else ?content.2: ( old_content | plain | specialchars )* @@ -81,6 +81,8 @@ scheduled: "[" [ content ":" ] content ":" numpar "]" // extra network tags extranetworktag: "<" /(?!ppp:)\w+:/ inside_content ">" +?commandinnercontent.3: content + // command: stn (send to negative) commandstn: "" content_negtag "" commandstni: "" @@ -90,7 +92,7 @@ commandif.2: commandif_if commandif_elif* commandif_else? "" commandif_if: "" ifvalue commandif_elif: "" ifvalue commandif_else: "" ifvalue -ifvalue: content +ifvalue: commandinnercontent // conditions ?condition: grouped_condition | ungrouped_condition @@ -104,14 +106,24 @@ operation_not: "not" ( ( _WHITESPACE ungrouped_condition ) | ( _WHITESPACE? grou truthy_operand: IDENTIFIER comparison_simple_value: IDENTIFIER _WHITESPACE ( /not/ _WHITESPACE )? /eq|ne|gt|lt|ge|le|contains/ _WHITESPACE SIMPLEVALUE comparison_list_value: IDENTIFIER _WHITESPACE ( /not/ _WHITESPACE )? /contains|in/ _WHITESPACE listvalue -listvalue: "(" _WHITESPACE? SIMPLEVALUE ( _WHITESPACE? "," _WHITESPACE? SIMPLEVALUE )* _WHITESPACE? ")" +listvalue.9: "(" _WHITESPACE? SIMPLEVALUE ( _WHITESPACE? "," _WHITESPACE? SIMPLEVALUE )* _WHITESPACE? ")" // command: set -commandset: "" content "" +commandset: "" commandsetcontent "" commandsetmodifiers: (_WHITESPACE /evaluate|ifundefined|add/ )+ +?commandsetcontent: commandinnercontent // command: echo -commandecho: "" [ content "" ] +commandecho: "" [ commandechodefault "" ] +?commandechodefault: commandinnercontent + +// command: ext +commandext: "" [ commandexttriggers "" ] +commandexttype: [/\$/] IDENTIFIER +?commandextid: STRING | CNAME +?commandextparams: STRING | SIGNED_NUMBER +?commandextif: "if" _WHITESPACE condition +?commandexttriggers: commandinnercontent // variable set variableset.2: "${" _WHITESPACE? IDENTIFIER [ variablesetmodifiers ] _WHITESPACE? "=" [ /!/ ] varvalue "}" diff --git a/ppp.py b/ppp.py index ae07f0a..79d170f 100644 --- a/ppp.py +++ b/ppp.py @@ -14,6 +14,7 @@ import numpy as np from ppp_hosts import SUPPORTED_APPS # pylint: disable=import-error from ppp_logging import DEBUG_LEVEL # pylint: disable=import-error from ppp_wildcards import PPPWildcard, PPPWildcards # pylint: disable=import-error +from ppp_enmappings import PPPENMappingVariant, PPPExtraNetworkMappings # pylint: disable=import-error class PPPInterrupt(Exception): @@ -92,6 +93,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in options: Optional[dict[str, Any]] = None, grammar_content: Optional[str] = None, wildcards_obj: PPPWildcards = None, + extranetwork_mappings_obj: PPPExtraNetworkMappings = None, ): """ Initializes the PPP object. @@ -103,6 +105,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in options: Optional. The options dictionary for configuring PPP behavior. grammar_content: Optional. The grammar content to be used for parsing. wildcards_obj: Optional. The wildcards object to be used for processing wildcards. + extranetwork_mappings_obj: Optional. The extranetwork mappings object to be used for processing. """ self.logger = logger self.rng = np.random.default_rng() # gets seeded on each process prompt call @@ -110,6 +113,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.options = options self.env_info = env_info self.wildcard_obj = wildcards_obj + self.extranetwork_mappings_obj = extranetwork_mappings_obj # General options self.debug_level = DEBUG_LEVEL(options.get("debug_level", DEBUG_LEVEL.none.value)) @@ -266,7 +270,12 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in propagate_positions=True, ) - # Partial parsers for wildcards and choices + # Partial parsers + self.parser_content = lark.Lark( + grammar_content_full, + propagate_positions=True, + start="content", + ) self.parser_choice = lark.Lark( grammar_content_full, propagate_positions=True, @@ -841,12 +850,13 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.logger.debug(self.format_output(f"Parsing {prompt_description}: '{prompt}'")) parsed_prompt = parser.parse(prompt) # we store the contents so we can use them later even if the meta position is not valid anymore - for n in parsed_prompt.iter_subtrees(): - if isinstance(n, lark.Tree): - if n.meta.empty: - n.meta.content = "" - else: - n.meta.content = prompt[n.meta.start_pos : n.meta.end_pos] + if isinstance(parsed_prompt, lark.Tree): + for n in parsed_prompt.iter_subtrees(): + if isinstance(n, lark.Tree): + if n.meta.empty: + n.meta.content = "" + else: + n.meta.content = prompt[n.meta.start_pos : n.meta.end_pos] except lark.exceptions.UnexpectedInput: if raise_parsing_error: raise @@ -855,7 +865,17 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if self.debug_level == DEBUG_LEVEL.full: self.logger.debug(f"Parse {prompt_description} time: {(t2 - t1) / 1_000_000_000:.3f} seconds") if parsed_prompt: - self.logger.debug("Tree:\n" + textwrap.indent(re.sub(r"\n$", "", parsed_prompt.pretty()), " ")) + self.logger.debug( + "Tree:\n" + + textwrap.indent( + re.sub( + r"\n$", + "", + parsed_prompt.pretty() if isinstance(parsed_prompt, lark.Tree) else parsed_prompt, + ), + " ", + ) + ) return parsed_prompt class TreeProcessor(lark.visitors.Interpreter): @@ -951,7 +971,11 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.visit(node) elif isinstance(node, lark.Token): self.result += node - added_result = self.result[len(backup_result) :] + 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: @@ -1515,6 +1539,135 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in 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.__ppp.rem_removeextranetworktags: + extnet_type: str = (tree.children[0].children[0] or "") + tree.children[0].children[1] + is_mapping = extnet_type.startswith("$") + if is_mapping: + extnet_type = extnet_type[1:] + extnet_id: str = tree.children[1].value + 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 = tree.children[2].value + 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_mappings: list[PPPENMappingVariant] = [] + else_mapping = None + if self.__ppp.extranetwork_mappings_obj: + enmapping = self.__ppp.extranetwork_mappings_obj.extranetwork_mappings.get(extnet_id, None) + if enmapping: + for v in enmapping.variants: + if v.condition: + try: + cnd = self.__ppp.parse_prompt( + "condition", v.condition, self.__ppp.parser_condition, True + ) + except lark.exceptions.UnexpectedInput as e: + self.warn_or_stop( + f"Error parsing condition '{v.condition}' in extranetwork mapping '{extnet_id}'! : {e.__class__.__name__}", + e, + ) + cnd = None + else: + cnd = "True" + if cnd is not None and (cnd == "True" or self.__eval_condition(cnd)): + if v.condition: + found_mappings.append(v) + else: + else_mapping = v + if found_mappings: + found = found_mappings[ + self.__ppp.rng.choice( + len(found_mappings), + p=[v.weight or 1 for v in found_mappings], + ) + ] + else: + found = else_mapping + if found: + if found.name: + if self.__ppp.debug_level != DEBUG_LEVEL.none: + self.__ppp.logger.info( + f"Mapping extranetwork '{extnet_id}' to '{extnet_type}:{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: + self.__ppp.logger.info(f"Mapping extranetwork '{extnet_id}' to just triggers") + extnet_id = None + else: + self.__ppp.logger.info(f"Mapping extranetwork '{extnet_id}' to nothing") + extnet_id = None + if found.triggers: + extra_triggers = ", ".join(found.triggers) + try: + compiled_extra_triggers = self.__ppp.parse_prompt( + "triggers", extra_triggers, self.__ppp.parser_content, True + ) + except lark.exceptions.UnexpectedInput as e: + self.warn_or_stop( + f"Error parsing triggers '{extra_triggers}' in extranetwork mapping '{extnet_id}'! : {e.__class__.__name__}", + e, + ) + compiled_extra_triggers = None + else: + self.warn_or_stop(f"Extranetwork mapping '{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.__ppp.cup_extranetworktags: + 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 extranetworktag(self, tree: lark.Tree): """ Process an extra network construct in the tree. @@ -1558,9 +1711,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in from_value: int = options.get("from", 1) to_value: int = options.get("to", 1) separator: str = options.get("separator", self.__ppp.wil_choice_separator) + msg_where = f"wildcard '{wildcard_key}'" if wildcard_key else "choices" if sampler != "~": - msg = f"wildcard '{wildcard_key}'" if wildcard_key else "choices" - self.warn_or_stop(f"Unsupported sampler '{sampler}' in {msg} options!") + self.warn_or_stop(f"Unsupported sampler '{sampler}' in {msg_where} options!") sampler = "~" if filter_specifier is not None: filtered_choice_values = [] @@ -1587,15 +1740,36 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) else: filtered_choice_values = choice_values.copy() - if filtered_choice_values: + available_choices: list[dict] = [] + weights = [] + included_choices = 0 + excluded_choices = 0 + excluded_weights_sum = 0 + for i, c in enumerate(filtered_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(filtered_choice_values): - from_value = len(filtered_choice_values) + elif from_value > len(available_choices): + from_value = len(available_choices) if to_value < 1: to_value = 1 - elif (to_value > len(filtered_choice_values) and not repeating) or from_value > to_value: - to_value = len(filtered_choice_values) + elif (to_value > len(available_choices) and not repeating) or from_value > to_value: + to_value = len(available_choices) num_choices = ( self.__ppp.rng.integers(from_value, to_value, endpoint=True) if from_value < to_value @@ -1603,6 +1777,10 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) else: num_choices = 0 + if from_value != 0: + self.__ppp.logger.warning( + f"No available choices could be selected for {msg_where}!" + ) if num_choices < 2: repeating = False if self.__ppp.debug_level == DEBUG_LEVEL.full: @@ -1613,30 +1791,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ) ) if num_choices > 0: - available_choices: list[dict] = [] - weights = [] - included_choices = 0 - excluded_choices = 0 - excluded_weights_sum = 0 - for i, c in enumerate(filtered_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 selected_choices: list[dict] = list( self.__ppp.rng.choice(available_choices, size=num_choices, p=weights, replace=repeating) - ) + ) if available_choices else [] if self.__ppp.wil_keep_choices_order: selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"]) selected_choices_text = [] @@ -1681,7 +1838,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in wildcard_key: str = None, ) -> str: r = self.__get_choices_internal(options, choice_values, filter_specifier, wildcard_key) - return r[0] + r[2].join(r[1]) + r[3] + if r[1]: + return r[0] + r[2].join(r[1]) + r[3] + return "" def __convert_choices_options(self, options: Optional[lark.Tree]) -> dict: """ diff --git a/ppp_comfyui.py b/ppp_comfyui.py index 801fd66..6d8433d 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -2,12 +2,13 @@ import os # pylint: disable=import-error import folder_paths # type: ignore -import nodes # type: ignore +import nodes from .ppp import PromptPostProcessor from .ppp_hosts import SUPPORTED_APPS from .ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory from .ppp_wildcards import PPPWildcards +from .ppp_enmappings import PPPExtraNetworkMappings if __name__ == "__main__": raise SystemExit("This script must be run from ComfyUI") @@ -27,6 +28,7 @@ class PromptPostProcessorComfyUINode: with open(grammar_filename, "r", encoding="utf-8") as file: self.grammar_content = file.read() self.wildcards_obj = PPPWildcards(lf.log) + self.extranetwork_mappings_obj = PPPExtraNetworkMappings(lf.log) self.logger.info(f"{PromptPostProcessor.NAME} {PromptPostProcessor.VERSION} initialized") class SmartType(str): @@ -281,6 +283,24 @@ class PromptPostProcessorComfyUINode: "label_off": "No", }, ), + "en_mappings_folders": ( + "STRING", + { + "default": "", + "tooltip": "Comma separated list of extranetwork mappings folders", + "dynamicPrompts": False, + }, + ), + "en_mappings_input": ( + "STRING", + { + "default": "", + "multiline": True, + "placeholder": "extranetwork mappings definitions", + "tooltip": "Extranetwork mappings definitions in yaml format", + "dynamicPrompts": False, + }, + ), }, } @@ -343,6 +363,8 @@ class PromptPostProcessorComfyUINode: cleanup_extranetwork_tags, cleanup_merge_attention, remove_extranetwork_tags, + en_mappings_folders, + en_mappings_input, ): if wc_process_wildcards: return float( @@ -376,6 +398,8 @@ class PromptPostProcessorComfyUINode: "cleanup_extranetwork_tags": cleanup_extranetwork_tags, "cleanup_merge_attention": cleanup_merge_attention, "remove_extranetwork_tags": remove_extranetwork_tags, + "en_mappings_folders": en_mappings_folders, + "en_mappings_input": en_mappings_input, } return new_run.__hash__ # return float("NaN") @@ -410,6 +434,8 @@ class PromptPostProcessorComfyUINode: cleanup_extranetwork_tags, cleanup_merge_attention, remove_extranetwork_tags, + en_mappings_folders, + en_mappings_input, ): modelclass = ( model.model.model_config.__class__.__name__ if model is not None and not isinstance(model, str) else model @@ -447,7 +473,15 @@ class PromptPostProcessorComfyUINode: # Also supported: SVD_img2vid, SVD3D_u, SVD3_p, Stable_Zero123, SD_X4Upscaler, Stable_Cascade_C, Stable_Cascade_B, StableAudio if wc_wildcards_folders == "": - wc_wildcards_folders = ",".join(folder_paths.get_folder_paths("wildcards") or []) + try: + fp1 = folder_paths.get_folder_paths("ppp_wildcards") + except Exception: # pylint: disable=W0718 + fp1 = None + try: + fp2 = folder_paths.get_folder_paths("wildcards") + except Exception: # pylint: disable=W0718 + fp2 = None + wc_wildcards_folders = ",".join(fp1 or fp2 or []) if wc_wildcards_folders == "": wc_wildcards_folders = os.getenv("WILDCARD_DIR", PPPWildcards.DEFAULT_WILDCARDS_FOLDER) wildcards_folders = [ @@ -455,6 +489,22 @@ class PromptPostProcessorComfyUINode: for f in wc_wildcards_folders.split(",") if f.strip() != "" ] + if en_mappings_folders == "": + try: + fp3 = folder_paths.get_folder_paths("ppp_extranetworkmappings") + except Exception: # pylint: disable=W0718 + fp3 = None + en_mappings_folders = ",".join(fp3 or []) + if en_mappings_folders == "": + en_mappings_folders = os.getenv( + "EXTRANETWORKMAPPINGS_DIR", PPPExtraNetworkMappings.DEFAULT_ENMAPPINGS_FOLDER + ) + enmappings_folders = [ + (f if os.path.isabs(f) else os.path.abspath(os.path.join(folder_paths.models_dir, f))) + for f in en_mappings_folders.split(",") + if f.strip() != "" + ] + if variants_definitions != "" and not "=" in variants_definitions: # mainly to warn about the old format raise ValueError("Invalid variants_definitions format") options = { @@ -485,8 +535,19 @@ class PromptPostProcessorComfyUINode: wildcards_folders if options["process_wildcards"] else None, wc_wildcards_input, ) + self.extranetwork_mappings_obj.refresh_extranetwork_mappings( + debug_level, + enmappings_folders, + en_mappings_input, + ) ppp = PromptPostProcessor( - self.logger, self.interrupt, env_info, options, self.grammar_content, self.wildcards_obj + self.logger, + self.interrupt, + env_info, + options, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_mappings_obj, ) pos_prompt, neg_prompt, variables = ppp.process_prompt(pos_prompt, neg_prompt, seed if seed is not None else 1) return ( diff --git a/ppp_enmappings.py b/ppp_enmappings.py new file mode 100644 index 0000000..c24e6d0 --- /dev/null +++ b/ppp_enmappings.py @@ -0,0 +1,257 @@ +import os +from typing import Optional +import logging +import yaml + +from ppp_logging import DEBUG_LEVEL # pylint: disable=import-error +from ppp_utils import deep_freeze # pylint: disable=import-error + + +class PPPENMappingVariant: + """ + A class to represent a variant of an extra network mapping. + + Attributes: + condition (str): The condition for the variant. + name (str): The name of the variant. + parameters (float|str): The parameters for the variant. + triggers (list[str]): The triggers for the variant. + weight (float): The weight for the variant when multiple variants apply. + """ + + def __init__(self, condition: str, name: str, parameters: float | str, triggers: list[str], weight: float): + self.condition: str = condition + self.name: str = name + self.parameters: float | str = parameters + self.triggers: list[str] = triggers + self.weight: float = weight + + +class PPPENMapping: + """ + A extra network mapping object. + + Attributes: + kind (str): The kind of the extra network. + name (str): The name of the extra network mapping. + file (str): The path to the file where the extranetwork mapping is defined. + variants (list[PPPENMappingVariant]): The processed variants of the extranetwork mapping. + """ + + def __init__(self, fullpath: str, kind: str, name: str, variants: list[dict]): + self.file: str = fullpath + self.kind: str = kind + self.name: str = name + self.variants: list[PPPENMappingVariant] = [ + PPPENMappingVariant(**{**{"condition": None, "name": None, "parameters": None, "triggers": None, "weight": 1.0}, **v}) + for v in variants + ] + + def __hash__(self) -> int: + t = (self.kind, self.name, deep_freeze(self.variants)) + return hash(t) + + def __sizeof__(self): + return self.kind.__sizeof__() + self.name.__sizeof__() + self.file.__sizeof__() + self.variants.__sizeof__() + + +class PPPExtraNetworkMappings: + """ + A class to manage extra network mappings. + + Attributes: + extranetwork_maps (dict[str, PPPENMapping]): The extra network mappings. + """ + + DEFAULT_ENMAPPINGS_FOLDER = "extranetworkmappings" + LOCALINPUT_FILENAME = "#INPUT" + + def __init__(self, logger): + self.__logger: logging.Logger = logger + self.__debug_level = DEBUG_LEVEL.none + self.__enmappings_folders = [] + self.__enmappings_files = {} + self.extranetwork_mappings: dict[str, PPPENMapping] = {} + + def __hash__(self) -> int: + return hash(deep_freeze(self.extranetwork_mappings)) + + def __sizeof__(self): + return ( + self.extranetwork_mappings.__sizeof__() + self.__enmappings_folders.__sizeof__() + self.__enmappings_files.__sizeof__() + ) + + def refresh_extranetwork_mappings( + self, debug_level: DEBUG_LEVEL, enmappings_folders: Optional[list[str]], enmappings_input: str = None + ): + """ + Initialize the extra network mappings. + """ + self.__debug_level = debug_level + self.__enmappings_folders = enmappings_folders or [] + # if self.__debug_level != DEBUG_LEVEL.none: + # self.__logger.info("Refreshing extra network mappings...") + # t1 = time.monotonic_ns() + for fullpath in list(self.__enmappings_files.keys()): + if fullpath != self.LOCALINPUT_FILENAME: + path = os.path.dirname(fullpath) + if not os.path.exists(fullpath) or not any( + os.path.commonpath([path, folder]) == folder for folder in self.__enmappings_folders + ): + self.__remove_extranetwork_mappings_from_path(fullpath) + elif enmappings_input is None: + self.__remove_extranetwork_mappings_from_path(fullpath) + if enmappings_folders is not None or enmappings_input is not None: + if enmappings_folders is not None: + for f in self.__enmappings_folders: + self.__get_extranetwork_mappings_in_directory(f) + if enmappings_input is not None: + self.__get_extranetwork_mappings_in_input(enmappings_input) + else: + self.extranetwork_mappings = {} + self.__enmappings_files = {} + # t2 = time.monotonic_ns() + # if self.__debug_level != DEBUG_LEVEL.none: + # self.__logger.info(f"Extra network mappings refresh time: {(t2 - t1) / 1_000_000_000:.3f} seconds") + + # def get_extranetwork_mappings(self, key: str) -> list[PPPENMapping]: + # """ + # Get all extra network mappings that match a key. + # + # Args: + # key (str): The key to match (kind:name). + # + # Returns: + # list: A list of all extra network mappings that match the key. + # """ + # keys = sorted(fnmatch.filter(self.extranetwork_mappings.keys(), key)) + # return [self.extranetwork_mappings[k] for k in keys] + + def __remove_extranetwork_mappings_from_path(self, full_path: str, debug=True): + """ + Clear all extra network mappings in a file. + + Args: + full_path (str): The path to the file. + debug (bool): Whether to print debug messages or not. + """ + last_modified_cached = self.__enmappings_files.get(full_path, None) # a time or a hash + if debug and last_modified_cached is not None and self.__debug_level != DEBUG_LEVEL.none: + if full_path == self.LOCALINPUT_FILENAME: + self.__logger.debug("Removing extra network mappings from input") + else: + self.__logger.debug(f"Removing extra network mappings from file: {full_path}") + if full_path in self.__enmappings_files.keys(): + del self.__enmappings_files[full_path] + for key in list(self.extranetwork_mappings.keys()): + if self.extranetwork_mappings[key].file == full_path: + del self.extranetwork_mappings[key] + + def __get_extranetwork_mappings_in_file(self, full_path: str): + """ + Get all extra network mappings in a file. + + Args: + full_path (str): The path to the file. + """ + last_modified = os.path.getmtime(full_path) + last_modified_cached = self.__enmappings_files.get(full_path, None) + if last_modified_cached is not None and last_modified == self.__enmappings_files[full_path]: + return + filename = os.path.basename(full_path) + _, extension = os.path.splitext(filename) + if extension not in (".yaml", ".yml", ".json"): + return + self.__remove_extranetwork_mappings_from_path(full_path, False) + if last_modified_cached is not None and self.__debug_level != DEBUG_LEVEL.none: + self.__logger.debug(f"Updating extra network mappings from file: {full_path}") + self.__get_extranetwork_mappings_in_structured_file(full_path) + self.__enmappings_files[full_path] = last_modified + + def __get_extranetwork_mappings_in_input(self, enmappings_input: str): + """ + Get all extra network mappings in the string. + + Args: + enmappings_input (str): The input string containing extra network mappings in yaml format. + """ + new_h = hash(enmappings_input) + h = self.__enmappings_files.get(self.LOCALINPUT_FILENAME, None) + if h == new_h: + return + self.__remove_extranetwork_mappings_from_path(self.LOCALINPUT_FILENAME, False) + if h is not None and self.__debug_level != DEBUG_LEVEL.none: + self.__logger.debug("Updating extra network mappings from input") + enmappings_input = enmappings_input.strip() + if enmappings_input != "": + try: + content = yaml.safe_load(enmappings_input) + except yaml.YAMLError as e: + self.__logger.warning(f"Invalid format for input extra network mappings: {e}") + return + if content is not None: + self.__add_extranetwork_mapping(content, self.LOCALINPUT_FILENAME) + self.__enmappings_files[self.LOCALINPUT_FILENAME] = new_h + + def __add_extranetwork_mapping(self, content: dict[str, dict[str, list[dict]]], full_path: str): + """ + Add an extra network mapping to the extra network mappings dictionary. + + Args: + content (object): The content of the extra network mapping. + full_path (str): The path to the file that contains it. + """ + if not isinstance(content, dict): + self.__logger.warning(f"Invalid extra network mapping in file '{full_path}'!") + return + for kind, maps in content.items(): + if not isinstance(maps, dict): + self.__logger.warning(f"Invalid extra network mapping definition for '{kind}:*' in file '{full_path}'!") + else: + for name, variants in maps.items(): + key = f"{kind}:{name}" + if not isinstance(variants, list): + self.__logger.warning( + f"Invalid extra network mapping definition for '{key}' in file '{full_path}'!" + ) + elif self.extranetwork_mappings.get(key, None) is not None: + self.__logger.warning( + f"Duplicate extra network mapping '{key}' in file '{full_path}' and '{self.extranetwork_mappings[key].file}'!" + ) + elif not isinstance(variants, list) or not all(isinstance(v, dict) for v in variants): + self.__logger.warning( + f"Invalid extra network mapping definition for '{key}' in file '{full_path}'!" + ) + else: + self.extranetwork_mappings[key] = PPPENMapping(full_path, kind, name, variants) + + def __get_extranetwork_mappings_in_structured_file(self, full_path): + """ + Get all extra network mappings in a structured file. + + Args: + full_path (str): The path to the file. + base (str): The base path for the extra network mappings. + """ + with open(full_path, "r", encoding="utf-8") as file: + content = yaml.safe_load(file) + self.__add_extranetwork_mapping(content, full_path) + + def __get_extranetwork_mappings_in_directory(self, directory: str): + """ + Get all extra network mappings in a directory. + + Args: + directory (str): The path to the directory. + """ + if not os.path.exists(directory): + self.__logger.warning(f"Extra network mappings directory '{directory}' does not exist!") + return + for filename in os.listdir(directory): + full_path = os.path.abspath(os.path.join(directory, filename)) + if os.path.basename(full_path).startswith("."): + continue + if os.path.isdir(full_path): + self.__get_extranetwork_mappings_in_directory(full_path) + elif os.path.isfile(full_path): + self.__get_extranetwork_mappings_in_file(full_path) diff --git a/ppp_utils.py b/ppp_utils.py new file mode 100644 index 0000000..0fd10b1 --- /dev/null +++ b/ppp_utils.py @@ -0,0 +1,17 @@ +def deep_freeze(obj): + """ + Deep freeze an object. + + Args: + obj (object): The object to freeze. + + Returns: + object: The frozen object. + """ + if isinstance(obj, dict): + return tuple((k, deep_freeze(v)) for k, v in sorted(obj.items())) + if isinstance(obj, list): + return tuple(deep_freeze(i) for i in obj) + if isinstance(obj, set): + return tuple(deep_freeze(i) for i in sorted(obj)) + return obj diff --git a/ppp_wildcards.py b/ppp_wildcards.py index 7a0d180..cd78fc5 100644 --- a/ppp_wildcards.py +++ b/ppp_wildcards.py @@ -5,25 +5,7 @@ import logging import yaml from ppp_logging import DEBUG_LEVEL # pylint: disable=import-error - - -def deep_freeze(obj): - """ - Deep freeze an object. - - Args: - obj (object): The object to freeze. - - Returns: - object: The frozen object. - """ - if isinstance(obj, dict): - return tuple((k, deep_freeze(v)) for k, v in sorted(obj.items())) - if isinstance(obj, list): - return tuple(deep_freeze(i) for i in obj) - if isinstance(obj, set): - return tuple(deep_freeze(i) for i in sorted(obj)) - return obj +from ppp_utils import deep_freeze # pylint: disable=import-error class PPPWildcard: diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py index 9699c40..cecfecd 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -19,6 +19,7 @@ from ppp_hosts import SUPPORTED_APPS, SUPPORTED_APPS_NAMES # pylint: disable=im from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # pylint: disable=import-error from ppp_cache import PPPLRUCache # pylint: disable=import-error from ppp_wildcards import PPPWildcards # pylint: disable=import-error +from ppp_enmappings import PPPExtraNetworkMappings # pylint: disable=import-error class PromptPostProcessorA1111Script(scripts.Script): @@ -68,6 +69,7 @@ class PromptPostProcessorA1111Script(scripts.Script): self.ppp_debug_level = DEBUG_LEVEL.none.value self.lru_cache = None self.wildcards_obj = None + self.extranetwork_mappings_obj = None def title(self): """ @@ -179,6 +181,7 @@ class PromptPostProcessorA1111Script(scripts.Script): self.ppp_debug_level = DEBUG_LEVEL(getattr(opts, "ppp_gen_debug_level", DEBUG_LEVEL.none.value)) self.lru_cache = PPPLRUCache(1000, logger=self.ppp_logger, debug_level=self.ppp_debug_level) self.wildcards_obj = PPPWildcards(self.ppp_logger) + self.extranetwork_mappings_obj = PPPExtraNetworkMappings(self.ppp_logger) self.ppp_logger.info( f"{PromptPostProcessor.NAME} {PromptPostProcessor.VERSION} initialized, running on {SUPPORTED_APPS_NAMES[app]}" ) @@ -292,6 +295,14 @@ class PromptPostProcessorA1111Script(scripts.Script): for f in wc_wildcards_folders.split(",") if f.strip() != "" ] + en_mappings_folders = getattr(opts, "ppp_en_mappingsfolders", "") + if en_mappings_folders == "": + en_mappings_folders = os.getenv("EXTRANETWORKMAPPINGS_DIR", PPPExtraNetworkMappings.DEFAULT_ENMAPPINGS_FOLDER) + enmappings_folders = [ + (f if os.path.isabs(f) else os.path.abspath(os.path.join(models_path, f))) + for f in en_mappings_folders.split(",") + if f.strip() != "" + ] options = { "debug_level": getattr(opts, "ppp_gen_debug_level", DEBUG_LEVEL.none.value), "on_warning": getattr(opts, "ppp_gen_onwarning", PromptPostProcessor.ONWARNING_CHOICES.warn.value), @@ -321,8 +332,15 @@ class PromptPostProcessorA1111Script(scripts.Script): self.wildcards_obj.refresh_wildcards( self.ppp_debug_level, wildcards_folders if options["process_wildcards"] else None ) + self.extranetwork_mappings_obj.refresh_extranetwork_mappings(self.ppp_debug_level, enmappings_folders) ppp = PromptPostProcessor( - self.ppp_logger, self.ppp_interrupt, env_info, options, self.grammar_content, self.wildcards_obj + self.ppp_logger, + self.ppp_interrupt, + env_info, + options, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_mappings_obj, ) prompts_list = [] @@ -571,6 +589,16 @@ def on_ui_settings(): ), ) + shared.opts.add_option( + key="ppp_en_mappingsfolders", + info=shared.OptionInfo( + PPPExtraNetworkMappings.DEFAULT_ENMAPPINGS_FOLDER, + label="Extranetwork Mappings folders", + comment_after='(absolute or relative to the models folder)', + section=section, + ), + ) + # wildcard settings shared.opts.add_option( key="ppp_wil_sep", diff --git a/tests/enmappings/mappings.yaml b/tests/enmappings/mappings.yaml new file mode 100644 index 0000000..41d0cc6 --- /dev/null +++ b/tests/enmappings/mappings.yaml @@ -0,0 +1,11 @@ +lora: + lora1: + - condition: _is_pony + name: lorapony + parameters: 0.8 + triggers: ["triggerpony1", "triggerpony2"] + - condition: _is_illustrious + name: loraillustrious + parameters: "0.9:0.8" + triggers: ["triggerillustrious1", "triggerillustrious2"] + - triggers: ["triggergeneric1", "triggergeneric2", "{one|two}"] diff --git a/tests/tests.py b/tests/tests.py index 1d06649..983e2bf 100644 --- a/tests/tests.py +++ b/tests/tests.py @@ -3,6 +3,7 @@ import logging from typing import NamedTuple, Optional import unittest +from ppp_enmappings import PPPExtraNetworkMappings # pylint: disable=import-error from ppp_wildcards import PPPWildcards # pylint: disable=import-error from ppp import PromptPostProcessor # pylint: disable=import-error from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory # pylint: disable=import-error @@ -68,6 +69,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): } self.interrupted = False self.wildcards_obj = PPPWildcards(self.lf.log) + self.extranetwork_maps_obj = PPPExtraNetworkMappings(self.lf.log) self.wildcards_obj.refresh_wildcards( DEBUG_LEVEL.full, [ @@ -82,6 +84,12 @@ class TestPromptPostProcessorBase(unittest.TestCase): - choice3 """, ) + self.extranetwork_maps_obj.refresh_extranetwork_mappings( + DEBUG_LEVEL.full, + [os.path.abspath(os.path.join(os.path.dirname(__file__), "enmappings"))], + """ + """, + ) grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "../grammar.lark") with open(grammar_filename, "r", encoding="utf-8") as file: self.grammar_content = file.read() @@ -92,6 +100,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): self.defopts, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ) self.nocupppp = PromptPostProcessor( self.ppp_logger, @@ -113,6 +122,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): }, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ) self.comfyuippp = PromptPostProcessor( self.ppp_logger, @@ -125,6 +135,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): self.defopts, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ) def interrupt(self): @@ -290,8 +301,8 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): def test_cl_simple(self): # simple cleanup self.process( - PromptPair(" this is a ((test ), , , (), , [] ( , test ,:2.0):1.5) (red:1.5) ", " normal quality "), - PromptPair("this is a ((test), (test,:2):1.5) (red:1.5)", "normal quality"), + PromptPair(" this is a ((test ), , , (), , [] ( , test ,:2.0):1.5), (red:1.5) ", " normal quality "), + PromptPair("this is a ((test), (test,:2):1.5), (red:1.5)", "normal quality"), ) def test_cl_complex(self): # complex cleanup @@ -317,6 +328,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): {**self.defopts, "remove_extranetwork_tags": True}, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ), ) @@ -335,6 +347,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): }, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ), ) @@ -372,6 +385,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): }, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ), ) @@ -443,6 +457,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): self.defopts, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ), ) @@ -452,6 +467,12 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): PromptPair("this test is OK", ""), ) + def test_cmd_set_empty(self): # set to empty + self.process( + PromptPair("${v2=}this test is not OKOK", ""), + PromptPair("this test is OK", ""), + ) + def test_cmd_set_eval_if(self): # set and if commands self.process( PromptPair("valuethis test is OKnot OK", ""), @@ -530,7 +551,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): def test_cmd_set_if2(self): # set and more complex if commands self.process( PromptPair( - "First: value1this test is OKnot OK\nSecond: value3this test is OKnot OK", + "First: value1this test is OKOK2not OK\nSecond: value3this test is OKnot OK", "", ), PromptPair("First: this test is OK\nSecond: this test is OK", ""), @@ -613,6 +634,111 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): PromptPair("this test is OK", ""), ) + def test_cmd_ext(self): # ext + self.process( + PromptPair( + "trigger1trigger2trigger4", + "", + ), + PromptPair( + "trigger1,trigger2,trigger4", + "", + ), + ) + + def test_cmd_ext_map1(self): # ext mapping, no lora + self.process( + PromptPair( + "inlinetrigger", + "", + ), + PromptPair("inlinetrigger, triggergeneric1, triggergeneric2, two", ""), + ) + + def test_cmd_ext_map2(self): # ext mapping, lora with weight + self.process( + PromptPair( + "inlinetrigger", + "", + ), + PromptPair("inlinetrigger, triggerpony1, triggerpony2", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cmd_ext_map3(self): # ext mapping, lora with weight adjusted + self.process( + PromptPair( + "inlinetrigger", + "", + ), + PromptPair("inlinetrigger, triggerpony1, triggerpony2", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cmd_ext_map4(self): # ext mapping, lora with parameters + self.process( + PromptPair( + "inlinetrigger", + "", + ), + PromptPair("inlinetrigger, triggerpony1, triggerpony2", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + + def test_cmd_ext_map5(self): # ext mapping, lora with no parameters + self.process( + PromptPair( + "inlinetrigger", + "", + ), + PromptPair("inlinetrigger, triggerillustrious1, triggerillustrious2", ""), + ppp=PromptPostProcessor( + self.ppp_logger, + self.interrupt, + { + **self.def_env_info, + "model_filename": "./webui/models/Stable-diffusion/ilxlmodel.safetensors", + }, + self.defopts, + self.grammar_content, + self.wildcards_obj, + self.extranetwork_maps_obj, + ), + ) + # Choices tests def test_ch_choices(self): # simple choices with weights @@ -689,6 +815,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): {**self.defopts, "remove_extranetwork_tags": True}, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ), ) @@ -709,6 +836,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): }, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ), ) @@ -733,6 +861,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): }, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ), ) @@ -751,6 +880,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): }, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ), ) @@ -772,6 +902,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): }, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ), interrupted=True, ) @@ -791,6 +922,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): }, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ), ) @@ -1039,6 +1171,7 @@ class TestPromptPostProcessor(TestPromptPostProcessorBase): }, self.grammar_content, self.wildcards_obj, + self.extranetwork_maps_obj, ), )