Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3fc61c8dd1 |
@@ -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.
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user