Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e8643680f9 | ||
|
|
3fc61c8dd1 | ||
|
|
77a74a8088 |
@@ -1,45 +1,54 @@
|
||||
# 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 or has it's own syntax expressions). Extensions load by their folder name in alphanumeric order.
|
||||
|
||||
With the ["Dynamic Prompts" extension](https://github.com/adieyal/sd-dynamic-prompts)
|
||||
this happens by default due to default folder names for both extensions. But if
|
||||
this is not the case, you can just rename 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\]
|
||||
* **Extra networks**: \<kind:model...\>
|
||||
* **BREAK**: prompt1 BREAK prompt2
|
||||
* **Composable Diffusion**: prompt1 AND prompt2
|
||||
|
||||
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. Since it should run after other extensions that apply to the prompt, the content should have already been processed by them and there should't be any non recognized syntax anymore.
|
||||
4. It does not create *AND/BREAK* constructs when moving content to the negative prompt.
|
||||
|
||||
## 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 +63,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 +79,41 @@ 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).
|
||||
|
||||
## Notes
|
||||
### Clean up settings
|
||||
|
||||
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.
|
||||
* **Apply in img2img**: check if you want to do this processing in img2img processes.
|
||||
* **Remove empty constructs**: removes attention/scheduling/alternation constructs when they are invalid.
|
||||
* **Remove extra separators**: removes unnecessary separators. This applies to the configured separator and regular commas.
|
||||
* **Clean up around BREAKs**: removes consecutive BREAKs and unnecessary commas and space around them.
|
||||
* **Clean up around ANDs**: removes consecutive ANDs and unnecessary commas and space around them.
|
||||
* **Clean up around extra network tags**: removes spaces around them.
|
||||
* **Remove extra spaces**: removes other unnecessary spaces.
|
||||
|
||||
## Notes on negative tags
|
||||
|
||||
Positional insertion tags have less priority that start/end tags, so even if they are at the start or end of the negative prompt, they will end up inside any start/end (and default position) tags.
|
||||
|
||||
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 +127,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.
|
||||
|
||||
|
||||
@@ -0,0 +1,903 @@
|
||||
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.3.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.cup_ands = getattr(opts, "ppp_cup_ands", True) if opts is not None else True
|
||||
self.cup_extranetworktags = getattr(opts, "ppp_cup_extranetworktags", False) if opts is not None else False
|
||||
|
||||
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: (promptcomp | specialchars)* // BUG: sometimes it chooses specialchars instead of promptcomp and fails!!!
|
||||
// prompt composition with AND
|
||||
promptcomp: promptcomppart ([":" numpar] (/\bAND\b/ promptcomppart [":" numpar])+)?
|
||||
promptcomppart: prompt
|
||||
// prompt scheduling and alternation
|
||||
alternate: "[" alternateoption ("|" alternateoption)+ "]"
|
||||
alternateoption: prompt
|
||||
scheduled: "[" [prompt ":"] prompt ":" numpar "]"
|
||||
// wildcard extension support
|
||||
wildcard: "__" /(?:(?!__)\S)+/ "__"
|
||||
choices: "{" choice ("|" choice)* "}"
|
||||
// we ignore weight and any other parameters in each choice
|
||||
choice: prompt
|
||||
// simple prompts
|
||||
prompt: (emphasized | deemphasized | scheduled | alternate | extranetworktag | negtag | wildcard | choices | plain)*
|
||||
nonegprompt: (emphasized | deemphasized | scheduled | alternate | extranetworktag | wildcard | choices | plain)*
|
||||
// attention modifiers
|
||||
emphasized: "(" prompt [":" numpar] ")"
|
||||
deemphasized: "[" prompt "]"
|
||||
// extra network tags
|
||||
extranetworktag: "<" /(?!!)[^>]+/ ">"
|
||||
// negative tags
|
||||
negtag: "<!" [negtagparameters] nonegprompt "!>"
|
||||
negtagparameters: "!" /s|e|[ip]\d/ "!"
|
||||
// plain text and weights
|
||||
numpar: WHITESPACE* NUMBER WHITESPACE*
|
||||
WHITESPACE: /\s+/
|
||||
plain: /((?!__|\bAND\b)[^\\[\]():|<>!{}]|\\.)+/s
|
||||
specialchars: /[\]():|<>!{}]|\bAND\b/+
|
||||
%import common.SIGNED_NUMBER -> NUMBER
|
||||
""",
|
||||
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", "data", "position"])
|
||||
AccumulatedShell = self.AccumulatedShell
|
||||
self.__shell: list[AccumulatedShell] = []
|
||||
self.NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"])
|
||||
NegTag = self.NegTag
|
||||
self.__negtags: list[NegTag] = []
|
||||
self.__already_processed = []
|
||||
self.add_at = add_at
|
||||
self.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
|
||||
"""
|
||||
treemetaposition = (
|
||||
[tree.meta.start_pos, tree.meta.end_pos] if hasattr(tree, "meta") and not tree.meta.empty else 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", pos, treemetaposition))
|
||||
if before is not None and hasattr(before, "data"):
|
||||
if self.__ppp.debug:
|
||||
before_metaposition = (
|
||||
[before.meta.start_pos, before.meta.end_pos]
|
||||
if hasattr(before, "meta") and not before.meta.empty
|
||||
else "?"
|
||||
)
|
||||
self.__ppp.logger.info(f"Shell scheduled before at {before_metaposition} with position {pos}")
|
||||
self.__shell.append(self.AccumulatedShell("scb", pos, treemetaposition))
|
||||
self.visit(before)
|
||||
self.__shell.pop()
|
||||
if hasattr(after, "data"):
|
||||
if self.__ppp.debug:
|
||||
after_metaposition = (
|
||||
[after.meta.start_pos, after.meta.end_pos]
|
||||
if hasattr(after, "meta") and not after.meta.empty
|
||||
else "?"
|
||||
)
|
||||
self.__ppp.logger.info(f"Shell scheduled after at {after_metaposition} with position {pos}")
|
||||
self.__shell.append(self.AccumulatedShell("sca", pos, treemetaposition))
|
||||
self.visit(after)
|
||||
self.__shell.pop()
|
||||
# self.__shell.pop()
|
||||
|
||||
def alternate(self, 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
|
||||
"""
|
||||
treemetaposition = (
|
||||
[tree.meta.start_pos, tree.meta.end_pos] if hasattr(tree, "meta") and not tree.meta.empty else None
|
||||
)
|
||||
# self.__shell.append(self.AccumulatedShell("al", len(tree.children), treemetaposition))
|
||||
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", {"pos": i + 1, "len": len(tree.children)}, treemetaposition)
|
||||
)
|
||||
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
|
||||
"""
|
||||
treemetaposition = (
|
||||
[tree.meta.start_pos, tree.meta.end_pos] if hasattr(tree, "meta") and not tree.meta.empty else None
|
||||
)
|
||||
numpar = tree.children[-1]
|
||||
weight = self.__get_numpar_value(numpar) if numpar is not None else 1.1
|
||||
if self.__ppp.debug:
|
||||
self.__ppp.logger.info(f"Shell attention at {treemetaposition or '?'} with weight {weight}")
|
||||
self.__shell.append(self.AccumulatedShell("at", weight, treemetaposition))
|
||||
self.visit_children(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
|
||||
treemetaposition = (
|
||||
[tree.meta.start_pos, tree.meta.end_pos] if hasattr(tree, "meta") and not tree.meta.empty else None
|
||||
)
|
||||
if self.__ppp.debug:
|
||||
self.__ppp.logger.info(f"Shell attention at {treemetaposition or '?'} with weight {weight}")
|
||||
self.__shell.append(self.AccumulatedShell("at", weight, treemetaposition))
|
||||
self.visit_children(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
|
||||
"""
|
||||
treemetaposition = (
|
||||
[tree.meta.start_pos, tree.meta.end_pos] if hasattr(tree, "meta") and not tree.meta.empty else 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:
|
||||
self.__ppp.logger.info(
|
||||
f"Negative tag at {treemetaposition or '?'}: {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].data * negtag.shell[i].data)
|
||||
/ 100, # we limit the new weight to two decimals
|
||||
negtag.shell[i - 1].position,
|
||||
)
|
||||
negtag.shell.pop(i)
|
||||
start = ""
|
||||
end = ""
|
||||
for s in negtag.shell:
|
||||
match s.type:
|
||||
case "at":
|
||||
if s.data == 0.9:
|
||||
start += "["
|
||||
end = "]" + end
|
||||
elif s.data == 1.1:
|
||||
start += "("
|
||||
end = ")" + end
|
||||
else:
|
||||
start += "("
|
||||
end = f":{s.data})" + end
|
||||
# case "sc":
|
||||
case "scb":
|
||||
start += "["
|
||||
end = f"::{s.data}]" + end
|
||||
case "sca":
|
||||
start += "["
|
||||
end = f":{s.data}]" + end
|
||||
# case "al":
|
||||
case "alo":
|
||||
start += "[" + ("|" * int(s.data["pos"] - 1))
|
||||
end = ("|" * int(s.data["len"] - s.data["pos"])) + "]" + end
|
||||
content = start + negtag.content + end
|
||||
position = negtag.parameters or "s"
|
||||
if 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)
|
||||
|
||||
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:
|
||||
promptcomp(tree): Replicates prompt composition constructs.
|
||||
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.
|
||||
extranetworktag(tree): Replicates extra network 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 promptcomp(self, tree):
|
||||
r = tree[0]
|
||||
if len(tree) > 1:
|
||||
if tree[1] is not None:
|
||||
r += f":{tree[1]}"
|
||||
for i in range(2, len(tree), 3):
|
||||
if self.__phase == "cleanup" and self.__ppp.cup_ands:
|
||||
r = re.sub(r"[, ]+$", " ", r)
|
||||
if r[-1:].isalnum(): # add space if needed
|
||||
r += " "
|
||||
r += "AND"
|
||||
t = tree[i + 1]
|
||||
if self.__phase == "cleanup" and self.__ppp.cup_ands:
|
||||
t = re.sub(r"^[, ]+", " ", t)
|
||||
if t[0:1].isalnum(): # add space if needed
|
||||
r += " "
|
||||
r += t
|
||||
if tree[i + 2] is not None:
|
||||
r += f":{tree[i+2]}"
|
||||
return r
|
||||
|
||||
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 extranetworktag(self, tree):
|
||||
return f"<{tree[0]}>" # replicate extra network 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":
|
||||
# take care of cleaning the joints only if there are no constructs that can be affected
|
||||
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. This is called for plain text only or when there are no constructs.
|
||||
|
||||
Args:
|
||||
text (str): The text to be cleaned up.
|
||||
|
||||
Returns:
|
||||
str: The cleaned up text.
|
||||
"""
|
||||
# NOTE: we can't use start/end of line regex since the text might only be a part of a larger line due to the parser
|
||||
if self.cup_extraseparators:
|
||||
#
|
||||
# sendtonegative separator
|
||||
#
|
||||
escapedSeparator = re.escape(self.stn_separator)
|
||||
# collapse separators
|
||||
text = re.sub(r"(?:\s*" + escapedSeparator + r"\s*){2,}", self.stn_separator, text)
|
||||
#
|
||||
# regular comma separator
|
||||
#
|
||||
# collapse separators
|
||||
text = re.sub(r"(?:\s*,\s*){2,}", ", ", text)
|
||||
if self.cup_breaks:
|
||||
# collapse separators and commas before BREAK
|
||||
text = re.sub(r"[, ]+BREAK\b", " BREAK", text)
|
||||
# collapse separators and commas after BREAK
|
||||
text = re.sub(r"\bBREAK[, ]+", "BREAK ", text)
|
||||
# collapse separators and commas around BREAK
|
||||
text = re.sub(r"[, ]+BREAK[, ]+", " BREAK ", text)
|
||||
# collapse BREAKs
|
||||
text = re.sub(r"\bBREAK(?:\s+BREAK)+\b", " BREAK ", text)
|
||||
if self.cup_extraspaces:
|
||||
# remove spaces before comma
|
||||
text = re.sub(r"[ ]+,", ",", text)
|
||||
# collapse spaces
|
||||
text = re.sub(r"[ ]{2,}", " ", text)
|
||||
return text
|
||||
|
||||
def trim_text(self, text):
|
||||
"""
|
||||
Trims the given text based on the specified cleanup options. This is only called for the reconstructed prompt.
|
||||
|
||||
Args:
|
||||
text (str): The text to be trimmed.
|
||||
|
||||
Returns:
|
||||
str: The trimmed text.
|
||||
"""
|
||||
# NOTE: here we can only do cleanups that can be done on the whole text, including inside constructs and around them
|
||||
if self.cup_extraseparators:
|
||||
#
|
||||
# sendtonegative separator
|
||||
#
|
||||
escapedSeparator = re.escape(self.stn_separator)
|
||||
# remove duplicate separator after starting parenthesis or bracket
|
||||
text = re.sub(r"(\s*" + escapedSeparator + r"\s*[([])\s*" + escapedSeparator + r"\s*", r"\1", text)
|
||||
# remove before colon or ending parenthesis or bracket
|
||||
text = re.sub(r"\s*" + escapedSeparator + r"\s*([:)\]]\s*" + escapedSeparator + r"\s*)", r"\1", text)
|
||||
# remove at start of prompt or line
|
||||
text = re.sub(r"^(?:\s*" + escapedSeparator + r"\s*)", "", text, flags=re.MULTILINE)
|
||||
# remove at end of prompt or line
|
||||
text = re.sub(r"(?:\s*" + escapedSeparator + r"\s*)$", "", text, flags=re.MULTILINE)
|
||||
#
|
||||
# regular comma separator
|
||||
#
|
||||
# remove duplicate separators after starting parenthesis or bracket
|
||||
text = re.sub(r"(\s*,\s*[([])\s*,\s*", r"\1", text)
|
||||
# remove duplicate separators before colon or ending parenthesis or bracket
|
||||
text = re.sub(r"\s*,\s*([:)\]]\s*,\s*)", r"\1", text)
|
||||
# remove at start of prompt or line
|
||||
text = re.sub(r"^\s*,\s*", "", text, flags=re.MULTILINE)
|
||||
# remove at end of prompt or line
|
||||
text = re.sub(r"\s*,\s*$", "", text, flags=re.MULTILINE)
|
||||
if self.cup_breaks:
|
||||
# remove spaces between start of line and BREAK
|
||||
text = re.sub(r"^[ ]+BREAK\b", "BREAK", text, flags=re.MULTILINE)
|
||||
# remove spaces between BREAK and end of line
|
||||
text = re.sub(r"\bBREAK[ ]+$", "BREAK", text, flags=re.MULTILINE)
|
||||
# remove at start of prompt
|
||||
text = re.sub(r"\ABREAK\b", "", text)
|
||||
# remove at end of prompt
|
||||
text = re.sub(r"\bBREAK\Z", "", text)
|
||||
if self.cup_ands:
|
||||
# collapse ANDs with space after
|
||||
text = re.sub(r"\bAND(?:\s+AND)+\s+", "AND ", text)
|
||||
# collapse ANDs without space after
|
||||
text = re.sub(r"\bAND(?:\s+AND)+\b", "AND", text)
|
||||
# collapse separators and spaces before ANDs
|
||||
text = re.sub(r"[, ]+AND\b", " AND", text)
|
||||
# collapse separators and spaces after ANDs
|
||||
text = re.sub(r"\bAND[, ]+", "AND ", text)
|
||||
# remove at start of prompt
|
||||
text = re.sub(r"\AAND\b", "", text)
|
||||
# remove at end of prompt
|
||||
text = re.sub(r"\bAND\Z", "", text)
|
||||
if self.cup_extranetworktags:
|
||||
#
|
||||
# all cases since we can't find them inside plain text
|
||||
#
|
||||
# remove spaces before <
|
||||
text = re.sub(r"\B\s+<(?!!)", "<", text)
|
||||
# remove spaces after >
|
||||
text = re.sub(r"(?<!!)>\s+\B", ">", text)
|
||||
if self.cup_extraspaces:
|
||||
# remove extra spaces after starting parenthesis or bracket
|
||||
text = re.sub(r"([,\.;\s]+[([])\s+", r"\1", text)
|
||||
# remove extra spaces before ending parenthesis or bracket
|
||||
text = re.sub(r"\s+([)\]][,\.;\s]+)", r"\1", text)
|
||||
# collapse spaces
|
||||
# text = re.sub(r"[ ]{2,}", " ", text)
|
||||
# remove spaces at start and end
|
||||
text = text.strip()
|
||||
return text
|
||||
|
||||
def __findwildcards(self, prompt, negative_prompt):
|
||||
"""
|
||||
Find and process wildcards in the prompt and negative_prompt strings.
|
||||
|
||||
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)
|
||||
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)
|
||||
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
|
||||
if hasattr(self.script, "ppp_interrupt"):
|
||||
self.script.ppp_interrupt()
|
||||
if self.debug:
|
||||
self.logger.info(f"prompt after wildcards: {self.formatOutput(prompt)}")
|
||||
self.logger.info(f"negative_prompt after wildcards: {self.formatOutput(negative_prompt)}")
|
||||
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)}")
|
||||
p_tree = self.__parser_complete.parse(prompt)
|
||||
self.logger.info(f"Tree from prompt:\n{p_tree.pretty()}")
|
||||
self.logger.info(f"Input negative_prompt: {self.formatOutput(negative_prompt)}")
|
||||
np_tree = self.__parser_complete.parse(negative_prompt)
|
||||
self.logger.info(f"Tree from negative prompt:\n{np_tree.pretty()}")
|
||||
|
||||
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
|
||||
or self.cup_ands
|
||||
or self.cup_extranetworktags
|
||||
):
|
||||
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
|
||||
@@ -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
|
||||
@@ -1,79 +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 SendToNegative.NAME
|
||||
|
||||
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]
|
||||
)
|
||||
# make it compatible with A1111 hires fix
|
||||
if hasattr(p, "all_hr_prompts") and hasattr(p, "all_hr_negative_prompts"):
|
||||
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] = stn.process_prompt(
|
||||
p.all_hr_prompts[i], p.all_hr_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,
|
||||
),
|
||||
)
|
||||
@@ -0,0 +1,244 @@
|
||||
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_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,
|
||||
),
|
||||
)
|
||||
shared.opts.add_option(
|
||||
key="ppp_cup_ands",
|
||||
info=shared.OptionInfo(
|
||||
True,
|
||||
label="Clean up around ANDs",
|
||||
section=section,
|
||||
),
|
||||
)
|
||||
shared.opts.add_option(
|
||||
key="ppp_cup_extranetworktags",
|
||||
info=shared.OptionInfo(
|
||||
False,
|
||||
label="Clean up around extra network tags",
|
||||
section=section,
|
||||
),
|
||||
)
|
||||
shared.opts.add_option(
|
||||
key="ppp_cup_extraspaces",
|
||||
info=shared.OptionInfo(
|
||||
True,
|
||||
label="Remove extra spaces",
|
||||
section=section,
|
||||
),
|
||||
)
|
||||
@@ -1,332 +0,0 @@
|
||||
from collections import namedtuple
|
||||
import re
|
||||
import math
|
||||
import lark
|
||||
|
||||
|
||||
class SendToNegative: # pylint: disable=too-few-public-methods
|
||||
NAME = "Send to Negative"
|
||||
VERSION = "2.1.4"
|
||||
|
||||
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 __get_numpar_value(self, numpar):
|
||||
return float(next(x for x in numpar.children if x.type == "NUMBER").value)
|
||||
|
||||
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 = 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"):
|
||||
self.__logger.debug(
|
||||
f"Shell scheduled before at {[before.meta.start_pos,before.meta.end_pos] if hasattr(before,'meta') and not before.meta.empty 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') and not after.meta.empty 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') and not opt.meta.empty 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 = self.__get_numpar_value(numpar) 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') and not tree.meta.empty 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') and not tree.meta.empty 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") 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())
|
||||
)
|
||||
self.__logger.debug(
|
||||
f"Negative tag at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') and not tree.meta.empty 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",
|
||||
math.floor(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
|
||||
@@ -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
|
||||
+192
-112
@@ -5,16 +5,73 @@ 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,
|
||||
"ppp_cup_ands": True,
|
||||
"ppp_cup_extranetworktags": True,
|
||||
}
|
||||
)
|
||||
self.__nocupopts = 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": False,
|
||||
"ppp_cup_extraseparators": False,
|
||||
"ppp_cup_extraspaces": False,
|
||||
"ppp_cup_breaks": False,
|
||||
"ppp_cup_ands": False,
|
||||
"ppp_cup_extranetworktags": False,
|
||||
}
|
||||
)
|
||||
self.defppp = PromptPostProcessor(self, self.__defopts)
|
||||
self.nocupppp = PromptPostProcessor(self, self.__nocupopts)
|
||||
|
||||
def process(
|
||||
self,
|
||||
@@ -22,9 +79,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,120 +103,42 @@ class TestSendToNegative(unittest.TestCase):
|
||||
f"Negative Prompt should be '{expected_negative_prompt}'",
|
||||
)
|
||||
|
||||
def test_tag_default(self):
|
||||
# Send To Negative tests
|
||||
|
||||
def test_nt_simple(self): # negtags with different parameters and separations
|
||||
self.process(
|
||||
"flowers<!red!>",
|
||||
"normal quality, worse quality",
|
||||
"flowers<!red!>, <!!s!green!>, <!!e!blue!><!!p0!yellow!>, <!!p1!purple!><!!p2!black!>",
|
||||
"<!!i0!!>normal quality<!!i1!!>, worse quality<!!i2!!>",
|
||||
"flowers",
|
||||
"red, normal quality, worse quality",
|
||||
"red, green, yellow, normal quality, purple, worse quality, black, blue",
|
||||
)
|
||||
|
||||
def test_tag_start(self):
|
||||
def test_nt_complex(self): # complex negtags
|
||||
self.process(
|
||||
"flowers<!!s!red!>",
|
||||
"normal quality, worse quality",
|
||||
"flowers",
|
||||
"red, normal quality, worse quality",
|
||||
)
|
||||
|
||||
def test_tag_end(self):
|
||||
self.process(
|
||||
"flowers<!!e!red!>",
|
||||
"normal quality, worse quality",
|
||||
"flowers",
|
||||
"normal quality, worse quality, red",
|
||||
)
|
||||
|
||||
def test_tag_insertion_mid_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"normal quality, <!!i0!!>, worse quality",
|
||||
"flowers",
|
||||
"normal quality, red, worse quality",
|
||||
)
|
||||
|
||||
def test_tag_insertion_mid_no_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"normal quality<!!i0!!>worse quality",
|
||||
"flowers",
|
||||
"normal quality, red, worse quality",
|
||||
)
|
||||
|
||||
def test_tag_insertion_start_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"<!!i0!!>, normal quality, worse quality",
|
||||
"flowers",
|
||||
"red, normal quality, worse quality",
|
||||
)
|
||||
|
||||
def test_tag_insertion_start_no_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"<!!i0!!>normal quality, worse quality",
|
||||
"flowers",
|
||||
"red, normal quality, worse quality",
|
||||
)
|
||||
|
||||
def test_tag_insertion_end_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"normal quality, worse quality, <!!i0!!>",
|
||||
"flowers",
|
||||
"normal quality, worse quality, red",
|
||||
)
|
||||
|
||||
def test_tag_insertion_end_no_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"normal quality, worse quality<!!i0!!>",
|
||||
"flowers",
|
||||
"normal quality, worse quality, red",
|
||||
)
|
||||
|
||||
def test_complex(self):
|
||||
self.process(
|
||||
"<!red!> (<!!s!pink!>), flowers <!!e!purple!>, <!!e!blue!>, <!!p0!yellow!> <!!p1!green!>",
|
||||
"<!red!> ((<!!s!pink!>)), flowers <!!e!purple!>, <!!p0!mauve!><!!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",
|
||||
"red, (pink:1.21), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue",
|
||||
)
|
||||
|
||||
def test_complex_no_cleanup(self):
|
||||
def test_nt_complex_nocleanup(self): # complex negtags with no cleanup
|
||||
self.process(
|
||||
"<!red!> (<!!s!pink!>), flowers <!!e!purple!>, <!!e!blue!>, <!!p0!yellow!> <!!p1!green!>",
|
||||
"<!red!> ((<!!s!pink!>)), flowers <!!e!purple!>, <!!p0!mauve!><!!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),
|
||||
" (()), flowers , , ",
|
||||
"red, (pink:1.21), normal quality, mauve, yellow, bad quality, green, worse quality, purple, blue",
|
||||
self.nocupppp,
|
||||
)
|
||||
|
||||
def test_inside_attention1(self):
|
||||
def test_nt_inside_attention(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 ) (red<![square]!>:1.5)",
|
||||
"normal quality",
|
||||
"this is a ((test) (test:2.0): 1.5 )",
|
||||
"[neg1], normal quality, (neg2:1.65)",
|
||||
"this is a ((test) (test:2.0):1.5) (red:1.5)",
|
||||
"[neg1], ([square]:1.5), normal quality, (neg2:1.65)",
|
||||
)
|
||||
|
||||
def test_inside_attention2(self):
|
||||
self.process(
|
||||
"(red<![square]!>:1.5)",
|
||||
"",
|
||||
"(red:1.5)",
|
||||
"([square]:1.5)",
|
||||
)
|
||||
|
||||
def test_inside_alternation1(self):
|
||||
self.process(
|
||||
"this is a (([complex|simple<!neg1!>|regular] test)(test:2.0):1.5)",
|
||||
"normal quality",
|
||||
"this is a (([complex|simple|regular] test)(test:2.0):1.5)",
|
||||
"([|neg1|]:1.65), normal quality",
|
||||
)
|
||||
|
||||
def test_inside_alternation2(self):
|
||||
def test_nt_inside_alternation(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 +146,7 @@ class TestSendToNegative(unittest.TestCase):
|
||||
"([neg1||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality",
|
||||
)
|
||||
|
||||
def test_inside_alternation3(self):
|
||||
def test_nt_inside_alternation_recursive(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 +154,110 @@ class TestSendToNegative(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_nt_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_nt_complex_features(self): # complex negtags with AND, BREAK and other 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]:0.5 AND loraword <lora:xxx:1> AND AND hypernetword <hypernet:yyy>:0.3",
|
||||
"normal quality, <!!i0!!>",
|
||||
"this is: a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK with [abc:def:5] <lora:xxx:1>",
|
||||
"this \\(is\\): a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK with [abc:def:5]:0.5 AND loraword <lora:xxx:1> AND hypernetword <hypernet:yyy>:0.3",
|
||||
"[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]",
|
||||
)
|
||||
|
||||
# Wildcard tests
|
||||
|
||||
def test_wc_ignore(self): # wildcards with ignore option
|
||||
self.process(
|
||||
"__bad_wildcard__",
|
||||
"{option1|option2}",
|
||||
"__bad_wildcard__",
|
||||
"{option1|option2}",
|
||||
PromptPostProcessor(
|
||||
self,
|
||||
DictToObj(
|
||||
{
|
||||
**self.__defopts.__dict__,
|
||||
"ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["ignore"],
|
||||
}
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
def test_wc_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_wc_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_wc_stop(self): # wildcards with stop option
|
||||
self.process(
|
||||
"__bad_wildcard__",
|
||||
"{option1|option2}",
|
||||
PromptPostProcessor.WILDCARD_STOP + "__bad_wildcard__",
|
||||
PromptPostProcessor.WILDCARD_STOP + "{option1|option2}",
|
||||
PromptPostProcessor(
|
||||
self,
|
||||
DictToObj(
|
||||
{
|
||||
**self.__defopts.__dict__,
|
||||
"ppp_gen_ifwildcards": PromptPostProcessor.IFWILDCARDS_CHOICES["stop"],
|
||||
}
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
# Cleanup tests
|
||||
|
||||
def test_cl_simple(self): # simple cleanup
|
||||
self.process(
|
||||
" this is a ((test ), , () [] ( , test ,:2.0):1.5) (red:1.5) ",
|
||||
" normal quality ",
|
||||
"this is a ((test), (test,:2.0):1.5) (red:1.5)",
|
||||
"normal quality",
|
||||
)
|
||||
|
||||
def test_cl_complex(self): # complex cleanup
|
||||
self.process(
|
||||
" this is BREAKABLE a ((test), ,AND AND() [] <lora:test> ANDERSON (test:2.0):1.5) :o BREAK \n BREAK (red:1.5) ",
|
||||
" [:hands, feet, :0.15]normal quality ",
|
||||
"this is BREAKABLE a ((test) AND <lora:test> ANDERSON (test:2.0):1.5) :o BREAK (red:1.5)",
|
||||
"[:hands, feet, :0.15]normal quality",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user