Compare commits

...
6 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
Antonio Cordero Balcazar 77a74a8088 * Renamed the extension to "Prompt Post-Processor".
* Improved logging.
* Fix compatibility with A1111 hiresfix.
2023-12-02 11:12:00 +01:00
Antonio Cordero Balcazar f59b3e51a5 Fix processing with A1111 hires fix. (#4)
# Pull Request

## Description

Fixes processing A1111 hr prompts.

Fixes #3 

## Type of change

- [x] Bug fix (non-breaking change which fixes an issue)

## How Has This Been Tested?

**Test Configuration**:

- A1111 v1.6
- SD.Next

## Checklist

- [x] My code follows the style guidelines of this project
- [x] I have performed a self-review of my own code
- [x] I have commented my code, particularly in hard-to-understand areas
- [ ] I have made corresponding changes to the documentation
- [x] My changes generate no new warnings
- [ ] I have added tests that prove my fix is effective or that my
feature works
- [x] New and existing unit tests pass locally with my changes
- [ ] Any dependent changes have been merged and published in downstream
modules
- [x] I have checked my code and corrected any misspellings
2023-10-14 14:29:26 +02:00
Antonio Cordero Balcazar 496117e004 * Fix processing with A1111 hires fix.
* Remove version from title.
2023-10-14 13:51:09 +02:00
Antonio Cordero Balcazar 5a87292a18 * Fix processing of weights/steps with spaces around. 2023-08-24 18:22:40 +02:00
Antonio Cordero Balcazar a0d22862ca Fix weight calculation. 2023-08-20 23:49:09 +02:00
8 changed files with 1331 additions and 508 deletions
+57 -38
View File
@@ -1,45 +1,51 @@
# Send to Negative for Stable Diffusion WebUI
# Prompt Postprocessor for Stable Diffusion WebUI
Extension for the [AUTOMATIC1111 Stable Diffusion WebUI](https://github.com/AUTOMATIC1111/stable-diffusion-webui) or compatible UIs.
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).
## Purpose
Currently this extension has these functions:
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.
* 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.
Note: The extension must be loaded after the installed wildcards extension. Extensions
load by their folder in alphanumeric order.
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.
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 the extension folder so the ordering
works out.
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 "Send to Negative".
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 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
1. Go to Extensions > Install from URL
2. Paste <https://github.com/acorderob/sd-webui-sendtonegative> in the URL for extension's git repository text field
2. Paste <https://github.com/acorderob/sd-webui-prompt-postprocessor> in the URL for extension's git repository text field
3. Click the Install button
4. Restart the webui
## 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:
```text
@@ -54,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
#### 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
@@ -72,22 +76,37 @@ 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
The content of the negative tags is not processed and is copied as is to the negative prompt. Other modifiers around the tags are processed in the following way.
The content of the negative tags is not processed and is copied as-is to the negative prompt. Other modifiers around the tags are processed in the following way.
### Attention modifiers (weights)
@@ -101,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.
+789
View File
@@ -0,0 +1,789 @@
from collections import namedtuple
import re
import math
import lark
class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-instance-attributes
"""
The PromptPostProcessor class is responsible for processing and manipulating prompt strings.
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,
script,
opts=None,
is_i2i=False,
):
"""
Initializes the PPP object.
Args:
script: The script object.
opts: Optional. The options object for configuring PPP behavior.
"""
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.ifwildcards = (
getattr(opts, "ppp_gen_ifwildcards", self.IFWILDCARDS_CHOICES["ignore"])
if opts is not None
else self.IFWILDCARDS_CHOICES["ignore"]
)
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
)
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.__parser_complete = lark.Lark(
r"""
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 "]"
alternate: "[" alternateoption ("|" alternateoption)+ "]"
alternateoption: prompt
negtag: "<!" [negtagparameters] nonegprompt "!>"
negtagparameters: "!" /s|e|[ip]\d/ "!"
modeltag: "<" /(?!!)[^>]+/ ">"
numpar: WHITESPACE* NUMBER WHITESPACE*
WHITESPACE: /\s+/
plain: /((?!__)[^\\[\]():|<>!{}]|\\.)+/s
%import common.SIGNED_NUMBER -> NUMBER
""", # prompt, nonegprompt, plain with ?
propagate_positions=True,
)
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.__ppp = ppp
self.__prompt = prompt
self.AccumulatedShell = namedtuple("AccumulatedShell", ["type", "info1", "info2"])
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.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:
before = None
after = tree.children[-2]
numpar = tree.children[-1]
pos = self.__get_numpar_value(numpar)
if pos >= 1:
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.__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.__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.__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)
self.__shell.pop()
# 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.__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.__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 = []
for x in tree.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
)
content = "".join(rest)
self.__negtags.append(
self.NegTag(tree.meta.start_pos, tree.meta.end_pos, content, parameters, self.__shell.copy())
)
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 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].info1 * negtag.shell[i].info1)
/ 100, # we limit the new weight to two decimals
None,
)
negtag.shell.pop(i)
start = ""
end = ""
for s in negtag.shell:
match s.type:
case "at":
if s.info1 == 0.9:
start += "["
end = "]" + end
elif s.info1 == 1.1:
start += "("
end = ")" + end
else:
start += "("
end = f":{s.info1})" + end
# case "sc":
case "scb":
start += "["
end = f"::{s.info1}]" + end
case "sca":
start += "["
end = f":{s.info1}]" + end
# case "al":
case "alo":
start += "[" + ("|" * int(s.info1 - 1))
end = ("|" * int(s.info2 - s.info1)) + "]" + end
content = start + negtag.content + end
position = negtag.parameters or "s"
if 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):
"""
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": []}
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] :]
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.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.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.stn_separator.join(add_at_insertion_point[n])
else:
ipp = 0
if negative_prompt.startswith(self.stn_separator):
ipp = len(self.stn_separator)
add_at_insertion_point[n].append(negative_prompt[ipp:])
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.stn_separator):
ipp = len(self.stn_separator) # adjust for existing end separator
add_at_start.append(negative_prompt[ipp:])
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.stn_separator):
ipl -= len(self.stn_separator) # adjust for existing start separator
add_at_end.insert(0, negative_prompt[:ipl])
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
+87
View File
@@ -0,0 +1,87 @@
import logging
import sys
import copy
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
"WARNING": "\033[0;33m", # YELLOW
"ERROR": "\033[0;31m", # RED
"CRITICAL": "\033[0;37;41m", # WHITE ON RED
"RESET": "\033[0m", # RESET COLOR
}
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"])
colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}"
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:
handler = logging.StreamHandler(sys.stdout)
handler.setFormatter(self.ColoredFormatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s"))
ppplog.addHandler(handler)
ppplog.setLevel(logging.INFO)
self.log = PromptPostProcessorLogCustomAdapter(ppplog)
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
-73
View File
@@ -1,73 +0,0 @@
if __name__ == "__main__":
raise SystemExit("This script must be run from a Stable Diffusion WebUI")
import sys
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
from sendtonegative import SendToNegative
from stnlogging import SendToNegativeLogFactory
class SendToNegativeScript(scripts.Script):
def __init__(self):
if not hasattr(self, "callbacks_added"):
lf = SendToNegativeLogFactory()
self.__logstn = lf.log
script_callbacks.on_ui_settings(self.__on_ui_settings)
self.callbacks_added = True
def title(self):
return f"{SendToNegative.NAME} v{SendToNegative.VERSION}"
def show(self, is_img2img):
return scripts.AlwaysVisible
def process(self, p: StableDiffusionProcessing, *args, **kwargs):
stn = SendToNegative(self.__logstn, opts=opts)
for i in range(len(p.all_prompts)): # pylint: disable=consider-using-enumerate
p.all_prompts[i], p.all_negative_prompts[i] = stn.process_prompt(
p.all_prompts[i], p.all_negative_prompts[i]
)
def __on_ui_settings(self):
section = ("send-to-negative", SendToNegative.NAME)
shared.opts.add_option(
key="stn_separator",
info=shared.OptionInfo(
SendToNegative.DEFAULT_SEPARATOR,
label="Separator used when adding to the negative prompt",
section=section,
),
)
shared.opts.add_option(
key="stn_ignorerepeats",
info=shared.OptionInfo(
True,
label="Ignore tags with repeated content",
section=section,
),
)
shared.opts.add_option(
key="stn_joinattention",
info=shared.OptionInfo(
True,
label="Join attention modifiers (weights) when possible",
section=section,
),
)
shared.opts.add_option(
key="stn_cleanup",
info=shared.OptionInfo(
True,
label="Try to clean-up the prompt after processing (removes extra spaces or the configured separator)",
section=section,
),
)
+228
View File
@@ -0,0 +1,228 @@
if __name__ == "__main__":
raise SystemExit("This script must be run from a Stable Diffusion WebUI")
import sys
import os
sys.path.insert(1, os.path.join(sys.path[0], ".."))
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):
"""
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.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
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): # 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")
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)
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_gen_sep", info=shared.OptionInfo("<h2>General settings</h2>", "", 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,
),
)
# 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,
),
)
shared.opts.add_option(
key="ppp_stn_ignorerepeats",
info=shared.OptionInfo(
True,
label="Ignore tags with 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("<h2>Clean-up settings</h2>", "", 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_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,
),
)
-323
View File
@@ -1,323 +0,0 @@
from collections import namedtuple
import re
import lark
class SendToNegative: # pylint: disable=too-few-public-methods
NAME = "Send to Negative"
VERSION = "2.1.1"
DEFAULT_SEPARATOR = ", "
def __init__(
self,
log,
separator=None,
ignore_repeats=None,
join_attention=None,
cleanup=None,
opts=None,
):
"""
Format for the tag:
<!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.
"""
self.__logger = log
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, "stn_ignorerepeats", True)
)
self.__join_attention = (
join_attention
if join_attention is not None
else getattr(opts, "stn_joinattention", True)
if opts is not None
else True
)
self.__cleanup = (
cleanup if cleanup is not None else getattr(opts, "stn_cleanup", True) if opts is not None else True
)
self.__separator = (
separator
if separator is not None
else getattr(opts, "stn_separator", self.DEFAULT_SEPARATOR)
if opts is not None
else self.DEFAULT_SEPARATOR
)
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(
r"""
start: (prompt | /[\][():|<>!]/+)*
?prompt: (emphasized | deemphasized | scheduled | alternate | modeltag | negtag | plain)*
?nonegprompt: (emphasized | deemphasized | scheduled | alternate | modeltag | plain)*
emphasized: "(" prompt [":" numpar] ")"
deemphasized: "[" prompt "]"
scheduled: "[" [prompt ":"] prompt ":" numpar "]"
alternate: "[" alternateoption ("|" alternateoption)+ "]"
alternateoption: prompt
negtag: "<!" [negtagparameters] nonegprompt "!>"
negtagparameters: "!" /s|e|[ip]\d/ "!"
modeltag: "<" /(?!!)[^>]+/ ">"
numpar: WHITESPACE* NUMBER WHITESPACE*
WHITESPACE: /\s+/
?plain: /([^\\[\]():|<>!]|\\.)+/s
%import common.SIGNED_NUMBER -> NUMBER
""",
propagate_positions=True,
)
class ReadTree(lark.visitors.Interpreter):
def __init__(self, logger, ignorerepeats, joinattention, prompt, add_at):
super().__init__()
self.__logger = logger
self.__ignore_repeats = ignorerepeats
self.__join_attention = joinattention
self.__prompt = prompt
self.AccumulatedShell = namedtuple("AccumulatedShell", ["type", "info1", "info2"])
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.remove = []
def scheduled(self, tree):
if len(tree.children) > 2: # before & after
before = tree.children[0]
else:
before = None
after = tree.children[-2]
numpar = tree.children[-1]
pos = float(numpar.children[0].value)
if pos >= 1:
pos = int(pos)
# self.__shell.append(self.AccumulatedShell("sc", tree.meta.start_pos, pos))
if before is not None and hasattr(before, "data"):
self.__logger.debug(
f"Shell scheduled before at {[before.meta.start_pos,before.meta.end_pos] if hasattr(before,'meta') else '?'} : {pos}"
)
self.__shell.append(self.AccumulatedShell("scb", pos, None))
self.visit(before)
self.__shell.pop()
if hasattr(after, "data"):
self.__logger.debug(
f"Shell scheduled after at {[after.meta.start_pos,after.meta.end_pos] if hasattr(after,'meta') else '?'} : {pos}"
)
self.__shell.append(self.AccumulatedShell("sca", pos, None))
self.visit(after)
self.__shell.pop()
# self.__shell.pop()
def alternate(self, tree):
# self.__shell.append(self.AccumulatedShell("al", tree.meta.start_pos, len(tree.children)))
for i, opt in enumerate(tree.children):
self.__logger.debug(
f"Shell alternate at {[opt.meta.start_pos,opt.meta.end_pos] if hasattr(opt,'meta') else '?'} : {i+1}"
)
if hasattr(opt, "data"):
self.__shell.append(self.AccumulatedShell("alo", i + 1, len(tree.children)))
self.visit(opt)
self.__shell.pop()
# self.__shell.pop()
def emphasized(self, tree):
numpar = tree.children[-1]
weight = float(numpar.children[0].value) if numpar is not None else 1.1
self.__logger.debug(
f"Shell attention at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') else '?'}: {weight}"
)
self.__shell.append(self.AccumulatedShell("at", weight, None))
self.visit_children(tree)
self.__shell.pop()
def deemphasized(self, tree):
weight = 0.9
self.__logger.debug(
f"Shell attention at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') else '?'}: {weight}"
)
self.__shell.append(self.AccumulatedShell("at", weight, None))
self.visit_children(tree)
self.__shell.pop()
def negtag(self, tree):
negtagparameters = tree.children[0]
parameters = negtagparameters.children[0].value if negtagparameters is not None else ""
rest = []
for x in tree.children[1::]:
rest.append(self.__prompt[x.meta.start_pos : x.meta.end_pos] if hasattr(x, "meta") else x.value)
content = "".join(rest)
self.__negtags.append(
self.NegTag(tree.meta.start_pos, tree.meta.end_pos, content, parameters, self.__shell.copy())
)
self.__logger.debug(
f"Negative tag at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') else '?'}: {parameters}: {content.encode('unicode_escape').decode('utf-8')}"
)
def start(self, tree):
self.visit_children(tree)
# process the found negtags
for nt in self.__negtags:
if self.__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(
"at",
(100 * nt.shell[i - 1].info1 * nt.shell[i].info1) / 100, # we limit to two decimals
None,
)
nt.shell.pop(i)
start = ""
end = ""
for s in nt.shell:
match s.type:
case "at":
if s.info1 == 0.9:
start += "["
end = "]" + end
elif s.info1 == 1.1:
start += "("
end = ")" + end
else:
start += "("
end = f":{s.info1})" + end
# case "sc":
case "scb":
start += "["
end = f"::{s.info1}]" + end
case "sca":
start += "["
end = f":{s.info1}]" + end
# case "al":
case "alo":
start += "[" + ("|" * int(s.info1 - 1))
end = ("|" * int(s.info2 - s.info1)) + "]" + end
content = start + nt.content + end
position = nt.parameters or "s"
if len(content) > 0:
if content not in self.__already_processed:
if self.__ignore_repeats:
self.__already_processed.append(content)
self.__logger.debug(
f"Adding content at position {position}: {content.encode('unicode_escape').decode('utf-8')}"
)
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.__logger.warning(
f"Ignoring repeated content: {content.encode('unicode_escape').decode('utf-8')}"
)
# 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.__logger.debug(f"Input prompt: {prompt.encode('unicode_escape').decode('utf-8')}")
self.__logger.debug(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"])
self.__logger.debug(f"Output prompt: {prompt.encode('unicode_escape').decode('utf-8')}")
self.__logger.debug(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
def __find_tags(self, prompt):
add_at = {"start": [], "insertion_point": [[] for x in range(10)], "end": []}
tree = self.__schedule_parser.parse(prompt)
self.__logger.debug(f"Initial tree:\n{tree.pretty()}")
readtree = self.ReadTree(self.__logger, self.__ignore_repeats, self.__join_attention, prompt, add_at)
readtree.visit(tree)
for r in readtree.remove[::-1]:
prompt = prompt[: r[0]] + prompt[r[1] :]
if self.__cleanup:
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
# clean up whitespace and extra separators
prompt = (
prompt.replace(" ", " ")
.replace(self.__separator + self.__separator, self.__separator)
.replace(" " + self.__separator, self.__separator)
.removeprefix(self.__separator)
.removesuffix(self.__separator)
.strip()
)
add_at = readtree.add_at
self.__logger.debug(f"New negative additions: {add_at}")
return prompt, add_at
def __add_to_insertion_points(self, negative_prompt, add_at_insertion_point):
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)
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
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])
else:
ipp = 0
if negative_prompt.startswith(self.__separator):
ipp = len(self.__separator)
add_at_insertion_point[n].append(negative_prompt[ipp:])
negative_prompt = self.__separator.join(add_at_insertion_point[n])
return negative_prompt
def __add_to_start(self, negative_prompt, add_at_start):
if len(negative_prompt) > 0:
ipp = 0
if negative_prompt.startswith(self.__separator):
ipp = len(self.__separator) # adjust for existing end separator
add_at_start.append(negative_prompt[ipp:])
negative_prompt = self.__separator.join(add_at_start)
return negative_prompt
def __add_to_end(self, negative_prompt, add_at_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
add_at_end.insert(0, negative_prompt[:ipl])
negative_prompt = self.__separator.join(add_at_end)
return negative_prompt
-40
View File
@@ -1,40 +0,0 @@
import sys
import copy
import logging
class SendToNegativeLogFactory:
class ColoredFormatter(logging.Formatter):
COLORS = {
"DEBUG": "\033[0;36m", # CYAN
"INFO": "\033[0;32m", # GREEN
"WARNING": "\033[0;33m", # YELLOW
"ERROR": "\033[0;31m", # RED
"CRITICAL": "\033[0;37;41m", # WHITE ON RED
"RESET": "\033[0m", # RESET COLOR
}
def format(self, record):
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']}"
return super().format(colored_record)
def __init__(self):
logsd = logging.getLogger("sd")
stnlog = logging.getLogger("SendToNegative")
stnlog.setLevel(logging.INFO)
stnlog.handlers = logsd.handlers
if not stnlog.handlers:
handler = logging.StreamHandler(sys.stdout)
handler.setFormatter(self.ColoredFormatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s"))
stnlog.addHandler(handler)
self.log = stnlog
else:
self.log = SendToNegativeLogCustomAdapter(stnlog)
class SendToNegativeLogCustomAdapter(logging.LoggerAdapter):
def process(self, msg, kwargs):
return f"[SendToNegative] {msg}", kwargs
+170 -34
View File
@@ -5,16 +5,56 @@ import os
sys.path.insert(1, os.path.join(sys.path[0], ".."))
from sendtonegative import SendToNegative # pylint: disable=import-error
from stnlogging import SendToNegativeLogFactory
from ppp import PromptPostProcessor
from ppp_logging import PromptPostProcessorLogFactory
class TestSendToNegative(unittest.TestCase):
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):
lf = SendToNegativeLogFactory()
self.__log = lf.log
self.__log.setLevel(logging.DEBUG)
self.defstn = SendToNegative(self.__log, separator=", ", ignore_repeats=True, join_attention=True, cleanup=True)
"""
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_extraspaces": True,
"ppp_cup_breaks": True,
}
)
self.defppp = PromptPostProcessor(self, self.__defopts)
def ppp_interrupt(self):
pass # fake interrupt
def process(
self,
@@ -22,9 +62,22 @@ class TestSendToNegative(unittest.TestCase):
negative_prompt,
expected_prompt,
expected_negative_prompt,
stn=None,
ppp=None,
):
the_obj = self.defstn if stn is None else stn
"""
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}'")
self.assertEqual(
@@ -33,7 +86,9 @@ class TestSendToNegative(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 TestSendToNegative(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 TestSendToNegative(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 TestSendToNegative(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 TestSendToNegative(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 TestSendToNegative(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 TestSendToNegative(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 TestSendToNegative(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 TestSendToNegative(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 TestSendToNegative(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 TestSendToNegative(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",
SendToNegative(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)",
"[<!neg1!>] this is a ((test<!!e!neg2!>) (test:2.0): 1.5 )",
"normal quality",
"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 TestSendToNegative(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 TestSendToNegative(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,30 +220,100 @@ class TestSendToNegative(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)",
"this is a (([complex<!neg1!>[one|two<!neg12!>||three|four(<!neg14!>)]|simple<!neg2!>|regular<!neg3!>] 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",
"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_inside_scheduling(self):
def test_inside_scheduling(self): # negtag inside scheduling
self.process(
"this is [abc<!neg1!>:def<!!e!neg2!>:5]",
"this is [abc<!neg1!>:def<!!e!neg2!>: 5 ]",
"normal quality",
"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()