diff --git a/.vscode/launch.json b/.vscode/launch.json index a000936..d36805c 100644 --- a/.vscode/launch.json +++ b/.vscode/launch.json @@ -6,7 +6,7 @@ "configurations": [ { "name": "Tests", - "type": "python", + "type": "debugpy", "request": "launch", "program": "tests/tests.py", "console": "integratedTerminal", diff --git a/.vscode/settings.json b/.vscode/settings.json index 3fc9ddc..d52f55c 100644 --- a/.vscode/settings.json +++ b/.vscode/settings.json @@ -9,7 +9,7 @@ ], "python.testing.pytestEnabled": false, "python.testing.unittestEnabled": true, - "python.analysis.typeCheckingMode": "basic", + "python.analysis.typeCheckingMode": "off", "black-formatter.args": [ "--line-length=120" ] diff --git a/README.md b/README.md index 88d7c7b..8fac366 100644 --- a/README.md +++ b/README.md @@ -1,57 +1,60 @@ -# Prompt Postprocessor for Stable Diffusion WebUI +# Prompt Postprocessor for Stable Diffusion WebUI and ComfyUI -The Prompt Postprocessor for Stable Diffusion WebUI, formerly known as "sd-webui-sendtonegative", is an extension designed to process the prompt after other extensions have potentially modified it. This extension is compatible with the [AUTOMATIC1111 Stable Diffusion WebUI](https://github.com/AUTOMATIC1111/stable-diffusion-webui) and [SD.Next](https://github.com/vladmandic/automatic). +The Prompt Postprocessor, formerly known as "sd-webui-sendtonegative", is an extension designed to process the prompt, possibly after other extensions have modified it. This extension is compatible with: + +* [AUTOMATIC1111 Stable Diffusion WebUI](https://github.com/AUTOMATIC1111/stable-diffusion-webui) +* [SD.Next](https://github.com/vladmandic/automatic). +* [Forge](https://github.com/lllyasviel/stable-diffusion-webui-forge) +* [reForge](https://github.com/Panchovix/stable-diffusion-webui-reForge) +* ...and probably other forks +* [ComfyUI](https://github.com/comfyanonymous/ComfyUI) Currently this extension has these functions: -* Allows marking parts of the prompt and moves them to the negative prompt. This allows for useful tricks when using a wildcard extension since you can add negative content from choices made in the positive prompt. -* Set values to local variables. -* Filter content based on the loaded SD model version or a set variable. -* Detect invalid wildcards and act on them. +* Sending parts of the prompt to the negative prompt. This allows for useful tricks when using wildcards since you can add negative content from choices made in the positive prompt. +* 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. * Clean up the prompt and negative prompt. -Note: The extension must be loaded after the installed wildcards extension (or any other that modifies the prompt or has it's own syntax expressions). Extensions load by their folder name in alphanumeric order. - -With the ["Dynamic Prompts" extension](https://github.com/adieyal/sd-dynamic-prompts) this happens by default due to default folder names for both extensions. But if this is not the case, you can just rename this extension's folder so the ordering works out. - -With the ["AUTOMATIC1111 Wildcards" extension](https://github.com/AUTOMATIC1111/stable-diffusion-webui-wildcards) you will have to rename one of the folders, so that it loads before than this extension. - -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. +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. Notes: -1. It only recognizes regular A1111 prompt formats. So: +1. Other than its own commands, it only recognizes regular A1111 prompt formats. So: * **Attention**: `\[prompt\] (prompt) (prompt:weight)` * **Alternation**: `\[prompt1|prompt2|...\]` * **Scheduling**: `\[prompt1:prompt2:step\]` * **Extra networks**: `\` * **BREAK**: `prompt1 BREAK prompt2` - * **Composable Diffusion**: `prompt1 AND prompt2` + * **Composable Diffusion**: `prompt1:weight1 AND prompt2:weight2` In SD.Next that means only the *A1111* or *Full* parsers. It will warn you if you use the *Compel* parser. -2. It only recognizes wildcards in the *\_\_wildcard\_\_* and *{choice|choice}* formats. -3. Since it should run after other extensions that apply to the prompt, the content should have already been processed by them and there should't be any non recognized syntax anymore. -4. It does not create *AND/BREAK* constructs when moving content to the negative prompt. +2. It recognizes wildcards in the *\_\_wildcard\_\_* and *{choice|choice}* formats (and anything that [Dynamic Prompts](https://github.com/adieyal/sd-dynamic-prompts) supports). +3. It does not create *AND/BREAK* constructs when moving content to the negative prompt. ## Installation +On A1111 compatible webuis: + 1. Go to Extensions > Install from URL 2. Paste in the URL for extension's git repository text field 3. Click the Install button 4. Restart the webui +On ComfyUI: + +1. Go to Manager > Custom Nodes Manager +2. Install through ComfyUI Manager +3. Click Install via Git URL and enter +4. Restart + ## Usage -### Detection of remaining wildcards - -This extension should run after any wildcard extensions, so any remaining wildcards present in the prompt or negative_prompt at this point of processing must be invalid. Usually you might not notice this problem until you check the image metadata, so this option gives you some ways to detect and treat the problem. - -If you choose to not ignore wildcards, the extension will look for any *\_\_wildcard\_\_* or *{choice|choice}* constructs and act as configured. - ### Commands -The extension uses now a new format for its commands. The format is similar to an extranetwork, but it has a "ppp:" prefix followed by the command, and then a space and any parameters (if any). +The extension uses a format for its commands similar to an extranetwork, but it has a "ppp:" prefix followed by the command, and then a space and any parameters (if any). ```text @@ -63,7 +66,79 @@ When a command is associated with any content, it will be between an opening and content ``` -The `set` and `if` commands are the first to be processed. +For wildcards and choices it uses the formats from the Dynamic Prompts extension, but sometimes with some additional options for more functionality. + +### Choices + +The generic format is: + +```text +{parameters$$opt1::choice1|opt2::choice2|opt3::choice3} +``` + +Both the construct parameters (up to the '$$') and the individual choice options (up to the '::') are optional. + +There is also a format where instead of "parameters$$" you just put the sampler, for compatibility with Dynamic Prompts. + +The construct parameters can be written with the following options (all are optional): + +* "**~**" or "**@**": sampler (for compatibility with Dynamic Prompts), but only "**~**" (random) is allowed. +* "**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. + +The choice options are as follows: + +* "**n**": weight of the choice (default 1) +* "**if condition**": filters out the choice if the condition is false (this is an extension to the Dynamic Prompts syntax). Same conditions as in the `if` command. +* "**::**": end of choice options + +Whitespace is allowed between parameters. + +These are examples of formats you can use to insert a choice construct: + +```text +{opt1|5::opt2|3::opt3} # select 1 choice, two have weights +{3$$opt1|5 if _is_sd1::opt2|opt3} # select 3 choices, one has a weight and a condition +{2-3$$opt1|opt2|opt3} # select 2 to 3 choices +{r2-3$$opt1|opt2|opt3} # select 2 to 3 choices allowing repetition +{2-3$$ / $$opt1|opt2|opt3} # select 2 to 3 choices with separator " / " +``` + +Notes: + +* The Dynamic Prompts format `{2$$__flavours__}` does not work as expected. It will only output one value. You can write is as `{r2$$__flavours__}` to get two values, but they may repeat since the evaluation of the wildcard is independent of the choices selection. +* Whitespace in the choices is not ignored like in Dynamic Prompts, but will be cleaned up if the appropiate settings are checked. + +### Wildcards + +The generic format is: + +```text +__parameters$$path/to/wildcard(var=value)__ +``` + +The parameters and the setting of a variable are optional. The parameters follow the same format as for the choices. The variable value only applies during the evaluation of the selected choices and is discarded afterward (the variable keeps its original value if there was one). + +In the wildcard definition (which supports the text, json and yaml formats), if the first choice follows the format of these parameters, it will be used as default parameters for that wildcard (see examples in the tests folder). The choices of the wildcard follow the same format as in the choices construct. If using the object format for a choice you can use a new "if" property for the condition in addition to the standard "weight" and "text"/"content". + +Wildcards can contain just one choice. In json and yaml formats this allows the use of a string value for the keys, rather than an array. + +These are examples of formats you can use to insert a wildcard: + +```text +__path/wildcard__ # select 1 choice +__3$$path/wildcard__ # select 3 choices +__2-3$$path/wildcard__ # select 2 to 3 choices +__r2-3$$path/wildcard__ # select 2 to 3 choices allowing repetition +__2-3$$ / $$path/wildcard__ # select 2 to 3 choices with separator " / " +__path/wildcard(var=value)__ # select 1 choice using the specified variable value in the evaluation. +``` + +#### Detection of remaining wildcards + +This extension should run after any other wildcard extensions, so if you don't use the internal wildcards processing, any remaining wildcards present in the prompt or negative_prompt at this point must be invalid. Usually you might not notice this problem until you check the image metadata, so this option gives you some ways to detect and treat the problem. ### Set command @@ -73,6 +148,27 @@ The format is: ```text value +value +value +value +``` + +The `evaluate` parameter makes it so the value of the variable is evaluated at this moment, instead of when it is used. + +With the `add` parameter the value is added to the current value of the variable. It does not force an immediate evaluation of the old nor the added value. + +The Dynamic Prompts format also works: + +```text +${var=value} +${var=!value} # immmediate evaluation +``` + +If also supports the addition as an extension of the Dynamic Prompts format: + +```text +${var+=value} +${var+=!value} ``` ### Echo command @@ -83,23 +179,57 @@ The format is: ```text +default +``` + +The Dynamic Prompts format is: + +```text +${var} +${var:default} ``` ### If command This command allows you to filter content based on conditions. -The format is: +The full format is: ```text content onecontent twoother content ``` -The *conditionN* compares a variable with a value. The operation can be `eq`, `ne`, `gt`, `lt`, `ge`, `le` and the value can be a quoted string or an integer. +The *conditionN* compares a variable with a value or a list of values. The allowed formats are: -The variable can be one set with the `set` command or special variables like: +```text +[not] variable +[not] variable operation value +variable [not] operation value +[not] variable operation (value1,value2...) +variable [not] operation (value1,value2...) +``` + +When there is no value it will check if the variable is truthy. + +For a simple value the allowed operations are `eq`, `ne`, `gt`, `lt`, `ge`, `le`, `contains` and the value can be a quoted string or an integer. + +For a list of values the allowed operations are `contains`, `in` and the value of the variable is checked against all the elements of the list until one matches. + +The variable can be one set with the `set` or `add` commands or you can use internal variables like these (names starting with an underscore are reserved): * `_sd` : the loaded model version (`"sd1"`, `"sd2"`, `"sdxl"`) +* `_sdname` : the loaded model filename (without path) +* `_sdfullname`: the loaded model filename (with path) +* `_is_sd`: true if the loaded model version is any version of SD +* `_is_sd1`: true if the loaded model version is SD 1.x +* `_is_sd2`: true if the loaded model version is SD 2.x +* `_is_sdxl`: true if the loaded model version is SDXL (includes Pony models) +* `_is_ssd`: true if the loaded model version is SSD (Segmind Stable Diffusion 1B). Note that for an SSD model `_is_sdxl` will also be true. +* `_is_sdxl_no_ssd`: true if the loaded model version is SDXL and not an SSD model. +* `_is_pony`: true if the loaded model version is SDXL and a Pony model (based on its filename). Note that for a pony model `_is_sdxl` will also be true. +* `_is_sdxl_no_pony`: true if the loaded model version is SDXL and not a Pony model. +* `_is_sd3`: true if the loaded model version is SD 3.x +* `_is_flux`: true if the loaded model is Flux Any `elif`s (there can be multiple) and the `else` are optional. @@ -108,9 +238,9 @@ Any `elif`s (there can be multiple) and the `else` are optional. (multiline to be easier to read) ```text - test sd1 - test sd2 - test sdxl + test sd1x + test pony + test sdxl unknown model ``` @@ -153,23 +283,7 @@ Then, if that option is chosen this extension will process it later and move tha #### Old format -The old format is still supported (for now) and is like this: - -```text - -``` - -And an optional position can be specified like this: - -```text - -``` - -With the insertion point like this: - -```text - -``` +The old format (``) is not supported anymore. ### Notes on negative commands @@ -198,35 +312,44 @@ This should still work as intended, and the only negative point i see is the unn ### General settings -* **Debug**: writes debugging information to the console. +* **Debug level**: what to write to the console. Note: in SD.Next debug messages only show if you launch it with the --debug argument. +* **Pony substrings**: list of substrings to detect a Pony model. +* **Apply in img2img**: check if you want to do the processing in img2img processes (does not apply to ComfyUI node). + +### Wildcard settings + +* **Process wildcards**: you can choose to process them 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. * **What to do with remaining wildcards?**: select what do you want to do with any found wildcards. * **Ignore**: do not try to detect wildcards. * **Remove**: detect wildcards and remove them. * **Add visible warning**: detect wildcards and add a warning text to the prompt, that hopefully produces a noticeable generation. * **Stop the generation**: detect wildcards and stop the generation. - -### Content removal settings - -* **Remove extra network tags**: removes all extra network tags. +* **Default separator used when adding multiple choices**: what do you want to use by default to separate multiple choices when the options allow it (by default it's ", "). +* **Keep the order of selected choices**: if checked, a multiple choice construct will return them in the order they are in the construct. ### Send to negative prompt settings -* **Apply in img2img**: check if you want to do this processing in img2img processes. * **Separator used when adding to the negative prompt**: you can specify the separator used when adding to the negative prompt (by default it's ", "). * **Ignore repeated content**: it ignores repeated content to avoid repetitions in the negative prompt. * **Join attention modifiers (weights) when possible**: it joins attention modifiers when possible (joins into one, multipliying their values). ### Clean up settings -* **Apply in img2img**: check if you want to do this processing in img2img processes. * **Remove empty constructs**: removes attention/scheduling/alternation constructs when they are invalid. * **Remove extra separators**: removes unnecessary separators. This applies to the configured separator and regular commas. * **Remove additional extra separators**: removes unnecessary separators at start or end of lines. This applies to the configured separator and regular commas. * **Clean up around BREAKs**: removes consecutive BREAKs and unnecessary commas and space around them. +* **Use EOL instead of Space before BREAKs**: add a newline before BREAKs. * **Clean up around ANDs**: removes consecutive ANDs and unnecessary commas and space around them. +* **Use EOL instead of Space before ANDs**: add a newline before ANDs. * **Clean up around extra network tags**: removes spaces around them. * **Remove extra spaces**: removes other unnecessary spaces. +### Content removal settings + +* **Remove extra network tags**: removes all extra network tags. + ## License MIT diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..9fe3f43 --- /dev/null +++ b/__init__.py @@ -0,0 +1,28 @@ +""" +@author: ACB +@title: Prompt Post Processor +@nickname: ACB PPP +@description: Node for processing prompts. Includes the following options: send to negative prompt, set variables, if/elif/else command for conditional content, wildcards and choices. +""" + +import sys +import os + +sys.path.append(os.path.dirname(os.path.abspath(__file__))) + +from .ppp_comfyui import PromptPostProcessorComfyUINode + +NODE_CLASS_MAPPINGS = {"ACBPromptPostProcessor": PromptPostProcessorComfyUINode} + +NODE_DISPLAY_NAME_MAPPINGS = {"ACBPromptPostProcessor": "ACB Prompt Post Processor"} + +MANIFEST = { + "name": "ACB Prompt Post Processor", + "version": PromptPostProcessorComfyUINode.VERSION, + "author": "ACB", + "project": "https://github.com/acorderob/sd-webui-prompt-postprocessor", + "description": "Node for processing prompts", + "license": "MIT", +} + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/grammar.lark b/grammar.lark new file mode 100644 index 0000000..af23df2 --- /dev/null +++ b/grammar.lark @@ -0,0 +1,98 @@ +%import common (LETTER, DIGIT, INT, CNAME, SIGNED_NUMBER, NUMBER) + +_WHITESPACE: /\s+/ +STRING: /("(?!"").*?(?${]|\\.)+/s // exclude only the starting ones +?plain_choice: /((?!__|\bAND\b|\${|\$\$)[^\\()\[\]:<>${|}~@]|\\.)+/s // add the specific internal choice ones +?plain_alternate: /((?!__|\bAND\b|\${)[^\\()\[\]:<>${|]|\\.)+/s // add the specific internal alternate ones +?plain_var: /((?!__|\bAND\b|\${)[^\\()\[\]:<>${}]|\\.)+/s // add the specific internal var ones +?specialchars: /[_{()\[\]:<>]|\$(?![{$])/ // include only the starting ones +?specialchars_negtag: /[_{()\[\]:<>!|}]|\$(?![{$])/ // add the internal negtag ones +?specialchars_alternate: /[_{()\[\]:<>|]|\$(?![{$])/ // add the internal alternate ones +?specialchars_choice: /[_{()\[\]:<>|}]|\$(?![{$])/ // add the internal choice ones +?specialchars_var: /[_{()\[\]:<>}]|\$(?![{$])/ // add the internal var ones +?numpar: _WHITESPACE? SIGNED_NUMBER _WHITESPACE? + +start: promptcomp | content + +// prompt composition with AND +promptcomp.4: promptcomppart ([":" numpar] (/\bAND\b/ promptcomppart [":" numpar])+)+ +promptcomppart: content + +// simple prompts +?content.2: (old_content | new_content | plain | specialchars)* +?content_choice.2: (old_content | new_content | plain_choice | specialchars_choice)* +?content_var.2: (old_content | new_content | plain_var | specialchars_var)* +?content_negtag.2: (old_content | new_content_negtag | plain | specialchars_negtag)* +?content_alternate.2: (old_content | new_content | plain_alternate | specialchars_alternate)* +?old_content.2: (emphasized | deemphasized | scheduled | alternate | extranetworktag)+ +?new_content.3: (variableset | variableuse | commandstn | commandstni | commandset | commandecho | commandif | wildcard | choices)+ +?new_content_negtag.3: (variableset | variableuse | commandset | commandecho | commandif | wildcard | choices)+ + +// attention modifiers +emphasized: "(" content [":" numpar] ")" +deemphasized: "[" content "]" + +// prompt scheduling and alternation +alternate: "[" alternateoption ("|" alternateoption)+ "]" +alternateoption: content_alternate +scheduled: "[" [content ":"] content ":" numpar "]" + +// extra network tags +extranetworktag: "<" /(?!ppp:)[^>]+/ ">" + +// command: stn (send to negative) +commandstn: "" content_negtag "" +commandstni: "" + +// command: if +commandif.2: commandif_if commandif_elif* commandif_else? "" +commandif_if: "" ifvalue +commandif_elif: "" ifvalue +commandif_else: "" ifvalue +ifvalue: content + +// conditions +condition: conditionsimplevalue | conditionlistvalue | conditionnocomparison +conditionnocomparison: (/not/ _WHITESPACE)? IDENTIFIER +conditionsimplevalue: (/not/ _WHITESPACE)? IDENTIFIER _WHITESPACE (/not/ _WHITESPACE)? /eq|ne|gt|lt|ge|le|contains/ _WHITESPACE SIMPLEVALUE +conditionlistvalue: (/not/ _WHITESPACE)? IDENTIFIER _WHITESPACE (/not/ _WHITESPACE)? /contains|in/ _WHITESPACE listvalue +IDENTIFIER: CNAME +SIMPLEVALUE: STRING | INT | BOOLEAN +listvalue: "(" _WHITESPACE? SIMPLEVALUE (_WHITESPACE? "," _WHITESPACE? SIMPLEVALUE)* _WHITESPACE? ")" + +// command: set +commandset: "" content "" + +// command: echo +commandecho: "" [ content "" ] + +// variable set +variableset.2: "${" _WHITESPACE? IDENTIFIER _WHITESPACE? [/\+/] "=" [/!/] varvalue "}" + +// variable use +variableuse.2: "${" _WHITESPACE? IDENTIFIER _WHITESPACE? [":" varvalue] "}" +varvalue: content_var + +// wildcards +wildcard.2: "__" [choicesoptions_sampler | (choicesoptions _WHITESPACE? "$$")] /(?:(?!__|\$\$|\()\S)+/ [ wildcard_var ] "__" +wildcard_var: "(" _WHITESPACE? IDENTIFIER _WHITESPACE? "=" varvalue ")" + +// choices +choices.2: "{" [choicesoptions_sampler | (choicesoptions _WHITESPACE? "$$")] choice ("|" choice)* "}" + +choicesoptions: [choicesoptions_sampler] [_WHITESPACE? choicesoptions_rep] (([_WHITESPACE? choicesoptions_from] "-" [_WHITESPACE? choicesoptions_to]) | [_WHITESPACE? choicesoptions_num] ) [_WHITESPACE? choicesoptions_sep] +choicesoptions_sampler: /[~@]/ // ~ for random, @ for cyclical +choicesoptions_rep: /r/ +choicesoptions_num: INT +choicesoptions_from: INT +choicesoptions_to: INT +choicesoptions_sep: "$$" plain + +choice: [[_WHITESPACE? choiceweight] [_WHITESPACE? choiceif] _WHITESPACE? "::"] choicevalue +choiceweight: NUMBER +choiceif: "if" _WHITESPACE condition +choicevalue: content_choice diff --git a/install.py b/install.py new file mode 100644 index 0000000..5fe51f0 --- /dev/null +++ b/install.py @@ -0,0 +1,5 @@ +import os +import launch + +requirements_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "requirements.txt") +launch.run_pip(f'install -r "{requirements_filename}"', "requirements for Prompt Post-Processor") diff --git a/metadata.ini b/metadata.ini new file mode 100644 index 0000000..1b61970 --- /dev/null +++ b/metadata.ini @@ -0,0 +1,5 @@ +[Extension] +Name = sd-webui-prompt-postprocessor + +[Scripts] +After = sd-dynamic-prompts, stable-diffusion-webui-wildcards diff --git a/ppp.py b/ppp.py index dc7660d..0d875ae 100644 --- a/ppp.py +++ b/ppp.py @@ -1,169 +1,127 @@ -from collections import namedtuple -import re +import fnmatch +import logging import math +import os +import re +import textwrap +import time +from collections import namedtuple +from enum import Enum +from typing import Callable, Optional + import lark +import lark.parsers +import numpy as np + +from ppp_logging import DEBUG_LEVEL +from ppp_wildcards import PPPWildcards class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-instance-attributes """ The PromptPostProcessor class is responsible for processing and manipulating prompt strings. - - Attributes: - NAME (str): The name of the prompt post-processor. - VERSION (str): The version of the prompt post-processor. - DEFAULT_STN_SEPARATOR (str): The default separator used for content sent to the negative prompt. - - Methods: - __init__(self, script, opts=None, is_i2i=False): Initializes the PromptPostProcessor instance. - formatOutput(self, text: str) -> str: Formats the output text by encoding and decoding it. - STNTree: A nested class for interpreting and processing a tree generated by the prompt parser. - __find_tags(self, prompt): Finds tags in the given prompt and returns the modified prompt and tag information. - __add_to_insertion_points(self, negative_prompt, add_at_insertion_point): Adds the negative prompt to the insertion points. - __add_to_start(self, negative_prompt, add_at_start): Adds the elements in `add_at_start` list to the start of the `negative_prompt` string. - __add_to_end(self, negative_prompt, add_at_end): Adds the elements in `add_at_end` list to the end of `negative_prompt` string. - CleanupTree: A nested class for cleaning up a prompt parsed into a tree. - __cleanup(self, prompt, negative_prompt): Cleans up the prompt and negative prompt by removing extra spaces, empty constructs, and extra separators. - cleanup_text(self, text): Cleans up the given text by removing extra separators, breaks, and spaces. - trim_text(self, text): Trims the given text based on the specified cleanup options. - process_prompt(self, original_prompt, original_negative_prompt): Process the prompt and negative prompt by moving content to the negative prompt, and cleaning up. """ NAME = "Prompt Post-Processor" - VERSION = "2.4.0" + VERSION = (2, 5, 0) + + class IFWILDCARDS_CHOICES(Enum): + ignore = "ignore" + remove = "remove" + warn = "warn" + stop = "stop" DEFAULT_STN_SEPARATOR = ", " - IFWILDCARDS_CHOICES = { - "ignore": "Ignore", - "remove": "Remove", - "warn": "Add visible warning", - "stop": "Stop the generation", - } + DEFAULT_PONY_SUBSTRINGS = ",".join(["pony", "pny", "pdxl"]) + DEFAULT_CHOICE_SEPARATOR = ", " WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK ' - WILDCARD_STOP = "INVALID WILDCARD!\nBREAK " + WILDCARD_STOP = "INVALID WILDCARD! {0}\nBREAK " UNPROCESSED_STOP = "UNPROCESSED CONSTRUCTS!\nBREAK " def __init__( self, - script, - processing, - state, - opts=None, - is_i2i=False, + logger: logging.Logger, + interrupt: Optional[Callable], + model_info: dict[str, any], + options: Optional[dict[str, any]] = None, + grammar_content: Optional[str] = None, + wildcards_obj: PPPWildcards = None, ): """ Initializes the PPP object. Args: - script: The script object. - opts: Optional. The options object for configuring PPP behavior. + logger: The logger object. + interrupt: The interrupt function. + model_info: A dictionary with information for the loaded model. + 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. """ - self.script = script - self.logger = script.ppp_logger - self.opts = opts - self.is_i2i = is_i2i - self.processing = processing - self.state = state - self.variables = {} - self.debug = getattr(self.opts, "ppp_gen_debug", False) if opts is not None else False - if opts is not None and getattr(opts, "prompt_attention", "") == "Compel parser": - self.logger.warning("Compel parser is not supported!") - self.ifwildcards = ( - getattr(opts, "ppp_gen_ifwildcards", self.IFWILDCARDS_CHOICES["ignore"]) - if opts is not None - else self.IFWILDCARDS_CHOICES["ignore"] - ) + self.logger = logger + self.rng = np.random.default_rng() # gets seeded on each process prompt call + self.the_interrupt = interrupt + self.options = options + self.model_info = model_info + self.wildcard_obj = wildcards_obj - self.stn_doi2i = getattr(opts, "ppp_stn_doi2i", False) if opts is not None else False - self.stn_ignore_repeats = getattr(opts, "ppp_stn_ignorerepeats", True) if opts is not None else True - self.stn_join_attention = getattr(opts, "ppp_stn_joinattention", True) if opts is not None else True - self.stn_separator = ( - getattr(opts, "ppp_stn_separator", self.DEFAULT_STN_SEPARATOR) - if opts is not None - else self.DEFAULT_STN_SEPARATOR + # General options + self.debug_level = DEBUG_LEVEL(options.get("debug_level", DEBUG_LEVEL.none.value)) + self.pony_substrings = list( + x.strip() for x in (str(options.get("pony_substrings", self.DEFAULT_PONY_SUBSTRINGS))).split(",") ) + # Wildcards options + self.wil_process_wildcards = options.get("process_wildcards", True) + self.wil_keep_choices_order = options.get("keep_choices_order", False) + self.wil_choice_separator = options.get("choice_separator", self.DEFAULT_CHOICE_SEPARATOR) + self.wil_ifwildcards = self.IFWILDCARDS_CHOICES(options.get("if_wildcards", self.IFWILDCARDS_CHOICES.ignore.value)) + # Send to negative options + self.stn_ignore_repeats = options.get("stn_ignore_repeats", True) + self.stn_join_attention = options.get("stn_join_attention", True) + self.stn_separator = options.get("stn_separator", self.DEFAULT_STN_SEPARATOR) + # Cleanup options + self.cup_extraspaces = options.get("cleanup_extra_spaces", True) + self.cup_emptyconstructs = options.get("cleanup_empty_constructs", True) + self.cup_extraseparators = options.get("cleanup_extra_separators", True) + self.cup_extraseparators2 = options.get("cleanup_extra_separators2", True) + self.cup_breaks = options.get("cleanup_breaks", True) + self.cup_breaks_eol = options.get("cleanup_breaks_eol", False) + self.cup_ands = options.get("cleanup_ands", True) + self.cup_ands_eol = options.get("cleanup_ands_eol", False) + self.cup_extranetworktags = options.get("cleanup_extranetwork_tags", False) + # Remove options + self.rem_removeextranetworktags = options.get("remove_extranetwork_tags", False) + + # if self.debug_level != DEBUG_LEVEL.none: + # self.logger.info(f"Detected model info: {model_info}") - self.cup_doi2i = getattr(opts, "ppp_cup_doi2i", False) if opts is not None else False - self.cup_extraspaces = getattr(opts, "ppp_cup_extraspaces", True) if opts is not None else True - self.cup_emptyconstructs = getattr(opts, "ppp_cup_emptyconstructs", True) if opts is not None else True - self.cup_extraseparators = getattr(opts, "ppp_cup_extraseparators", True) if opts is not None else True - self.cup_extraseparators2 = getattr(opts, "ppp_cup_extraseparators2", True) if opts is not None else True - self.cup_breaks = getattr(opts, "ppp_cup_breaks", True) if opts is not None else True - self.cup_breaks_eol = getattr(opts, "ppp_cup_breaks_eol", False) if opts is not None else False - self.cup_ands = getattr(opts, "ppp_cup_ands", True) if opts is not None else True - self.cup_ands_eol = getattr(opts, "ppp_cup_ands_eol", False) if opts is not None else False - self.cup_extranetworktags = getattr(opts, "ppp_cup_extranetworktags", False) if opts is not None else False - self.rem_removeextranetworktags = ( - getattr(opts, "ppp_rem_removeextranetworktags", False) if opts is not None else False - ) # Process with lark (debug with https://www.lark-parser.org/ide/) - self.__parser_complete = lark.Lark( - r""" - start: (promptcomp | specialchars)* - - // prompt composition with AND - promptcomp: promptcomppart ([":" numpar] (/\bAND\b/ promptcomppart [":" numpar])+)? - promptcomppart: prompt - - // prompt scheduling and alternation - alternate: "[" alternateoption ("|" alternateoption)+ "]" - alternateoption: prompt - scheduled: "[" [prompt ":"] prompt ":" numpar "]" - - // wildcard extension support - wildcard: "__" /(?:(?!__)\S)+/ "__" - choices: "{" choice ("|" choice)* "}" - // we ignore weight and any other parameters in each choice - choice: prompt - - // simple prompts - prompt: (emphasized | deemphasized | scheduled | alternate | extranetworktag | commandstn | commandstni | negtag | commandset | commandif | commandecho | wildcard | choices | plain)* - - nonegprompt: (emphasized | deemphasized | scheduled | alternate | extranetworktag | commandset | commandif | commandecho | wildcard | choices | plain)* - - // attention modifiers - emphasized: "(" prompt [":" numpar] ")" - deemphasized: "[" prompt "]" - - // extra network tags - extranetworktag: "<" /(?!!|ppp:)[^>]+/ ">" - - // negative tags - negtag: "" - negtagparameters: "!" /s|e|[ip]\d/ "!" - - // command: stn (send to negative) - commandstn: "" nonegprompt "" - commandstni: "" - - // command: if - commandif: commandif_if commandif_elif* commandif_else? "" - commandif_if: "" prompt - commandif_elif: "" prompt - commandif_else: "" prompt - condition: IDENTIFIER WHITESPACE /eq|ne|gt|lt|ge|le/ WHITESPACE VALUE - IDENTIFIER: CNAME - VALUE: STRING | INT - - // command: set - commandset: "" prompt "" - - // command: echo - commandecho: "" - - // plain text and weights - numpar: WHITESPACE? NUMBER WHITESPACE? - WHITESPACE: /\s+/ - STRING: /("(?!"").*?(?!{}]|\\.)+/s - specialchars: /[\]():|<>!{}]|\bAND\b/+ - %import common.SIGNED_NUMBER -> NUMBER - %import common.CNAME -> CNAME - %import common.INT -> INT - """, + if grammar_content is None: + grammar_filename = os.path.join(os.path.dirname(os.path.realpath(__file__)), "grammar.lark") + with open(grammar_filename, "r", encoding="utf-8") as file: + grammar_content = file.read() + self.parser_complete = lark.Lark( + grammar_content, propagate_positions=True, ) + self.parser_choice = lark.Lark( + grammar_content, + propagate_positions=True, + start="choice", + ) + self.parser_choicesoptions = lark.Lark( + grammar_content, + propagate_positions=True, + start="choicesoptions", + ) + self.__init_sysvars() + self.user_variables = {} - def formatOutput(self, text: str): + def interrupt(self): + if self.the_interrupt is not None: + self.the_interrupt() + + def formatOutput(self, text: str) -> str: """ Formats the output text by encoding it using unicode_escape and decoding it using utf-8. @@ -175,469 +133,36 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in """ return text.encode("unicode_escape").decode("utf-8") - def eval_condition(self, cond_var, cond_comp, cond_value): - """ - Evaluate a condition based on the given variable, comparison, and value. + def __init_sysvars(self): + self.system_variables = {} + sdchecks = { + "sd1": self.model_info.get("is_sd1", False), + "sd2": self.model_info.get("is_sd2", False), + "sdxl": self.model_info.get("is_sdxl", False), + "sd3": self.model_info.get("is_sd3", False), + "flux": self.model_info.get("is_flux", False), + "": True, + } + self.system_variables["_sd"] = [k for k, v in sdchecks.items() if v][0] + model_filename = self.model_info.get("model_filename", "") + is_pony = any(s in model_filename.lower() for s in self.pony_substrings) + is_ssd = self.model_info.get("is_ssd", False) + self.system_variables["_sdfullname"] = model_filename + self.system_variables["_sdname"] = os.path.basename(model_filename) + self.system_variables["_is_sd1"] = sdchecks["sd1"] + self.system_variables["_is_sd2"] = sdchecks["sd2"] + self.system_variables["_is_sdxl"] = sdchecks["sdxl"] + self.system_variables["_is_ssd"] = is_ssd + self.system_variables["_is_sdxl_no_ssd"] = sdchecks["sdxl"] and not is_ssd + self.system_variables["_is_pony"] = sdchecks["sdxl"] and is_pony + self.system_variables["_is_sdxl_no_pony"] = sdchecks["sdxl"] and not is_pony + self.system_variables["_is_sd3"] = sdchecks["sd3"] + self.system_variables["_is_sd"] = sdchecks["sd1"] or sdchecks["sd2"] or sdchecks["sdxl"] or sdchecks["sd3"] + self.system_variables["_is_flux"] = sdchecks["flux"] - Args: - cond_var (str): The variable to be compared. - cond_comp (str): The comparison operator. - cond_value (str): The value to be compared with. - - Returns: - bool: The result of the condition evaluation. - """ - result = False - var_value = self.variables.get(cond_var, "") - value_is_int = False - if cond_value.startswith('"') or cond_value.startswith("'"): - cond_value = cond_value[1:-1] - else: - value_is_int = True - if cond_comp == "eq": - result = var_value == cond_value - elif cond_comp == "ne": - result = var_value != cond_value - elif cond_comp == "gt": - if value_is_int: - result = int(var_value) > int(cond_value) - else: - result = var_value > cond_value - elif cond_comp == "lt": - if value_is_int: - result = int(var_value) < int(cond_value) - else: - result = var_value < cond_value - elif cond_comp == "ge": - if value_is_int: - result = int(var_value) >= int(cond_value) - else: - result = var_value >= cond_value - elif cond_comp == "le": - if value_is_int: - result = int(var_value) <= int(cond_value) - else: - result = var_value <= cond_value - return result - - class VARTree(lark.visitors.Interpreter): - """ - A class for interpreting and processing a tree generated by the prompt parser in the context of initializing variables. - - Attributes: - __ppp (object): The instance of the parent class. - variables (dict): The dictionary to store the detected variables. - - Methods: - commandset(node): Process a set command in the tree and add it to the dictionary of variables. - commandif(node): Process an if command in the tree. - start(node): Process the given tree and perform necessary operations on the found variables. - """ - - def __init__(self, ppp, prompt): - super().__init__() - self.__ppp = ppp - self.__prompt = prompt - - def commandset(self, node): - """ - Process a set command in the tree and add it to the dictionary of variables. - - Args: - node (Node): The tree node representing the set command. - - Returns: - None - """ - variable = node.children[1] - content = node.children[2] - value = self.__prompt[content.meta.start_pos : content.meta.end_pos] - self.__ppp.variables[variable.value] = value - if self.__ppp.debug: - self.__ppp.logger.info(f"Setting variable {variable} to '{value}'") - - def commandif(self, node): - """ - Process an if command in the tree. - - Args: - node (Node): The tree node representing the if command. - - Returns: - None - """ - for n in node.children: - if len(n.children) >= 3 and n.children[1].data.type != "WHITESPACE": - # has a condition - condition = n.children[1] - cond_var = condition.children[0].value - cond_comp = condition.children[2].value - cond_value = condition.children[4].value - # there could be a whitespace node here - content = n.children[-1] - if self.__ppp.eval_condition(cond_var, cond_comp, cond_value): - if self.__ppp.debug: - conditiontext = self.__prompt[condition.meta.start_pos : condition.meta.end_pos] - contenttext = self.__prompt[content.meta.start_pos : content.meta.end_pos] - self.__ppp.logger.info(f"Applying if condition ({conditiontext}): {contenttext}") - self.visit(content) - return - else: - # its an else - content = n.children[-1] - if self.__ppp.debug: - contenttext = self.__prompt[content.meta.start_pos : content.meta.end_pos] - self.__ppp.logger.info(f"Applying if else: {contenttext}") - self.visit(content) - - def start(self, node): - self.visit_children(node) - - class STNTree(lark.visitors.Interpreter): - """ - A class for interpreting and processing a tree generated by the prompt parser in the context of processing send to negative prompt commands. - - Attributes: - __ppp (object): The instance of the parent class. - __prompt (str): The prompt string. - __shell (list): The list of accumulated shell elements. - __negtags (list): The list of negative tags. - __already_processed (list): The list of already processed content. - add_at (dict): The dictionary to store the content to be added at different positions. - remove (list): The list of content to be removed from the prompt. - - Methods: - __get_numpar_value(numpar): Get the numerical value from a numpar object. - scheduled(tree): Process a scheduling construct in the tree and add it to the accumulated shell. - alternate(tree): Process an alternation construct in the tree and add it to the accumulated shell. - emphasized(tree): Process an attention change construct in the tree and add it to the accumulated shell. - deemphasized(tree): Process a decrease attention construct in the tree and add it to the accumulated shell. - negtag(tree): Process a negative tag in the tree and add it to the list of negative tags. - commandstn(tree): Process a send to negative command in the tree and add it to the list of negative tags. - commandstni(tree): Process a send to negative command in the tree and add it to the list of negative tags. - start(tree): Process the given tree and perform necessary operations on the found negative tags. - """ - - def __init__(self, ppp, prompt, add_at, insertion_at): - super().__init__() - self.__ppp = ppp - self.__prompt = prompt - self.AccumulatedShell = namedtuple("AccumulatedShell", ["type", "data", "position"]) - AccumulatedShell = self.AccumulatedShell - self.__shell: list[AccumulatedShell] = [] - self.NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"]) - NegTag = self.NegTag - self.__negtags: list[NegTag] = [] - self.__already_processed = [] - self.add_at = add_at - self.insertion_at = insertion_at - self.remove = [] - - def __get_numpar_value(self, numpar): - """ - Get the numerical value from a numpar object. - - Args: - numpar (object): The numpar object to extract the value from. - - Returns: - float: The numerical value extracted from the numpar object. - """ - return float(next(x for x in numpar.children if x.type == "NUMBER").value) - - def scheduled(self, node): - """ - Process a scheduling construct in the tree and add it to the accumulated shell. - - Args: - node (Node): The tree node representing the scheduling construct. - - Returns: - None - """ - treemetaposition = ( - [node.meta.start_pos, node.meta.end_pos] if hasattr(node, "meta") and not node.meta.empty else None - ) - if len(node.children) > 2: # before & after - before = node.children[0] - else: - before = None - after = node.children[-2] - numpar = node.children[-1] - pos = self.__get_numpar_value(numpar) - if pos >= 1: - pos = int(pos) - # self.__shell.append(self.AccumulatedShell("sc", pos, treemetaposition)) - if before is not None and hasattr(before, "data"): - if self.__ppp.debug: - before_metaposition = ( - [before.meta.start_pos, before.meta.end_pos] - if hasattr(before, "meta") and not before.meta.empty - else "?" - ) - self.__ppp.logger.info(f"Shell scheduled before at {before_metaposition} with position {pos}") - self.__shell.append(self.AccumulatedShell("scb", pos, treemetaposition)) - self.visit(before) - self.__shell.pop() - if hasattr(after, "data"): - if self.__ppp.debug: - after_metaposition = ( - [after.meta.start_pos, after.meta.end_pos] - if hasattr(after, "meta") and not after.meta.empty - else "?" - ) - self.__ppp.logger.info(f"Shell scheduled after at {after_metaposition} with position {pos}") - self.__shell.append(self.AccumulatedShell("sca", pos, treemetaposition)) - self.visit(after) - self.__shell.pop() - # self.__shell.pop() - - def alternate(self, node): - """ - Process an alternation construct in the tree and add it to the accumulated shell. - - Args: - node (Node): The tree node representing the alternation construct. - - Returns: - None - """ - treemetaposition = ( - [node.meta.start_pos, node.meta.end_pos] if hasattr(node, "meta") and not node.meta.empty else None - ) - # self.__shell.append(self.AccumulatedShell("al", len(tree.children), treemetaposition)) - for i, opt in enumerate(node.children): - if self.__ppp.debug: - metaposition = ( - [opt.meta.start_pos, opt.meta.end_pos] if hasattr(opt, "meta") and not opt.meta.empty else "?" - ) - self.__ppp.logger.info(f"Shell alternate at {metaposition} option {i+1}") - if hasattr(opt, "data"): - self.__shell.append( - self.AccumulatedShell("alo", {"pos": i + 1, "len": len(node.children)}, treemetaposition) - ) - self.visit(opt) - self.__shell.pop() - # self.__shell.pop() - - def emphasized(self, node): - """ - Process a attention change construct in the tree and add it to the accumulated shell. - - Args: - node (Node): The tree node representing the attention construct. - - Returns: - None - """ - treemetaposition = ( - [node.meta.start_pos, node.meta.end_pos] if hasattr(node, "meta") and not node.meta.empty else None - ) - numpar = node.children[-1] - weight = self.__get_numpar_value(numpar) if numpar is not None else 1.1 - if self.__ppp.debug: - self.__ppp.logger.info(f"Shell attention at {treemetaposition or '?'} with weight {weight}") - self.__shell.append(self.AccumulatedShell("at", weight, treemetaposition)) - self.visit_children(node) - self.__shell.pop() - - def deemphasized(self, node): - """ - Process a decrease attention construct in the tree and add it to the accumulated shell. - - Args: - node (Node): The tree node representing the decreased attention construct. - - Returns: - None - """ - weight = 0.9 - treemetaposition = ( - [node.meta.start_pos, node.meta.end_pos] if hasattr(node, "meta") and not node.meta.empty else None - ) - if self.__ppp.debug: - self.__ppp.logger.info(f"Shell attention at {treemetaposition or '?'} with weight {weight}") - self.__shell.append(self.AccumulatedShell("at", weight, treemetaposition)) - self.visit_children(node) - self.__shell.pop() - - def negtag(self, node): - """ - Process a negative tag in the tree and add it to the list of negative tags. - - Args: - node (Node): The tree node representing the negative tag. - - Returns: - None - """ - treemetaposition = ( - [node.meta.start_pos, node.meta.end_pos] if hasattr(node, "meta") and not node.meta.empty else None - ) - negtagparameters = node.children[0] - parameters = negtagparameters.children[0].value if negtagparameters is not None else "" - rest = [] - for x in node.children[1::]: - rest.append( - self.__prompt[x.meta.start_pos : x.meta.end_pos] - if hasattr(x, "meta") and not x.meta.empty - else x.value if hasattr(x, "value") else "" - ) - content = "".join(rest) - self.__negtags.append( - self.NegTag(node.meta.start_pos, node.meta.end_pos, content, parameters, self.__shell.copy()) - ) - if self.__ppp.debug: - self.__ppp.logger.info( - f"Negative tag at {treemetaposition or '?'}: {parameters or 'with no parameters'} : {self.__ppp.formatOutput(content)}" - ) - - def commandstn(self, node): - """ - Process a send to negative command in the tree and add it to the list of negative tags. - - Args: - node (Node): The tree node representing the stn command. - - Returns: - None - """ - treemetaposition = ( - [node.meta.start_pos, node.meta.end_pos] if hasattr(node, "meta") and not node.meta.empty else None - ) - negtagparameters = node.children[1] - if negtagparameters is not None: - if hasattr(negtagparameters, "children"): - parameters = negtagparameters.children[0].value - else: - parameters = negtagparameters.value - else: - parameters = "" - rest = [] - for x in node.children[2::]: - rest.append( - self.__prompt[x.meta.start_pos : x.meta.end_pos] - if hasattr(x, "meta") and not x.meta.empty - else x.value if hasattr(x, "value") else "" - ) - content = "".join(rest) - self.__negtags.append( - self.NegTag(node.meta.start_pos, node.meta.end_pos, content, parameters, self.__shell.copy()) - ) - if self.__ppp.debug: - self.__ppp.logger.info( - f"Negative tag at {treemetaposition or '?'}: {parameters or 'with no parameters'} : {self.__ppp.formatOutput(content)}" - ) - - def commandstni(self, node): - self.commandstn(node) - - def start(self, node): - """ - Process the given tree and perform necessary operations on the found negative tags. - - Args: - node: The tree to be processed. - - Returns: - None - """ - self.visit_children(node) - # process the found negative tags - for negtag in self.__negtags: - if self.__ppp.stn_join_attention: - # join consecutive attention elements - for i in range(len(negtag.shell) - 1, 0, -1): - if negtag.shell[i].type == "at" and negtag.shell[i - 1].type == "at": - negtag.shell[i - 1] = self.AccumulatedShell( - "at", - math.floor(100 * negtag.shell[i - 1].data * negtag.shell[i].data) - / 100, # we limit the new weight to two decimals - negtag.shell[i - 1].position, - ) - negtag.shell.pop(i) - start = "" - end = "" - for s in negtag.shell: - match s.type: - case "at": - if s.data == 0.9: - start += "[" - end = "]" + end - elif s.data == 1.1: - start += "(" - end = ")" + end - else: - start += "(" - end = f":{s.data})" + end - # case "sc": - case "scb": - start += "[" - end = f"::{s.data}]" + end - case "sca": - start += "[" - end = f":{s.data}]" + end - # case "al": - case "alo": - start += "[" + ("|" * int(s.data["pos"] - 1)) - end = ("|" * int(s.data["len"] - s.data["pos"])) + "]" + end - content = start + negtag.content + end - position = negtag.parameters or "s" - if position.startswith("i"): - n = int(position[1]) - self.insertion_at[n] = [negtag.start, negtag.end] - elif len(content) > 0: - if content not in self.__already_processed: - if self.__ppp.stn_ignore_repeats: - self.__already_processed.append(content) - if self.__ppp.debug: - self.__ppp.logger.info( - f"Adding content at position {position}: {self.__ppp.formatOutput(content)}" - ) - if position == "e": - self.add_at["end"].append(content) - elif position.startswith("p"): - n = int(position[1]) - self.add_at["insertion_point"][n].append(content) - else: # position == "s" or invalid - self.add_at["start"].append(content) - else: - self.__ppp.logger.warning(f"Ignoring repeated content: {self.__ppp.formatOutput(content)}") - # remove from prompt - self.remove.append([negtag.start, negtag.end]) - - def __find_tags(self, prompt, negative_prompt): - """ - Finds tags in the given prompt and returns the modified prompt and tag information. - - Args: - prompt (str): The input prompt. - - Returns: - tuple: A tuple containing the modified prompt and tag information. - - """ - add_at = {"start": [], "insertion_point": [[] for x in range(10)], "end": []} - insertion_at = [None for x in range(10)] - - tree = self.__parser_complete.parse(prompt) - readtree = self.STNTree(self, prompt, add_at, insertion_at) - readtree.visit(tree) - for r in readtree.remove[::-1]: - prompt = prompt[: r[0]] + prompt[r[1] :] - add_at = readtree.add_at # we only use the additions from the prompt - - tree = self.__parser_complete.parse(negative_prompt) - readtree = self.STNTree(self, negative_prompt, add_at, insertion_at) - readtree.visit(tree) - insertion_at = readtree.insertion_at # we only use the insertion positions from the negative prompt - - if self.debug: - self.logger.info(f"New negative additions: {add_at}") - self.logger.info(f"New negative indexes: {insertion_at}") - return prompt, negative_prompt, add_at, insertion_at - - def __add_to_insertion_points(self, negative_prompt, add_at_insertion_point, insertion_at): + def __add_to_insertion_points( + self, negative_prompt: str, add_at_insertion_point: list[str], insertion_at: list[tuple[int, int]] + ) -> str: """ Adds the negative prompt to the insertion points. @@ -674,7 +199,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in negative_prompt = self.stn_separator.join(add_at_insertion_point[n]) return negative_prompt - def __add_to_start(self, negative_prompt, add_at_start): + def __add_to_start(self, negative_prompt: str, add_at_start: list[str]) -> str: """ Adds the elements in `add_at_start` list to the start of the `negative_prompt` string. @@ -693,7 +218,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in negative_prompt = self.stn_separator.join(add_at_start) return negative_prompt - def __add_to_end(self, negative_prompt, add_at_end): + def __add_to_end(self, negative_prompt: str, add_at_end: list[str]) -> str: """ Adds the elements in `add_at_end` list to the end of `negative_prompt` string. @@ -712,400 +237,50 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in negative_prompt = self.stn_separator.join(add_at_end) return negative_prompt - def __sendtonegative(self, prompt, negative_prompt): + def __cleanup(self, text: str) -> str: """ - Modifies the prompt and negative_prompt by moving content from the prompt to the negative prompt. - - Args: - prompt (str): The prompt. - negative_prompt (str): The negative prompt. - - Returns: - tuple: A tuple containing the modified prompt and negative_prompt. - """ - if self.debug: - self.logger.info("Doing send-to-negative") - prompt, negative_prompt, add_at, insertion_at = self.__find_tags(prompt, negative_prompt) - negative_prompt = self.__add_to_insertion_points(negative_prompt, add_at["insertion_point"], insertion_at) - if len(add_at["start"]) > 0: - negative_prompt = self.__add_to_start(negative_prompt, add_at["start"]) - if len(add_at["end"]) > 0: - negative_prompt = self.__add_to_end(negative_prompt, add_at["end"]) - if self.debug: - self.logger.info(f"prompt after send-to-negative: {self.formatOutput(prompt)}") - self.logger.info(f"negative_prompt after send-to-negative: {self.formatOutput(negative_prompt)}") - return prompt, negative_prompt - - class TransformerTree(lark.visitors.Transformer_NonRecursive): - """ - Transformer class for detecting wildcards and/or cleaning up a prompt parsed into a tree. - - This class provides methods for transforming different constructs in a parse tree - based on certain conditions. It is used for detecting wildcards and cleaning up invalid - or empty constructs in the parse tree. - - Args: - ppp (object): An instance of the parent class `ppp`. - phase (str): The phase of the transformation. - - Attributes: - __ppp (object): An instance of the parent class `ppp`. - __phase (str): The phase of the transformation. - detectedWildcards (list): A list of detected wildcards. - - Methods: - promptcomp(children): Replicates prompt composition constructs. - scheduled(children): Replicates or removes scheduling constructs based on conditions. - alternate(children): Replicates or removes alternation constructs based on conditions. - emphasized(children): Replicates or removes attention constructs based on conditions. - deemphasized(children): Replicates or removes attention constructs based on conditions. - extranetworktag(children): Replicates extra network constructs. - numpar(children): Cleans up number parameter. - negtag(children): Replicates or removes negative tag constructs based on conditions. - commandstn(children): Replicates or removes send to negative constructs based on conditions. - commandstni(children): Replicates or removes send to negative constructs based on conditions. - commandif(children): Replicates or removes if constructs based on conditions. - commandset(children): Replicates or removes set constructs based on conditions. - commandecho(children): Replicates or removes echo constructs based on conditions. - wildcard(children): Replicates or removes wildcard constructs based on conditions. - choices(children): Replicates or removes choices constructs based on conditions. - choice(children): Replicates choices. - plain(children): Cleans up plain text based on conditions. - __default__(data, children, meta): Default method for joining children and cleaning up text based on conditions. - """ - - def __init__(self, ppp, phase="cleanup"): - super().__init__(visit_tokens=True) - self.__ppp = ppp - self.__phase = phase - self.detectedWildcards = [] - - def promptcomp(self, children): - r = children[0] - if len(children) > 1: - if children[1] is not None: - r += f":{children[1]}" - for i in range(2, len(children), 3): - if self.__phase == "cleanup" and self.__ppp.cup_ands: - r = re.sub(r"[, ]+$", "\n" if self.__ppp.cup_ands_eol else " ", r) - if r[-1:].isalnum(): # add space if needed - r += " " - r += "AND" - t = children[i + 1] - if self.__phase == "cleanup" and self.__ppp.cup_ands: - t = re.sub(r"^[, ]+", " ", t) - if t[0:1].isalnum(): # add space if needed - r += " " - r += t - if children[i + 2] is not None: - r += f":{children[i+2]}" - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.promptcomp: {r}") - return r - - def scheduled(self, children): - if len(children) == 0 and self.__phase == "cleanup" and self.__ppp.cup_emptyconstructs: - r = "" # remove invalid scheduling construct (probably this is not reachable) - else: - # replicate scheduling construct - if len(children) > 0 and children[0] is None: - r = f"[{':'.join(children[1:])}]" - else: - r = f"[{':'.join(children)}]" - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.scheduled: {r}") - return r - - def alternate(self, children): - if len(children) == 0 and self.__phase == "cleanup" and self.__ppp.cup_emptyconstructs: - r = "" # remove invalid alternation construct (probably this is not reachable) - else: - r = f"[{'|'.join(children)}]" # replicate alternation construct - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.alternate: {r}") - return r - - def emphasized(self, children): - if ( - (len(children) == 0 or children[0] == "") - and self.__phase == "cleanup" - and self.__ppp.cup_emptyconstructs - ): - r = "" # remove empty attention construct - else: - if len(children) > 1 and children[1] is not None: - r = f"({children[0]}:{children[1]})" # replicate attention construct with weight - else: - r = f"({children[0]})" # replicate attention construct without weight - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.emphasized: {r}") - return r - - def deemphasized(self, children): - if ( - (len(children) == 0 or children[0] == "") - and self.__phase == "cleanup" - and self.__ppp.cup_emptyconstructs - ): - r = "" # remove empty attention construct (invalid scheduling or alternation constructs end up here too?) - else: - r = f"[{children[0]}]" # replicate attention construct - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.deemphasized: {r}") - return r - - def extranetworktag(self, children): - if self.__phase == "removecontent" and self.__ppp.rem_removeextranetworktags: - r = "" # remove extra network construct - else: - r = f"<{children[0]}>" # replicate extra network construct - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.deemphasized: {r}") - return r - - def numpar(self, children): - r = next(x for x in children if x.type == "NUMBER").value.strip() # clean up number parameter - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.numpar: {r}") - return r - - def negtag(self, children): - if self.__phase == "cleanup": - r = "" # remove negative tag construct (there shouldn't be any at this point) - else: - parameters = f"!{children[0]}!" if children[0] is not None else "" - content = "".join(children[1::]) - r = f"" # replicate negative tag construct - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.negtag: {r}") - return r - - def commandstn(self, children): - if self.__phase == "cleanup": - r = "" # remove send to negative command construct (there shouldn't be any at this point) - else: - parameters = " " + children[1] if children[1] is not None else "" - content = "".join(children[2::]) - r = f"{content}" # replicate send to negative command construct - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.commandstn: {r}") - return r - - def commandstni(self, children): - if self.__phase == "cleanup": - r = "" # remove send to negative command construct (there shouldn't be any at this point) - else: - parameters = " " + children[1] if children[1] is not None else "" - r = f"" # replicate send to negative command construct - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.commandstni: {r}") - return r - - def commandset(self, children): - if self.__phase == "removecontent": - r = "" - else: - r = f"{children[2]}" # replicate set construct - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.commandset: {r}") - return r - - def commandecho(self, children): - if self.__phase == "removecontent": - r = self.__ppp.variables.get(children[1], "") - else: - r = f"{children[1]}" - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.commandecho: {r}") - return r - - def commandif(self, children): - if self.__phase == "removecontent": - selectedcontent = "" - for n in children: - if len(n) >= 3 and isinstance(n[1], list): - # has a condition - condition = n[1] - cond_var = condition[0].value - cond_comp = condition[2].value - cond_value = condition[4].value - # there could be a whitespace node here - content = n[-1] - if self.__ppp.eval_condition(cond_var, cond_comp, cond_value): - selectedcontent = content - break - else: - # its an else - selectedcontent = n[-1] - break - r = selectedcontent - else: - r = f"{''.join(children)}" - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.commandif: {r}") - return r - - def commandif_if(self, children): - if self.__phase == "removecontent": - return children - return f"{children[2]}" # replicate if construct - - def commandif_elif(self, children): - if self.__phase == "removecontent": - return children - return f"{children[2]}" # replicate elif construct - - def commandif_else(self, children): - if self.__phase == "removecontent": - return children - return f"{children[1]}" # replicate else construct - - def condition(self, children): - if self.__phase == "removecontent": - return children - return "".join(children) # replicate condition rule - - def wildcard(self, children): - r = f"__{children[0]}__" # replicate wildcard construct - self.detectedWildcards.append(r) - if self.__phase == "wildcards" and self.__ppp.ifwildcards == self.__ppp.IFWILDCARDS_CHOICES["remove"]: - r = "" - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.wildcard: {r}") - return r - - def choices(self, children): - r = "{" + "|".join(children) + "}" # replicate wildcard choices construct - self.detectedWildcards.append(r) - if self.__phase == "wildcards" and self.__ppp.ifwildcards == self.__ppp.IFWILDCARDS_CHOICES["remove"]: - r = "" - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.choices: {r}") - return r - - def choice(self, children): - r = children[0] # replicate choice - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.choice: {r}") - return r - - def plain(self, children): - r = children[0].value - if self.__phase == "cleanup": - r = self.__ppp.cleanup_text(r) # clean up plain text - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.plain: {r}") - return r - - def __default__(self, data, children, meta): - joined = "".join(children) # join all children - if self.__phase == "cleanup" and not re.match(r"[([<{]", joined): - # take care of cleaning the joints only if there are no constructs that can be affected - joined = self.__ppp.cleanup_text(joined) - if self.__ppp.debug: - self.__ppp.logger.info(f"{self.__phase} TransformerTree.default: {joined}") - return joined - - def __removecontent(self, prompt, negative_prompt): - """ - Removes content in the prompt and negative prompt. - - Args: - prompt (str): The original prompt. - negative_prompt (str): The negative prompt. - - Returns: - tuple: A tuple containing the processed prompt and negative prompt. - """ - - self.variables = {} - self.variables["_sd"] = ( - "sd1" - if self.processing.sd_model.is_sd1 - else "sd2" if self.processing.sd_model.is_sd2 else "sdxl" if self.processing.sd_model.is_sdxl else "" - ) - - transformtree = self.TransformerTree(self, phase="removecontent") - try: - if self.debug: - self.logger.info("Removing content in the prompt") - prompt_tree = self.__parser_complete.parse(prompt) - vartree = self.VARTree(self, prompt) - vartree.visit(prompt_tree) - prompt = transformtree.transform(prompt_tree) - except Exception as e: # pylint: disable=broad-except - self.logger.warning("parsing failed on prompt!: %s", e) - try: - if self.debug: - self.logger.info("Removing content in the negative prompt") - negativeprompt_tree = self.__parser_complete.parse(negative_prompt) - vartree = self.VARTree(self, negative_prompt) - vartree.visit(negativeprompt_tree) - negative_prompt = transformtree.transform(negativeprompt_tree) - except Exception as e: # pylint: disable=broad-except - self.logger.warning("parsing failed on negative prompt!: %s", e) - - if self.debug: - self.logger.info(f"prompt after content removal: {self.formatOutput(prompt)}") - self.logger.info(f"negative_prompt after content removal: {self.formatOutput(negative_prompt)}") - return prompt, negative_prompt - - def __cleanup(self, prompt, negative_prompt): - """ - Cleans up the prompt and negative prompt by removing extra spaces, empty constructs, and extra separators. - - Args: - prompt (str): The original prompt. - negative_prompt (str): The negative prompt. - - Returns: - tuple: A tuple containing the cleaned up prompt and negative prompt. - """ - - transformtree = self.TransformerTree(self, phase="cleanup") - try: - if self.debug: - self.logger.info("Cleaning the prompt") - prompt_tree = self.__parser_complete.parse(prompt) - prompt = self.trim_text(transformtree.transform(prompt_tree)) - except Exception as e: # pylint: disable=broad-except - self.logger.warning("Parsing failed on prompt!: %s", e) - try: - if self.debug: - self.logger.info("Cleaning the negative prompt") - negativeprompt_tree = self.__parser_complete.parse(negative_prompt) - negative_prompt = self.trim_text(transformtree.transform(negativeprompt_tree)) - except Exception as e: # pylint: disable=broad-except - self.logger.warning("Parsing failed on negative prompt!: %s", e) - - if self.debug: - self.logger.info(f"prompt after cleanup: {self.formatOutput(prompt)}") - self.logger.info(f"negative_prompt after cleanup: {self.formatOutput(negative_prompt)}") - return prompt, negative_prompt - - def cleanup_text(self, text): - """ - Cleans up the given text by removing extra separators, breaks, and spaces. This is called for plain text only or when there are no constructs. + Trims the given text based on the specified cleanup options. Args: text (str): The text to be cleaned up. Returns: - str: The cleaned up text. + str: The resulting text. """ - # NOTE: we can't use start/end of line regex since the text might only be a part of a larger line due to the parser + escapedSeparator = re.escape(self.stn_separator) if self.cup_extraseparators: # # sendtonegative separator # - escapedSeparator = re.escape(self.stn_separator) # collapse separators text = re.sub(r"(?:\s*" + escapedSeparator + r"\s*){2,}", self.stn_separator, text) + # remove separator after starting parenthesis or bracket + text = re.sub(r"(\s*" + escapedSeparator + r"\s*[([])(?:\s*" + escapedSeparator + r"\s*)+", r"\1", text) + # remove before colon or ending parenthesis or bracket + text = re.sub(r"(?:\s*" + escapedSeparator + r"\s*)+([:)\]]\s*" + escapedSeparator + r"\s*)", r"\1", text) + if self.cup_extraseparators2: + # remove at start of prompt or line + text = re.sub(r"^(?:\s*" + escapedSeparator + r"\s*)+", "", text, flags=re.MULTILINE) + # remove at end of prompt or line + text = re.sub(r"(?:\s*" + escapedSeparator + r"\s*)+$", "", text, flags=re.MULTILINE) + if self.cup_extraseparators: # # regular comma separator # # collapse separators text = re.sub(r"(?:\s*,\s*){2,}", ", ", text) + # remove separators after starting parenthesis or bracket + text = re.sub(r"(\s*,\s*[([])(?:\s*,\s*)+", r"\1", text) + # remove separators before colon or ending parenthesis or bracket + text = re.sub(r"(?:\s*,\s*)+([:)\]]\s*,\s*)", r"\1", text) + if self.cup_extraseparators2: + # remove at start of prompt or line + text = re.sub(r"^(?:\s*,\s*)+", "", text, flags=re.MULTILINE) + # remove at end of prompt or line + text = re.sub(r"(?:\s*,\s*)+$", "", text, flags=re.MULTILINE) + if self.cup_breaks_eol: + # replace spaces before break with EOL + text = re.sub(r"[, ]+BREAK\b", "\nBREAK", text) if self.cup_breaks: # collapse separators and commas before BREAK text = re.sub(r"[, ]+BREAK\b", " BREAK", text) @@ -1115,63 +290,14 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in text = re.sub(r"[, ]+BREAK[, ]+", " BREAK ", text) # collapse BREAKs text = re.sub(r"\bBREAK(?:\s+BREAK)+\b", " BREAK ", text) - if self.cup_extraspaces: - # remove spaces before comma - text = re.sub(r"[ ]+,", ",", text) - # collapse spaces - text = re.sub(r"[ ]{2,}", " ", text) - return text - - def trim_text(self, text): - """ - Trims the given text based on the specified cleanup options. This is only called for the reconstructed prompt. - - Args: - text (str): The text to be trimmed. - - Returns: - str: The trimmed text. - """ - # NOTE: here we can only do cleanups that can be done on the whole text, including inside constructs and around them - if self.cup_extraseparators: - # - # sendtonegative separator - # - escapedSeparator = re.escape(self.stn_separator) - # remove duplicate separator after starting parenthesis or bracket - text = re.sub(r"(\s*" + escapedSeparator + r"\s*[([])\s*" + escapedSeparator + r"\s*", r"\1", text) - # remove before colon or ending parenthesis or bracket - text = re.sub(r"\s*" + escapedSeparator + r"\s*([:)\]]\s*" + escapedSeparator + r"\s*)", r"\1", text) - if self.cup_extraseparators2: - # remove at start of prompt or line - text = re.sub(r"^(?:\s*" + escapedSeparator + r"\s*)", "", text, flags=re.MULTILINE) - # remove at end of prompt or line - text = re.sub(r"(?:\s*" + escapedSeparator + r"\s*)$", "", text, flags=re.MULTILINE) - if self.cup_extraseparators: - # - # regular comma separator - # - # remove duplicate separators after starting parenthesis or bracket - text = re.sub(r"(\s*,\s*[([])\s*,\s*", r"\1", text) - # remove duplicate separators before colon or ending parenthesis or bracket - text = re.sub(r"\s*,\s*([:)\]]\s*,\s*)", r"\1", text) - if self.cup_extraseparators2: - # remove at start of prompt or line - text = re.sub(r"^\s*,\s*", "", text, flags=re.MULTILINE) - # remove at end of prompt or line - text = re.sub(r"\s*,\s*$", "", text, flags=re.MULTILINE) - if self.cup_breaks_eol: - # replace spaces before break with EOL - text = re.sub(r"[, ]+BREAK\b", "\nBREAK", text) - if self.cup_breaks: # remove spaces between start of line and BREAK text = re.sub(r"^[ ]+BREAK\b", "BREAK", text, flags=re.MULTILINE) # remove spaces between BREAK and end of line text = re.sub(r"\bBREAK[ ]+$", "BREAK", text, flags=re.MULTILINE) # remove at start of prompt - text = re.sub(r"\ABREAK\b", "", text) + text = re.sub(r"\A(?:\s*BREAK\b\s*)+", "", text) # remove at end of prompt - text = re.sub(r"\bBREAK\Z", "", text) + text = re.sub(r"(?:\s*\bBREAK\s*)+\Z", "", text) if self.cup_ands: # collapse ANDs with space after text = re.sub(r"\bAND(?:\s+AND)+\s+", "AND ", text) @@ -1182,147 +308,985 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in # collapse separators and spaces after ANDs text = re.sub(r"\bAND[, ]+", "AND ", text) # remove at start of prompt - text = re.sub(r"\AAND\b", "", text) + text = re.sub(r"\A(?:AND\b\s*)+", "", text) # remove at end of prompt - text = re.sub(r"\bAND\Z", "", text) + text = re.sub(r"(\s*\bAND)+\Z", "", text) if self.cup_extranetworktags: - # - # all cases since we can't find them inside plain text - # # remove spaces before < text = re.sub(r"\B\s+<(?!!)", "<", text) # remove spaces after > - text = re.sub(r"(?\s+\B", ">", text) + text = re.sub(r">\s+\B", ">", text) if self.cup_extraspaces: + # remove spaces before comma + text = re.sub(r"[ ]+,", ",", text) + # remove spaces at end of line + text = re.sub(r"[ ]+$", "", text, flags=re.MULTILINE) + # remove spaces at start of line + text = re.sub(r"^[ ]+", "", text, flags=re.MULTILINE) # remove extra whitespace after starting parenthesis or bracket text = re.sub(r"([,\.;\s]+[([])\s+", r"\1", text) # remove extra whitespace before ending parenthesis or bracket text = re.sub(r"\s+([)\]][,\.;\s]+)", r"\1", text) + # remove empty lines + text = re.sub(r"(?:^|\n)[ ]*\n", "\n", text) + text = re.sub(r"\n[ ]*\n$", "\n", text) # collapse spaces - # text = re.sub(r"[ ]{2,}", " ", text) + text = re.sub(r"[ ]{2,}", " ", text) # remove spaces at start and end text = text.strip() return text - def __findwildcards(self, prompt, negative_prompt): - """ - Find and process wildcards in the prompt and negative_prompt strings. + def __processprompts(self, prompt, negative_prompt): + self.user_variables = {} - Args: - prompt (str): The prompt string. - negative_prompt (str): The negative prompt string. + # Process prompt + p_processor = self.TreeProcessor(self) + p_parsed = self.parse_prompt("prompt", prompt, self.parser_complete) + prompt = p_processor.start_visit("prompt", p_parsed, False) - Returns: - tuple: A tuple containing the processed prompt and negative_prompt strings. - """ + # Process negative prompt + n_processor = self.TreeProcessor(self) + n_parsed = self.parse_prompt("negative prompt", negative_prompt, self.parser_complete) + negative_prompt = n_processor.start_visit("negative prompt", n_parsed, True) - p_transformtree = self.TransformerTree(self, phase="wildcards") - try: - if self.debug: - self.logger.info("Wildcard processing in the prompt") - p_tree = self.__parser_complete.parse(prompt) - prompt = p_transformtree.transform(p_tree) - except Exception as e: # pylint: disable=broad-except - self.logger.warning("Parsing failed in prompt!: %s", e) + # Insertions in the negative prompt + if self.debug_level == DEBUG_LEVEL.full: + self.logger.debug(self.formatOutput(f"New negative additions: {p_processor.add_at}")) + self.logger.debug(self.formatOutput(f"New negative indexes: {n_processor.insertion_at}")) + negative_prompt = self.__add_to_insertion_points( + negative_prompt, p_processor.add_at["insertion_point"], n_processor.insertion_at + ) + if len(p_processor.add_at["start"]) > 0: + negative_prompt = self.__add_to_start(negative_prompt, p_processor.add_at["start"]) + if len(p_processor.add_at["end"]) > 0: + negative_prompt = self.__add_to_end(negative_prompt, p_processor.add_at["end"]) - np_transformtree = self.TransformerTree(self, phase="wildcards") - try: - if self.debug: - self.logger.info("Wildcard processing in the negative prompt") - np_tree = self.__parser_complete.parse(negative_prompt) - negative_prompt = np_transformtree.transform(np_tree) - except Exception as e: # pylint: disable=broad-except - self.logger.warning("Parsing failed in negative prompt!: %s", e) - - foundP = False - foundNP = False - if len(p_transformtree.detectedWildcards) > 0: - foundP = True - self.logger.info(f"Found wildcards in prompt: {p_transformtree.detectedWildcards}") - if len(np_transformtree.detectedWildcards) > 0: - foundNP = True - self.logger.info(f"Found wildcards in negative prompt: {np_transformtree.detectedWildcards}") + # Clean up + prompt = self.__cleanup(prompt) + negative_prompt = self.__cleanup(negative_prompt) + # Check for wildcards not processed + foundP = len(p_processor.detectedWildcards) > 0 + foundNP = len(n_processor.detectedWildcards) > 0 if foundP or foundNP: - if self.ifwildcards == self.IFWILDCARDS_CHOICES["warn"]: + self.logger.error("Found unprocessed wildcards!") + ppwl = ", ".join(p_processor.detectedWildcards) + npwl = ", ".join(n_processor.detectedWildcards) + if foundP: + self.logger.error(self.formatOutput(f"In the positive prompt: {ppwl}")) + if foundNP: + self.logger.error(self.formatOutput(f"In the negative prompt: {npwl}")) + if self.wil_ifwildcards == self.IFWILDCARDS_CHOICES.warn: prompt = self.WILDCARD_WARNING + prompt - elif self.ifwildcards == self.IFWILDCARDS_CHOICES["stop"]: - self.logger.error("Found unprocessed wildcards! Stopping the generation.") + elif self.wil_ifwildcards == self.IFWILDCARDS_CHOICES.stop: + self.logger.error("Stopping the generation.") if foundP: - prompt = self.WILDCARD_STOP + prompt + prompt = self.WILDCARD_STOP.format(ppwl) + prompt if foundNP: - negative_prompt = self.WILDCARD_STOP + negative_prompt - if hasattr(self.script, "ppp_interrupt"): - self.script.ppp_interrupt() - if self.debug: - self.logger.info(f"prompt after wildcards: {self.formatOutput(prompt)}") - self.logger.info(f"negative_prompt after wildcards: {self.formatOutput(negative_prompt)}") + negative_prompt = self.WILDCARD_STOP.format(npwl) + negative_prompt + self.interrupt() return prompt, negative_prompt - def process_prompt(self, original_prompt, original_negative_prompt): + def process_prompt( + self, + original_prompt: str, + original_negative_prompt: str, + seed: int = 0, + ): """ Process the prompt and negative prompt by moving content to the negative prompt, and cleaning up. Args: original_prompt (str): The original prompt. original_negative_prompt (str): The original negative prompt. + seed (int): The seed. Returns: tuple: A tuple containing the processed prompt and negative prompt. """ try: - self.variables = {} + self.rng = np.random.default_rng(seed & 0xFFFFFFFF) prompt = original_prompt negative_prompt = original_negative_prompt - self.debug = getattr(self.opts, "ppp_gen_debug", False) - if not self.is_i2i or self.stn_doi2i or self.cup_doi2i: - if self.debug: - self.logger.info(f"Input prompt: {self.formatOutput(prompt)}") - p_tree = self.__parser_complete.parse(prompt) - self.logger.info(f"Tree from prompt:\n{p_tree.pretty()}") - self.logger.info(f"Input negative_prompt: {self.formatOutput(negative_prompt)}") - np_tree = self.__parser_complete.parse(negative_prompt) - self.logger.info(f"Tree from negative prompt:\n{np_tree.pretty()}") - - prompt, negative_prompt = self.__removecontent(prompt, negative_prompt) - - if self.ifwildcards != self.IFWILDCARDS_CHOICES["ignore"]: - prompt, negative_prompt = self.__findwildcards(prompt, negative_prompt) - - if not self.is_i2i or self.stn_doi2i: - prompt, negative_prompt = self.__sendtonegative(prompt, negative_prompt) - - # pylint: disable-next=too-many-boolean-expressions - if (not self.is_i2i or self.cup_doi2i) and ( - self.cup_extraspaces - or self.cup_emptyconstructs - or self.cup_extraseparators - or self.cup_extraseparators2 - or self.cup_breaks - or self.cup_breaks_eol - or self.cup_ands - or self.cup_ands_eol - or self.cup_extranetworktags - ): - prompt, negative_prompt = self.__cleanup(prompt, negative_prompt) - - # Check for constructs not processed due to parsing problems - if ( - prompt.find("= 0 - or negative_prompt.find("= 0 - or prompt.find("= 0 - or negative_prompt.find("= 0 - ): - self.logger.error( - "Found unprocessed constructs in prompt or negative prompt! Stopping the generation." - ) - self.logger.info(f"prompt: {self.formatOutput(prompt)}") - self.logger.info(f"negative_prompt: {self.formatOutput(negative_prompt)}") - prompt = self.UNPROCESSED_STOP + prompt - if hasattr(self.script, "ppp_interrupt"): - self.script.ppp_interrupt() + self.debug_level = DEBUG_LEVEL(self.options.get("debug_level", DEBUG_LEVEL.none.value)) + if self.debug_level != DEBUG_LEVEL.none: + self.logger.info(f"System variables: {self.system_variables}") + self.logger.info(f"Input seed: {seed}") + self.logger.info(self.formatOutput(f"Input prompt: {prompt}")) + self.logger.info(self.formatOutput(f"Input negative_prompt: {negative_prompt}")) + t1 = time.time() + prompt, negative_prompt = self.__processprompts(prompt, negative_prompt) + t2 = time.time() + if self.debug_level != DEBUG_LEVEL.none: + self.logger.info(self.formatOutput(f"Result prompt: {prompt}")) + self.logger.info(self.formatOutput(f"Result negative_prompt: {negative_prompt}")) + self.logger.info(f"Process prompt pair time: {t2 - t1:.3f} seconds") + # Check for constructs not processed due to parsing problems + fullcontent: str = prompt + negative_prompt + if fullcontent.find("= 0: + self.logger.error("Found unprocessed constructs in prompt or negative prompt! Stopping the generation.") + prompt = self.UNPROCESSED_STOP + prompt + self.interrupt() return prompt, negative_prompt except Exception as e: # pylint: disable=broad-exception-caught self.logger.exception(e) return original_prompt, original_negative_prompt + + def parse_prompt(self, prompt_description: str, prompt: str, parser: lark.Lark, raise_parsing_error: bool = False): + t1 = time.time() + try: + if self.debug_level == DEBUG_LEVEL.full: + self.logger.debug(self.formatOutput(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] + except lark.exceptions.UnexpectedInput: + if raise_parsing_error: + raise + self.logger.exception(self.formatOutput(f"Parsing failed on prompt!: {prompt}")) + t2 = time.time() + if self.debug_level == DEBUG_LEVEL.full: + self.logger.debug("Tree:\n" + textwrap.indent(re.sub(r"\n$", "", parsed_prompt.pretty()), " ")) + self.logger.debug(f"Parse {prompt_description} time: {t2 - t1:.3f} seconds") + return parsed_prompt + + class TreeProcessor(lark.visitors.Interpreter): + """ + A class for interpreting and processing a tree generated by the prompt parser. + + Args: + ppp (PromptPostProcessor): The PromptPostProcessor object. + + Attributes: + add_at (dict): The dictionary to store the content to be added at different positions of the negative prompt. + insertion_at (list): The list of insertion points in the negative prompt. + detectedWildcards (list): The list of detected invalid wildcards or choices. + result (str): The final processed prompt. + """ + + def __init__(self, ppp: "PromptPostProcessor"): + super().__init__() + self.__ppp = ppp + self.AccumulatedShell = namedtuple("AccumulatedShell", ["type", "data"]) + self.NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"]) + self.__shell: list[self.AccumulatedShell] = [] + self.__negtags: list[self.NegTag] = [] + self.__already_processed: list[str] = [] + self.__is_negative = False + self.add_at: dict = {"start": [], "insertion_point": [[] for x in range(10)], "end": []} + self.insertion_at: list[tuple[int, int]] = [None for x in range(10)] + self.detectedWildcards: list[str] = [] + self.result = "" + + def start_visit(self, prompt_description: str, parsed_prompt: lark.Tree, is_negative: bool = False) -> str: + """ + Start the visit process. + + Args: + prompt_description (str): The description of the prompt. + parsed_prompt (Tree): The parsed prompt. + is_negative (bool): Whether the prompt is negative or not. + + Returns: + str: The processed prompt. + """ + t1 = time.time() + self.__is_negative = is_negative + if self.__ppp.debug_level != DEBUG_LEVEL.none: + self.__ppp.logger.info(f"Processing {prompt_description}...") + self.visit(parsed_prompt) + t2 = time.time() + if self.__ppp.debug_level != DEBUG_LEVEL.none: + self.__ppp.logger.info(f"Process {prompt_description} time: {t2 - t1:.3f} seconds") + return self.result + + def __visit( + self, + node: lark.Tree | lark.Token | list[lark.Tree | lark.Token] | None, + restore_state: bool = False, + discard_content: bool = False, + ) -> str: + """ + Visit a node in the tree and process it or accumulate its value if it is a Token. + + Args: + node (Tree|Token|list): The node or list of nodes to visit. + restore_state (bool): Whether to restore the state after visiting the node. + discard_content (bool): Whether to discard the content of the node. + + Returns: + str: The result of the visit. + """ + backup_result = self.result + if restore_state: + backup_shell = self.__shell.copy() + backup_negtags = self.__negtags.copy() + backup_already_processed = self.__already_processed.copy() + backup_add_at = self.add_at.copy() + backup_insertion_at = self.insertion_at.copy() + backup_detectedwildcards = self.detectedWildcards.copy() + if node is not None: + if isinstance(node, list): + for child in node: + self.__visit(child) + elif isinstance(node, lark.Tree): + self.visit(node) + elif isinstance(node, lark.Token): + self.result += node + added_result = self.result[len(backup_result) :] + if discard_content or restore_state: + self.result = backup_result + if restore_state: + self.__shell = backup_shell + self.__negtags = backup_negtags + self.__already_processed = backup_already_processed + self.add_at = backup_add_at + self.insertion_at = backup_insertion_at + self.detectedWildcards = backup_detectedwildcards + return added_result + + def __get_original_node_content(self, node: lark.Tree | lark.Token, default=None) -> str: + return ( + node.meta.content + if hasattr(node, "meta") and node.meta is not None and not node.meta.empty + else default + ) + + def __get_user_variable_value(self, name: str, default="", evaluate=True) -> str: + if evaluate: + v = self.__ppp.user_variables.get(name, default) + if isinstance(v, lark.Tree): + v = self.__visit(v, True) + else: + v = ( + self.__ppp.user_variables[name] + if isinstance(self.__ppp.user_variables[name], str) + else self.__get_original_node_content( + self.__ppp.user_variables[name], default or "not evaluated yet" + ) + ) + return v + + def __set_user_variable_value(self, name: str, value: str): + self.__ppp.user_variables[name] = value + + def __remove_user_variable(self, name: str): + if name in self.__ppp.user_variables: + del self.__ppp.user_variables[name] + + def __debug_end(self, construct: str, start_result: str, duration: float, info=None): + if self.__ppp.debug_level == DEBUG_LEVEL.full: + info = f"({info}) " if info is not None and info != "" else "" + output = self.result[len(start_result) :] + if output != "": + output = f" >> '{output}'" + self.__ppp.logger.debug( + self.__ppp.formatOutput(f"TreeProcessor.{construct} {info}({duration:.3f} seconds){output}") + ) + + def __eval_condition(self, cond_var: str, cond_comp: str, cond_value: str | list[str]) -> bool: + """ + Evaluate a condition based on the given variable, comparison, and value. + + Args: + cond_var (str): The variable to be compared. + cond_comp (str): The comparison operator. + cond_value (str or list[str]): The value to be compared with. + + Returns: + bool: The result of the condition evaluation. + """ + var_value = self.__ppp.system_variables.get(cond_var, self.__get_user_variable_value(cond_var, None)) + if var_value is None: + var_value = "" + self.__ppp.logger.warning(f"Unknown variable {cond_var}") + if isinstance(var_value, str): + var_value = var_value.lower() + if isinstance(cond_value, list): + comp_ops = { + "contains": lambda x, y: y in x, + "in": lambda x, y: x == y, + } + else: + cond_value = [cond_value] + comp_ops = { + "eq": lambda x, y: x == y, + "ne": lambda x, y: x != y, + "gt": lambda x, y: x > y, + "lt": lambda x, y: x < y, + "ge": lambda x, y: x >= y, + "le": lambda x, y: x <= y, + "contains": lambda x, y: y in x, + "truthy": lambda x, y: bool(x), + } + if cond_comp not in comp_ops: + return False + cond_value_adjusted = list( + ( + c[1:-1].lower() + if c.startswith('"') or c.startswith("'") + else True if c.lower() == "true" else False if c.lower() == "false" else int(c) + ) + for c in cond_value + ) + result = False + for c in cond_value_adjusted: + var_value_adjusted = ( + var_value + if isinstance(c, str) + else ( + True + if isinstance(c, bool) and var_value != "false" and var_value is not False + else ( + False + if isinstance(c, bool) and (var_value != "true" or var_value is False) + else int(var_value) + ) + ) + ) + result = comp_ops[cond_comp](var_value_adjusted, c) + if result: + break + return result + + def __evaluate_if(self, condition: lark.Tree) -> bool: + """ + Evaluate an if condition based on the given condition tree. + + Args: + condition (Node): The condition tree to be evaluated. + + Returns: + bool: The result of the if condition evaluation. + """ + get_value = lambda n: n.value # pylint: disable=unnecessary-lambda-assignment + # if hasattr(condition, "children"): + get_children = lambda n: n.children # pylint: disable=unnecessary-lambda-assignment + # else: + # get_children = lambda n: n # pylint: disable=unnecessary-lambda-assignment + # get_value = lambda n: n # pylint: disable=unnecessary-lambda-assignment + individualcondition = get_children(condition)[0] + # we get the name of the variable and check for a preceding not + invert = False + first = get_value(get_children(individualcondition)[0]) + if first == "not": + invert = True + cond_var = get_value(get_children(individualcondition)[1]) + poscomp = 2 + else: + cond_var = first + poscomp = 1 + if poscomp >= len(get_children(individualcondition)): + # no condition, just a variable + cond_comp = "truthy" + cond_value = "true" + else: + # we get the comparison (with possible not) and the value + cond_comp = get_value(get_children(individualcondition)[poscomp]) + if cond_comp == "not": + invert = not invert + poscomp += 1 + cond_comp = get_value(get_children(individualcondition)[poscomp]) + poscomp += 1 + cond_value_node = get_children(individualcondition)[poscomp] + cond_value = ( + list(get_value(v) for v in get_children(cond_value_node)) + if isinstance(cond_value_node, (lark.Tree, list)) + else cond_value_node.value if isinstance(cond_value_node, lark.Token) else cond_value_node + ) + condresult = self.__eval_condition(cond_var, cond_comp, cond_value) + if invert: + condresult = not condresult + return condresult + + def promptcomp(self, tree: lark.Tree): + """ + Process a prompt composition construct in the tree. + """ + start_result = self.result + t1 = time.time() + self.__visit(tree.children[0]) + if len(tree.children) > 1: + if tree.children[1] is not None: + self.result += f":{tree.children[1]}" + for i in range(2, len(tree.children), 3): + if self.__ppp.cup_ands: + self.result = re.sub(r"[, ]+$", "\n" if self.__ppp.cup_ands_eol else " ", self.result) + if self.result[-1:].isalnum(): # add space if needed + self.result += " " + self.result += "AND" + added_result = self.__visit(tree.children[i + 1], False, True) + if self.__ppp.cup_ands: + added_result = re.sub(r"^[, ]+", " ", added_result) + if added_result[0:1].isalnum(): # add space if needed + added_result = " " + added_result + self.result += added_result + if tree.children[i + 2] is not None: + self.result += f":{tree.children[i+2]}" + t2 = time.time() + self.__debug_end("promptcomp", start_result, t2 - t1) + + def scheduled(self, tree: lark.Tree): + """ + Process a scheduling construct in the tree and add it to the accumulated shell. + """ + start_result = self.result + t1 = time.time() + before = tree.children[0] + after = tree.children[-2] + pos_str = tree.children[-1] + pos = float(pos_str) + if pos >= 1: + pos = int(pos) + # self.__shell.append(self.AccumulatedShell("sc", pos)) + self.result += "[" + if before is not None: + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug(f"Shell scheduled before with position {pos}") + self.__shell.append(self.AccumulatedShell("scb", pos)) + self.__visit(before) + self.__shell.pop() + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug(f"Shell scheduled after with position {pos}") + self.__shell.append(self.AccumulatedShell("sca", pos)) + self.result += ":" + self.__visit(after) + self.__shell.pop() + if self.__ppp.cup_emptyconstructs and self.result == start_result + "[:": + self.result = start_result + else: + self.result += f":{pos_str}]" + # self.__shell.pop() + t2 = time.time() + self.__debug_end("scheduled", start_result, t2 - t1, pos_str) + + def alternate(self, tree: lark.Tree): + """ + Process an alternation construct in the tree and add it to the accumulated shell. + """ + start_result = self.result + t1 = time.time() + # self.__shell.append(self.AccumulatedShell("al", len(tree.children))) + self.result += "[" + for i, opt in enumerate(tree.children): + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug(f"Shell alternate option {i+1}") + self.__shell.append(self.AccumulatedShell("alo", {"pos": i + 1, "len": len(tree.children)})) + if i > 0: + self.result += "|" + self.__visit(opt) + self.__shell.pop() + self.result += "]" + if self.__ppp.cup_emptyconstructs and self.result == start_result + "[]": + self.result = start_result + # self.__shell.pop() + t2 = time.time() + self.__debug_end("alternate", start_result, t2 - t1) + + def emphasized(self, tree: lark.Tree): + """ + Process a attention change construct in the tree and add it to the accumulated shell. + """ + start_result = self.result + t1 = time.time() + weight_str = tree.children[-1] + if weight_str is not None: + weight = float(weight_str) + else: + weight_str = "" + weight = 1.1 + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug(f"Shell attention with weight {weight}") + self.__shell.append(self.AccumulatedShell("at", weight)) + self.result += "(" + self.__visit(tree.children[:-1]) + if self.__ppp.cup_emptyconstructs and self.result == start_result + "(": + self.result = start_result + else: + if weight_str != "": + self.result += f":{weight_str}" + self.result += ")" + self.__shell.pop() + t2 = time.time() + self.__debug_end("emphasized", start_result, t2 - t1, weight_str) + + def deemphasized(self, tree: lark.Tree): + """ + Process a decrease attention construct in the tree and add it to the accumulated shell. + """ + start_result = self.result + t1 = time.time() + weight = 0.9 + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug(f"Shell attention with weight {weight}") + self.__shell.append(self.AccumulatedShell("at", weight)) + self.result += "[" + self.__visit(tree.children) + if self.__ppp.cup_emptyconstructs and self.result == start_result + "[": + self.result = start_result + else: + self.result += "]" + self.__shell.pop() + t2 = time.time() + self.__debug_end("deemphasized", start_result, t2 - t1) + + def commandstn(self, tree: lark.Tree): + """ + Process a send to negative command in the tree and add it to the list of negative tags. + """ + start_result = self.result + info = None + t1 = time.time() + if not self.__is_negative: + negtagparameters = tree.children[0] + if negtagparameters is not None: + parameters = negtagparameters.value + else: + parameters = "" + content = self.__visit(tree.children[1::], False, True) + self.__negtags.append( + self.NegTag(len(self.result), len(self.result), content, parameters, self.__shell.copy()) + ) + info = f"with {parameters or 'no parameters'} : {content}" + else: + self.__ppp.logger.warning("Ignored negative command in negative prompt") + self.__visit(tree.children[1::]) + t2 = time.time() + self.__debug_end("commandstn", start_result, t2 - t1, info) + + def commandstni(self, tree: lark.Tree): + """ + Process a send to negative insertion point command in the tree and add it to the list of negative tags. + """ + start_result = self.result + info = None + t1 = time.time() + if self.__is_negative: + negtagparameters = tree.children[0] + if negtagparameters is not None: + parameters = negtagparameters.value + else: + parameters = "" + self.__negtags.append( + self.NegTag(len(self.result), len(self.result), "", parameters, self.__shell.copy()) + ) + info = f"with {parameters or 'no parameters'}" + else: + self.__ppp.logger.warning("Ignored negative insertion point command in positive prompt") + t2 = time.time() + self.__debug_end("commandstni", start_result, t2 - t1, info) + + def __varset( + self, + command: str, + variable: str, + immediateevaluation: str | None, + adding: str | None, + content: lark.Tree | None, + ): + """ + Process a generic set command in the tree. + """ + t1 = time.time() + start_result = self.result + if variable.startswith("_"): + self.__ppp.logger.warning(f"Invalid variable name '{variable}' detected!") + self.__ppp.interrupt() + return + info = variable + value_description = self.__get_original_node_content(content, None) + value = content + if adding is not None: + info += f" += '{value_description}'" + raw_oldvalue = self.__ppp.user_variables.get(variable, None) + if raw_oldvalue is None: + newvalue = value + self.__ppp.logger.warning(f"Unknown variable {variable}") + elif isinstance(raw_oldvalue, str): + newvalue = lark.Tree( + lark.Token("RULE", "varvalue"), + [lark.Token("plain", raw_oldvalue), value], + # Meta should be {"content": raw_oldvalue + value}, + ) + else: + newvalue = lark.Tree( + lark.Token("RULE", "varvalue"), + [raw_oldvalue, value], + # Meta should be {"content": raw_oldvalue.meta.content + value.meta.content}, + ) + else: + newvalue = value + if immediateevaluation is not None: + newvalue = self.__visit(newvalue, False, True) + info += " =! " + else: + info += " = " + self.__set_user_variable_value(variable, newvalue) + currentvalue = self.__get_user_variable_value(variable, None, False) + if currentvalue is None: + info += "not evaluated yet" + else: + info += f"'{currentvalue}'" + t2 = time.time() + self.__debug_end(command, start_result, t2 - t1, info) + + def variableset(self, tree: lark.Tree): + """ + Process a DP set variable command in the tree and add it to the dictionary of variables. + """ + self.__varset("variableset", tree.children[0], tree.children[2], tree.children[1], tree.children[3]) + + def commandset(self, tree: lark.Tree): + """ + Process a set command in the tree and add it to the dictionary of variables. + """ + self.__varset("commandset", tree.children[0], tree.children[1], tree.children[2], tree.children[3]) + + def __varecho(self, command: str, variable: str, default: lark.Tree | None): + """ + Process a generic echo command in the tree. + """ + t1 = time.time() + start_result = self.result + value = self.__get_user_variable_value(variable, None) + if default is not None: + default_value = self.__visit(default, True) # for log + if value is None: + if default is not None: + value = self.__visit(default, False, True) + else: + value = "" + self.__ppp.logger.warning(f"Unknown variable {variable}") + self.result += value + t2 = time.time() + info = variable + if default is not None: + info += f" with default '{default_value}'" + self.__debug_end(command, start_result, t2 - t1, info) + + def variableuse(self, tree: lark.Tree): + """ + Process a DP use variable command in the tree. + """ + self.__varecho("variableuse", tree.children[0], tree.children[1]) + + def commandecho(self, tree: lark.Tree): + """ + Process an echo command in the tree. + """ + self.__varecho("commandecho", tree.children[0], tree.children[1]) + + def commandif(self, tree: lark.Tree): + """ + Process an if command in the tree. + """ + t1 = time.time() + start_result = self.result + for i, n in enumerate(tree.children): + content = n.children[-1] + if len(n.children) == 2: # its not an else + # has a condition + condition = n.children[0] + c = self.__get_original_node_content(condition, f"condition {i}") + if self.__evaluate_if(condition): + self.__visit(content) + t2 = time.time() + self.__debug_end("commandif", start_result, t2 - t1, c) + return + else: # its an else + self.__visit(content) + t2 = time.time() + self.__debug_end("commandif", start_result, t2 - t1, "else") + return + + def extranetworktag(self, tree: lark.Tree): + """ + Process an extra network construct in the tree. + """ + t1 = time.time() + start_result = self.result + if not self.__ppp.rem_removeextranetworktags: + # keep extra network construct + self.result += self.__get_original_node_content(tree, f"<{tree.children[0]}>") + t2 = time.time() + self.__debug_end("extranetworktag", start_result, t2 - t1) + + def __get_choices(self, options: lark.Tree | None, choice_values: list[lark.Tree]) -> str: + """ + Select choices based on the options. + + Args: + is_wildcard (bool): A flag indicating whether the choices are from a wildcard. + options (Tree): The tree object representing the options construct. + choice_values (list[Tree]): A list of choice tree objects. + + Returns: + str: The selected choice. + """ + sampler: str = "~" + repeating: bool = False + from_value: int = 1 + to_value: int = 1 + separator: str = self.__ppp.wil_choice_separator + if options is not None: + if len(options.children) == 1: + sampler = options.children[0] if options.children[0] is not None else "~" + else: + sampler = options.children[0].children[0] if options.children[0] is not None else "~" + repeating = options.children[1].children[0] == "r" if options.children[1] is not None else False + if len(options.children) == 4: + ifrom = 2 + ito = 2 + isep = 3 + else: # 6 + ifrom = 2 + ito = 3 + isep = 4 + from_value = int(options.children[ifrom].children[0]) if options.children[ifrom] is not None else 1 + to_value = int(options.children[ito].children[0]) if options.children[ito] is not None else 1 + separator = ( + self.__visit(options.children[isep], False, True) + if options.children[isep] is not None + else self.__ppp.wil_choice_separator + ) + if sampler != "~": + self.__ppp.logger.warning(f"Unsupported sampler '{sampler}' in wildcard/choices options!") + self.__ppp.interrupt() + return "" + if from_value < 0: + from_value = 1 + elif from_value > len(choice_values): + from_value = len(choice_values) + if to_value < 1: + to_value = 1 + elif (to_value > len(choice_values) and not repeating) or from_value > to_value: + to_value = len(choice_values) + num_choices = ( + self.__ppp.rng.integers(from_value, to_value, endpoint=True) if from_value < to_value else from_value + ) + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug( + self.__ppp.formatOutput( + f"Selecting {'repeating ' if repeating else ''}{num_choices} choices and separating with '{separator}'" + ) + ) + if num_choices > 0: + weights = [] + included_choices = 0 + excluded_choices = 0 + excluded_weights_sum = 0 + for i, c in enumerate(choice_values): + c.choice_index = i # we index them to later sort the results + w = float(c.children[0].children[0]) if c.children[0] is not None else 1.0 + if w > 0 and (c.children[1] is None or self.__evaluate_if(c.children[1].children[0])): + weights.append(w) + included_choices += 1 + else: + weights.append(-1) + excluded_choices += 1 + excluded_weights_sum += w + if excluded_choices > 0: # we need to redistribute the excluded weights + weights = [w + excluded_weights_sum / included_choices if w >= 0 else 0.0 for w in weights] + weights = np.array(weights) + weights /= weights.sum() # normalize weights + selected_choices: list[lark.Tree] = list( + self.__ppp.rng.choice(choice_values, size=num_choices, p=weights, replace=repeating) + ) + if self.__ppp.wil_keep_choices_order: + selected_choices = sorted(selected_choices, key=lambda x: x.choice_index) + selected_choices_text = [] + for i, c in enumerate(selected_choices): + t1 = time.time() + choice_content = self.__visit(c.children[2], False, True) + t2 = time.time() + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug( + f"Adding choice {i+1} ({t2-t1:.3f} seconds):\n" + + textwrap.indent(re.sub(r"\n$", "", c.pretty()), " ") + ) + selected_choices_text.append(choice_content) + # remove comments + results = [re.sub(r"\s*#[^\n]*(?:\n|$)", "", r, flags=re.DOTALL) for r in selected_choices_text] + return separator.join(results) + return "" + + def wildcard(self, tree: lark.Tree): + """ + Process a wildcard construct in the tree. + """ + t1 = time.time() + start_result = self.result + options = tree.children[0] + wildcard_key = tree.children[1].value + wc = self.__get_original_node_content(tree, f"?__{wildcard_key}__") + if self.__ppp.wil_process_wildcards: + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug(f"Processing wildcard: {wildcard_key}") + wildcard_keys = fnmatch.filter(self.__ppp.wildcard_obj.wildcards.keys(), wildcard_key) + if len(wildcard_keys) == 0: + self.detectedWildcards.append(wc) + self.result += wc + t2 = time.time() + self.__debug_end("wildcard", start_result, t2 - t1, wc) + return + variablename = None + if tree.children[2] is not None: + variablename = tree.children[2].children[0] # should be a token + variablevalue = self.__visit(tree.children[2].children[1], False, True) + variablebackup = self.__ppp.user_variables.get(variablename, None) + self.__remove_user_variable(variablename) + self.__set_user_variable_value(variablename, variablevalue) + choice_values_obj_all = [] + for key in wildcard_keys: + wildcard = self.__ppp.wildcard_obj.wildcards.get(key, None) + if wildcard is None: + self.detectedWildcards.append(wc) + self.result += wc + t2 = time.time() + self.__debug_end("wildcard", start_result, t2 - t1, wc) + return + choice_values_obj = wildcard.get("choices_obj", None) + options_obj = wildcard.get("options_obj", None) + if choice_values_obj is None: + t1 = time.time() + choice_values_obj = [] + choices = wildcard["choices"] + try: + options_obj = self.__ppp.parse_prompt( + "as choices options", choices[0], self.__ppp.parser_choicesoptions, True + ) + n = 1 + except lark.exceptions.UnexpectedInput: + options_obj = None + n = 0 + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug("Does not have options") + wildcard["options_obj"] = options_obj + for cv in choices[n:]: + try: + choice_values_obj.append( + self.__ppp.parse_prompt("choice", cv, self.__ppp.parser_choice, True) + ) + except lark.exceptions.UnexpectedInput as e: + self.__ppp.logger.warning( + f"Error parsing choice '{cv}' in wildcard '{key}'! : {e.__class__.__name__}" + ) + wildcard["choices_obj"] = choice_values_obj + t2 = time.time() + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug(f"Processed choices for wildcard '{key}' ({t2-t1:.3f} seconds)") + if options_obj is not None: + if options is None: + options = options_obj + else: + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug(f"Options for wildcard '{key}' are ignored!") + choice_values_obj_all += choice_values_obj + self.result += self.__get_choices(options, choice_values_obj_all) + if variablename is not None: + self.__remove_user_variable(variablename) + if variablebackup is not None: + self.__ppp.user_variables[variablename] = variablebackup + elif self.__ppp.wil_ifwildcards != self.__ppp.IFWILDCARDS_CHOICES.remove: + self.detectedWildcards.append(wc) + self.result += wc + t2 = time.time() + self.__debug_end("wildcard", start_result, t2 - t1, f"'{wc}'") + + def choices(self, tree: lark.Tree): + """ + Process a choices construct in the tree. + """ + t1 = time.time() + start_result = self.result + options = tree.children[0] + choice_values = tree.children[1::] + ch = self.__get_original_node_content(tree, "?{...}") + if self.__ppp.wil_process_wildcards: + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug("Processing choices:") + self.result += self.__get_choices(options, choice_values) + elif self.__ppp.wil_ifwildcards != self.__ppp.IFWILDCARDS_CHOICES.remove: + self.detectedWildcards.append(ch) + self.result += ch + t2 = time.time() + self.__debug_end("choices", start_result, t2 - t1, f"'{ch}'") + + def __default__(self, tree): + t1 = time.time() + start_result = self.result + self.__visit(tree.children) + t2 = time.time() + self.__debug_end(tree.data.value, start_result, t2 - t1) + + def start(self, tree): + self.result = "" + t1 = time.time() + self.__visit(tree.children) + # process the found negative tags + for negtag in self.__negtags: + if self.__ppp.stn_join_attention: + # join consecutive attention elements + for i in range(len(negtag.shell) - 1, 0, -1): + if negtag.shell[i].type == "at" and negtag.shell[i - 1].type == "at": + negtag.shell[i - 1] = self.AccumulatedShell( + "at", + math.floor(100 * negtag.shell[i - 1].data * negtag.shell[i].data) + / 100, # we limit the new weight to two decimals + ) + negtag.shell.pop(i) + start = "" + end = "" + for s in negtag.shell: + match s.type: + case "at": + if s.data == 0.9: + start += "[" + end = "]" + end + elif s.data == 1.1: + start += "(" + end = ")" + end + else: + start += "(" + end = f":{s.data})" + end + # case "sc": + case "scb": + start += "[" + end = f"::{s.data}]" + end + case "sca": + start += "[" + end = f":{s.data}]" + end + # case "al": + case "alo": + start += "[" + ("|" * int(s.data["pos"] - 1)) + end = ("|" * int(s.data["len"] - s.data["pos"])) + "]" + end + content = start + negtag.content + end + position = negtag.parameters or "s" + if position.startswith("i"): + n = int(position[1]) + self.insertion_at[n] = [negtag.start, negtag.end] + elif len(content) > 0: + if content not in self.__already_processed: + if self.__ppp.stn_ignore_repeats: + self.__already_processed.append(content) + if self.__ppp.debug_level == DEBUG_LEVEL.full: + self.__ppp.logger.debug( + self.__ppp.formatOutput(f"Adding content at position {position}: {content}") + ) + if position == "e": + self.add_at["end"].append(content) + elif position.startswith("p"): + n = int(position[1]) + self.add_at["insertion_point"][n].append(content) + else: # position == "s" or invalid + self.add_at["start"].append(content) + else: + self.__ppp.logger.warning(self.__ppp.formatOutput(f"Ignoring repeated content: {content}")) + t2 = time.time() + self.__debug_end("start", "", t2 - t1) diff --git a/ppp_cache.py b/ppp_cache.py new file mode 100644 index 0000000..d88473c --- /dev/null +++ b/ppp_cache.py @@ -0,0 +1,24 @@ +from collections import OrderedDict +from typing import Tuple + + +class PPPLRUCache: + + ProcessInput = Tuple[int, str, str] + ProcessResult = Tuple[str, str] + + def __init__(self, capacity: int): + self.cache = OrderedDict() + self.capacity = capacity + + def get(self, key: ProcessInput) -> ProcessResult: + if key not in self.cache: + return None + self.cache.move_to_end(key) + return self.cache[key] + + def put(self, key: ProcessInput, value: ProcessResult) -> None: + self.cache[key] = value + self.cache.move_to_end(key) + if len(self.cache) > self.capacity: + self.cache.popitem(last=False) diff --git a/ppp_comfyui.py b/ppp_comfyui.py new file mode 100644 index 0000000..80f4e8d --- /dev/null +++ b/ppp_comfyui.py @@ -0,0 +1,406 @@ +# pylint: disable=missing-module-docstring, missing-class-docstring, missing-function-docstring, invalid-name + +import os + +# pylint: disable=import-error +import folder_paths # type: ignore +import nodes # type: ignore + +from .ppp import PromptPostProcessor +from .ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory +from .ppp_wildcards import PPPWildcards + +if __name__ == "__main__": + raise SystemExit("This script must be run from ComfyUI") + + +class PromptPostProcessorComfyUINode: + + VERSION = PromptPostProcessor.VERSION + + logger = None + + def __init__(self): + lf = PromptPostProcessorLogFactory() + self.logger = lf.log + 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() + self.wildcards_obj = PPPWildcards(lf.log) + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ( + "MODEL", + { + "forceInput": True, + }, + ), + "modelname": ( + "STRING", + { + "default": "", + "forceInput": True, + }, + ), + "seed": ( + "INT", + { + "default": None, + "forceInput": False, + }, + ), + "pos_prompt": ( + "STRING", + { + "multiline": True, + "default": "", + "forceInput": True, + }, + ), + "neg_prompt": ( + "STRING", + { + "multiline": True, + "default": "", + "forceInput": True, + }, + ), + }, + "optional": { + "debug_level": ( + [e.value for e in DEBUG_LEVEL], + { + "default": DEBUG_LEVEL.minimal.value, + "tooltip": "Debug level", + }, + ), + "pony_substrings": ( + "STRING", + { + "default": PromptPostProcessor.DEFAULT_PONY_SUBSTRINGS, + "placeholder": "comma separated list", + "tooltip": "Comma separated list of substrings to look for in the modelname to determine if the model is a pony model", + }, + ), + "wc_process_wildcards": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Process wildcards in the prompt", + "label_on": "Yes", + "label_off": "No", + }, + ), + "wc_wildcards_folders": ( + "STRING", + { + "default": "", + "tooltip": "Comma separated list of wildcards folders", + }, + ), + "wc_if_wildcards": ( + [e.value for e in PromptPostProcessor.IFWILDCARDS_CHOICES], + { + "default": PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value, + "tooltip": "How to handle invalid wildcards in the prompt", + }, + ), + "wc_choice_separator": ( + "STRING", + { + "default": PromptPostProcessor.DEFAULT_CHOICE_SEPARATOR, + "tooltip": "Default separator for selected choices", + }, + ), + "wc_keep_choices_order": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Keep the order of the choices in the prompt", + "label_on": "Yes", + "label_off": "No", + }, + ), + "stn_separator": ( + "STRING", + { + "default": PromptPostProcessor.DEFAULT_STN_SEPARATOR, + "tooltip": "Separator for the content added to the negative prompt", + }, + ), + "stn_ignore_repeats": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Ignore repeated content added to the negative prompt", + "label_on": "Yes", + "label_off": "No", + }, + ), + "stn_join_attention": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Merge attention in the content added to the negative prompt", + "label_on": "Yes", + "label_off": "No", + }, + ), + "cleanup_extra_spaces": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Remove extra spaces", + "label_on": "Yes", + "label_off": "No", + }, + ), + "cleanup_empty_constructs": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Remove empty constructs", + "label_on": "Yes", + "label_off": "No", + }, + ), + "cleanup_extra_separators": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Remove extra separators", + "label_on": "Yes", + "label_off": "No", + }, + ), + "cleanup_extra_separators2": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Remove extra separators (additional cases)", + "label_on": "Yes", + "label_off": "No", + }, + ), + "cleanup_breaks": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Cleanup around BREAKs", + "label_on": "Yes", + "label_off": "No", + }, + ), + "cleanup_breaks_eol": ( + "BOOLEAN", + { + "default": False, + "tooltip": "Set BREAKs in their own line", + "label_on": "Yes", + "label_off": "No", + }, + ), + "cleanup_ands": ( + "BOOLEAN", + { + "default": True, + "tooltip": "Cleanup around ANDs", + "label_on": "Yes", + "label_off": "No", + }, + ), + "cleanup_ands_eol": ( + "BOOLEAN", + { + "default": False, + "tooltip": "Set ANDs in their own line", + "label_on": "Yes", + "label_off": "No", + }, + ), + "cleanup_extranetwork_tags": ( + "BOOLEAN", + { + "default": False, + "tooltip": "Clean up around extra network tags", + "label_on": "Yes", + "label_off": "No", + }, + ), + "remove_extranetwork_tags": ( + "BOOLEAN", + { + "default": False, + "tooltip": "Remove extra network tags", + "label_on": "Yes", + "label_off": "No", + }, + ), + }, + } + + RETURN_TYPES = ( + "STRING", + "STRING", + ) + RETURN_NAMES = ( + "pos_prompt", + "neg_prompt", + ) + + FUNCTION = "process" + + CATEGORY = "ACB" + + @classmethod + def IS_CHANGED( + cls, + model, + modelname, + pos_prompt, + neg_prompt, + seed, + debug_level, # pylint: disable=unused-argument + pony_substrings, + wc_process_wildcards, + wc_wildcards_folders, + wc_if_wildcards, + wc_choice_separator, + wc_keep_choices_order, + stn_separator, + stn_ignore_repeats, + stn_join_attention, + cleanup_extra_spaces, + cleanup_empty_constructs, + cleanup_extra_separators, + cleanup_extra_separators2, + cleanup_breaks, + cleanup_breaks_eol, + cleanup_ands, + cleanup_ands_eol, + cleanup_extranetwork_tags, + remove_extranetwork_tags, + ): + new_run = { + "model": model, + "modelname": modelname, + "pos_prompt": pos_prompt, + "neg_prompt": neg_prompt, + "seed": seed, + "pony_substrings": pony_substrings, + "process_wildcards": wc_process_wildcards, + "wildcards_folders": wc_wildcards_folders, + "if_wildcards": wc_if_wildcards, + "choice_separator": wc_choice_separator, + "keep_choices_order": wc_keep_choices_order, + "stn_separator": stn_separator, + "stn_ignore_repeats": stn_ignore_repeats, + "stn_join_attention": stn_join_attention, + "cleanup_extra_spaces": cleanup_extra_spaces, + "cleanup_empty_constructs": cleanup_empty_constructs, + "cleanup_extra_separators": cleanup_extra_separators, + "cleanup_extra_separators2": cleanup_extra_separators2, + "cleanup_breaks": cleanup_breaks, + "cleanup_breaks_eol": cleanup_breaks_eol, + "cleanup_ands": cleanup_ands, + "cleanup_ands_eol": cleanup_ands_eol, + "cleanup_extranetwork_tags": cleanup_extranetwork_tags, + "remove_extranetwork_tags": remove_extranetwork_tags, + } + return new_run.__hash__ + # return float("NaN") + + def process( + self, + model, + modelname, + pos_prompt, + neg_prompt, + seed, + debug_level, + pony_substrings, + wc_process_wildcards, + wc_wildcards_folders, + wc_if_wildcards, + wc_choice_separator, + wc_keep_choices_order, + stn_separator, + stn_ignore_repeats, + stn_join_attention, + cleanup_extra_spaces, + cleanup_empty_constructs, + cleanup_extra_separators, + cleanup_extra_separators2, + cleanup_breaks, + cleanup_breaks_eol, + cleanup_ands, + cleanup_ands_eol, + cleanup_extranetwork_tags, + remove_extranetwork_tags, + ): + model_info = { + "models_path": folder_paths.models_dir, + "model_filename": modelname, # path is relative to checkpoints folder + "is_sd1": model.model.model_config.__class__.__name__ in ("SD15", "SD15_instructpix2pix"), + "is_sd2": model.model.model_config.__class__.__name__ in ("SD20", "SD21UnclipL", "SD21UnclipH"), + "is_sdxl": model.model.model_config.__class__.__name__ + in ( + "SDXL", + "SDXLRefiner", + "SDXL_instructpix2pix", + "Segmind_Vega", + "KOALA_700M", + "KOALA_1B", + ), + "is_ssd": model.model.model_config.__class__.__name__ in ("SSD1B"), + "is_sd3": model.model.model_config.__class__.__name__ in ("SD3"), + "is_flux": model.model.model_config.__class__.__name__ in ("Flux"), + } + # 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 []) + if wc_wildcards_folders == "": + wc_wildcards_folders = os.getenv("WILDCARD_DIR", PPPWildcards.DEFAULT_WILDCARDS_FOLDER) + wildcards_folders = [ + (f if os.path.isabs(f) else os.path.abspath(os.path.join(folder_paths.models_dir, f))) + for f in wc_wildcards_folders.split(",") + if f.strip() != "" + ] + options = { + "debug_level": debug_level, + "pony_substrings": pony_substrings, + "process_wildcards": wc_process_wildcards, + "if_wildcards": wc_if_wildcards, + "choice_separator": wc_choice_separator, + "keep_choices_order": wc_keep_choices_order, + "stn_separator": stn_separator, + "stn_ignore_repeats": stn_ignore_repeats, + "stn_join_attention": stn_join_attention, + "cleanup_extra_spaces": cleanup_extra_spaces, + "cleanup_empty_constructs": cleanup_empty_constructs, + "cleanup_extra_separators": cleanup_extra_separators, + "cleanup_extra_separators2": cleanup_extra_separators2, + "cleanup_breaks": cleanup_breaks, + "cleanup_breaks_eol": cleanup_breaks_eol, + "cleanup_ands": cleanup_ands, + "cleanup_ands_eol": cleanup_ands_eol, + "cleanup_extranetwork_tags": cleanup_extranetwork_tags, + "remove_extranetwork_tags": remove_extranetwork_tags, + } + self.wildcards_obj.refresh_wildcards(debug_level, wildcards_folders if options["process_wildcards"] else None) + ppp = PromptPostProcessor( + self.logger, self.interrupt, model_info, options, self.grammar_content, self.wildcards_obj + ) + pos_prompt, neg_prompt = ppp.process_prompt(pos_prompt, neg_prompt, seed if seed is not None else 1) + return ( + pos_prompt, + neg_prompt, + ) + + def interrupt(self): + nodes.interrupt_processing(True) diff --git a/ppp_logging.py b/ppp_logging.py index 4ae86ce..36d1298 100644 --- a/ppp_logging.py +++ b/ppp_logging.py @@ -1,8 +1,15 @@ +from enum import Enum import logging import sys import copy +class DEBUG_LEVEL(Enum): + none = "none" + minimal = "minimal" + full = "full" + + class PromptPostProcessorLogFactory: # pylint: disable=too-few-public-methods """ Factory class for creating loggers for the PromptPostProcessor module. @@ -42,7 +49,7 @@ class PromptPostProcessorLogFactory: # pylint: disable=too-few-public-methods colored_record = copy.copy(record) levelname = colored_record.levelname seq = self.COLORS.get(levelname, self.COLORS["RESET"]) - colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}" + colored_record.levelname = f"{seq}{levelname:8s}{self.COLORS['RESET']}" return super().format(colored_record) def __init__(self): @@ -61,9 +68,9 @@ class PromptPostProcessorLogFactory: # pylint: disable=too-few-public-methods ppplog.propagate = False if not ppplog.handlers: handler = logging.StreamHandler(sys.stdout) - handler.setFormatter(self.ColoredFormatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s")) + handler.setFormatter(self.ColoredFormatter("%(asctime)s %(levelname)s %(message)s")) # Used in A1111 / Forge / reForge / ComfyUI, but not in SD.Next ppplog.addHandler(handler) - ppplog.setLevel(logging.INFO) + ppplog.setLevel(logging.DEBUG) self.log = PromptPostProcessorLogCustomAdapter(ppplog) @@ -76,12 +83,10 @@ class PromptPostProcessorLogCustomAdapter(logging.LoggerAdapter): def process(self, msg, kwargs): """ Process the log message and keyword arguments. - Args: msg (str): The log message. kwargs (dict): The keyword arguments. - Returns: tuple: A tuple containing the processed log message and keyword arguments. """ - return f"[PromptPostProcessor] {msg}", kwargs + return f"[PPP] {msg}", kwargs diff --git a/ppp_wildcards.py b/ppp_wildcards.py new file mode 100644 index 0000000..16cbb7b --- /dev/null +++ b/ppp_wildcards.py @@ -0,0 +1,209 @@ +import os +import json +from typing import Optional +import yaml + +from ppp_logging import DEBUG_LEVEL + + +class PPPWildcards: + + DEFAULT_WILDCARDS_FOLDER = "wildcards" + + def __init__(self, logger): + self.logger = logger + self.debug_level = DEBUG_LEVEL.none + self.wildcards_folders = [] + self.wildcards = {} + self.wildcard_files = {} + + def refresh_wildcards(self, debug_level: DEBUG_LEVEL, wildcards_folders: Optional[list[str]]): + """ + Initialize the wildcards. + """ + self.debug_level = debug_level + self.wildcards_folders = wildcards_folders + if wildcards_folders is not None: + # if self.debug_level != DEBUG_LEVEL.none: + # self.logger.info("Initializing wildcards...") + # t1 = time.time() + for fullpath in self.wildcard_files.keys(): + path = os.path.dirname(fullpath) + if not os.path.exists(fullpath) or not any( + os.path.commonpath([path, folder]) == folder for folder in self.wildcards_folders + ): + self.__remove_wildcards_from_file(fullpath) + for f in self.wildcards_folders: + self.__get_wildcards_in_directory(f, f) + # t2 = time.time() + # if self.debug_level != DEBUG_LEVEL.none: + # self.logger.info(f"Wildcards init time: {t2 - t1:.3f} seconds") + else: + self.wildcards_folders = [] + self.wildcards = {} + self.wildcard_files = {} + + def __get_keys_in_dict(self, dictionary: dict, prefix="") -> list[str]: + """ + Get all keys in a dictionary. + + Args: + dictionary (dict): The dictionary to check. + prefix (str): The prefix for the current key. + + Returns: + list: A list of all keys in the dictionary, including nested keys. + """ + keys = [] + for key in dictionary.keys(): + if isinstance(dictionary[key], dict): + keys.extend(self.__get_keys_in_dict(dictionary[key], prefix + key + "/")) + else: + keys.append(prefix + str(key)) + return keys + + def __get_nested(self, dictionary: dict, keys: str) -> object: + """ + Get a nested value from a dictionary. + + Args: + dictionary (dict): The dictionary to check. + keys (str): The keys to get the value from. + + Returns: + object: The value of the nested keys in the dictionary. + """ + keys = keys.split("/") + current_dict = dictionary + for key in keys: + current_dict = current_dict.get(key) + if current_dict is None: + return None + return current_dict + + def __remove_wildcards_from_file(self, full_path: str): + """ + Clear all wildcards in a file. + + Args: + full_path (str): The path to the file. + """ + last_modified_cached = self.wildcard_files.get(full_path, None) + if last_modified_cached is not None and self.debug_level != DEBUG_LEVEL.none: + self.logger.debug(f"Removing wildcards from file: {full_path}") + if full_path in self.wildcard_files.keys(): + del self.wildcard_files[full_path] + for key in list(self.wildcards.keys()): + if self.wildcards[key]["file"] == full_path: + del self.wildcards[key] + + def __get_wildcards_in_file(self, base, full_path: str): + """ + Get all wildcards in a file. + + Args: + base (str): The base path for the wildcards. + full_path (str): The path to the file. + """ + last_modified = os.path.getmtime(full_path) + last_modified_cached = self.wildcard_files.get(full_path, None) + if last_modified_cached is not None and last_modified == self.wildcard_files[full_path]: + return + filename = os.path.basename(full_path) + name, extension = os.path.splitext(filename) + if extension not in (".txt", ".json", ".yaml", ".yml"): + return + self.__remove_wildcards_from_file(full_path) + if last_modified_cached is not None and self.debug_level != DEBUG_LEVEL.none: + self.logger.debug(f"Updating wildcards from file: {full_path}") + relfolders = os.path.relpath(os.path.dirname(full_path), base) + if relfolders == ".": + relfolders = "" + elif relfolders != "": + relfolders += "/" + if extension == ".txt": + self.__get_wildcards_in_text_file(full_path, name, relfolders) + elif extension in (".json", ".yaml", ".yml"): + self.__get_wildcards_in_structured_file(full_path, extension, relfolders) + self.wildcard_files[full_path] = last_modified + + def __get_wildcards_in_structured_file(self, full_path, extension, relfolders): + with open(full_path, "r", encoding="utf-8") as file: + if extension == ".json": + content = json.loads(file.read()) + else: + content = yaml.safe_load(file) + keys = self.__get_keys_in_dict(content) + for key in keys: + fullkey = f"{relfolders}{key}" + if self.wildcards.get(fullkey) is not None: + self.logger.warning( + f"Duplicate wildcard '{fullkey}' in file '{full_path}' and '{self.wildcards[fullkey]['file']}'!" + ) + else: + obj = self.__get_nested(content, key) + if obj is not None: + if isinstance(obj, str): + choices = [obj] + elif isinstance(obj, (int, float, bool)): + choices = [str(obj)] + elif isinstance(obj, list) and len(obj) > 0: + choices = [] + for c in obj: + if isinstance(c, str): + choices.append(c) + elif isinstance(c, dict): # we convert the dict to a string + d = "" + if "weight" in c.keys(): + d += str(c["weight"]) + if "if" in c.keys(): + d += f" if {c['if']}" + if d != "": + d += "::" + if "text" in c.keys(): + d += c["text"] + elif "content" in c.keys(): + d += c["content"] + choices.append(d) + else: + obj = None + if obj is None: + self.logger.warning(f"Invalid wildcard '{fullkey}' in file '{full_path}'!") + else: + self.wildcards[fullkey] = {"file": full_path, "choices": choices} + + def __get_wildcards_in_text_file(self, full_path, name, relfolders): + with open(full_path, "r", encoding="utf-8") as file: + text_content = map(lambda x: x.strip("\n\r"), file.readlines()) + text_content = list(filter(lambda x: x.strip() != "" and not x.strip().startswith("#"), text_content)) + text_content = [x.split("#")[0].rstrip() if len(x.split("#")) > 1 else x for x in text_content] + fullkey = f"{relfolders}{name}" + if self.wildcards.get(fullkey) is not None: + self.logger.warning( + f"Duplicate wildcard '{fullkey}' in file '{full_path}' and '{self.wildcards[fullkey]['file']}'!" + ) + else: + if len(text_content) == 0: + self.logger.warning(f"Invalid wildcard in file '{full_path}'!") + else: + self.wildcards[fullkey] = {"file": full_path, "choices": text_content} + + def __get_wildcards_in_directory(self, base: str, directory: str): + """ + Get all wildcards in a directory. + + Args: + base (str): The base path for the wildcards. + directory (str): The path to the directory. + """ + if not os.path.exists(directory): + self.logger.warning(f"Wildcard 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_wildcards_in_directory(base, full_path) + elif os.path.isfile(full_path): + self.__get_wildcards_in_file(base, full_path) diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..76597fa --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +lark diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py index 95c32ad..b7d0212 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -3,19 +3,23 @@ if __name__ == "__main__": import sys import os +import time -sys.path.insert(1, os.path.join(sys.path[0], "..")) +sys.path.append(os.path.join(sys.path[0], "..")) from modules import scripts, shared, script_callbacks from modules.processing import StableDiffusionProcessing from modules.shared import opts +from modules.paths import models_path import gradio as gr from ppp import PromptPostProcessor -from ppp_logging import PromptPostProcessorLogFactory +from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory +from ppp_cache import PPPLRUCache +from ppp_wildcards import PPPWildcards -class PromptPostProcessorScript(scripts.Script): +class PromptPostProcessorA1111Script(scripts.Script): """ This class represents a script for prompt post-processing. It is responsible for processing prompts and applying various settings and cleanup operations. @@ -28,6 +32,7 @@ class PromptPostProcessorScript(scripts.Script): title(): Returns the title of the script. show(is_img2img): Determines whether the script should be shown based on the input type. process(p, *args, **kwargs): Processes the prompts and applies post-processing operations. + ppp_interrupt(): Interrupts the generation. __on_ui_settings(): Callback function for UI settings. """ @@ -43,12 +48,15 @@ class PromptPostProcessorScript(scripts.Script): Returns: None """ - if not hasattr(self, "ppp_callbacks_added"): - lf = PromptPostProcessorLogFactory() - self.ppp_logger = lf.log - self.ppp_debug = getattr(opts, "ppp_gen_debug", False) if opts is not None else False - script_callbacks.on_ui_settings(self.__on_ui_settings) - self.ppp_callbacks_added = True + lf = PromptPostProcessorLogFactory() + self.name = PromptPostProcessor.NAME + self.ppp_logger = lf.log + self.ppp_debug_level = DEBUG_LEVEL(getattr(opts, "ppp_gen_debug_level", DEBUG_LEVEL.none.value)) + self.lru_cache = PPPLRUCache(1000) + 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() + self.wildcards_obj = PPPWildcards(lf.log) def title(self): """ @@ -81,29 +89,155 @@ class PromptPostProcessorScript(scripts.Script): Returns: None """ + t1 = time.time() + if getattr(opts, "prompt_attention", "") == "Compel parser": + self.ppp_logger.warning("Compel parser is not supported!") is_i2i = getattr(p, "init_images", [None])[0] is not None - self.ppp_debug = getattr(opts, "ppp_gen_debug", False) if opts is not None else False - if self.ppp_debug: - self.ppp_logger.info(f"Post-processing prompts ({'i2i' if is_i2i else 't2i'} mode)") - ppp = PromptPostProcessor(self, p, shared.state, opts, is_i2i) - # processes regular prompts - if ( - hasattr(p, "all_prompts") - and p.all_prompts is not None - and hasattr(p, "all_negative_prompts") - and p.all_negative_prompts is not None - ): - for i, (prompt, negative_prompt) in enumerate(zip(p.all_prompts, p.all_negative_prompts)): - p.all_prompts[i], p.all_negative_prompts[i] = ppp.process_prompt(prompt, negative_prompt) + self.ppp_debug_level = DEBUG_LEVEL(getattr(opts, "ppp_gen_debug_level", DEBUG_LEVEL.none.value)) + do_i2i = getattr(opts, "ppp_gen_doi2i", False) + if is_i2i and not do_i2i: + if self.ppp_debug_level != DEBUG_LEVEL.none: + self.ppp_logger.info("Not processing the prompt for i2i") + return + if self.ppp_debug_level != DEBUG_LEVEL.none: + self.ppp_logger.info(f"Post-processing prompts ({'i2i' if is_i2i else 't2i'})") + model_info = { + "models_path": models_path, + "model_filename": getattr(p.sd_model.sd_checkpoint_info, "filename", ""), # path is absolute + "is_sd1": False, # Stable Diffusion 1 + "is_sd2": False, # Stable Diffusion 2 + "is_sdxl": False, # Stable Diffusion XL + "is_ssd": False, # Segmind Stable Diffusion 1B + "is_sd3": False, # Stable Diffusion 3 + "is_flux": False, # Flux + } + app = ( + "forge" + if hasattr(p.sd_model, "model_config") + else "sdnext" if hasattr(p.sd_model, "is_sdxl") and not hasattr(p.sd_model, "is_ssd") else "a1111" + ) + if app == "sdnext": + # cannot differenciate SD1 and SD2, we set True to both + # LatentDiffusion is for the original backend, StableDiffusionPipeline is for the diffusers backend + model_info["is_sd1"] = p.sd_model.__class__.__name__ in ("LatentDiffusion", "StableDiffusionPipeline") + model_info["is_sd2"] = p.sd_model.__class__.__name__ in ("LatentDiffusion", "StableDiffusionPipeline") + model_info["is_sdxl"] = p.sd_model.__class__.__name__ == "StableDiffusionXLPipeline" + model_info["is_ssd"] = False # ? + model_info["is_sd3"] = p.sd_model.__class__.__name__ == "StableDiffusion3Pipeline" + model_info["is_flux"] = False + elif app == "forge": + model_info["is_sd1"] = getattr(p.sd_model, "is_sd1", False) + model_info["is_sd2"] = getattr(p.sd_model, "is_sd2", False) + model_info["is_sdxl"] = getattr(p.sd_model, "is_sdxl", False) + model_info["is_ssd"] = False # ? + model_info["is_sd3"] = getattr(p.sd_model, "is_sd3", False) + model_info["is_flux"] = p.sd_model.model_config.__class__.__name__ == "Flux" + else: # assume A1111 compatible (p.sd_model.__class__.__name__=="DiffusionEngine") + model_info["is_sd1"] = getattr(p.sd_model, "is_sd1", False) + model_info["is_sd2"] = getattr(p.sd_model, "is_sd2", False) + model_info["is_sdxl"] = getattr(p.sd_model, "is_sdxl", False) + model_info["is_ssd"] = getattr(p.sd_model, "is_ssd", False) + model_info["is_sd3"] = getattr(p.sd_model, "is_sd3", False) + model_info["is_flux"] = False + wc_wildcards_folders = getattr(opts, "ppp_wil_wildcardsfolders", "") + if wc_wildcards_folders == "": + wc_wildcards_folders = os.getenv("WILDCARD_DIR", PPPWildcards.DEFAULT_WILDCARDS_FOLDER) + wildcards_folders = [ + (f if os.path.isabs(f) else os.path.abspath(os.path.join(models_path, f))) + for f in wc_wildcards_folders.split(",") + if f.strip() != "" + ] + options = { + "debug_level": getattr(opts, "ppp_gen_debug_level", DEBUG_LEVEL.none.value), + "pony_substrings": getattr(opts, "ppp_gen_ponysubstrings", PromptPostProcessor.DEFAULT_PONY_SUBSTRINGS), + "process_wildcards": getattr(opts, "ppp_wil_process_wildcards", True), + "if_wildcards": getattr(opts, "ppp_wil_ifwildcards", PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value), + "choice_separator": getattr(opts, "ppp_wil_choice_separator", PromptPostProcessor.DEFAULT_CHOICE_SEPARATOR), + "keep_choices_order": getattr(opts, "ppp_wil_keep_choices_order", False), + "stn_separator": getattr(opts, "ppp_stn_separator", PromptPostProcessor.DEFAULT_STN_SEPARATOR), + "stn_ignore_repeats": getattr(opts, "ppp_stn_ignorerepeats", True), + "stn_join_attention": getattr(opts, "ppp_stn_joinattention", True), + "cleanup_extra_spaces": getattr(opts, "ppp_cup_extraspaces", True), + "cleanup_empty_constructs": getattr(opts, "ppp_cup_emptyconstructs", True), + "cleanup_extra_separators": getattr(opts, "ppp_cup_extraseparators", True), + "cleanup_extra_separators2": getattr(opts, "ppp_cup_extraseparators2", True), + "cleanup_breaks": getattr(opts, "ppp_cup_breaks", True), + "cleanup_breaks_eol": getattr(opts, "ppp_cup_breaks_eol", False), + "cleanup_ands": getattr(opts, "ppp_cup_ands", True), + "cleanup_ands_eol": getattr(opts, "ppp_cup_ands_eol", False), + "cleanup_extranetwork_tags": getattr(opts, "ppp_cup_extranetworktags", False), + "remove_extranetwork_tags": getattr(opts, "ppp_rem_removeextranetworktags", False), + } + self.wildcards_obj.refresh_wildcards( + self.ppp_debug_level, wildcards_folders if options["process_wildcards"] else None + ) + ppp = PromptPostProcessor( + self.ppp_logger, self.ppp_interrupt, model_info, options, self.grammar_content, self.wildcards_obj + ) + prompts_list = [] + + seeds = getattr(p, "all_seeds", []) + subseeds = getattr(p, "all_subseeds", []) + subseed_strength = getattr(p, "subseed_strength", 0.0) + if subseed_strength > 0: + calculated_seeds = [ + int(subseed * subseed_strength + seed * (1 - subseed_strength)) + for seed, subseed in zip(seeds, subseeds) + ] + else: + calculated_seeds = seeds + if len(set(calculated_seeds)) < len(calculated_seeds): + self.ppp_logger.info("Adjusting seeds because some are equal.") + calculated_seeds = [seed + i for i, seed in enumerate(calculated_seeds)] + + # adds regular prompts + rpr = getattr(p, "all_prompts", None) + rnr = 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 - if ( - hasattr(p, "all_hr_prompts") - and p.all_hr_prompts is not None - and hasattr(p, "all_hr_negative_prompts") - and p.all_hr_negative_prompts is not None - ): - for i, (hr_prompt, hr_negative_prompt) in enumerate(zip(p.all_hr_prompts, p.all_hr_negative_prompts)): - p.all_hr_prompts[i], p.all_hr_negative_prompts[i] = ppp.process_prompt(hr_prompt, hr_negative_prompt) + rph = getattr(p, "all_hr_prompts", None) + rnh = getattr(p, "all_hr_negative_prompts", None) + if rph is not None and rnh is not None: + 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 + ] + + # processes prompts + for i, (prompttype, seed, prompt, negative_prompt) in enumerate(prompts_list): + if self.ppp_debug_level != DEBUG_LEVEL.none: + self.ppp_logger.info(f"processing prompts[{i+1}] ({prompttype})") + if self.lru_cache.get((seed, prompt, negative_prompt)) is None: + pp, np = ppp.process_prompt(prompt, negative_prompt, seed) + self.lru_cache.put((seed, prompt, negative_prompt), (pp, np)) + # adds also the result so i2i doesn't process it unnecessarily + self.lru_cache.put((seed, pp, np), (pp, np)) + elif self.ppp_debug_level != DEBUG_LEVEL.none: + self.ppp_logger.info("result already in cache") + + # updates the prompts + if rpr is not None and rnr is not None: + for i, (seed, prompt, negative_prompt) in enumerate(zip(calculated_seeds, rpr, rnr)): + found = self.lru_cache.get((seed, prompt, negative_prompt)) + if found is not None: + rpr[i] = found[0] + rnr[i] = found[1] + if rph is not None and rnh is not None: + for i, (seed, prompt, negative_prompt) in enumerate(zip(calculated_seeds, rph, rnh)): + found = self.lru_cache.get((seed, prompt, negative_prompt)) + if found is not None: + rph[i] = found[0] + rnh[i] = found[1] + + t2 = time.time() + if self.ppp_debug_level != DEBUG_LEVEL.none: + self.ppp_logger.info(f"process time: {t2 - t1:.3f} seconds") def ppp_interrupt(self): """ @@ -114,171 +248,266 @@ class PromptPostProcessorScript(scripts.Script): """ shared.state.interrupted = True - def __on_ui_settings(self): - """ - Callback function for UI settings. - Returns: - None - """ - # general settings - section = ("prompt-post-processor", PromptPostProcessor.NAME) - shared.opts.add_option( - key="ppp_gen_sep", info=shared.OptionInfo("

General settings

", "", gr.HTML, section=section) - ) - shared.opts.add_option( - key="ppp_gen_debug", - info=shared.OptionInfo( - False, - label="Debug", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_gen_ifwildcards", - info=shared.OptionInfo( - default=PromptPostProcessor.IFWILDCARDS_CHOICES["ignore"], - label="What to do with remaining wildcards?", - component=gr.Radio, - component_args={"choices": PromptPostProcessor.IFWILDCARDS_CHOICES.values()}, - section=section, - ), - ) +def on_ui_settings(): + """ + Callback function for UI settings. - # content removal settings - shared.opts.add_option( - key="ppp_rem_sep", info=shared.OptionInfo("

Content removal settings

", "", gr.HTML, section=section) - ) - shared.opts.add_option( - key="ppp_rem_removeextranetworktags", - info=shared.OptionInfo( - False, - label="Remove extra network tags", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_rem_if", info=shared.OptionInfo("

* Parsing of the 'if' commands cannot be disabled

", "", gr.HTML, section=section) - ) + Returns: + None + """ - # send to negative settings - shared.opts.add_option( - key="ppp_stn_sep", - info=shared.OptionInfo("

