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**
A clear and concise description of what the issue is.
## Prerequisites
**To Reproduce**
Steps to reproduce the behavior
Please answer the following questions for yourself before submitting an issue. **YOU MAY DELETE THE PREREQUISITES SECTION.**
**Expected behavior**
A clear and concise description of what you expected to happen.
- [ ] I am running the latest version
- [ ] 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**
If applicable, add screenshots to help explain your problem.
## Current Behavior
**Commit where the problem happens**
Installed commit or version tag
What is the current behavior?
**Additional information**
Add any other context about the problem here.
## Expected Behavior
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.**
A clear and concise description of what the problem is. Ex. I'm always frustrated when [...]
## Prerequisites
**Describe the solution you'd like**
A clear and concise description of what you want to happen.
Please answer the following questions for yourself before submitting an issue. **YOU MAY DELETE THE PREREQUISITES SECTION.**
**Describe alternatives you've considered**
A clear and concise description of any alternative solutions or features you've considered.
- [ ] I am running the latest version
- [ ] 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.
+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__
.vscode/**/*
!.vscode/settings.json
!.vscode/launch.json
+1 -2
View File
@@ -1,4 +1,5 @@
{
"python.analysis.extraPaths": ["../.."],
"python.testing.unittestArgs": [
"-v",
"-s",
@@ -9,8 +10,6 @@
"python.testing.pytestEnabled": false,
"python.testing.unittestEnabled": true,
"python.analysis.typeCheckingMode": "basic",
"python.linting.pylintEnabled": true,
"python.linting.enabled": true,
"black-formatter.args": [
"--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.
@@ -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
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)
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 ["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
The format of the tags is like this:
@@ -54,11 +77,39 @@ that part to the negative prompt.
## Configuration
The extension settings allow you to change the format of the tag in case there
is some incompatibility with another extension.
Separator used when adding to the negative prompt: You can specify the separator used when adding to the negative prompt (by default it's ", ").
You can also specify the separator added to the negative prompt which by
default is ", ".
Ignore tags with repeated content: by default it ignores repeated content to avoid repetitions in the negative prompt.
By default it ignores repeated content and also tries to clean up the prompt
after removing the tags, but these can also be changed in the settings.
Join attention modifiers (weights) when possible: by default it joins attention modifiers when possible (joins into one, multipliying their values).
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.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
@@ -28,7 +31,7 @@ class SendToNegativeScript(scripts.Script):
return scripts.AlwaysVisible
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
p.all_prompts[i], p.all_negative_prompts[i] = stn.process_prompt(
p.all_prompts[i], p.all_negative_prompts[i]
@@ -36,38 +39,6 @@ class SendToNegativeScript(scripts.Script):
def __on_ui_settings(self):
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(
key="stn_separator",
info=shared.OptionInfo(
@@ -84,6 +55,14 @@ class SendToNegativeScript(scripts.Script):
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(
+217 -92
View File
@@ -1,30 +1,26 @@
from collections import namedtuple
import logging
import re
import lark
class SendToNegative: # pylint: disable=too-few-public-methods
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 = ", "
def __init__(
self,
tag_start=None,
tag_end=None,
tag_param_start=None,
tag_param_end=None,
log,
separator=None,
ignore_repeats=None,
join_attention=None,
cleanup=None,
opts=None,
):
"""
Default format for the tag:
Format for the tag:
<!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.
"""
self.__logger = logging.getLogger(__name__)
str_start = (
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.__logger = log
if opts is not None and 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
)
@@ -82,24 +57,180 @@ class SendToNegative: # pylint: disable=too-few-public-methods
if opts is not None
else self.DEFAULT_SEPARATOR
)
self.__insertion_point_tags = [
(str_start + str_param_start + "i" + str(x) + str_param_end + str_end) for x in range(10)
]
self.__regex = re.compile(
"("
+ escape_sequence
+ re.escape(str_start)
+ "(?:"
+ re.escape(str_param_start)
+ "([se]|(?:[pi][0-9]))"
+ re.escape(str_param_end)
+ ")?(.*?)"
+ escape_sequence
+ re.escape(str_end)
+ ")",
re.S,
self.__insertion_point_tags = [f"<!!i{x}!!>" for x in range(10)]
# Process with lark (debug with https://www.lark-parser.org/ide/)
self.__schedule_parser = lark.Lark(
r"""
start: (prompt | /[\][():|<>!]/+)*
?prompt: (emphasized | deemphasized | scheduled | alternate | modeltag | negtag | plain)*
?nonegprompt: (emphasized | deemphasized | scheduled | alternate | modeltag | plain)*
emphasized: "(" prompt [":" numpar] ")"
deemphasized: "[" prompt "]"
scheduled: "[" [prompt ":"] prompt ":" numpar "]"
alternate: "[" alternateoption ("|" alternateoption)+ "]"
alternateoption: prompt
negtag: "<!" [negtagparameters] nonegprompt "!>"
negtagparameters: "!" /s|e|[ip]\d/ "!"
modeltag: "<" /(?!!)[^>]+/ ">"
numpar: WHITESPACE* NUMBER WHITESPACE*
WHITESPACE: /\s+/
?plain: /([^\\[\]():|<>!]|\\.)+/s
%import common.SIGNED_NUMBER -> NUMBER
""",
propagate_positions=True,
)
class ReadTree(lark.visitors.Interpreter):
def __init__(self, logger, ignorerepeats, joinattention, prompt, add_at):
super().__init__()
self.__logger = logger
self.__ignore_repeats = ignorerepeats
self.__join_attention = joinattention
self.__prompt = prompt
self.AccumulatedShell = namedtuple("AccumulatedShell", ["type", "info1", "info2"])
AccumulatedShell = self.AccumulatedShell
self.__shell: list[AccumulatedShell] = []
self.NegTag = namedtuple("NegTag", ["start", "end", "content", "parameters", "shell"])
NegTag = self.NegTag
self.__negtags: list[NegTag] = []
self.__already_processed = []
self.add_at = add_at
self.remove = []
def scheduled(self, tree):
if len(tree.children) > 2: # before & after
before = tree.children[0]
else:
before = None
after = tree.children[-2]
numpar = tree.children[-1]
pos = float(numpar.children[0].value)
if pos >= 1:
pos = int(pos)
# self.__shell.append(self.AccumulatedShell("sc", tree.meta.start_pos, pos))
if before is not None and hasattr(before, "data"):
self.__logger.debug(
f"Shell scheduled before at {[before.meta.start_pos,before.meta.end_pos] if hasattr(before,'meta') else '?'} : {pos}"
)
self.__shell.append(self.AccumulatedShell("scb", pos, None))
self.visit(before)
self.__shell.pop()
if hasattr(after, "data"):
self.__logger.debug(
f"Shell scheduled after at {[after.meta.start_pos,after.meta.end_pos] if hasattr(after,'meta') else '?'} : {pos}"
)
self.__shell.append(self.AccumulatedShell("sca", pos, None))
self.visit(after)
self.__shell.pop()
# self.__shell.pop()
def alternate(self, tree):
# self.__shell.append(self.AccumulatedShell("al", tree.meta.start_pos, len(tree.children)))
for i, opt in enumerate(tree.children):
self.__logger.debug(
f"Shell alternate at {[opt.meta.start_pos,opt.meta.end_pos] if hasattr(opt,'meta') else '?'} : {i+1}"
)
if hasattr(opt, "data"):
self.__shell.append(self.AccumulatedShell("alo", i + 1, len(tree.children)))
self.visit(opt)
self.__shell.pop()
# self.__shell.pop()
def emphasized(self, tree):
numpar = tree.children[-1]
weight = float(numpar.children[0].value) if numpar is not None else 1.1
self.__logger.debug(
f"Shell attention at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') else '?'}: {weight}"
)
self.__shell.append(self.AccumulatedShell("at", weight, None))
self.visit_children(tree)
self.__shell.pop()
def deemphasized(self, tree):
weight = 0.9
self.__logger.debug(
f"Shell attention at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') else '?'}: {weight}"
)
self.__shell.append(self.AccumulatedShell("at", weight, None))
self.visit_children(tree)
self.__shell.pop()
def negtag(self, tree):
negtagparameters = tree.children[0]
parameters = negtagparameters.children[0].value if negtagparameters is not None else ""
rest = []
for x in tree.children[1::]:
rest.append(self.__prompt[x.meta.start_pos : x.meta.end_pos] if hasattr(x, "meta") else x.value)
content = "".join(rest)
self.__negtags.append(
self.NegTag(tree.meta.start_pos, tree.meta.end_pos, content, parameters, self.__shell.copy())
)
self.__logger.debug(
f"Negative tag at {[tree.meta.start_pos,tree.meta.end_pos] if hasattr(tree,'meta') else '?'}: {parameters}: {content.encode('unicode_escape').decode('utf-8')}"
)
def start(self, tree):
self.visit_children(tree)
# process the found negtags
for nt in self.__negtags:
if self.__join_attention:
# join consecutive attention elements
for i in range(len(nt.shell) - 1, 0, -1):
if nt.shell[i].type == "at" and nt.shell[i - 1].type == "at":
nt.shell[i - 1] = self.AccumulatedShell(
"at",
(100 * nt.shell[i - 1].info1 * nt.shell[i].info1) / 100, # we limit to two decimals
None,
)
nt.shell.pop(i)
start = ""
end = ""
for s in nt.shell:
match s.type:
case "at":
if s.info1 == 0.9:
start += "["
end = "]" + end
elif s.info1 == 1.1:
start += "("
end = ")" + end
else:
start += "("
end = f":{s.info1})" + end
# case "sc":
case "scb":
start += "["
end = f"::{s.info1}]" + end
case "sca":
start += "["
end = f":{s.info1}]" + end
# case "al":
case "alo":
start += "[" + ("|" * int(s.info1 - 1))
end = ("|" * int(s.info2 - s.info1)) + "]" + end
content = start + nt.content + end
position = nt.parameters or "s"
if len(content) > 0:
if content not in self.__already_processed:
if self.__ignore_repeats:
self.__already_processed.append(content)
self.__logger.debug(
f"Adding content at position {position}: {content.encode('unicode_escape').decode('utf-8')}"
)
if position == "e":
self.add_at["end"].append(content)
elif position.startswith("p"):
n = int(position[1])
self.add_at["insertion_point"][n].append(content)
else: # position == "s" or invalid
self.add_at["start"].append(content)
else:
self.__logger.warning(
f"Ignoring repeated content: {content.encode('unicode_escape').decode('utf-8')}"
)
# remove from prompt
self.remove.append([nt.start, nt.end])
def process_prompt(self, original_prompt, original_negative_prompt):
"""
Extract from the prompt the tagged parts and add them to the negative prompt
@@ -107,54 +238,48 @@ class SendToNegative: # pylint: disable=too-few-public-methods
try:
prompt = original_prompt
negative_prompt = original_negative_prompt
self.__logger.debug(f"Input prompt: {prompt}")
self.__logger.debug(f"Input negative_prompt: {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}")
self.__logger.debug(f"Output negative_prompt: {negative_prompt}")
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):
already_processed = []
add_at = {"start": [], "insertion_point": [[] for x in range(10)], "end": []}
# process tags in prompt
matches = self.__regex.findall(prompt)
for match in matches:
position = match[1] or "s"
content = match[2]
if len(content) > 0:
if content not in already_processed:
if self.__ignore_repeats:
already_processed.append(content)
self.__logger.debug(f"Processing content at position {position}: {content}")
if position == "e":
add_at["end"].append(content)
elif position.startswith("p"):
n = int(position[1])
add_at["insertion_point"][n].append(content)
else: # position == "s" or invalid
add_at["start"].append(content)
else:
self.__logger.warning(f"Ignoring repeated content: {content}")
# clean-up
prompt = prompt.replace(match[0], "")
if self.__cleanup:
prompt = (
prompt.replace(" ", " ")
.replace(self.__separator + self.__separator, self.__separator)
.replace(" " + self.__separator, self.__separator)
.removeprefix(self.__separator)
.removesuffix(self.__separator)
.strip()
)
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):
+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 sys
import os
@@ -5,19 +6,15 @@ import os
sys.path.insert(1, os.path.join(sys.path[0], ".."))
from sendtonegative import SendToNegative # pylint: disable=import-error
from stnlogging import SendToNegativeLogFactory
class TestSendToNegative(unittest.TestCase):
def setUp(self):
self.defstn = SendToNegative(
tag_start="<!",
tag_end="!>",
tag_param_start="!",
tag_param_end="!",
separator=", ",
ignore_repeats=True,
cleanup=True,
)
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)
def process(
self,
@@ -110,27 +107,75 @@ class TestSendToNegative(unittest.TestCase):
def test_complex(self):
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",
"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):
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",
" , flowers , , ",
"red, pink, normal quality, yellow, bad quality, green, worse quality, purple, blue",
SendToNegative(
tag_start="<!",
tag_end="!>",
tag_param_start="!",
tag_param_end="!",
separator=", ",
ignore_repeats=True,
cleanup=False,
),
" (), 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),
)
def test_inside_attention1(self):
self.process(
"[<!neg1!>] this is a ((test<!!e!neg2!>) (test:2.0):1.5)",
"normal quality",
"this is a ((test) (test:2.0):1.5)",
"[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]",
)