Compare commits

...
1 Commits
Author SHA1 Message Date
Antonio Cordero Balcazar 3fc61c8dd1 * Improved documentation.
* New options to choose whether to process in img2img.
* Option to detect and do something with unwanted wildcards.
* Cleanup processing rewritten and separated in multiple options.
2023-12-10 12:16:55 +01:00
5 changed files with 1021 additions and 250 deletions
+51 -38
View File
@@ -1,39 +1,33 @@
# Prompt Postprocessor for Stable Diffusion WebUI
(formerly known as "sd-webui-sendtonegative")
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).
Extension for [AUTOMATIC1111 Stable Diffusion WebUI](https://github.com/AUTOMATIC1111/stable-diffusion-webui). Compatible with [SD.Next](https://github.com/vladmandic/automatic).
Currently this extension has these functions:
## Purpose
* Allows the tagging of 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.
* Detect invalid wildcards and act on them.
* Clean up the prompt and negative prompt.
The purpose of this extension is to process the prompt after other extensions have possibly modified it.
Note: The extension must be loaded after the installed wildcards extension (or any other that modifies the prompt). Extensions load by their folder name in alphanumeric order.
Currently this extension allows the tagging of parts of the prompt and moves them to the
negative prompt. This allows useful tricks when using a wildcard extension
since you can add negative content from choices made in the positive prompt.
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.
Note: The extension must be loaded after the installed wildcards extension. 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.
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: it only recognizes regular A1111 prompt formats. So:
Notes:
* Attention: \[prompt\] (prompt) (prompt:weight)
* Alternation: \[prompt1|prompt2|...\]
* Scheduling: \[prompt1:prompt2:step\]
1. It only recognizes regular A1111 prompt formats. So:
In SD.Next that means only the A1111 or Full parsers.
* **Attention**: \[prompt\] (prompt) (prompt:weight)
* **Alternation**: \[prompt1|prompt2|...\]
* **Scheduling**: \[prompt1:prompt2:step\]
* **Models**: \<model\>
It does not build equivalent AND/BREAK separations into the negative prompt.
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. It does not translate equivalent *AND/BREAK* separations into the negative prompt.
## Installation
@@ -44,6 +38,12 @@ It does not build equivalent AND/BREAK separations into the negative prompt.
## 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.
### Sending content to the negative prompt
The format of the tags is like this:
@@ -60,17 +60,15 @@ And an optional position in the negative prompt can be specified like this:
Where position can be:
* s: at the start (the default)
* e: at the end
* pN: at the position of the insertion point "<!!iN!!>" with N being 0-9
* **s**: at the start (the default)
* **e**: at the end
* **pN**: at the position of the insertion point "**<!!iN!!>**" with N being 0-9
If the insertion point is not found it inserts at the start.
The insertion point of course must be in the negative prompt. If the insertion point is not found it inserts at the start.
#### Example
You have a wildcard for hair colors (\_\_haircolors\_\_) with one being
strawberry blonde, but you don't want strawberries. So in that option you add a
tag to add to the negative prompt, like so:
You have a wildcard for hair colors (\_\_haircolors\_\_) with one being strawberry blonde, but you don't want strawberries. So in that option you add a tag to add to the negative prompt, like so:
```text
blonde
@@ -78,18 +76,33 @@ strawberry blonde <!strawberry!>
brunette
```
Then, if that option is chosen this extension will process it later and move
that part to the negative prompt.
Then, if that option is chosen this extension will process it later and move that part to the negative prompt.
## Configuration
Separator used when adding to the negative prompt: You can specify the separator used when adding to the negative prompt (by default it's ", ").
### General settings
Ignore tags with repeated content: by default it ignores repeated content to avoid repetitions in the negative prompt.
* **Debug**: writes debugging information to the console.
* **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.
Join attention modifiers (weights) when possible: by default it joins attention modifiers when possible (joins into one, multipliying their values).
### Send to negative prompt settings
Try to clean-up the prompt after processing: by default cleans up the positive prompt after processing, removing extra spaces and separators.
* **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 tags with 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 unnecesary separators. This applies to the configured separator and regular commas.
* **Clean up around BREAKs**: removes consecutive BREAKs and unnecesary commas and space around them.
* **Remove extra spaces**: removes unnecesary spaces.
## Notes
@@ -107,8 +120,8 @@ They will be translated to the negative prompt. For example:
Negative tags inside such constructs will copy the construct to the negative prompt, but separating its elements. For example:
* Alternation: `[red<!square!>|blue<!circle!>]` will end up as `[square|], [|circle]` in the negative prompt, instead of `[square|circle]`
* Scheduling: `[red<!square!>:blue<!circle!>:0.5]` will end up as `[square::0.5], [:circle:0.5]` instead of `[square:circle:0.5]`
* **Alternation**: `[red<!square!>|blue<!circle!>]` will end up as `[square|], [|circle]` in the negative prompt, instead of `[square|circle]`
* **Scheduling**: `[red<!square!>:blue<!circle!>:0.5]` will end up as `[square::0.5], [:circle:0.5]` instead of `[square:circle:0.5]`
This should still work as intended, and the only negative point i see is the unnecessary separators.
+595 -159
View File
@@ -4,68 +4,105 @@ import math
import lark
class PromptPostProcessor: # pylint: disable=too-few-public-methods
NAME = "Prompt Post-Processor"
VERSION = "2.1.5"
class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-instance-attributes
"""
The PromptPostProcessor class is responsible for processing and manipulating prompt strings.
DEFAULT_SEPARATOR = ", "
The format for the negative tags is:
<!content!>
<!!x!content!>
with x being:
s - content is added at the start of the negative prompt. This is the default if no parameter exists.
e - content is added at the end of the negative prompt.
pN - content is added where the insertion point N is in the negative prompt or at the start if it does not exist. N can be 0 to 9.
iN - tags the position of insertion point N. Used only in the negative prompt and does not accept content. N can be 0 to 9.
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.2.0"
DEFAULT_STN_SEPARATOR = ", "
IFWILDCARDS_CHOICES = {
"ignore": "Ignore",
"remove": "Remove",
"warn": "Add visible warning",
"stop": "Stop the generation",
}
WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5) BREAK\n'
WILDCARD_STOP = "INVALID WILDCARD! BREAK\n"
def __init__(
self,
log,
separator=None,
ignore_repeats=None,
join_attention=None,
cleanup=None,
script,
opts=None,
is_i2i=False,
):
"""
Format for the tag:
<!content!>
Initializes the PPP object.
<!!x!content!>
with x being:
s - content is added at the start of the negative prompt. This is the default if no parameter exists.
e - content is added at the end of the negative prompt.
pN - content is added where the insertion point N is in the negative prompt or at the start if it does not exist. N can be 0 to 9.
iN - tags the position of insertion point N. Used only in the negative prompt and does not accept content. N can be 0 to 9.
Args:
script: The script object.
opts: Optional. The options object for configuring PPP behavior.
"""
self.__opts = opts
self.__logger = log
self.__debug = getattr(self.__opts, "ppp_debug", False)
self.script = script
self.logger = script.ppp_logger
self.opts = opts
self.is_i2i = is_i2i
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.__ignore_repeats = (
ignore_repeats if ignore_repeats is not None else getattr(opts, "ppp_stn_ignorerepeats", True)
)
self.__join_attention = (
join_attention
if join_attention is not None
else getattr(opts, "ppp_stn_joinattention", True)
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 True
else self.IFWILDCARDS_CHOICES["ignore"]
)
self.__cleanup = (
cleanup if cleanup is not None else getattr(opts, "ppp_cleanup", True) if opts is not None else True
)
self.__separator = (
separator
if separator is not None
else getattr(opts, "ppp_separator", self.DEFAULT_SEPARATOR)
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_SEPARATOR
else self.DEFAULT_STN_SEPARATOR
)
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_breaks = getattr(opts, "ppp_cup_breaks", True) if opts is not None else True
self.__insertion_point_tags = [f"<!!i{x}!!>" for x in range(10)]
# Process with lark (debug with https://www.lark-parser.org/ide/)
self.__schedule_parser = lark.Lark(
self.__parser_complete = lark.Lark(
r"""
start: (prompt | /[\][():|<>!]/+)*
?prompt: (emphasized | deemphasized | scheduled | alternate | modeltag | negtag | plain)*
?nonegprompt: (emphasized | deemphasized | scheduled | alternate | modeltag | plain)*
start: (prompt | /[\][():|<>!{}]/+)*
prompt: (emphasized | deemphasized | scheduled | alternate | modeltag | negtag | wildcard | choices | plain)*
nonegprompt: (emphasized | deemphasized | scheduled | alternate | modeltag | wildcard | choices | plain)*
wildcard: "__" /(?:(?!__)\S)+/ "__"
choices: "{" choice ("|" choice)* "}"
choice: prompt # we ignore weight and any other parameters
emphasized: "(" prompt [":" numpar] ")"
deemphasized: "[" prompt "]"
scheduled: "[" [prompt ":"] prompt ":" numpar "]"
@@ -76,19 +113,50 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods
modeltag: "<" /(?!!)[^>]+/ ">"
numpar: WHITESPACE* NUMBER WHITESPACE*
WHITESPACE: /\s+/
?plain: /([^\\[\]():|<>!]|\\.)+/s
plain: /((?!__)[^\\[\]():|<>!{}]|\\.)+/s
%import common.SIGNED_NUMBER -> NUMBER
""",
""", # prompt, nonegprompt, plain with ?
propagate_positions=True,
)
class ReadTree(lark.visitors.Interpreter):
def __init__(self, logger, debug, ignorerepeats, joinattention, prompt, add_at):
def formatOutput(self, text: str):
"""
Formats the output text by encoding it using unicode_escape and decoding it using utf-8.
Args:
text (str): The input text to be formatted.
Returns:
str: The formatted output text.
"""
return text.encode("unicode_escape").decode("utf-8")
class STNTree(lark.visitors.Interpreter):
"""
A class for interpreting and processing a tree generated by the prompt parser.
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.
start(tree): Process the given tree and perform necessary operations on the found negative tags.
"""
def __init__(self, ppp, prompt, add_at):
super().__init__()
self.__logger = logger
self.__debug = debug
self.__ignore_repeats = ignorerepeats
self.__join_attention = joinattention
self.__ppp = ppp
self.__prompt = prompt
self.AccumulatedShell = namedtuple("AccumulatedShell", ["type", "info1", "info2"])
AccumulatedShell = self.AccumulatedShell
@@ -101,9 +169,27 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods
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, tree):
"""
Process a scheduling construct in the tree and add it to the accumulated shell.
Args:
tree (Node): The tree node representing the scheduling construct.
Returns:
None
"""
if len(tree.children) > 2: # before & after
before = tree.children[0]
else:
@@ -115,30 +201,46 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods
pos = int(pos)
# self.__shell.append(self.AccumulatedShell("sc", tree.meta.start_pos, pos))
if before is not None and hasattr(before, "data"):
if self.__debug:
self.__logger.info(
f"Shell scheduled before at {[before.meta.start_pos,before.meta.end_pos] if hasattr(before,'meta') and not before.meta.empty else '?'} with position {pos}"
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, None))
self.visit(before)
self.__shell.pop()
if hasattr(after, "data"):
if self.__debug:
self.__logger.info(
f"Shell scheduled after at {[after.meta.start_pos,after.meta.end_pos] if hasattr(after,'meta') and not after.meta.empty else '?'} with position {pos}"
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, None))
self.visit(after)
self.__shell.pop()
# self.__shell.pop()
def alternate(self, tree):
"""
Process an alternation construct in the tree and add it to the accumulated shell.
Args:
tree (Node): The tree node representing the alternation construct.
Returns:
None
"""
# self.__shell.append(self.AccumulatedShell("al", tree.meta.start_pos, len(tree.children)))
for i, opt in enumerate(tree.children):
if self.__debug:
self.__logger.info(
f"Shell alternate at {[opt.meta.start_pos,opt.meta.end_pos] if hasattr(opt,'meta') and not opt.meta.empty else '?'} option {i+1}"
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", i + 1, len(tree.children)))
self.visit(opt)
@@ -146,27 +248,56 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods
# self.__shell.pop()
def emphasized(self, tree):
"""
Process a attention change construct in the tree and add it to the accumulated shell.
Args:
tree (Node): The tree node representing the attention construct.
Returns:
None
"""
numpar = tree.children[-1]
weight = self.__get_numpar_value(numpar) if numpar is not None else 1.1
if self.__debug:
self.__logger.info(
f"Shell attention at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') and not tree.meta.empty else '?'} with weight {weight}"
if self.__ppp.debug:
metaposition = (
[tree.meta.start_pos, tree.meta.end_pos] if hasattr(tree, "meta") and not tree.meta.empty else "?"
)
self.__ppp.logger.info(f"Shell attention at {metaposition} with weight {weight}")
self.__shell.append(self.AccumulatedShell("at", weight, None))
self.visit_children(tree)
self.__shell.pop()
def deemphasized(self, tree):
"""
Process a decrease attention construct in the tree and add it to the accumulated shell.
Args:
tree (Node): The tree node representing the decreased attention construct.
Returns:
None
"""
weight = 0.9
if self.__debug:
self.__logger.info(
f"Shell attention at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') and not tree.meta.empty else '?'} with weight {weight}"
if self.__ppp.debug:
metaposition = (
[tree.meta.start_pos, tree.meta.end_pos] if hasattr(tree, "meta") and not tree.meta.empty else "?"
)
self.__ppp.logger.info(f"Shell attention at {metaposition} with weight {weight}")
self.__shell.append(self.AccumulatedShell("at", weight, None))
self.visit_children(tree)
self.__shell.pop()
def negtag(self, tree):
"""
Process a negative tag in the tree and add it to the list of negative tags.
Args:
tree (Node): The tree node representing the negative tag.
Returns:
None
"""
negtagparameters = tree.children[0]
parameters = negtagparameters.children[0].value if negtagparameters is not None else ""
rest = []
@@ -180,29 +311,41 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods
self.__negtags.append(
self.NegTag(tree.meta.start_pos, tree.meta.end_pos, content, parameters, self.__shell.copy())
)
if self.__debug:
self.__logger.info(
f"Negative tag at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') and not tree.meta.empty else '?'}: {parameters or 'with no parameters :'} {content.encode('unicode_escape').decode('utf-8')}"
if self.__ppp.debug:
metaposition = (
[tree.meta.start_pos, tree.meta.end_pos] if hasattr(tree, "meta") and not tree.meta.empty else "?"
)
self.__ppp.logger.info(
f"Negative tag at {metaposition}: {parameters or 'with no parameters :'} {self.__ppp.formatOutput(content)}"
)
def start(self, tree):
"""
Process the given tree and perform necessary operations on the found negative tags.
Args:
tree: The tree to be processed.
Returns:
None
"""
self.visit_children(tree)
# process the found negtags
for nt in self.__negtags:
if self.__join_attention:
# 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(nt.shell) - 1, 0, -1):
if nt.shell[i].type == "at" and nt.shell[i - 1].type == "at":
nt.shell[i - 1] = self.AccumulatedShell(
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 * nt.shell[i - 1].info1 * nt.shell[i].info1)
/ 100, # we limit to two decimals
math.floor(100 * negtag.shell[i - 1].info1 * negtag.shell[i].info1)
/ 100, # we limit the new weight to two decimals
None,
)
nt.shell.pop(i)
negtag.shell.pop(i)
start = ""
end = ""
for s in nt.shell:
for s in negtag.shell:
match s.type:
case "at":
if s.info1 == 0.9:
@@ -225,15 +368,15 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods
case "alo":
start += "[" + ("|" * int(s.info1 - 1))
end = ("|" * int(s.info2 - s.info1)) + "]" + end
content = start + nt.content + end
position = nt.parameters or "s"
content = start + negtag.content + end
position = negtag.parameters or "s"
if len(content) > 0:
if content not in self.__already_processed:
if self.__ignore_repeats:
if self.__ppp.stn_ignore_repeats:
self.__already_processed.append(content)
if self.__debug:
self.__logger.info(
f"Adding content at position {position}: {content.encode('unicode_escape').decode('utf-8')}"
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)
@@ -243,111 +386,404 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods
else: # position == "s" or invalid
self.add_at["start"].append(content)
else:
self.__logger.warning(
f"Ignoring repeated content: {content.encode('unicode_escape').decode('utf-8')}"
)
self.__ppp.logger.warning(f"Ignoring repeated content: {self.__ppp.formatOutput(content)}")
# remove from prompt
self.remove.append([nt.start, nt.end])
def process_prompt(self, original_prompt, original_negative_prompt):
"""
Extract from the prompt the tagged parts and add them to the negative prompt
"""
try:
prompt = original_prompt
negative_prompt = original_negative_prompt
self.__debug = getattr(self.__opts, "ppp_debug", False)
if self.__debug:
self.__logger.info(f"Input prompt: {prompt.encode('unicode_escape').decode('utf-8')}")
self.__logger.info(f"Input negative_prompt: {negative_prompt.encode('unicode_escape').decode('utf-8')}")
prompt, add_at = self.__find_tags(prompt)
negative_prompt = self.__add_to_insertion_points(negative_prompt, add_at["insertion_point"])
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"Output prompt: {prompt.encode('unicode_escape').decode('utf-8')}")
self.__logger.info(
f"Output negative_prompt: {negative_prompt.encode('unicode_escape').decode('utf-8')}"
)
return prompt, negative_prompt
except Exception as e: # pylint: disable=broad-exception-caught
self.__logger.exception(e)
return original_prompt, original_negative_prompt
self.remove.append([negtag.start, negtag.end])
def __find_tags(self, prompt):
add_at = {"start": [], "insertion_point": [[] for x in range(10)], "end": []}
tree = self.__schedule_parser.parse(prompt)
# if self.__debug:
# self.__logger.info(f"Initial tree:\n{tree.pretty()}")
"""
Finds tags in the given prompt and returns the modified prompt and tag information.
readtree = self.ReadTree(
self.__logger, self.__debug, self.__ignore_repeats, self.__join_attention, prompt, add_at
)
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": []}
tree = self.__parser_complete.parse(prompt)
# self.logger.info(f"tree from prompt:\n{tree.pretty()}")
readtree = self.STNTree(self, prompt, add_at)
readtree.visit(tree)
for r in readtree.remove[::-1]:
prompt = prompt[: r[0]] + prompt[r[1] :]
if self.__cleanup:
if self.__debug:
self.__logger.info(f"Prompt before cleanup: {prompt.encode('unicode_escape').decode('utf-8')}")
prompt = re.sub(r"\((?::[+-]?[\d\.]+)?\)", "", prompt) # clean up empty attention
prompt = re.sub(r"\[\]", "", prompt) # clean up empty attention
prompt = re.sub(r"\[:?:[+-]?[\d\.]+\]", "", prompt) # clean up empty scheduling
prompt = re.sub(r"\[\|+\]", "", prompt) # clean up empty alternation
prompt = re.sub(r"[ ]{2,}", " ", prompt) # collapse spaces
# clean up extra separators
prompt = (
prompt.replace(self.__separator + self.__separator, self.__separator)
.replace(" " + self.__separator, self.__separator)
.removeprefix(self.__separator)
.removesuffix(self.__separator)
.strip()
)
add_at = readtree.add_at
if self.__debug:
self.__logger.info(f"New negative additions: {add_at}")
add_at = readtree.add_at
if self.debug:
self.logger.info(f"New negative additions: {add_at}")
return prompt, add_at
def __add_to_insertion_points(self, negative_prompt, add_at_insertion_point):
"""
Adds the negative prompt to the insertion points.
Args:
negative_prompt (str): The negative prompt to be added.
add_at_insertion_point (list): A list of insertion points.
Returns:
str: The modified negative prompt.
"""
for n in range(10):
ipp = negative_prompt.find(self.__insertion_point_tags[n])
if ipp >= 0:
ipl = len(self.__insertion_point_tags[n])
if negative_prompt[ipp - len(self.__separator) : ipp] == self.__separator:
ipp -= len(self.__separator) # adjust for existing start separator
ipl += len(self.__separator)
if negative_prompt[ipp - len(self.stn_separator) : ipp] == self.stn_separator:
ipp -= len(self.stn_separator) # adjust for existing start separator
ipl += len(self.stn_separator)
add_at_insertion_point[n].insert(0, negative_prompt[:ipp])
if negative_prompt[ipp + ipl : ipp + ipl + len(self.__separator)] == self.__separator:
ipl += len(self.__separator) # adjust for existing end separator
if negative_prompt[ipp + ipl : ipp + ipl + len(self.stn_separator)] == self.stn_separator:
ipl += len(self.stn_separator) # adjust for existing end separator
endPart = negative_prompt[ipp + ipl :]
if len(endPart) > 0:
add_at_insertion_point[n].append(endPart)
negative_prompt = self.__separator.join(add_at_insertion_point[n])
negative_prompt = self.stn_separator.join(add_at_insertion_point[n])
else:
ipp = 0
if negative_prompt.startswith(self.__separator):
ipp = len(self.__separator)
if negative_prompt.startswith(self.stn_separator):
ipp = len(self.stn_separator)
add_at_insertion_point[n].append(negative_prompt[ipp:])
negative_prompt = self.__separator.join(add_at_insertion_point[n])
negative_prompt = self.stn_separator.join(add_at_insertion_point[n])
return negative_prompt
def __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.
Args:
negative_prompt (str): The original negative prompt string.
add_at_start (list): The list of elements to be added at the start of the negative prompt.
Returns:
str: The updated negative prompt string with the elements added at the start.
"""
if len(negative_prompt) > 0:
ipp = 0
if negative_prompt.startswith(self.__separator):
ipp = len(self.__separator) # adjust for existing end separator
if negative_prompt.startswith(self.stn_separator):
ipp = len(self.stn_separator) # adjust for existing end separator
add_at_start.append(negative_prompt[ipp:])
negative_prompt = self.__separator.join(add_at_start)
negative_prompt = self.stn_separator.join(add_at_start)
return negative_prompt
def __add_to_end(self, negative_prompt, add_at_end):
"""
Adds the elements in `add_at_end` list to the end of `negative_prompt` string.
Args:
negative_prompt (str): The original negative prompt string.
add_at_end (list): The list of elements to be added at the end of `negative_prompt`.
Returns:
str: The updated negative prompt string with elements added at the end.
"""
if len(negative_prompt) > 0:
ipl = len(negative_prompt)
if negative_prompt.endswith(self.__separator):
ipl -= len(self.__separator) # adjust for existing start separator
if negative_prompt.endswith(self.stn_separator):
ipl -= len(self.stn_separator) # adjust for existing start separator
add_at_end.insert(0, negative_prompt[:ipl])
negative_prompt = self.__separator.join(add_at_end)
negative_prompt = self.stn_separator.join(add_at_end)
return negative_prompt
def __sendtonegative(self, prompt, negative_prompt):
"""
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, add_at = self.__find_tags(prompt)
negative_prompt = self.__add_to_insertion_points(negative_prompt, add_at["insertion_point"])
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`.
Attributes:
__ppp (object): An instance of the parent class `ppp`.
Methods:
scheduled(tree): Replicates or removes scheduling constructs based on conditions.
alternate(tree): Replicates or removes alternation constructs based on conditions.
emphasized(tree): Replicates or removes attention constructs based on conditions.
deemphasized(tree): Replicates or removes attention constructs based on conditions.
modeltag(tree): Replicates model constructs.
numpar(tree): Cleans up number parameter.
negtag(tree): Replicates or removes negative tag constructs based on conditions.
wildcard(tree): Replicates or removes wildcard constructs based on conditions.
choices(tree): Replicates or removes choices constructs based on conditions.
choice(tree): Replicates choices.
plain(tree): 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 scheduled(self, tree):
if len(tree) == 0 and self.__phase == "cleanup" and self.__ppp.cup_emptyconstructs:
return "" # remove invalid scheduling construct (probably this is not reachable)
# replicate scheduling construct
if len(tree) > 0 and tree[0] is None:
return f"[{':'.join(tree[1:])}]"
return f"[{':'.join(tree)}]"
def alternate(self, tree):
if len(tree) == 0 and self.__phase == "cleanup" and self.__ppp.cup_emptyconstructs:
return "" # remove invalid alternation construct (probably this is not reachable)
return f"[{'|'.join(tree)}]" # replicate alternation construct
def emphasized(self, tree):
if (len(tree) == 0 or tree[0] == "") and self.__phase == "cleanup" and self.__ppp.cup_emptyconstructs:
return "" # remove empty attention construct
if len(tree) > 1 and tree[1] is not None:
return f"({tree[0]}:{tree[1]})" # replicate attention construct with weight
return f"({tree[0]})" # replicate attention construct without weight
def deemphasized(self, tree):
if (len(tree) == 0 or tree[0] == "") and self.__phase == "cleanup" and self.__ppp.cup_emptyconstructs:
return "" # remove empty attention construct (invalid scheduling or alternation constructs end up here too?)
return f"[{tree[0]}]" # replicate attention construct
def modeltag(self, tree):
return f"<{tree[0]}>" # replicate model construct
def numpar(self, tree):
return next(x for x in tree if x.type == "NUMBER").value.strip() # clean up number parameter
def negtag(self, tree):
if self.__phase == "cleanup":
return "" # remove negative tag construct (there shouldn't be any at this point)
parameters = "!" + tree[0] + "!" if tree[0] is not None else ""
content = "".join(tree[1::])
return f"<!{parameters}{content}!>" # replicate negative tag construct
def wildcard(self, tree):
content = f"__{tree[0]}__" # replicate wildcard construct
self.detectedWildcards.append(content)
if self.__phase == "wildcards" and self.__ppp.ifwildcards == self.__ppp.IFWILDCARDS_CHOICES["remove"]:
return ""
return content
def choices(self, tree):
content = "{" + "|".join(tree) + "}" # replicate wildcard choices construct
self.detectedWildcards.append(content)
if self.__phase == "wildcards" and self.__ppp.ifwildcards == self.__ppp.IFWILDCARDS_CHOICES["remove"]:
return ""
return content
def choice(self, tree):
return f"{tree[0]}" # replicate choice
def plain(self, tree):
if self.__phase == "cleanup":
return self.__ppp.cleanup_text(tree[0].value) # clean up plain text
return tree[0].value
def __default__(self, data, children, meta):
joined = "".join(children) # join all children
if self.__phase == "cleanup":
# clean up joined text if there are no constructs to take care of cleaning the joints
if not re.match(r"[([<{]", joined):
joined = self.__ppp.cleanup_text(joined)
return joined
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.
"""
if self.debug:
self.logger.info("Doing cleanup")
transformtree = self.TransformerTree(self, phase="cleanup")
try:
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("Cleanup parsing failed on prompt!: %s", e)
try:
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("Cleanup 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.
Args:
text (str): The text to be cleaned up.
Returns:
str: The cleaned up text.
"""
if self.cup_extraseparators:
# sendtonegative separator
escapedSeparator = re.escape(self.stn_separator)
text = re.sub(r"(?:\s*" + escapedSeparator + r"\s*){2,}", self.stn_separator, text)
# regular comma separator
text = re.sub(r"(?:\s*,\s*){2,}", ", ", text)
if self.cup_breaks:
text = re.sub(r"[, ]+BREAK[, ]+", " BREAK ", text)
text = re.sub(r"BREAK(?:\s+BREAK)+[ ]+", "BREAK ", text)
text = re.sub(r"[ ]+BREAK(?:\s+BREAK)+", " BREAK", text)
if self.cup_extraspaces:
text = re.sub(r"[ ]+,", ",", text) # remove spaces before comma
text = re.sub(r"[ ]{2,}", " ", text) # collapse spaces
return text
def trim_text(self, text):
"""
Trims the given text based on the specified cleanup options.
Args:
text (str): The text to be trimmed.
Returns:
str: The trimmed text.
"""
if self.cup_extraseparators:
# sendtonegative separator
escapedSeparator = re.escape(self.stn_separator)
text = re.sub(r"^(?:\s*" + escapedSeparator + r"\s*)", "", text)
text = re.sub(r"(?:\s*" + escapedSeparator + r"\s*)$", "", text)
# regular comma separator
text = re.sub(r"^\s*,\s*", "", text)
text = re.sub(r"\s*,\s*$", "", text)
if self.cup_breaks:
text = re.sub(r"^BREAK\s+", "", text)
text = re.sub(r"\s+BREAK$", "", text)
if self.cup_extraspaces:
text = text.strip()
return text
def __findwildcards(self, prompt, negative_prompt):
"""
Find and process wildcards in the prompt and negative_prompt strings.
Args:
prompt (str): The prompt string.
negative_prompt (str): The negative prompt string.
Returns:
tuple: A tuple containing the processed prompt and negative_prompt strings.
"""
if self.debug:
self.logger.info("Doing wildcard processing")
p_transformtree = self.TransformerTree(self, phase="wildcards")
try:
p_tree = self.__parser_complete.parse(prompt)
# self.logger.info(f"Wildcards tree from prompt:\n{p_tree.pretty()}")
prompt = p_transformtree.transform(p_tree)
except Exception as e: # pylint: disable=broad-except
self.logger.warning("Wildcards parsing failed in prompt!: %s", e)
np_transformtree = self.TransformerTree(self, phase="wildcards")
try:
np_tree = self.__parser_complete.parse(negative_prompt)
# self.logger.info(f"Wildcards tree from negative prompt:\n{np_tree.pretty()}")
negative_prompt = np_transformtree.transform(np_tree)
except Exception as e: # pylint: disable=broad-except
self.logger.warning("Wildcards 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}")
if foundP or foundNP:
if self.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.")
if foundP:
prompt = self.WILDCARD_STOP + prompt
if foundNP:
negative_prompt = self.WILDCARD_STOP + negative_prompt
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)}")
return prompt, negative_prompt
def process_prompt(self, original_prompt, original_negative_prompt):
"""
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.
Returns:
tuple: A tuple containing the processed prompt and negative prompt.
"""
try:
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)}")
self.logger.info(f"Input negative_prompt: {self.formatOutput(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_breaks
):
prompt, negative_prompt = self.__cleanup(prompt, negative_prompt)
return prompt, negative_prompt
except Exception as e: # pylint: disable=broad-exception-caught
self.logger.exception(e)
return original_prompt, original_negative_prompt
+52 -2
View File
@@ -1,10 +1,25 @@
import logging
import sys
import copy
import logging
class PromptPostProcessorLogFactory:
class PromptPostProcessorLogFactory: # pylint: disable=too-few-public-methods
"""
Factory class for creating loggers for the PromptPostProcessor module.
"""
class ColoredFormatter(logging.Formatter):
"""
A custom logging formatter that adds color to log records based on their level.
Attributes:
COLORS (dict): A dictionary mapping log levels to ANSI escape codes for colors.
Methods:
format(record): Formats the log record with color based on its level.
"""
COLORS = {
"DEBUG": "\033[0;36m", # CYAN
"INFO": "\033[0;32m", # GREEN
@@ -15,6 +30,15 @@ class PromptPostProcessorLogFactory:
}
def format(self, record):
"""
Formats the log record with color based on the log level.
Args:
record (LogRecord): The log record to be formatted.
Returns:
str: The formatted log record.
"""
colored_record = copy.copy(record)
levelname = colored_record.levelname
seq = self.COLORS.get(levelname, self.COLORS["RESET"])
@@ -22,6 +46,17 @@ class PromptPostProcessorLogFactory:
return super().format(colored_record)
def __init__(self):
"""
Initializes the PromptPostProcessor class.
This method sets up the logger for the PromptPostProcessor class and configures its log level and handlers.
Args:
None
Returns:
None
"""
ppplog = logging.getLogger("PromptPostProcessor")
ppplog.propagate = False
if not ppplog.handlers:
@@ -33,5 +68,20 @@ class PromptPostProcessorLogFactory:
class PromptPostProcessorLogCustomAdapter(logging.LoggerAdapter):
"""
Custom logger adapter for the PromptPostProcessor.
This adapter adds a prefix to log messages to indicate that they are related to the PromptPostProcessor.
"""
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
+161 -25
View File
@@ -7,35 +7,94 @@ import os
sys.path.insert(1, os.path.join(sys.path[0], ".."))
# pylint: disable=import-error
from modules import scripts, shared, script_callbacks
from modules.processing import StableDiffusionProcessing
from modules.shared import opts
import gradio as gr
from ppp import PromptPostProcessor
from ppp_logging import PromptPostProcessorLogFactory
class PromptPostProcessorScript(scripts.Script):
"""
This class represents a script for prompt post-processing.
It is responsible for processing prompts and applying various settings and cleanup operations.
Attributes:
callbacks_added (bool): Flag indicating whether the script callbacks have been added.
Methods:
__init__(): Initializes the PromptPostProcessorScript object.
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.
__on_ui_settings(): Callback function for UI settings.
"""
def __init__(self):
if not hasattr(self, "callbacks_added"):
"""
Initializes the PromptPostProcessor object.
This method adds callbacks for UI settings and initializes the logger.
Parameters:
None
Returns:
None
"""
if not hasattr(self, "ppp_callbacks_added"):
lf = PromptPostProcessorLogFactory()
self.__logppp = lf.log
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.callbacks_added = True
self.ppp_callbacks_added = True
def title(self):
"""
Returns the title of the script.
Returns:
str: The title of the script.
"""
return PromptPostProcessor.NAME
def show(self, is_img2img):
"""
Determines whether the script should be shown based on the kind of processing.
Args:
is_img2img (bool): Flag indicating whether the processing is image-to-image.
Returns:
scripts.Visibility: The visibility setting for the script.
"""
return scripts.AlwaysVisible
def process(self, p: StableDiffusionProcessing, *args, **kwargs):
ppp = PromptPostProcessor(self.__logppp, opts=opts)
for i in range(len(p.all_prompts)): # pylint: disable=consider-using-enumerate
p.all_prompts[i], p.all_negative_prompts[i] = ppp.process_prompt(
p.all_prompts[i], p.all_negative_prompts[i]
)
def process(self, p: StableDiffusionProcessing, *args, **kwargs): # pylint: disable=unused-argument
"""
Processes the prompts and applies post-processing operations.
Args:
p (StableDiffusionProcessing): The StableDiffusionProcessing object containing the prompts.
Returns:
None
"""
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, 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)
# make it compatible with A1111 hires fix
if (
hasattr(p, "all_hr_prompts")
@@ -43,17 +102,66 @@ class PromptPostProcessorScript(scripts.Script):
and hasattr(p, "all_hr_negative_prompts")
and p.all_hr_negative_prompts is not None
):
for i in range(len(p.all_hr_prompts)): # pylint: disable=consider-using-enumerate
p.all_hr_prompts[i], p.all_hr_negative_prompts[i] = ppp.process_prompt(
p.all_hr_prompts[i], p.all_hr_negative_prompts[i]
)
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)
def ppp_interrupt(self):
"""
Interrupts the generation.
Returns:
None
"""
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_separator",
key="ppp_gen_sep", info=shared.OptionInfo("<h2>General settings</h2>", "", gr.HTML, section=section)
)
shared.opts.add_option(
key="ppp_gen_debug",
info=shared.OptionInfo(
PromptPostProcessor.DEFAULT_SEPARATOR,
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,
),
)
# send to negative settings
shared.opts.add_option(
key="ppp_stn_sep",
info=shared.OptionInfo("<h2>Send to Negative settings</h2>", "", 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,
),
)
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,
),
@@ -74,19 +182,47 @@ class PromptPostProcessorScript(scripts.Script):
section=section,
),
)
# clean-up settings
shared.opts.add_option(
key="ppp_cleanup",
info=shared.OptionInfo(
True,
label="Try to clean-up the prompt after processing (removes extra spaces, empty attention, or the configured separator)",
section=section,
),
key="ppp_cup_sep", info=shared.OptionInfo("<h2>Clean-up settings</h2>", "", gr.HTML, section=section)
)
shared.opts.add_option(
key="ppp_debug",
key="ppp_cup_doi2i",
info=shared.OptionInfo(
False,
label="Debug",
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_extraspaces",
info=shared.OptionInfo(
True,
label="Remove extra spaces",
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_breaks",
info=shared.OptionInfo(
True,
label="Clean up around BREAKs",
section=section,
),
)
+162 -26
View File
@@ -5,16 +5,56 @@ import os
sys.path.insert(1, os.path.join(sys.path[0], ".."))
from ppp import PromptPostProcessor # pylint: disable=import-error
from ppp import PromptPostProcessor
from ppp_logging import 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)
class TestPromptPostProcessor(unittest.TestCase):
"""
A test case class for testing the PromptPostProcessor class.
"""
def setUp(self):
"""
Set up the test case by initializing the necessary objects and configurations.
"""
lf = PromptPostProcessorLogFactory()
self.__log = lf.log
self.__log.setLevel(logging.DEBUG)
self.defppp = PromptPostProcessor(self.__log, separator=", ", ignore_repeats=True, join_attention=True, cleanup=True)
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_extraspaces": True,
"ppp_cup_breaks": True,
}
)
self.defppp = PromptPostProcessor(self, self.__defopts)
def ppp_interrupt(self):
pass # fake interrupt
def process(
self,
@@ -24,6 +64,19 @@ class TestPromptPostProcessor(unittest.TestCase):
expected_negative_prompt,
ppp=None,
):
"""
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.
ppp (object, optional): The post-processor object. Defaults to None.
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}'")
@@ -33,7 +86,9 @@ class TestPromptPostProcessor(unittest.TestCase):
f"Negative Prompt should be '{expected_negative_prompt}'",
)
def test_tag_default(self):
# Send To Negative tests
def test_tag_default(self): # negtag with no parameters
self.process(
"flowers<!red!>",
"normal quality, worse quality",
@@ -41,7 +96,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"red, normal quality, worse quality",
)
def test_tag_start(self):
def test_tag_start(self): # negtag with s parameter
self.process(
"flowers<!!s!red!>",
"normal quality, worse quality",
@@ -49,7 +104,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"red, normal quality, worse quality",
)
def test_tag_end(self):
def test_tag_end(self): # negtag with e parameter
self.process(
"flowers<!!e!red!>",
"normal quality, worse quality",
@@ -57,7 +112,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"normal quality, worse quality, red",
)
def test_tag_insertion_mid_sep(self):
def test_tag_insertion_mid_sep(self): # negtag with p parameter and insertion in the middle
self.process(
"flowers<!!p0!red!>",
"normal quality, <!!i0!!>, worse quality",
@@ -65,7 +120,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"normal quality, red, worse quality",
)
def test_tag_insertion_mid_no_sep(self):
def test_tag_insertion_mid_no_sep(self): # negtag with p parameter and insertion in the middle without separator
self.process(
"flowers<!!p0!red!>",
"normal quality<!!i0!!>worse quality",
@@ -73,7 +128,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"normal quality, red, worse quality",
)
def test_tag_insertion_start_sep(self):
def test_tag_insertion_start_sep(self): # negtag with p parameter and insertion at the start
self.process(
"flowers<!!p0!red!>",
"<!!i0!!>, normal quality, worse quality",
@@ -81,7 +136,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"red, normal quality, worse quality",
)
def test_tag_insertion_start_no_sep(self):
def test_tag_insertion_start_no_sep(self): # negtag with p parameter and insertion at the start without separator
self.process(
"flowers<!!p0!red!>",
"<!!i0!!>normal quality, worse quality",
@@ -89,7 +144,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"red, normal quality, worse quality",
)
def test_tag_insertion_end_sep(self):
def test_tag_insertion_end_sep(self): # negtag with p parameter and insertion at the end
self.process(
"flowers<!!p0!red!>",
"normal quality, worse quality, <!!i0!!>",
@@ -97,7 +152,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"normal quality, worse quality, red",
)
def test_tag_insertion_end_no_sep(self):
def test_tag_insertion_end_no_sep(self): # negtag with p parameter and insertion at the end without separator
self.process(
"flowers<!!p0!red!>",
"normal quality, worse quality<!!i0!!>",
@@ -105,7 +160,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"normal quality, worse quality, red",
)
def test_complex(self):
def test_complex(self): # complex negtags
self.process(
"<!red!> (<!!s!pink!>), flowers <!!e!purple!>, <!!e!blue!>, <!!p0!yellow!> <!!p1!green!>",
"normal quality, <!!i0!!>, bad quality<!!i1!!>, worse quality",
@@ -113,24 +168,35 @@ class TestPromptPostProcessor(unittest.TestCase):
"red, (pink), normal quality, yellow, bad quality, green, worse quality, purple, blue",
)
def test_complex_no_cleanup(self):
def test_complex_no_cleanup(self): # complex negtags with no cleanup
self.process(
"<!red!> (<!!s!pink!>), flowers <!!e!purple!>, <!!e!blue!>, <!!p0!yellow!> <!!p1!green!>",
"normal quality, <!!i0!!>, bad quality<!!i1!!>, worse quality",
" (), flowers , , ",
"red, (pink), normal quality, yellow, bad quality, green, worse quality, purple, blue",
PromptPostProcessor(self.__log, separator=", ", ignore_repeats=True, join_attention=True, cleanup=False),
PromptPostProcessor(
self,
DictToObj(
{
**self.__defopts.__dict__,
"ppp_cup_emptyconstructs": False,
"ppp_cup_extraseparators": False,
"ppp_cup_extraspaces": False,
"ppp_cup_breaks": False,
}
),
),
)
def test_inside_attention1(self):
def test_inside_attention1(self): # negtag inside attention
self.process(
"[<!neg1!>] this is a ((test<!!e!neg2!>) (test:2.0): 1.5 )",
"normal quality",
"this is a ((test) (test:2.0): 1.5 )",
"this is a ((test) (test:2.0):1.5)",
"[neg1], normal quality, (neg2:1.65)",
)
def test_inside_attention2(self):
def test_inside_attention2(self): # negtag inside attention
self.process(
"(red<![square]!>:1.5)",
"",
@@ -138,7 +204,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"([square]:1.5)",
)
def test_inside_alternation1(self):
def test_inside_alternation1(self): # negtag inside alternation
self.process(
"this is a (([complex|simple<!neg1!>|regular] test)(test:2.0):1.5)",
"normal quality",
@@ -146,7 +212,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"([|neg1|]:1.65), normal quality",
)
def test_inside_alternation2(self):
def test_inside_alternation2(self): # negtag inside alternation
self.process(
"this is a (([complex<!neg1!>|simple<!neg2!>|regular<!neg3!>] test)(test:2.0):1.5)",
"normal quality",
@@ -154,7 +220,7 @@ class TestPromptPostProcessor(unittest.TestCase):
"([neg1||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality",
)
def test_inside_alternation3(self):
def test_inside_alternation3(self): # negtag inside alternation (recursive alternation)
self.process(
"this is a (([complex<!neg1!>[one|two<!neg12!>||three|four(<!neg14!>)]|simple<!neg2!>|regular<!neg3!>] test)(test:2.0):1.5)",
"normal quality",
@@ -162,22 +228,92 @@ class TestPromptPostProcessor(unittest.TestCase):
"([neg1||]:1.65), ([[|neg12|||]||]:1.65), ([[||||(neg14)]||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality",
)
def test_inside_scheduling(self):
def test_inside_scheduling(self): # negtag inside scheduling
self.process(
"this is [abc<!neg1!>:def<!!e!neg2!>: 5 ]",
"normal quality",
"this is [abc:def: 5 ]",
"this is [abc:def:5]",
"[neg1::5], normal quality, [neg2:5]",
)
def test_complex_features(self):
def test_complex_features(self): # complex negtags with features
self.process(
"[<!neg5!>] this is: a (([complex|simple<!neg6!>|regular] test<!neg1!>)(test:2.0):1.5) \nBREAK with [abc<!neg4!>:def<!!p0!neg2(neg3:1.6)!>:5] <lora:xxx:1>",
"[<!neg5!>] this is: a (([complex|simple<!neg6!>|regular] test<!neg1!>)(test:2.0):1.5) \nBREAK, BREAK with [abc<!neg4!>:def<!!p0!neg2(neg3:1.6)!>:5] <lora:xxx:1>",
"normal quality, <!!i0!!>",
"this is: a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK with [abc:def:5] <lora:xxx:1>",
"[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]",
)
# Wildcard tests
def test_wildcards_ignore(self): # wildcards with ignore option
self.process(
"__bad_wildcard__",
"{option1|option2}",
"__bad_wildcard__",
"{option1|option2}",
PromptPostProcessor(
self,
DictToObj(
{
**self.__defopts.__dict__,
"ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["ignore"],
}
),
),
)
def test_wildcards_remove(self): # wildcards with remove option
self.process(
"[<!neg5!>] this is: __bad_wildcard__ a (([complex|simple<!neg6!>|regular] test<!neg1!>)(test:2.0):1.5) \nBREAK, BREAK with [abc<!neg4!>:def<!!p0!neg2(neg3:1.6)!>:5] <lora:xxx:1>",
"normal quality, <!!i0!!> {option1|option2}",
"this is: a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK with [abc:def:5] <lora:xxx:1>",
"[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]",
PromptPostProcessor(
self,
DictToObj(
{
**self.__defopts.__dict__,
"ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["remove"],
}
),
),
)
def test_wildcards_warn(self): # wildcards with warn option
self.process(
"__bad_wildcard__",
"{option1|option2}",
PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__",
"{option1|option2}",
PromptPostProcessor(
self,
DictToObj(
{
**self.__defopts.__dict__,
"ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["warn"],
}
),
),
)
def test_wildcards_stop(self): # wildcards with stop option
self.process(
"__bad_wildcard__",
"{option1|option2}",
PromptPostProcessor.WILDCARD_STOP + "__bad_wildcard__",
PromptPostProcessor.WILDCARD_STOP + "{option1|option2}",
PromptPostProcessor(
self,
DictToObj(
{
**self.__defopts.__dict__,
"ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["stop"],
}
),
),
)
if __name__ == "__main__":
unittest.main()