Send to Negative settings

", "", gr.HTML, section=section), + section = ("prompt-post-processor", PromptPostProcessor.NAME) + + def import_old_settings(names, default): + for name in names: + if hasattr(opts, name): + return getattr(opts, name) + return default + + def import_bool_to_any(name, value_false, value_true, default): + if hasattr(opts, name): + return value_true if getattr(opts, name) else value_false + return default + + def new_html_title(title): + info = shared.OptionInfo( + title, + "", + gr.HTML, + section=section, ) - shared.opts.add_option( - key="ppp_stn_doi2i", - info=shared.OptionInfo( - False, - label="Apply in img2img (this includes any pass that contains an initial image, like refiner, hires fix, adetailer)", - section=section, + info.do_not_save = True + return info + + # general settings + shared.opts.add_option( + key="ppp_gen_sep", + info=new_html_title("

General settings

"), + ) + shared.opts.add_option( + key="ppp_gen_debug_level", + info=shared.OptionInfo( + default=import_bool_to_any( + "ppp_gen_debug", + DEBUG_LEVEL.minimal.value, + DEBUG_LEVEL.full.value, + DEBUG_LEVEL.minimal.value, ), - ) - shared.opts.add_option( - key="ppp_stn_separator", - info=shared.OptionInfo( - PromptPostProcessor.DEFAULT_STN_SEPARATOR, - label="Separator used when adding to the negative prompt", - section=section, + label="Debug level", + component=gr.Radio, + component_args={ + "choices": ( + ("None", DEBUG_LEVEL.none.value), + ("Minimal", DEBUG_LEVEL.minimal.value), + ("Full", DEBUG_LEVEL.full.value), + ), + }, + section=section, + ), + ) + shared.opts.add_option( + key="ppp_gen_ponysubstrings", + info=shared.OptionInfo( + PromptPostProcessor.DEFAULT_PONY_SUBSTRINGS, + label="Comma separated list of substrings to look for in the model full filename to flag it as Pony (case insensitive)", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_gen_doi2i", + info=shared.OptionInfo( + False, + label="Apply in img2img", + comment_after='(this includes any pass that contains an initial image, like adetailer)', + section=section, + ), + ) + + # wildcard settings + shared.opts.add_option( + key="ppp_wil_sep", + info=new_html_title('

