Compare commits

...
6 Commits
10 changed files with 485 additions and 177 deletions
+33 -12
View File
@@ -6,20 +6,41 @@ assignees: acorderob
--- ---
**Describe the bug** ## Prerequisites
A clear and concise description of what the issue is.
**To Reproduce** Please answer the following questions for yourself before submitting an issue. **YOU MAY DELETE THE PREREQUISITES SECTION.**
Steps to reproduce the behavior
**Expected behavior** - [ ] I am running the latest version
A clear and concise description of what you expected to happen. - [ ] I checked the documentation and found no answer
- [ ] I checked to make sure that this issue has not already been filed
- [ ] I'm reporting the issue to the correct repository (for multi-repository projects)
**Screenshots** ## Current Behavior
If applicable, add screenshots to help explain your problem.
**Commit where the problem happens** What is the current behavior?
Installed commit or version tag
**Additional information** ## Expected Behavior
Add any other context about the problem here.
Please describe the behavior you are expecting
## Failure Information (for bugs)
Please help provide information about the failure if this is a bug. If it is not a bug, please remove the rest of this template.
## Steps to Reproduce
Please provide detailed steps for reproducing the issue.
1. step 1
2. step 2
3. you get it...
## Context
Please provide any relevant information about your setup. This is important in case the issue is not reproducible except for under certain conditions.
- WebUI used and version:
## Failure Logs
Please include any relevant log snippets or files here.
+12 -7
View File
@@ -6,14 +6,19 @@ assignees: acorderob
--- ---
**Is your feature request related to a problem? Please describe.** ## Prerequisites
A clear and concise description of what the problem is. Ex. I'm always frustrated when [...]
**Describe the solution you'd like** Please answer the following questions for yourself before submitting an issue. **YOU MAY DELETE THE PREREQUISITES SECTION.**
A clear and concise description of what you want to happen.
**Describe alternatives you've considered** - [ ] I am running the latest version
A clear and concise description of any alternative solutions or features you've considered. - [ ] I checked the documentation and found no answer
- [ ] I checked to make sure that this issue has not already been filed
- [ ] I'm reporting the issue to the correct repository (for multi-repository projects)
## Description
A clear and concise description of the feature that you want.
## Additional context
**Additional context**
Add any other context or screenshots about the feature request here. Add any other context or screenshots about the feature request here.
+39
View File
@@ -0,0 +1,39 @@
# Pull Request
## Description
Please include a summary of the change and which issue is fixed. Please also include relevant motivation and context. List any dependencies that are required for this change.
Fixes # (issue)
## Type of change
Please delete options that are not relevant.
- [ ] Bug fix (non-breaking change which fixes an issue)
- [ ] New feature (non-breaking change which adds functionality)
- [ ] Breaking change (fix or feature that would cause existing functionality to not work as expected)
- [ ] This change requires a documentation update
## How Has This Been Tested?
Please describe the tests that you ran to verify your changes. Provide instructions so we can reproduce. Please also list any relevant details for your test configuration
- [ ] Test A
- [ ] Test B
**Test Configuration**:
- WebUI used and version:
## Checklist
- [ ] My code follows the style guidelines of this project
- [ ] I have performed a self-review of my own code
- [ ] I have commented my code, particularly in hard-to-understand areas
- [ ] I have made corresponding changes to the documentation
- [ ] My changes generate no new warnings
- [ ] I have added tests that prove my fix is effective or that my feature works
- [ ] New and existing unit tests pass locally with my changes
- [ ] Any dependent changes have been merged and published in downstream modules
- [ ] I have checked my code and corrected any misspellings
+4
View File
@@ -1 +1,5 @@
**/__pycache__ **/__pycache__
.vscode/**/*
!.vscode/settings.json
!.vscode/launch.json
+1 -2
View File
@@ -1,4 +1,5 @@
{ {
"python.analysis.extraPaths": ["../.."],
"python.testing.unittestArgs": [ "python.testing.unittestArgs": [
"-v", "-v",
"-s", "-s",
@@ -9,8 +10,6 @@
"python.testing.pytestEnabled": false, "python.testing.pytestEnabled": false,
"python.testing.unittestEnabled": true, "python.testing.unittestEnabled": true,
"python.analysis.typeCheckingMode": "basic", "python.analysis.typeCheckingMode": "basic",
"python.linting.pylintEnabled": true,
"python.linting.enabled": true,
"black-formatter.args": [ "black-formatter.args": [
"--line-length=120" "--line-length=120"
] ]
+59 -8
View File
@@ -1,4 +1,4 @@
# Send To Negative for Stable Diffusion WebUI # Send to Negative for Stable Diffusion WebUI
Extension for the [AUTOMATIC1111 Stable Diffusion WebUI](https://github.com/AUTOMATIC1111/stable-diffusion-webui) or compatible UIs. Extension for the [AUTOMATIC1111 Stable Diffusion WebUI](https://github.com/AUTOMATIC1111/stable-diffusion-webui) or compatible UIs.
@@ -8,13 +8,36 @@ 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 negative prompt. This allows useful tricks when using a wildcard extension
since you can add negative content from choices made in the positive prompt. since you can add negative content from choices made in the positive prompt.
Note: The extension must be loaded after the wildcard extension. Note: The extension must be loaded after the installed wildcards extension. Extensions
load by their folder in alphanumeric order.
With the ["Dynamic Prompts" extension](https://github.com/adieyal/sd-dynamic-prompts) 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 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 this is not the case, you can just rename the extension folder so the ordering
works out. 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".
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:
* Attention: \[prompt\] (prompt) (prompt:weight)
* Alternation: \[prompt1|prompt2|...\]
* Scheduling: \[prompt1:prompt2:step\]
In SD.Next that means only the A1111 or Full parsers.
It does not build AND/BREAK separations into the negative prompt.
## Installation
1. Go to Extensions > Install from URL
2. Paste <https://github.com/acorderob/sd-webui-sendtonegative> in the URL for extension's git repository text field
3. Click the Install button
4. Restart the webui
## Usage ## Usage
The format of the tags is like this: The format of the tags is like this:
@@ -54,11 +77,39 @@ that part to the negative prompt.
## Configuration ## Configuration
The extension settings allow you to change the format of the tag in case there Separator used when adding to the negative prompt: You can specify the separator used when adding to the negative prompt (by default it's ", ").
is some incompatibility with another extension.
You can also specify the separator added to the negative prompt which by Ignore tags with repeated content: by default it ignores repeated content to avoid repetitions in the negative prompt.
default is ", ".
By default it ignores repeated content and also tries to clean up the prompt Join attention modifiers (weights) when possible: by default it joins attention modifiers when possible (joins into one, multipliying their values).
after removing the tags, but these can also be changed in the settings.
Try to clean-up the prompt after processing: by default cleans up the positive prompt after processing, removing extra spaces and separators.
## Notes
The content of the negative tags is not processed and is copied as is to the negative prompt. Other modifiers around the tags are processed in the following way.
### Attention modifiers (weights)
They will be translated to the negative prompt. For example:
* `(red<!square!>:1.5)` will end up as `(square:1.5)` in the negative prompt
* `(red[<!square!>]:1.5)` will end up as `(square:1.35)` in the negative prompt (weight=1.5*0.9)
* However `(red<![square]!>:1.5)` will end up as `([square]:1.5)` in the negative prompt. The content of the negative tag is copied as is, and not joined with the surrounding modifier.
### Prompt editing constructs (alternation and scheduling)
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]`
This should still work as intended, and the only negative point i see is the unnecessary separators.
## License
MIT
## Contact
If you have any questions or concerns, please leave an issue, or start a thread in the discussions.
+12 -33
View File
@@ -13,11 +13,14 @@ from modules import scripts, shared, script_callbacks
from modules.processing import StableDiffusionProcessing from modules.processing import StableDiffusionProcessing
from modules.shared import opts from modules.shared import opts
from sendtonegative import SendToNegative from sendtonegative import SendToNegative
from stnlogging import SendToNegativeLogFactory
class SendToNegativeScript(scripts.Script): class SendToNegativeScript(scripts.Script):
def __init__(self): def __init__(self):
if not hasattr(self, "callbacks_added"): if not hasattr(self, "callbacks_added"):
lf = SendToNegativeLogFactory()
self.__logstn = lf.log
script_callbacks.on_ui_settings(self.__on_ui_settings) script_callbacks.on_ui_settings(self.__on_ui_settings)
self.callbacks_added = True self.callbacks_added = True
@@ -28,7 +31,7 @@ class SendToNegativeScript(scripts.Script):
return scripts.AlwaysVisible return scripts.AlwaysVisible
def process(self, p: StableDiffusionProcessing, *args, **kwargs): def process(self, p: StableDiffusionProcessing, *args, **kwargs):
stn = SendToNegative(opts=opts) stn = SendToNegative(self.__logstn, opts=opts)
for i in range(len(p.all_prompts)): # pylint: disable=consider-using-enumerate 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] = stn.process_prompt(
p.all_prompts[i], p.all_negative_prompts[i] p.all_prompts[i], p.all_negative_prompts[i]
@@ -36,38 +39,6 @@ class SendToNegativeScript(scripts.Script):
def __on_ui_settings(self): def __on_ui_settings(self):
section = ("send-to-negative", SendToNegative.NAME) section = ("send-to-negative", SendToNegative.NAME)
shared.opts.add_option(
key="stn_tagstart",
info=shared.OptionInfo(
SendToNegative.DEFAULT_TAG_START,
label="Tag start",
section=section,
),
)
shared.opts.add_option(
key="stn_tagend",
info=shared.OptionInfo(
SendToNegative.DEFAULT_TAG_END,
label="Tag end",
section=section,
),
)
shared.opts.add_option(
key="stn_tagparamstart",
info=shared.OptionInfo(
SendToNegative.DEFAULT_TAG_PARAM_START,
label="Tag parameter start",
section=section,
),
)
shared.opts.add_option(
key="stn_tagparamend",
info=shared.OptionInfo(
SendToNegative.DEFAULT_TAG_PARAM_END,
label="Tag parameter end",
section=section,
),
)
shared.opts.add_option( shared.opts.add_option(
key="stn_separator", key="stn_separator",
info=shared.OptionInfo( info=shared.OptionInfo(
@@ -84,6 +55,14 @@ class SendToNegativeScript(scripts.Script):
section=section, 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( shared.opts.add_option(
key="stn_cleanup", key="stn_cleanup",
info=shared.OptionInfo( info=shared.OptionInfo(
+217 -92
View File
@@ -1,30 +1,26 @@
from collections import namedtuple
import logging import logging
import re import re
import lark
class SendToNegative: # pylint: disable=too-few-public-methods class SendToNegative: # pylint: disable=too-few-public-methods
NAME = "Send to Negative" NAME = "Send to Negative"
VERSION = "1.1" VERSION = "2.1"
DEFAULT_TAG_START = "<!"
DEFAULT_TAG_END = "!>"
DEFAULT_TAG_PARAM_START = "!"
DEFAULT_TAG_PARAM_END = "!"
DEFAULT_SEPARATOR = ", " DEFAULT_SEPARATOR = ", "
def __init__( def __init__(
self, self,
tag_start=None, log,
tag_end=None,
tag_param_start=None,
tag_param_end=None,
separator=None, separator=None,
ignore_repeats=None, ignore_repeats=None,
join_attention=None,
cleanup=None, cleanup=None,
opts=None, opts=None,
): ):
""" """
Default format for the tag: Format for the tag:
<!content!> <!content!>
<!!x!content!> <!!x!content!>
@@ -38,40 +34,19 @@ class SendToNegative: # pylint: disable=too-few-public-methods
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. 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 = logging.getLogger(__name__) self.__logger = log
if opts is not None and opts.prompt_attention == "Compel parser":
str_start = ( self.__logger.warning("Compel parser is not supported!")
tag_start
if tag_start is not None
else getattr(opts, "stn_tagstart", self.DEFAULT_TAG_START)
if opts is not None
else self.DEFAULT_TAG_START
)
str_end = (
tag_end
if tag_end is not None
else getattr(opts, "stn_tagend", self.DEFAULT_TAG_END)
if opts is not None
else self.DEFAULT_TAG_END
)
str_param_start = (
tag_param_start
if tag_param_start is not None
else getattr(opts, "stn_tagparamstart", self.DEFAULT_TAG_PARAM_START)
if opts is not None
else self.DEFAULT_TAG_PARAM_START
)
str_param_end = (
tag_param_end
if tag_param_end is not None
else getattr(opts, "stn_tagparamend", self.DEFAULT_TAG_PARAM_END)
if opts is not None
else self.DEFAULT_TAG_PARAM_END
)
escape_sequence = r"(?<!\\)"
self.__ignore_repeats = ( self.__ignore_repeats = (
ignore_repeats if ignore_repeats is not None else getattr(opts, "stn_ignorerepeats", True) 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 = ( self.__cleanup = (
cleanup if cleanup is not None else getattr(opts, "stn_cleanup", True) if opts is not None else True cleanup if cleanup is not None else getattr(opts, "stn_cleanup", True) if opts is not None else True
) )
@@ -82,24 +57,180 @@ class SendToNegative: # pylint: disable=too-few-public-methods
if opts is not None if opts is not None
else self.DEFAULT_SEPARATOR else self.DEFAULT_SEPARATOR
) )
self.__insertion_point_tags = [ self.__insertion_point_tags = [f"<!!i{x}!!>" for x in range(10)]
(str_start + str_param_start + "i" + str(x) + str_param_end + str_end) for x in range(10) # Process with lark (debug with https://www.lark-parser.org/ide/)
] self.__schedule_parser = lark.Lark(
self.__regex = re.compile( r"""
"(" start: (prompt | /[\][():|<>!]/+)*
+ escape_sequence ?prompt: (emphasized | deemphasized | scheduled | alternate | modeltag | negtag | plain)*
+ re.escape(str_start) ?nonegprompt: (emphasized | deemphasized | scheduled | alternate | modeltag | plain)*
+ "(?:" emphasized: "(" prompt [":" numpar] ")"
+ re.escape(str_param_start) deemphasized: "[" prompt "]"
+ "([se]|(?:[pi][0-9]))" scheduled: "[" [prompt ":"] prompt ":" numpar "]"
+ re.escape(str_param_end) alternate: "[" alternateoption ("|" alternateoption)+ "]"
+ ")?(.*?)" alternateoption: prompt
+ escape_sequence negtag: "<!" [negtagparameters] nonegprompt "!>"
+ re.escape(str_end) negtagparameters: "!" /s|e|[ip]\d/ "!"
+ ")", modeltag: "<" /(?!!)[^>]+/ ">"
re.S, numpar: WHITESPACE* NUMBER WHITESPACE*
WHITESPACE: /\s+/
?plain: /([^\\[\]():|<>!]|\\.)+/s
%import common.SIGNED_NUMBER -> NUMBER
""",
propagate_positions=True,
) )
class ReadTree(lark.visitors.Interpreter):
def __init__(self, logger, ignorerepeats, joinattention, prompt, add_at):
super().__init__()
self.__logger = logger
self.__ignore_repeats = ignorerepeats
self.__join_attention = joinattention
self.__prompt = prompt
self.AccumulatedShell = namedtuple("AccumulatedShell", ["type", "info1", "info2"])
AccumulatedShell = self.AccumulatedShell
self.__shell: list[AccumulatedShell] = []
self.NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"])
NegTag = self.NegTag
self.__negtags: list[NegTag] = []
self.__already_processed = []
self.add_at = add_at
self.remove = []
def scheduled(self, tree):
if len(tree.children) > 2: # before & after
before = tree.children[0]
else:
before = None
after = tree.children[-2]
numpar = tree.children[-1]
pos = float(numpar.children[0].value)
if pos >= 1:
pos = int(pos)
# self.__shell.append(self.AccumulatedShell("sc", tree.meta.start_pos, pos))
if before is not None and hasattr(before, "data"):
self.__logger.debug(
f"Shell scheduled before at {[before.meta.start_pos,before.meta.end_pos] if hasattr(before,'meta') else '?'} : {pos}"
)
self.__shell.append(self.AccumulatedShell("scb", pos, None))
self.visit(before)
self.__shell.pop()
if hasattr(after, "data"):
self.__logger.debug(
f"Shell scheduled after at {[after.meta.start_pos,after.meta.end_pos] if hasattr(after,'meta') else '?'} : {pos}"
)
self.__shell.append(self.AccumulatedShell("sca", pos, None))
self.visit(after)
self.__shell.pop()
# self.__shell.pop()
def alternate(self, tree):
# self.__shell.append(self.AccumulatedShell("al", tree.meta.start_pos, len(tree.children)))
for i, opt in enumerate(tree.children):
self.__logger.debug(
f"Shell alternate at {[opt.meta.start_pos,opt.meta.end_pos] if hasattr(opt,'meta') else '?'} : {i+1}"
)
if hasattr(opt, "data"):
self.__shell.append(self.AccumulatedShell("alo", i + 1, len(tree.children)))
self.visit(opt)
self.__shell.pop()
# self.__shell.pop()
def emphasized(self, tree):
numpar = tree.children[-1]
weight = float(numpar.children[0].value) if numpar is not None else 1.1
self.__logger.debug(
f"Shell attention at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') else '?'}: {weight}"
)
self.__shell.append(self.AccumulatedShell("at", weight, None))
self.visit_children(tree)
self.__shell.pop()
def deemphasized(self, tree):
weight = 0.9
self.__logger.debug(
f"Shell attention at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') else '?'}: {weight}"
)
self.__shell.append(self.AccumulatedShell("at", weight, None))
self.visit_children(tree)
self.__shell.pop()
def negtag(self, tree):
negtagparameters = tree.children[0]
parameters = negtagparameters.children[0].value if negtagparameters is not None else ""
rest = []
for x in tree.children[1::]:
rest.append(self.__prompt[x.meta.start_pos : x.meta.end_pos] if hasattr(x, "meta") else x.value)
content = "".join(rest)
self.__negtags.append(
self.NegTag(tree.meta.start_pos, tree.meta.end_pos, content, parameters, self.__shell.copy())
)
self.__logger.debug(
f"Negative tag at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') else '?'}: {parameters}: {content.encode('unicode_escape').decode('utf-8')}"
)
def start(self, tree):
self.visit_children(tree)
# process the found negtags
for nt in self.__negtags:
if self.__join_attention:
# join consecutive attention elements
for i in range(len(nt.shell) - 1, 0, -1):
if nt.shell[i].type == "at" and nt.shell[i - 1].type == "at":
nt.shell[i - 1] = self.AccumulatedShell(
"at",
(100 * nt.shell[i - 1].info1 * nt.shell[i].info1) / 100, # we limit to two decimals
None,
)
nt.shell.pop(i)
start = ""
end = ""
for s in nt.shell:
match s.type:
case "at":
if s.info1 == 0.9:
start += "["
end = "]" + end
elif s.info1 == 1.1:
start += "("
end = ")" + end
else:
start += "("
end = f":{s.info1})" + end
# case "sc":
case "scb":
start += "["
end = f"::{s.info1}]" + end
case "sca":
start += "["
end = f":{s.info1}]" + end
# case "al":
case "alo":
start += "[" + ("|" * int(s.info1 - 1))
end = ("|" * int(s.info2 - s.info1)) + "]" + end
content = start + nt.content + end
position = nt.parameters or "s"
if len(content) > 0:
if content not in self.__already_processed:
if self.__ignore_repeats:
self.__already_processed.append(content)
self.__logger.debug(
f"Adding content at position {position}: {content.encode('unicode_escape').decode('utf-8')}"
)
if position == "e":
self.add_at["end"].append(content)
elif position.startswith("p"):
n = int(position[1])
self.add_at["insertion_point"][n].append(content)
else: # position == "s" or invalid
self.add_at["start"].append(content)
else:
self.__logger.warning(
f"Ignoring repeated content: {content.encode('unicode_escape').decode('utf-8')}"
)
# remove from prompt
self.remove.append([nt.start, nt.end])
def process_prompt(self, original_prompt, original_negative_prompt): def process_prompt(self, original_prompt, original_negative_prompt):
""" """
Extract from the prompt the tagged parts and add them to the negative prompt Extract from the prompt the tagged parts and add them to the negative prompt
@@ -107,54 +238,48 @@ class SendToNegative: # pylint: disable=too-few-public-methods
try: try:
prompt = original_prompt prompt = original_prompt
negative_prompt = original_negative_prompt negative_prompt = original_negative_prompt
self.__logger.debug(f"Input prompt: {prompt}") self.__logger.debug(f"Input prompt: {prompt.encode('unicode_escape').decode('utf-8')}")
self.__logger.debug(f"Input negative_prompt: {negative_prompt}") self.__logger.debug(f"Input negative_prompt: {negative_prompt.encode('unicode_escape').decode('utf-8')}")
prompt, add_at = self.__find_tags(prompt) prompt, add_at = self.__find_tags(prompt)
negative_prompt = self.__add_to_insertion_points(negative_prompt, add_at["insertion_point"]) negative_prompt = self.__add_to_insertion_points(negative_prompt, add_at["insertion_point"])
if len(add_at["start"]) > 0: if len(add_at["start"]) > 0:
negative_prompt = self.__add_to_start(negative_prompt, add_at["start"]) negative_prompt = self.__add_to_start(negative_prompt, add_at["start"])
if len(add_at["end"]) > 0: if len(add_at["end"]) > 0:
negative_prompt = self.__add_to_end(negative_prompt, add_at["end"]) negative_prompt = self.__add_to_end(negative_prompt, add_at["end"])
self.__logger.debug(f"Output prompt: {prompt}") self.__logger.debug(f"Output prompt: {prompt.encode('unicode_escape').decode('utf-8')}")
self.__logger.debug(f"Output negative_prompt: {negative_prompt}") self.__logger.debug(f"Output negative_prompt: {negative_prompt.encode('unicode_escape').decode('utf-8')}")
return prompt, negative_prompt return prompt, negative_prompt
except Exception as e: # pylint: disable=broad-exception-caught except Exception as e: # pylint: disable=broad-exception-caught
self.__logger.exception(e) self.__logger.exception(e)
return original_prompt, original_negative_prompt return original_prompt, original_negative_prompt
def __find_tags(self, prompt): def __find_tags(self, prompt):
already_processed = []
add_at = {"start": [], "insertion_point": [[] for x in range(10)], "end": []} add_at = {"start": [], "insertion_point": [[] for x in range(10)], "end": []}
# process tags in prompt tree = self.__schedule_parser.parse(prompt)
matches = self.__regex.findall(prompt) self.__logger.debug(f"Initial tree:\n{tree.pretty()}")
for match in matches:
position = match[1] or "s" readtree = self.ReadTree(self.__logger, self.__ignore_repeats, self.__join_attention, prompt, add_at)
content = match[2] readtree.visit(tree)
if len(content) > 0:
if content not in already_processed: for r in readtree.remove[::-1]:
if self.__ignore_repeats: prompt = prompt[: r[0]] + prompt[r[1] :]
already_processed.append(content) if self.__cleanup:
self.__logger.debug(f"Processing content at position {position}: {content}") prompt = re.sub(r"\((?::[+-]?[\d\.]+)?\)", "", prompt) # clean up empty attention
if position == "e": prompt = re.sub(r"\[\]", "", prompt) # clean up empty attention
add_at["end"].append(content) prompt = re.sub(r"\[:?:[+-]?[\d\.]+\]", "", prompt) # clean up empty scheduling
elif position.startswith("p"): prompt = re.sub(r"\[\|+\]", "", prompt) # clean up empty alternation
n = int(position[1]) # clean up whitespace and extra separators
add_at["insertion_point"][n].append(content) prompt = (
else: # position == "s" or invalid prompt.replace(" ", " ")
add_at["start"].append(content) .replace(self.__separator + self.__separator, self.__separator)
else: .replace(" " + self.__separator, self.__separator)
self.__logger.warning(f"Ignoring repeated content: {content}") .removeprefix(self.__separator)
# clean-up .removesuffix(self.__separator)
prompt = prompt.replace(match[0], "") .strip()
if self.__cleanup: )
prompt = ( add_at = readtree.add_at
prompt.replace(" ", " ") self.__logger.debug(f"New negative additions: {add_at}")
.replace(self.__separator + self.__separator, self.__separator)
.replace(" " + self.__separator, self.__separator)
.removeprefix(self.__separator)
.removesuffix(self.__separator)
.strip()
)
return prompt, add_at return prompt, add_at
def __add_to_insertion_points(self, negative_prompt, add_at_insertion_point): def __add_to_insertion_points(self, negative_prompt, add_at_insertion_point):
+40
View File
@@ -0,0 +1,40 @@
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
+68 -23
View File
@@ -1,3 +1,4 @@
import logging
import unittest import unittest
import sys import sys
import os import os
@@ -5,19 +6,15 @@ import os
sys.path.insert(1, os.path.join(sys.path[0], "..")) sys.path.insert(1, os.path.join(sys.path[0], ".."))
from sendtonegative import SendToNegative # pylint: disable=import-error from sendtonegative import SendToNegative # pylint: disable=import-error
from stnlogging import SendToNegativeLogFactory
class TestSendToNegative(unittest.TestCase): class TestSendToNegative(unittest.TestCase):
def setUp(self): def setUp(self):
self.defstn = SendToNegative( lf = SendToNegativeLogFactory()
tag_start="<!", self.__log = lf.log
tag_end="!>", self.__log.setLevel(logging.DEBUG)
tag_param_start="!", self.defstn = SendToNegative(self.__log, separator=", ", ignore_repeats=True, join_attention=True, cleanup=True)
tag_param_end="!",
separator=", ",
ignore_repeats=True,
cleanup=True,
)
def process( def process(
self, self,
@@ -110,27 +107,75 @@ class TestSendToNegative(unittest.TestCase):
def test_complex(self): def test_complex(self):
self.process( self.process(
"<!red!> <!!s!pink!>, flowers <!!e!purple!>, <!!e!blue!>, <!!p0!yellow!> <!!p1!green!>", "<!red!> (<!!s!pink!>), flowers <!!e!purple!>, <!!e!blue!>, <!!p0!yellow!> <!!p1!green!>",
"normal quality, <!!i0!!>, bad quality<!!i1!!>, worse quality", "normal quality, <!!i0!!>, bad quality<!!i1!!>, worse quality",
"flowers", "flowers",
"red, pink, normal quality, yellow, bad quality, green, worse quality, purple, blue", "red, (pink), normal quality, yellow, bad quality, green, worse quality, purple, blue",
) )
def test_complex_no_cleanup(self): def test_complex_no_cleanup(self):
self.process( self.process(
"<!red!> <!!s!pink!>, flowers <!!e!purple!>, <!!e!blue!>, <!!p0!yellow!> <!!p1!green!>", "<!red!> (<!!s!pink!>), flowers <!!e!purple!>, <!!e!blue!>, <!!p0!yellow!> <!!p1!green!>",
"normal quality, <!!i0!!>, bad quality<!!i1!!>, worse quality", "normal quality, <!!i0!!>, bad quality<!!i1!!>, worse quality",
" , flowers , , ", " (), flowers , , ",
"red, pink, normal quality, yellow, bad quality, green, worse quality, purple, blue", "red, (pink), normal quality, yellow, bad quality, green, worse quality, purple, blue",
SendToNegative( SendToNegative(self.__log, separator=", ", ignore_repeats=True, join_attention=True, cleanup=False),
tag_start="<!", )
tag_end="!>",
tag_param_start="!", def test_inside_attention1(self):
tag_param_end="!", self.process(
separator=", ", "[<!neg1!>] this is a ((test<!!e!neg2!>) (test:2.0):1.5)",
ignore_repeats=True, "normal quality",
cleanup=False, "this is a ((test) (test:2.0):1.5)",
), "[neg1], 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):
self.process(
"this is a (([complex<!neg1!>|simple<!neg2!>|regular<!neg3!>] test)(test:2.0):1.5)",
"normal quality",
"this is a (([complex|simple|regular] test)(test:2.0):1.5)",
"([neg1||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality",
)
def test_inside_alternation3(self):
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",
"this is a (([complex[one|two|three|four]|simple|regular] test)(test:2.0):1.5)",
"([neg1||]:1.65), ([[|neg12||]||]:1.65), ([[|||(neg14)]||]:1.65), ([|neg2|]:1.65), ([||neg3]:1.65), normal quality",
)
def test_inside_scheduling(self):
self.process(
"this is [abc<!neg1!>:def<!!e!neg2!>:5]",
"normal quality",
"this is [abc:def:5]",
"[neg1::5], normal quality, [neg2:5]",
)
def test_complex_features(self):
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>",
"normal quality, <!!i0!!>",
"this is: a (([complex|simple|regular] test)(test:2.0):1.5) \nBREAK with [abc:def:5] <lora:xxx:1>",
"[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]",
) )