improved
refactoring added tests
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user