Wildcard settings

'), + ) + shared.opts.add_option( + key="ppp_wil_processwildcards", + info=shared.OptionInfo( + True, + label="Process wildcards", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_wil_wildcardsfolders", + info=shared.OptionInfo( + PPPWildcards.DEFAULT_WILDCARDS_FOLDER, + label="Wildcards folders", + comment_after='(absolute or relative to the models folder)', + section=section, + ), + ) + shared.opts.add_option( + key="ppp_wil_ifwildcards", + info=shared.OptionInfo( + default=import_old_settings( + ["ppp_gen_ifwildcards", "ppp_ifwildcards"], + PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value, ), - ) - shared.opts.add_option( - key="ppp_stn_ignorerepeats", - info=shared.OptionInfo( - True, - label="Ignore repeated content", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_stn_joinattention", - info=shared.OptionInfo( - True, - label="Join attention modifiers (weights) when possible", - section=section, - ), - ) - # clean-up settings - shared.opts.add_option( - key="ppp_cup_sep", info=shared.OptionInfo("

Clean-up settings

", "", gr.HTML, section=section) - ) - shared.opts.add_option( - key="ppp_cup_doi2i", - info=shared.OptionInfo( - False, - label="Apply in img2img (this includes any pass that contains an initial image, like refiner, hires fix, adetailer)", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_cup_emptyconstructs", - info=shared.OptionInfo( - True, - label="Remove empty constructs (attention, alternation, scheduling)", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_cup_extraseparators", - info=shared.OptionInfo( - True, - label="Remove extra separators", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_cup_extraseparators2", - info=shared.OptionInfo( - True, - label="Remove additional extra separators", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_cup_breaks", - info=shared.OptionInfo( - True, - label="Clean up around BREAKs", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_cup_breaks_eol", - info=shared.OptionInfo( - False, - label="Use EOL instead of Space before BREAKs", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_cup_ands", - info=shared.OptionInfo( - True, - label="Clean up around ANDs", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_cup_ands_eol", - info=shared.OptionInfo( - False, - label="Use EOL instead of Space before ANDs", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_cup_extranetworktags", - info=shared.OptionInfo( - False, - label="Clean up around extra network tags", - section=section, - ), - ) - shared.opts.add_option( - key="ppp_cup_extraspaces", - info=shared.OptionInfo( - True, - label="Remove extra spaces", - section=section, - ), - ) + label="What to do with remaining/invalid wildcards?", + component=gr.Radio, + component_args={ + "choices": ( + ("Ignore", PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value), + ("Remove", PromptPostProcessor.IFWILDCARDS_CHOICES.remove.value), + ("Add visible warning", PromptPostProcessor.IFWILDCARDS_CHOICES.warn.value), + ("Stop the generation", PromptPostProcessor.IFWILDCARDS_CHOICES.stop.value), + ) + }, + section=section, + ), + ) + shared.opts.add_option( + key="ppp_wil_choice_separator", + info=shared.OptionInfo( + PromptPostProcessor.DEFAULT_CHOICE_SEPARATOR, + label="Default separator used when adding multiple choices", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_wil_keep_choices_order", + info=shared.OptionInfo( + False, + label="Keep the order of selected choices", + section=section, + ), + ) + + # content removal settings + shared.opts.add_option( + key="ppp_rem_sep", + info=new_html_title('

Content removal settings

'), + ) + shared.opts.add_option( + key="ppp_rem_removeextranetworktags", + info=shared.OptionInfo( + False, + label="Remove extra network tags", + section=section, + ), + ) + + # send to negative settings + shared.opts.add_option( + key="ppp_stn_sep", + info=new_html_title('

Send to Negative settings

'), + ) + shared.opts.add_option( + key="ppp_stn_separator", + info=shared.OptionInfo( + PromptPostProcessor.DEFAULT_STN_SEPARATOR, + label="Separator used when adding to the negative prompt", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_stn_ignorerepeats", + info=shared.OptionInfo( + True, + label="Ignore repeated content", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_stn_joinattention", + info=shared.OptionInfo( + True, + label="Join attention modifiers (weights) when possible", + section=section, + ), + ) + # clean-up settings + shared.opts.add_option( + key="ppp_cup_sep", + info=new_html_title('

Clean-up settings

'), + ) + shared.opts.add_option( + key="ppp_cup_emptyconstructs", + info=shared.OptionInfo( + True, + label="Remove empty constructs (attention, alternation, scheduling)", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_cup_extraseparators", + info=shared.OptionInfo( + True, + label="Remove extra separators", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_cup_extraseparators2", + info=shared.OptionInfo( + True, + label="Remove additional extra separators", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_cup_breaks", + info=shared.OptionInfo( + True, + label="Clean up around BREAKs", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_cup_breaks_eol", + info=shared.OptionInfo( + False, + label="Use EOL instead of Space before BREAKs", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_cup_ands", + info=shared.OptionInfo( + True, + label="Clean up around ANDs", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_cup_ands_eol", + info=shared.OptionInfo( + False, + label="Use EOL instead of Space before ANDs", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_cup_extranetworktags", + info=shared.OptionInfo( + False, + label="Clean up around extra network tags", + section=section, + ), + ) + shared.opts.add_option( + key="ppp_cup_extraspaces", + info=shared.OptionInfo( + True, + label="Remove extra spaces", + section=section, + ), + ) + + # Remove old settings + # for name in ["ppp_gen_ifwildcards", "ppp_ifwildcards", "ppp_gen_debug", "ppp_stn_doi2i", "ppp_cup_doi2i"]: + # if hasattr(opts, name): + # delattr(opts, name) + + +script_callbacks.on_ui_settings(on_ui_settings) diff --git a/tests/tests.py b/tests/tests.py index e27b2f7..804b1f0 100644 --- a/tests/tests.py +++ b/tests/tests.py @@ -1,27 +1,18 @@ +from collections import namedtuple import logging import unittest import sys import os -sys.path.insert(1, os.path.join(sys.path[0], "..")) +from ppp_wildcards import PPPWildcards + +sys.path.append(os.path.join(sys.path[0], "..")) from ppp import PromptPostProcessor -from ppp_logging import PromptPostProcessorLogFactory +from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory -class DictToObj: # pylint: disable=too-few-public-methods - """ - Converts a dictionary to an object with attribute access. - from https://joelmccune.com/python-dictionary-as-object/ - """ - - def __init__(self, in_dict: dict): - assert isinstance(in_dict, dict) - for key, val in in_dict.items(): - if isinstance(val, (list, tuple)): - setattr(self, key, [DictToObj(x) if isinstance(x, dict) else x for x in val]) - else: - setattr(self, key, DictToObj(val) if isinstance(val, dict) else val) +PromptPair = namedtuple("PromptPair", ["prompt", "negative_prompt"], defaults=["", ""]) class TestPromptPostProcessor(unittest.TestCase): @@ -34,334 +25,667 @@ class TestPromptPostProcessor(unittest.TestCase): Set up the test case by initializing the necessary objects and configurations. """ lf = PromptPostProcessorLogFactory() - self.ppp_logger = lf.log - self.ppp_logger.setLevel(logging.DEBUG) - self.__defopts = DictToObj( - { - "ppp_gen_debug": True, - "ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["ignore"], - "ppp_stn_doi2i": False, - "ppp_stn_separator": ", ", - "ppp_stn_ignore_repeats": True, - "ppp_stn_join_attention": True, - "ppp_cup_doi2i": False, - "ppp_cup_emptyconstructs": True, - "ppp_cup_extraseparators": True, - "ppp_cup_extraseparators2": True, - "ppp_cup_extraspaces": True, - "ppp_cup_breaks": True, - "ppp_cup_breaks_eol": False, - "ppp_cup_ands": True, - "ppp_cup_ands_eol": False, - "ppp_cup_extranetworktags": True, - "ppp_rem_removeextranetworktags": False, - } + self.__ppp_logger = lf.log + self.__ppp_logger.setLevel(logging.DEBUG) + self.__defopts = { + "debug_level": DEBUG_LEVEL.full.value, + "pony_substrings": PromptPostProcessor.DEFAULT_PONY_SUBSTRINGS, + "process_wildcards": True, + "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value, + "choice_separator": ", ", + "keep_choices_order": False, + "stn_separator": ", ", + "stn_ignore_repeats": True, + "stn_join_attention": True, + "cleanup_empty_constructs": True, + "cleanup_extra_separators": True, + "cleanup_extra_separators2": True, + "cleanup_extra_spaces": True, + "cleanup_breaks": True, + "cleanup_breaks_eol": False, + "cleanup_ands": True, + "cleanup_ands_eol": False, + "cleanup_extranetwork_tags": True, + "remove_extranetwork_tags": False, + } + self.__def_model_info = { + "is_sd1": False, + "is_sd2": False, + "is_sdxl": True, + "is_ssd": False, + "is_sd3": False, + "is_flux": False, + "models_path": "./webui/models", + "model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", + } + self.__interrupted = False + self.__wildcards_obj = PPPWildcards(lf.log) + self.__wildcards_obj.refresh_wildcards( + DEBUG_LEVEL.full, + [ + os.path.abspath(os.path.join(os.path.dirname(__file__), "wildcards")), + os.path.abspath(os.path.join(os.path.dirname(__file__), "wildcards2")), + ], ) - self.__nocupopts = DictToObj( + 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() + self.__defppp = PromptPostProcessor( + self.__ppp_logger, + self.__interrupt, + self.__def_model_info, + self.__defopts, + self.__grammar_content, + self.__wildcards_obj, + ) + self.__nocupppp = PromptPostProcessor( + self.__ppp_logger, + self.__interrupt, + self.__def_model_info, { - "ppp_gen_debug": True, - "ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["ignore"], - "ppp_stn_doi2i": False, - "ppp_stn_separator": ", ", - "ppp_stn_ignore_repeats": True, - "ppp_stn_join_attention": True, - "ppp_cup_doi2i": False, - "ppp_cup_emptyconstructs": False, - "ppp_cup_extraseparators": False, - "ppp_cup_extraseparators2": False, - "ppp_cup_extraspaces": False, - "ppp_cup_breaks": False, - "ppp_cup_breaks_eol": False, - "ppp_cup_ands": False, - "ppp_cup_ands_eol": False, - "ppp_cup_extranetworktags": False, - "ppp_rem_removeextranetworktags": False, - } + **self.__defopts, + "cleanup_empty_constructs": False, + "cleanup_extra_separators": False, + "cleanup_extra_separators2": False, + "cleanup_extra_spaces": False, + "cleanup_breaks": False, + "cleanup_breaks_eol": False, + "cleanup_ands": False, + "cleanup_ands_eol": False, + "cleanup_extranetwork_tags": False, + }, + self.__grammar_content, + self.__wildcards_obj, ) - self.__defprocessing = DictToObj({"sd_model": DictToObj({"is_sd1": False, "is_sd2": False, "is_sdxl": True})}) - self.__defstate = None - self.defppp = PromptPostProcessor(self, self.__defprocessing, self.__defstate, self.__defopts) - self.nocupppp = PromptPostProcessor(self, self.__defprocessing, self.__defstate, self.__nocupopts) - def process( + def __interrupt(self): + self.__interrupted = True + + def __process( self, - prompt, - negative_prompt, - expected_prompt, - expected_negative_prompt, + input_prompts: PromptPair, + expected_output_prompts: PromptPair | list[PromptPair], + seed: int = 1, ppp=None, + interrupted=False, ): """ Process the prompt and compare the results with the expected prompts. Args: - prompt (str): The input prompt. - negative_prompt (str): The input negative prompt. - expected_prompt (str): The expected output prompt. - expected_negative_prompt (str): The expected output negative prompt. + input_prompts (PromptPair): The input prompts. + expected_output_prompts (PromptPair | list[PromptPair]): The expected prompts. + seed (int, optional): The seed value. Defaults to 1. ppp (object, optional): The post-processor object. Defaults to None. + interrupted (bool, optional): The interrupted flag. Defaults to False. Returns: None """ - the_obj = self.defppp if ppp is None else ppp - result_prompt, result_negative_prompt = the_obj.process_prompt(prompt, negative_prompt) - self.assertEqual(result_prompt, expected_prompt, f"Prompt should be '{expected_prompt}'") - self.assertEqual( - result_negative_prompt, - expected_negative_prompt, - f"Negative Prompt should be '{expected_negative_prompt}'", - ) + the_obj = ppp or self.__defppp + out = expected_output_prompts if isinstance(expected_output_prompts, list) else [expected_output_prompts] + for eo in out: + result_prompt, result_negative_prompt = the_obj.process_prompt( + input_prompts.prompt, + input_prompts.negative_prompt, + seed, + ) + self.assertEqual(self.__interrupted, interrupted, "Interrupted flag is incorrect") + if not self.__interrupted: + self.assertEqual(result_prompt, eo.prompt, "Incorrect prompt") + self.assertEqual(result_negative_prompt, eo.negative_prompt, "Incorrect negative prompt") + seed += 1 # Send To Negative tests - def test_nt_simple_oldformat(self): # negtags with different parameters and separations - self.process( - "flowers, , , ", - "normal quality, worse quality", - "flowers", - "red, green, yellow, normal quality, purple, worse quality, black, blue", + def test_stn_simple(self): # negtags with different parameters and separations + self.__process( + PromptPair( + "flowersred, green, blueyellow, purpleblack", + "normal quality, worse quality", + ), + PromptPair("flowers", "red, green, yellow, normal quality, purple, worse quality, black, blue"), ) - def test_nt_simple(self): # negtags with different parameters and separations - self.process( - "flowersred, green, blueyellow, purpleblack", - "normal quality, worse quality", - "flowers", - "red, green, yellow, normal quality, purple, worse quality, black, blue", - ) - - def test_nt_complex(self): # complex negtags - self.process( - "red ((pink)), flowers purple, mauveblue, yellow green", - "normal quality, , bad quality, worse quality", - "flowers", - "red, (pink:1.21), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue", - ) - - def test_nt_complex_nocleanup(self): # complex negtags with no cleanup - self.process( - "red ((pink)), flowers purple, mauveblue, yellow green", - "normal quality, , bad quality, worse quality", - " (()), flowers , , ", - "red, (pink:1.21), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue", - self.nocupppp, - ) - - def test_nt_inside_attention(self): # negtag inside attention - self.process( - "[neg1] this is a ((testneg2) (test:2.0): 1.5 ) (red[square]:1.5)", - "normal quality", - "this is a ((test) (test:2.0):1.5) (red:1.5)", - "[neg1], ([square]:1.5), normal quality, (neg2:1.65)", - ) - - def test_nt_inside_alternation(self): # negtag inside alternation - self.process( - "this is a (([complexneg1|simpleneg2|regularneg3] test)(test:2.0):1.5)", - "normal quality", - "this is a (([complex|simple|regular] test)(test:2.0):1.5)", - "([neg1||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality", - ) - - def test_nt_inside_alternation_recursive(self): # negtag inside alternation (recursive alternation) - self.process( - "this is a (([complexneg1[one|twoneg12||three|four(neg14)]|simpleneg2|regularneg3] test)(test:2.0):1.5)", - "normal quality", - "this is a (([complex[one|two||three|four]|simple|regular] test)(test:2.0):1.5)", - "([neg1||]:1.65), ([[|neg12|||]||]:1.65), ([[||||(neg14)]||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality", - ) - - def test_nt_inside_scheduling(self): # negtag inside scheduling - self.process( - "this is [abcneg1:defneg2: 5 ]", - "normal quality", - "this is [abc:def:5]", - "[neg1::5], normal quality, [neg2:5]", - ) - - def test_nt_complex_features(self): # complex negtags with AND, BREAK and other features - self.process( - "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loraword AND AND hypernetword :0.3", - "normal quality, ", - "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK with [abc:def:5]:0.5 AND loraword AND hypernetword :0.3", - "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", - ) - - def test_nt_complex_features_newformat(self): # complex negtags with AND, BREAK and other features (new format) - self.process( - "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loraword AND AND hypernetword :0.3", - "normal quality, ", - "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK with [abc:def:5]:0.5 AND loraword AND hypernetword :0.3", - "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", - ) - - # Wildcard tests - - def test_wc_ignore(self): # wildcards with ignore option - self.process( - "__bad_wildcard__", - "{option1|option2}", - "__bad_wildcard__", - "{option1|option2}", - PromptPostProcessor( - self, - self.__defprocessing, - self.__defstate, - DictToObj( - { - **self.__defopts.__dict__, - "ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["ignore"], - } - ), + def test_stn_complex(self): # complex negtags + self.__process( + PromptPair( + "red ((pink)), flowers purple, mauveblue, yellow green", + "normal quality, , bad quality, worse quality", + ), + PromptPair( + "flowers", + "red, (pink:1.21), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue", ), ) - def test_wc_remove(self): # wildcards with remove option - self.process( - "[neg5] this is: __bad_wildcard__ a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5] ", - "normal quality, {option1|option2}", - "this is: a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK with [abc:def:5]", - "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", - PromptPostProcessor( - self, - self.__defprocessing, - self.__defstate, - DictToObj( - { - **self.__defopts.__dict__, - "ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["remove"], - } - ), + def test_stn_complex_nocleanup(self): # complex negtags with no cleanup + self.__process( + PromptPair( + "red ((pink)), flowers purple, mauveblue, yellow green", + "normal quality, , bad quality, worse quality", + ), + PromptPair( + " (()), flowers , , ", + "red, (pink:1.21), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue", + ), + ppp=self.__nocupppp, + ) + + def test_stn_inside_attention(self): # negtag inside attention + self.__process( + PromptPair( + "[neg1] this is a ((testneg2) (test:2.0): 1.5 ) (red[square]:1.5)", + "normal quality", + ), + PromptPair( + "this is a ((test) (test:2.0):1.5) (red:1.5)", "[neg1], ([square]:1.5), normal quality, (neg2:1.65)" ), ) - def test_wc_warn(self): # wildcards with warn option - self.process( - "__bad_wildcard__", - "{option1|option2}", - PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", - "{option1|option2}", - PromptPostProcessor( - self, - self.__defprocessing, - self.__defstate, - DictToObj( - { - **self.__defopts.__dict__, - "ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["warn"], - } - ), + def test_stn_inside_alternation(self): # negtag inside alternation + self.__process( + PromptPair( + "this is a (([complexneg1|simpleneg2|regularneg3] test)(test:2.0):1.5)", + "normal quality", + ), + PromptPair( + "this is a (([complex|simple|regular] test)(test:2.0):1.5)", + "([neg1||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality", ), ) - def test_wc_stop(self): # wildcards with stop option - self.process( - "__bad_wildcard__", - "{option1|option2}", - PromptPostProcessor.WILDCARD_STOP + "__bad_wildcard__", - PromptPostProcessor.WILDCARD_STOP + "{option1|option2}", - PromptPostProcessor( - self, - self.__defprocessing, - self.__defstate, - DictToObj( - { - **self.__defopts.__dict__, - "ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["stop"], - } - ), + def test_stn_inside_alternation_recursive(self): # negtag inside alternation (recursive alternation) + self.__process( + PromptPair( + "this is a (([complexneg1[one|twoneg12||three|four(neg14)]|simpleneg2|regularneg3] test)(test:2.0):1.5)", + "normal quality", + ), + PromptPair( + "this is a (([complex[one|two||three|four]|simple|regular] test)(test:2.0):1.5)", + "([neg1||]:1.65), ([[|neg12|||]||]:1.65), ([[||||(neg14)]||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality", + ), + ) + + def test_stn_inside_scheduling(self): # negtag inside scheduling + self.__process( + PromptPair("this is [abcneg1:defneg2: 5 ]", "normal quality"), + [PromptPair("this is [abc:def:5]", "[neg1::5], normal quality, [neg2:5]")], + ) + + def test_stn_complex_features(self): # complex negtags with AND, BREAK and other features + self.__process( + PromptPair( + "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loratrigger AND AND hypernettrigger :0.3", + "normal quality, ", + ), + PromptPair( + "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger AND hypernettrigger :0.3", + "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", + ), + ) + + def test_stn_complex_features_newformat(self): # complex negtags with AND, BREAK and other features (new format) + self.__process( + PromptPair( + "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loratrigger AND AND hypernettrigger :0.3", + "normal quality, ", + ), + PromptPair( + "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger AND hypernettrigger :0.3", + "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", ), ) # Cleanup tests def test_cl_simple(self): # simple cleanup - self.process( - " this is a ((test ), , () [] ( , test ,:2.0):1.5) (red:1.5) ", - " normal quality ", - "this is a ((test), (test,:2.0):1.5) (red:1.5)", - "normal quality", + self.__process( + PromptPair(" this is a ((test ), , , (), , [] ( , test ,:2.0):1.5) (red:1.5) ", " normal quality "), + PromptPair("this is a ((test), (test,:2.0):1.5) (red:1.5)", "normal quality"), ) def test_cl_complex(self): # complex cleanup - self.process( - " this is BREAKABLE a ((test), ,AND AND() [] ANDERSON (test:2.0):1.5) :o BREAK \n BREAK (red:1.5) ", - " [:hands, feet, :0.15]normal quality ", - "this is BREAKABLE a ((test) AND ANDERSON (test:2.0):1.5) :o BREAK (red:1.5)", - "[:hands, feet, :0.15]normal quality", + self.__process( + PromptPair( + " this is BREAKABLE a ((test)), ,AND AND(() [] ANDERSON (test:2.0):1.5) :o BREAK \n BREAK (red:1.5) ", + " [:hands, feet, :0.15]normal quality ", + ), + PromptPair( + "this is BREAKABLE a ((test)) AND( ANDERSON (test:2.0):1.5) :o BREAK (red:1.5)", + "[:hands, feet, :0.15]normal quality", + ), ) def test_cl_removenetworktags(self): # remove network tags - self.process( - "this is a test", - "", - "this is a test", - "", - PromptPostProcessor( - self, - self.__defprocessing, - self.__defstate, - DictToObj({**self.__defopts.__dict__, "ppp_rem_removeextranetworktags": True}), + self.__process( + PromptPair("this is a test", ""), + PromptPair("this is a test", ""), + ppp=PromptPostProcessor( + self.__ppp_logger, + self.__interrupt, + self.__def_model_info, + {**self.__defopts, "remove_extranetwork_tags": True}, ), ) def test_cl_dontremoveseparatorsoneol(self): # dont remove separators on eol - self.process( - "this is a test,\nsecond line", - "", - "this is a test,\nsecond line", - "", - PromptPostProcessor( - self, - self.__defprocessing, - self.__defstate, - DictToObj({**self.__defopts.__dict__, "ppp_cup_extraseparators2": False}), + self.__process( + PromptPair("this is a test,\nsecond line", ""), + PromptPair("this is a test,\nsecond line", ""), + ppp=PromptPostProcessor( + self.__ppp_logger, + self.__interrupt, + self.__def_model_info, + {**self.__defopts, "cleanup_extra_separators2": False}, ), ) # Command tests def test_cmd_stn_complex_features(self): # complex stn command with AND, BREAK and other features - self.process( - "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loraword AND AND hypernetword :0.3", - "normal quality, ", - "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK with [abc:def:5]:0.5 AND loraword AND hypernetword :0.3", - "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", + self.__process( + PromptPair( + "[neg5] this \\(is\\): a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5]:0.5 AND loratrigger AND AND hypernettrigger :0.3", + "normal quality, ", + ), + PromptPair( + "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5)\nBREAK with [abc:def:5]:0.5 AND loratrigger AND hypernettrigger :0.3", + "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", + ), ) def test_cmd_if_complex_features(self): # complex if command - self.process( - "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:def:5]:0.5 AND loraword hypernetword nothing:0.3", - "normal quality", - "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK :0.5 AND hypernetword :0.3", - "normal quality", + self.__process( + PromptPair( + "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:def:5]:0.5 AND loratrigger hypernettrigger nothing:0.3", + "normal quality", + ), + PromptPair( + "this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5)\nBREAK :0.5 AND hypernettrigger :0.3", + "normal quality", + ), ) def test_cmd_if_nested(self): # nested if command - self.process( - "this is SD1SDXLSD2", - "", - "this is SDXL", - "", + self.__process( + PromptPair( + "this is SD1PONYSD2", "" + ), + PromptPair("this is PONY", ""), + ppp=PromptPostProcessor( + self.__ppp_logger, + self.__interrupt, + { + **self.__def_model_info, + "model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", + }, + self.__defopts, + ), ) def test_cmd_set_if(self): # set and if commands - self.process( - "valuethis test is OKnot OK", - "", - "this test is OK", - "", + self.__process( + PromptPair("valuethis test is OKnot OK", ""), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_eval_if(self): # set and if commands + self.__process( + PromptPair("valuethis test is OKnot OK", ""), + PromptPair("this test is OK", ""), ) def test_cmd_set_if_echo_nested(self): # nested set, if and echo commands - self.process( - "1OKnot OK", - "", - "OK", - "", + self.__process( + PromptPair( + "1OKnot OK NOK OK", + "", + ), + PromptPair("OK OK OK", ""), ) + 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", + "", + ), + PromptPair("First: this test is OK\nSecond: this test is OK", ""), + ) + + def test_cmd_set_add_if(self): # set, add and if commands + self.__process( + PromptPair( + "value2this test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_add_DP_if(self): # set, add (DP format) and if commands + self.__process( + PromptPair( + "${v=value}${v+=2}this test is OKnot OK", + "", + ), + PromptPair("this test is OK", ""), + ) + + def test_cmd_set_immediateeval(self): # set (DP format) with mixed evaluation + self.__process( + PromptPair( + "${var=!__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var2:default}, ${var3=__yaml/wildcard1__}${var3}, ${var3}", + "", + ), + PromptPair("the choices are: choice2, choice2, default, choice3, choice1", ""), + ppp=self.__nocupppp, + ) + + def test_cmd_set_mixeval(self): # set and add (DP format) with mixed evaluation + self.__process( + PromptPair( + "${var=__yaml/wildcard1__}the choices are: ${var}, ${var}, ${var+=, __yaml/wildcard2__}${var}, ${var}, ${var+=!, __yaml/wildcard3__}${var}, ${var}", + "", + ), + PromptPair( + "the choices are: choice2, choice3, choice1, choice1- choice2 -choice3, choice2, choice2 -choice1-choice3, choice2, choice3-choice1- choice2 , choice1, choice2 , choice2, choice3-choice1- choice2 , choice1, choice2 ", + "", + ), + ppp=self.__nocupppp, + ) + + # Choices tests + + def test_ch_choices(self): # simple choices with weights + self.__process( + PromptPair("the choices are: {3::choice1|2::choice2|choice3}", ""), + PromptPair("the choices are: choice2", ""), + ppp=self.__nocupppp, + ) + + def test_ch_unsupportedsampler(self): # unsupported sampler + self.__process( + PromptPair("the choices are: {@choice1|choice2|choice3}", ""), + PromptPair("", ""), + ppp=self.__nocupppp, + interrupted=True, + ) + + def test_ch_choices_withcomments(self): # choices with comments and multiline + self.__process( + PromptPair( + "the choices are: {\n3::choice1 # this is option 1\n|2::choice2\n# this was option 2\n|choice3 # this is option 3\n}", + "", + ), + PromptPair("the choices are: choice2", ""), + ppp=self.__nocupppp, + ) + + def test_ch_choices_multiple(self): # choices with multiple selection + self.__process( + PromptPair("the choices are: {~2$$, $$3::choice1|2:: choice2 |choice3}", ""), + PromptPair("the choices are: choice2 , choice3", ""), + ppp=self.__nocupppp, + ) + + def test_ch_choices_if_multiple(self): # choices with if and multiple selection + self.__process( + PromptPair("the choices are: {2$$, $$3::choice1|2 if _is_sd1::choice2|choice3}", ""), + PromptPair("the choices are: choice1, choice3", ""), + ppp=self.__nocupppp, + ) + + def test_ch_choices_set_if_multiple(self): # choices with if user variable and multiple selection + self.__process( + PromptPair("${var=test}the choices are: {2$$, $$3::choice1|2 if not var eq 'test'::choice2|choice3}", ""), + PromptPair("the choices are: choice1, choice3", ""), + ppp=self.__nocupppp, + ) + + def test_ch_choices_set_if_nested(self): # nested choices with if user variable and multiple selection + self.__process( + PromptPair( + "${var=test}the choices are: {2$$, $$3::choice1${var2=test2} {if var2 eq 'test2'::choice11|choice12}|2 if not var eq 'test'::choice2|choice3}", + "", + ), + PromptPair("the choices are: choice1 choice11, choice3", ""), + ppp=self.__nocupppp, + ) + + # Wildcards tests + + def test_wc_ignore(self): # wildcards with ignore option + self.__process( + PromptPair("__bad_wildcard__", "{option1|option2}"), + PromptPair("__bad_wildcard__", "{option1|option2}"), + ppp=PromptPostProcessor( + self.__ppp_logger, + self.__interrupt, + self.__def_model_info, + { + **self.__defopts, + "process_wildcards": False, + "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value, + }, + ), + ) + + def test_wc_remove(self): # wildcards with remove option + self.__process( + PromptPair( + "[neg5] this is: __bad_wildcard__ a (([complex|simpleneg6|regular] testneg1)(test:2.0):1.5) \nBREAK, BREAK with [abcneg4:defneg2(neg3:1.6):5] ", + "normal quality, {option1|option2}", + ), + PromptPair( + "this is: a (([complex|simple|regular] test)(test:2.0):1.5)\nBREAK with [abc:def:5]", + "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", + ), + ppp=PromptPostProcessor( + self.__ppp_logger, + self.__interrupt, + self.__def_model_info, + { + **self.__defopts, + "process_wildcards": False, + "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.remove.value, + }, + ), + ) + + def test_wc_warn(self): # wildcards with warn option + self.__process( + PromptPair("__bad_wildcard__", "{option1|option2}"), + PromptPair(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", "{option1|option2}"), + ppp=PromptPostProcessor( + self.__ppp_logger, + self.__interrupt, + self.__def_model_info, + { + **self.__defopts, + "process_wildcards": False, + "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.warn.value, + }, + ), + ) + + def test_wc_stop(self): # wildcards with stop option + self.__process( + PromptPair("__bad_wildcard__", "{option1|option2}"), + PromptPair( + PromptPostProcessor.WILDCARD_STOP + "__bad_wildcard__", + PromptPostProcessor.WILDCARD_STOP + "{option1|option2}", + ), + ppp=PromptPostProcessor( + self.__ppp_logger, + self.__interrupt, + self.__def_model_info, + { + **self.__defopts, + "process_wildcards": False, + "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.stop.value, + }, + ), + interrupted=True, + ) + + def test_wc_wildcard1a_text(self): # simple text wildcard + self.__process( + PromptPair("the choices are: __text/wildcard1__", ""), + PromptPair("the choices are: choice2", ""), + ppp=self.__nocupppp, + ) + + def test_wc_wildcard1a_json(self): # simple json wildcard + self.__process( + PromptPair("the choices are: __json/wildcard1__", ""), + PromptPair("the choices are: choice2", ""), + ppp=self.__nocupppp, + ) + + def test_wc_wildcard1a_yaml(self): # simple yaml wildcard + self.__process( + PromptPair("the choices are: __yaml/wildcard1__", ""), + PromptPair("the choices are: choice2", ""), + ppp=self.__nocupppp, + ) + + def test_wc_wildcard1b_text(self): # simple text wildcard with multiple choices + self.__process( + PromptPair("the choices are: __2-$$text/wildcard1__", ""), + PromptPair("the choices are: choice3, choice1", ""), + ppp=self.__nocupppp, + ) + + def test_wc_wildcard1b_json(self): # simple json wildcard with multiple choices + self.__process( + PromptPair("the choices are: __2-$$json/wildcard1__", ""), + PromptPair("the choices are: choice3, choice1", ""), + ppp=self.__nocupppp, + ) + + def test_wc_wildcard1b_yaml(self): # simple yaml wildcard with multiple choices + self.__process( + PromptPair("the choices are: __2-$$yaml/wildcard1__", ""), + PromptPair("the choices are: choice3, choice1", ""), + ppp=self.__nocupppp, + ) + + def test_wc_wildcard2_text(self): # simple text wildcard with default options + self.__process( + PromptPair("the choices are: __text/wildcard2__", ""), + PromptPair("the choices are: choice3-choice1", ""), + ppp=self.__nocupppp, + ) + + def test_wc_wildcard2_json(self): # simple json wildcard with default options + self.__process( + PromptPair("the choices are: __json/wildcard2__", ""), + PromptPair("the choices are: choice3-choice1", ""), + ppp=self.__nocupppp, + ) + + def test_wc_wildcard2_yaml(self): # simple yaml wildcard with default options + self.__process( + PromptPair("the choices are: __yaml/wildcard2__", ""), + PromptPair("the choices are: choice3-choice1", ""), + ppp=self.__nocupppp, + ) + + def test_wc_nested_wildcard_text(self): # nested text wildcard with repeating multiple choices + self.__process( + PromptPair("the choices are: __r3$$-$$text/wildcard3__", ""), + PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""), + ppp=self.__nocupppp, + ) + + def test_wc_nested_wildcard_json(self): # nested json wildcard with repeating multiple choices + self.__process( + PromptPair("the choices are: __r3$$-$$json/wildcard3__", ""), + PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""), + ppp=self.__nocupppp, + ) + + def test_wc_nested_wildcard_yaml(self): # nested yaml wildcard with repeating multiple choices + self.__process( + PromptPair("the choices are: __r3$$-$$yaml/wildcard3__", ""), + PromptPair("the choices are: choice3,choice1- choice2 ,choice3", ""), + ppp=self.__nocupppp, + ) + + def test_wc_wildcard4_yaml(self): # simple yaml wildcard with one option + self.__process( + PromptPair("the choices are: __yaml/wildcard4__", ""), + PromptPair("the choices are: inline text", ""), + ppp=self.__nocupppp, + ) + + def test_wc_wildcard6_yaml(self): # simple yaml wildcard with object formatted choices + self.__process( + PromptPair("the choices are: __yaml/wildcard6__", ""), + PromptPair("the choices are: choice2", ""), + ppp=self.__nocupppp, + ) + + def test_wc_choice_wildcard_mix(self): # choices with wildcard mix + self.__process( + PromptPair("the choices are: {__~2$$yaml/wildcard2__|choice0}", ""), + [ + PromptPair("the choices are: choice0", ""), + PromptPair("the choices are: choice1, choice3", ""), + PromptPair("the choices are: choice1, choice3", ""), + ], + ppp=self.__nocupppp, + ) + + def test_wc_unsupportedsampler(self): # unsupported sampler + self.__process( + PromptPair("the choices are: __@yaml/wildcard2__", ""), + PromptPair("", ""), + ppp=self.__nocupppp, + interrupted=True, + ) + + def test_wc_wildcard_globbing(self): # wildcard with globbing + self.__process( + PromptPair("the choices are: __yaml/wildcard[12]__, __yaml/wildcard*__", ""), + PromptPair("the choices are: choice3-choice2, choice3-choice1- choice2 ", ""), + ppp=self.__nocupppp, + ) + + def test_wc_wildcardwithvar(self): # wildcard with inline variable + self.__process( + PromptPair("the choices are: __yaml/wildcard5(var=test)__, __yaml/wildcard5__", ""), + PromptPair("the choices are: inline test, inline default", ""), + ppp=self.__nocupppp, + ) + + # def test_mix(self): + # self.__process( + # PromptPair( + # "__text/wildcard1__ (__text/wildcard2__) (__text/wildcard3__:1.5) [__text/wildcard1__] [__text/wildcard2__:__text/wildcard3__:0.5] [__text/wildcard1__|__text/wildcard2__] # {opt1_1|opt1_2} ({opt2_1|opt2_2}) ({opt3_1|opt3_2}:1.5) [{opt4_1|opt4_2}] [{opt5_1|opt5_2}:{opt6_1|opt6_2}:0.5] [{opt7_1|opt7_2}|{opt8_1|opt8_2}] # {opt1_1|__text/wildcard1__} ({opt2_1|__text/wildcard2__}) ({opt3_1|__text/wildcard3__}:1.5) [{opt4_1|__text/wildcard1__}] [{opt5_1|__text/wildcard2__}# :{opt6_1|__text/wildcard3__}:0.5] [{opt7_1|__text/wildcard1__}|{opt8_1|__text/wildcard2__}] {|}", + # "", + # ), + # PromptPair( + # "choice2 ( choice2 -choice1) (choice1, choice2 :1.5) [choice1] [choice1-choice1:choice3, choice2 :0.5] [choice2| choice2 - choice2 ] opt1_1 # (opt2_2) (opt3_1:1.5) [opt4_2] [opt5_1:opt6_2:0.5] [opt7_1|opt8_1] choice3 (opt2_1) (choice1,choice3:1.5) [choice1] [choice3-choice3:opt6_1:0.5] [opt7_1|# opt8_1] ", + # "", + # ), + # ppp=self.__nocupppp, + # ) + + # def test_real(self): + # self.__wildcards_obj.refresh_wildcards( + # DEBUG_LEVEL.full, + # ["D:\\AI\\SD\\_configuraciones\\acb-wildcards\\wildcards"], + # ) + # self.__process( + # PromptPair( + # "${separator=()}, __quality/high__ __misc/sep__, photograph of a __character__", + # "__negatives/ng_generic__", + # ), + # PromptPair("", ""), + # ) + if __name__ == "__main__": unittest.main() diff --git a/tests/wildcards/test.json b/tests/wildcards/test.json new file mode 100644 index 0000000..111176c --- /dev/null +++ b/tests/wildcards/test.json @@ -0,0 +1,19 @@ +{ + "json": { + "wildcard1": [ + "choice1", + "choice2", + "choice3" + ], + "wildcard2": [ + "r2-3$$-", + "4::choice1", + "3:: choice2 ", + "2::choice3", + "5 if _is_sd1::choice4" + ], + "wildcard3": [ + "__2$$,$$json/wildcard2__" + ] + } +} \ No newline at end of file diff --git a/tests/wildcards/test.yaml b/tests/wildcards/test.yaml new file mode 100644 index 0000000..1e390ec --- /dev/null +++ b/tests/wildcards/test.yaml @@ -0,0 +1,25 @@ +yaml: + wildcard1: + - choice1 + - choice2 + - choice3 + + wildcard2: + - ~r2-3$$- + - 4::choice1 + - "3:: choice2 " + - 2::choice3 + - 5 if _is_sd1::choice4 + + wildcard3: + - __2$$,$$yaml/wildcard2__ + + wildcard4: inline text + + wildcard5: inline ${var:default} + + wildcard6: + - { weight: 2, text: choice1 } + - { weight: 3, content: choice2 } + - { text: choice3 } + - { weight: 4, if: "_is_ssd", text: choice4 } diff --git a/tests/wildcards2/text/wildcard1.txt b/tests/wildcards2/text/wildcard1.txt new file mode 100644 index 0000000..4df3c75 --- /dev/null +++ b/tests/wildcards2/text/wildcard1.txt @@ -0,0 +1,4 @@ +# wildcard1 +choice1 +choice2 +choice3 \ No newline at end of file diff --git a/tests/wildcards2/text/wildcard2.txt b/tests/wildcards2/text/wildcard2.txt new file mode 100644 index 0000000..fb9c0e7 --- /dev/null +++ b/tests/wildcards2/text/wildcard2.txt @@ -0,0 +1,6 @@ +# wildcard2 +r2-3$$- +4::choice1 +3:: choice2 +2::choice3 +5 if _is_sd1::choice4 \ No newline at end of file diff --git a/tests/wildcards2/text/wildcard3.txt b/tests/wildcards2/text/wildcard3.txt new file mode 100644 index 0000000..61b6647 --- /dev/null +++ b/tests/wildcards2/text/wildcard3.txt @@ -0,0 +1,2 @@ +# wildcard3 +__2$$,$$text/wildcard2__ \ No newline at end of file