From c2ced311ba76603ff1b571af8ad0338459fc816f Mon Sep 17 00:00:00 2001 From: Antonio Cordero Balcazar Date: Sat, 13 May 2023 15:10:14 +0200 Subject: [PATCH] improved refactoring added tests --- .vscode/launch.json | 16 ++ scripts/extension.py | 97 +++++++++++++ scripts/sendtonegative.py | 298 +++++++++++++++++++++----------------- scripts/tests.py | 136 +++++++++++++++++ 4 files changed, 417 insertions(+), 130 deletions(-) create mode 100644 .vscode/launch.json create mode 100644 scripts/extension.py create mode 100644 scripts/tests.py diff --git a/.vscode/launch.json b/.vscode/launch.json new file mode 100644 index 0000000..c047885 --- /dev/null +++ b/.vscode/launch.json @@ -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 + } + ] +} \ No newline at end of file diff --git a/scripts/extension.py b/scripts/extension.py new file mode 100644 index 0000000..00fc1f7 --- /dev/null +++ b/scripts/extension.py @@ -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 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, + ), + ) diff --git a/scripts/sendtonegative.py b/scripts/sendtonegative.py index 40acdc8..6536e37 100644 --- a/scripts/sendtonegative.py +++ b/scripts/sendtonegative.py @@ -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 = "" -strParamStart = "!" -strParamEnd = "!" -""" -Default format: or - 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"(? 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: + + + 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", "") + 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"(? 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 diff --git a/scripts/tests.py b/scripts/tests.py new file mode 100644 index 0000000..0270210 --- /dev/null +++ b/scripts/tests.py @@ -0,0 +1,136 @@ +import unittest +from sendtonegative import SendToNegative + + +class TestSendToNegative(unittest.TestCase): + def setUp(self): + self.defstn = SendToNegative( + tagStart="", + 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", + "normal quality, worse quality", + "flowers", + "red, normal quality, worse quality", + ) + + def test_tagStart(self): + self.process( + "flowers", + "normal quality, worse quality", + "flowers", + "red, normal quality, worse quality", + ) + + def test_tagEnd(self): + self.process( + "flowers", + "normal quality, worse quality", + "flowers", + "normal quality, worse quality, red", + ) + + def test_tagInsertion_midSep(self): + self.process( + "flowers", + "normal quality, , worse quality", + "flowers", + "normal quality, red, worse quality", + ) + + def test_tagInsertion_midNoSep(self): + self.process( + "flowers", + "normal qualityworse quality", + "flowers", + "normal quality, red, worse quality", + ) + + def test_tagInsertion_startSep(self): + self.process( + "flowers", + ", normal quality, worse quality", + "flowers", + "red, normal quality, worse quality", + ) + + def test_tagInsertion_startNoSep(self): + self.process( + "flowers", + "normal quality, worse quality", + "flowers", + "red, normal quality, worse quality", + ) + + def test_tagInsertion_endSep(self): + self.process( + "flowers", + "normal quality, worse quality, ", + "flowers", + "normal quality, worse quality, red", + ) + + def test_tagInsertion_endNoSep(self): + self.process( + "flowers", + "normal quality, worse quality", + "flowers", + "normal quality, worse quality, red", + ) + + def test_complex(self): + self.process( + " , flowers, , ", + "normal quality, , bad quality, worse quality", + "flowers", + "red, pink, normal quality, yellow, bad quality, green, worse quality, blue", + ) + + def test_complexNoCleanUp(self): + self.process( + " , flowers, , ", + "normal quality, , bad quality, worse quality", + " , flowers, , ", + "red, pink, normal quality, yellow, bad quality, green, worse quality, blue", + SendToNegative( + tagStart="", + tagParamStart="!", + tagParamEnd="!", + separator=", ", + ignoreRepeats=True, + cleanup=False, + ), + ) + + +if __name__ == "__main__": + unittest.main()