refactoring
added tests
This commit is contained in:
Antonio Cordero Balcazar
2023-05-13 15:10:14 +02:00
parent f2159d3df8
commit c2ced311ba
4 changed files with 417 additions and 130 deletions
+16
View File
@@ -0,0 +1,16 @@
{
// Use IntelliSense to learn about possible attributes.
// Hover to view descriptions of existing attributes.
// For more information, visit: https://go.microsoft.com/fwlink/?linkid=830387
"version": "0.2.0",
"configurations": [
{
"name": "Tests",
"type": "python",
"request": "launch",
"program": "scripts/tests.py",
"console": "integratedTerminal",
"justMyCode": true
}
]
}
+97
View File
@@ -0,0 +1,97 @@
import logging
from modules import scripts, shared, script_callbacks
from modules.processing import StableDiffusionProcessing
from modules.shared import opts
from sendtonegative import SendToNegative
if __name__ == "__main__":
raise SystemExit("This script must be run from a Stable Diffusion WebUI")
__all__ = ["SendToNegativeScript"]
class SendToNegativeScript(scripts.Script):
NAME = "Send to Negative"
VERSION = "0.3"
def __init__(self):
self.logger = logging.getLogger(__name__)
self.logger.setLevel(logging.INFO)
if getattr(opts, "is_debug", False):
self.logger.setLevel(logging.DEBUG)
if not hasattr(self, "callbacks_added"):
script_callbacks.on_ui_settings(on_ui_settings)
self.callbacks_added = True
def title(self):
return f"{self.NAME} v{self.VERSION}"
def show(self, is_img2img):
return scripts.AlwaysVisible
def process(self, p: StableDiffusionProcessing, *args, **kwargs):
stn = SendToNegative(opts=opts, logger=self.logger)
for i in range(len(p.all_prompts)):
p.all_prompts[i], p.all_negative_prompts[i] = stn.processPrompts(
p.all_prompts[i], p.all_negative_prompts[i]
)
def on_ui_settings():
section = ("send-to-negative", SendToNegativeScript.NAME)
shared.opts.add_option(
key="stn_tagstart",
info=shared.OptionInfo(
"<!",
label="Tag start",
section=section,
),
)
shared.opts.add_option(
key="stn_tagend",
info=shared.OptionInfo(
"!>",
label="Tag end",
section=section,
),
)
shared.opts.add_option(
key="stn_tagparamstart",
info=shared.OptionInfo(
"!",
label="Tag parameter start",
section=section,
),
)
shared.opts.add_option(
key="stn_tagparamend",
info=shared.OptionInfo(
"!",
label="Tag parameter end",
section=section,
),
)
shared.opts.add_option(
key="stn_separator",
info=shared.OptionInfo(
", ",
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_cleanup",
info=shared.OptionInfo(
True,
label="Try to clean-up the prompt after processing",
section=section,
),
)
+168 -130
View File
@@ -1,140 +1,178 @@
import logging
from modules.processing import StableDiffusionProcessing
import modules.scripts as scripts
import re
VERSION = "0.2"
# TODO : support marking in the negative prompt to move to the positive
# TODO : add separator and ignoreRepeats to ui
# TODO : add tests
# TODO : replace regex with proper parsing to detect recursion
logger = logging.getLogger(__name__)
logger.setLevel(logging.DEBUG)
strStart = "<!"
strEnd = "!>"
strParamStart = "!"
strParamEnd = "!"
"""
Default format: <!content!> or <!!x!content!>
with x = s (start position) e (end position) p (specified position) i (insertion point)
both p and i have to be followed by a number 0-9
the insertion point does not accept content
"""
escapeSequence = r"(?<!\\)"
ignoreRepeats = True
separator = ", "
insertionPointTags = [
(strStart + strParamStart + "i" + str(x) + strParamEnd + strEnd) for x in range(10)
]
find = (
"("
+ escapeSequence
+ re.escape(strStart)
+ "(?:"
+ re.escape(strParamStart)
+ "([se]|(?:[pi][0-9]))"
+ re.escape(strParamEnd)
+ ")?(.*?)"
+ escapeSequence
+ re.escape(strEnd)
+ ")"
)
regex = re.compile(find, re.S)
def processPrompts(original_prompt, original_negative_prompt):
class SendToNegative:
"""
Extract from the prompt the marked parts and add them to the negative prompt
This class contains the code independent of the webui interface
"""
try:
prompt = original_prompt
negative_prompt = original_negative_prompt
insertionPointPositions = [negative_prompt.find(x) for x in insertionPointTags]
alreadyProcessed = []
addAtStart = []
addAtEnd = []
addAtInsertionPoint = [[] for x in range(10)]
matches = regex.findall(prompt)
for match in matches:
position = match[1] or "s"
content = match[2]
if len(content) > 0:
if content not in alreadyProcessed:
if ignoreRepeats:
alreadyProcessed.append(content)
logger.debug("Processing content: %s", content)
if position is "e":
addAtEnd.append(content)
elif position.startswith("p"):
n = int(position[1])
addAtInsertionPoint[n].append(content)
else: # position is "s" or invalid
addAtStart.append(content)
def __init__(
self,
tagStart=None,
tagEnd=None,
tagParamStart=None,
tagParamEnd=None,
separator=None,
ignoreRepeats=None,
cleanup=None,
opts=None,
logger=None,
):
"""
Default format:
<!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 - marks the position of insertion point N. Used only in the negative prompt and does not accept content. N can be 0 to 9.
The tags will be removed from the prompt or negative prompt without considering neighboring whitespace or separators.
"""
if logger is None:
self.logger = logging.getLogger(__name__)
self.logger.setLevel(logging.INFO)
else:
self.logger = logger
strStart = (
tagStart if tagStart is not None else getattr(opts, "stn_tagstart", "<!")
)
strEnd = tagEnd if tagEnd is not None else getattr(opts, "stn_tagend", "!>")
strParamStart = (
tagParamStart
if tagParamStart is not None
else getattr(opts, "stn_tagparamstart", "!")
)
strParamEnd = (
tagParamEnd
if tagParamEnd is not None
else getattr(opts, "stn_tagparamend", "!")
)
escapeSequence = r"(?<!\\)"
self.ignoreRepeats = (
ignoreRepeats
if ignoreRepeats is not None
else getattr(opts, "stn_ignorerepeats", True)
)
self.cleanup = (
cleanup if cleanup is not None else getattr(opts, "stn_cleanup", True)
)
self.separator = (
separator if separator is not None else getattr(opts, "stn_separator", ", ")
)
self.insertionPointTags = [
(strStart + strParamStart + "i" + str(x) + strParamEnd + strEnd)
for x in range(10)
]
self.regex = re.compile(
"("
+ escapeSequence
+ re.escape(strStart)
+ "(?:"
+ re.escape(strParamStart)
+ "([se]|(?:[pi][0-9]))"
+ re.escape(strParamEnd)
+ ")?(.*?)"
+ escapeSequence
+ re.escape(strEnd)
+ ")",
re.S,
)
def processPrompts(self, original_prompt, original_negative_prompt):
"""
Extract from the prompt the marked parts and add them to the negative prompt
"""
try:
prompt = original_prompt
negative_prompt = original_negative_prompt
alreadyProcessed = []
addAtStart = []
addAtEnd = []
addAtInsertionPoint = [[] for x in range(10)]
# 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 alreadyProcessed:
if self.ignoreRepeats:
alreadyProcessed.append(content)
self.logger.debug("Processing content: %s", content)
if position == "e":
addAtEnd.append(content)
elif position.startswith("p"):
n = int(position[1])
addAtInsertionPoint[n].append(content)
else: # position == "s" or invalid
addAtStart.append(content)
else:
self.logger.warn("Ignoring repeated content: %s", content)
# clean-up
prompt = prompt.replace(match[0], "")
if self.cleanup:
prompt = (
prompt.replace(" ", " ")
.replace(self.separator + self.separator, self.separator)
.removeprefix(self.separator)
.removesuffix(self.separator)
.strip()
)
# Add content to insertion points
for n in range(10):
ipp = negative_prompt.find(self.insertionPointTags[n])
if ipp >= 0:
ipl = len(self.insertionPointTags[n])
if (
negative_prompt[ipp - len(self.separator) : ipp]
== self.separator
):
ipp -= len(
self.separator
) # adjust for existing start separator
ipl += len(self.separator)
addAtInsertionPoint[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:
addAtInsertionPoint[n].append(endPart)
negative_prompt = self.separator.join(addAtInsertionPoint[n])
else:
logger.warn("Ignoring repeated content: %s", content)
prompt = prompt.replace(match[0], "")
ipp = 0
if negative_prompt.startswith(self.separator):
ipp = len(self.separator)
addAtInsertionPoint[n].append(negative_prompt[ipp:])
negative_prompt = self.separator.join(addAtInsertionPoint[n])
# Add content to insertion points
for n in range(10):
if insertionPointPositions[n] >= 0:
ipp = insertionPointPositions[n]
ipl = len(insertionPointTags[n])
if negative_prompt[ipp - len(separator) : ipp] == separator:
ipp -= len(separator) # adjust for existing start separator
ipl += len(separator)
addAtInsertionPoint[n].insert(0, negative_prompt[:ipp])
if negative_prompt[ipp + ipl : ipp + ipl + len(separator)] == separator:
ipl += len(separator) # adjust for existing end separator
endPart = negative_prompt[ipp + ipl :]
if len(endPart) > 0:
addAtInsertionPoint[n].append(endPart)
negative_prompt = separator.join(addAtInsertionPoint[n])
else:
ipp = 0
if negative_prompt.startswith(separator):
ipp = len(separator)
addAtInsertionPoint[n].append(negative_prompt[ipp:])
negative_prompt = separator.join(addAtInsertionPoint[n])
# Add content to start
if len(addAtStart) > 0:
if len(negative_prompt) > 0:
ipp = 0
if negative_prompt.startswith(self.separator):
ipp = len(self.separator) # adjust for existing end separator
addAtStart.append(negative_prompt[ipp:])
negative_prompt = self.separator.join(addAtStart)
# Add content to start
if len(addAtStart) > 0:
if len(negative_prompt) > 0:
ipp = 0
if negative_prompt.startswith(separator):
ipp = len(separator) # adjust for existing end separator
addAtStart.append(negative_prompt[ipp:])
negative_prompt = separator.join(addAtStart)
# Add content to end
if len(addAtEnd) > 0:
if len(negative_prompt) > 0:
ipl = len(negative_prompt)
if negative_prompt.endswith(self.separator):
ipl -= len(
self.separator
) # adjust for existing start separator
addAtEnd.insert(0, negative_prompt[:ipl])
negative_prompt = self.separator.join(addAtEnd)
# Add content to end
if len(addAtEnd) > 0:
if len(negative_prompt) > 0:
ipl = len(negative_prompt)
if negative_prompt.endswith(separator):
ipl -= len(separator) # adjust for existing start separator
addAtEnd.insert(0, negative_prompt[:ipl])
negative_prompt = separator.join(addAtEnd)
return prompt, negative_prompt
except Exception as e:
logger.exception(e)
return original_prompt, original_negative_prompt
class Script(scripts.Script):
def __init__(self):
pass
def title(self):
return f"SendToNegative v{VERSION}"
def show(self, is_img2img):
return scripts.AlwaysVisible
def process(self, p: StableDiffusionProcessing, *args, **kwargs):
for i in range(len(p.all_prompts)):
p.all_prompts[i], p.all_negative_prompts[i] = processPrompts(
p.all_prompts[i], p.all_negative_prompts[i]
)
return prompt, negative_prompt
except Exception as e:
self.logger.exception(e)
return original_prompt, original_negative_prompt
+136
View File
@@ -0,0 +1,136 @@
import unittest
from sendtonegative import SendToNegative
class TestSendToNegative(unittest.TestCase):
def setUp(self):
self.defstn = SendToNegative(
tagStart="<!",
tagEnd="!>",
tagParamStart="!",
tagParamEnd="!",
separator=", ",
ignoreRepeats=True,
cleanup=True,
)
def process(
self,
prompt,
negative_prompt,
expected_prompt,
expected_negative_prompt,
stn=None,
):
result_prompt, result_negative_prompt = (self.defstn if stn is None else stn).processPrompts(
prompt, negative_prompt
)
self.assertEqual(
result_prompt, expected_prompt, f"Prompt should be '{expected_prompt}'"
)
self.assertEqual(
result_negative_prompt,
expected_negative_prompt,
f"Negative Prompt should be '{expected_negative_prompt}'",
)
def test_tagDefault(self):
self.process(
"flowers<!red!>",
"normal quality, worse quality",
"flowers",
"red, normal quality, worse quality",
)
def test_tagStart(self):
self.process(
"flowers<!!s!red!>",
"normal quality, worse quality",
"flowers",
"red, normal quality, worse quality",
)
def test_tagEnd(self):
self.process(
"flowers<!!e!red!>",
"normal quality, worse quality",
"flowers",
"normal quality, worse quality, red",
)
def test_tagInsertion_midSep(self):
self.process(
"flowers<!!p0!red!>",
"normal quality, <!!i0!!>, worse quality",
"flowers",
"normal quality, red, worse quality",
)
def test_tagInsertion_midNoSep(self):
self.process(
"flowers<!!p0!red!>",
"normal quality<!!i0!!>worse quality",
"flowers",
"normal quality, red, worse quality",
)
def test_tagInsertion_startSep(self):
self.process(
"flowers<!!p0!red!>",
"<!!i0!!>, normal quality, worse quality",
"flowers",
"red, normal quality, worse quality",
)
def test_tagInsertion_startNoSep(self):
self.process(
"flowers<!!p0!red!>",
"<!!i0!!>normal quality, worse quality",
"flowers",
"red, normal quality, worse quality",
)
def test_tagInsertion_endSep(self):
self.process(
"flowers<!!p0!red!>",
"normal quality, worse quality, <!!i0!!>",
"flowers",
"normal quality, worse quality, red",
)
def test_tagInsertion_endNoSep(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!blue!>, <!!p0!yellow!> <!!p1!green!>",
"normal quality, <!!i0!!>, bad quality<!!i1!!>, worse quality",
"flowers",
"red, pink, normal quality, yellow, bad quality, green, worse quality, blue",
)
def test_complexNoCleanUp(self):
self.process(
"<!red!> <!!s!pink!>, flowers, <!!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, blue",
SendToNegative(
tagStart="<!",
tagEnd="!>",
tagParamStart="!",
tagParamEnd="!",
separator=", ",
ignoreRepeats=True,
cleanup=False,
),
)
if __name__ == "__main__":
unittest.main()