refactoring and linting
removed unneeded code for logging
This commit is contained in:
@@ -0,0 +1,171 @@
|
||||
[MAIN]
|
||||
|
||||
[BASIC]
|
||||
argument-naming-style=snake_case
|
||||
attr-naming-style=snake_case
|
||||
class-attribute-naming-style=any
|
||||
class-const-naming-style=UPPER_CASE
|
||||
class-naming-style=PascalCase
|
||||
const-naming-style=snake_case
|
||||
docstring-min-length=-1
|
||||
function-naming-style=snake_case
|
||||
good-names=i,j,k,e,ex,ok,p
|
||||
good-names-rgxs=
|
||||
include-naming-hint=no
|
||||
inlinevar-naming-style=any
|
||||
method-naming-style=snake_case
|
||||
module-naming-style=snake_case
|
||||
name-group=
|
||||
no-docstring-rgx=^_
|
||||
property-classes=abc.abstractproperty
|
||||
variable-naming-style=snake_case
|
||||
|
||||
[CLASSES]
|
||||
check-protected-access-in-special-methods=no
|
||||
defining-attr-methods=__init__,
|
||||
__new__,
|
||||
setUp,
|
||||
asyncSetUp,
|
||||
__post_init__
|
||||
exclude-protected=_asdict,_fields,_replace,_source,_make,os._exit
|
||||
valid-classmethod-first-arg=cls
|
||||
valid-metaclass-classmethod-first-arg=mcs
|
||||
|
||||
[DESIGN]
|
||||
exclude-too-few-public-methods=
|
||||
ignored-parents=
|
||||
max-args=10
|
||||
max-attributes=7
|
||||
max-bool-expr=5
|
||||
max-branches=12
|
||||
max-locals=15
|
||||
max-parents=7
|
||||
max-public-methods=20
|
||||
max-returns=6
|
||||
max-statements=50
|
||||
min-public-methods=2
|
||||
|
||||
[EXCEPTIONS]
|
||||
overgeneral-exceptions=builtins.BaseException,builtins.Exception
|
||||
|
||||
[FORMAT]
|
||||
expected-line-ending-format=
|
||||
ignore-long-lines=^\s*(# )?<?https?://\S+>?$
|
||||
indent-after-paren=4
|
||||
indent-string=' '
|
||||
max-line-length=200
|
||||
max-module-lines=9999
|
||||
single-line-class-stmt=no
|
||||
single-line-if-stmt=no
|
||||
|
||||
[IMPORTS]
|
||||
allow-any-import-level=
|
||||
allow-reexport-from-package=no
|
||||
allow-wildcard-with-all=no
|
||||
deprecated-modules=
|
||||
ext-import-graph=
|
||||
import-graph=
|
||||
int-import-graph=
|
||||
known-standard-library=
|
||||
known-third-party=enchant
|
||||
preferred-modules=
|
||||
|
||||
[LOGGING]
|
||||
logging-format-style=new
|
||||
logging-modules=logging
|
||||
|
||||
[MESSAGES CONTROL]
|
||||
confidence=HIGH,
|
||||
CONTROL_FLOW,
|
||||
INFERENCE,
|
||||
INFERENCE_FAILURE,
|
||||
UNDEFINED
|
||||
# disable=C,R,W
|
||||
disable=raw-checker-failed,
|
||||
bad-inline-option,
|
||||
locally-disabled,
|
||||
file-ignored,
|
||||
suppressed-message,
|
||||
useless-suppression,
|
||||
deprecated-pragma,
|
||||
use-symbolic-message-instead,
|
||||
line-too-long,
|
||||
missing-function-docstring,
|
||||
missing-module-docstring,
|
||||
missing-class-docstring,
|
||||
logging-fstring-interpolation,
|
||||
import-outside-toplevel,
|
||||
consider-iterating-dictionary,
|
||||
wrong-import-position,
|
||||
unnecessary-lambda,
|
||||
consider-using-dict-items,
|
||||
dangerous-default-value,
|
||||
unnecessary-dunder-call,
|
||||
invalid-name,
|
||||
R0801,
|
||||
enable=c-extension-no-member
|
||||
|
||||
[METHOD_ARGS]
|
||||
timeout-methods=requests.api.delete,requests.api.get,requests.api.head,requests.api.options,requests.api.patch,requests.api.post,requests.api.put,requests.api.request
|
||||
|
||||
[MISCELLANEOUS]
|
||||
notes=FIXME,
|
||||
XXX,
|
||||
TODO
|
||||
notes-rgx=
|
||||
|
||||
[REFACTORING]
|
||||
max-nested-blocks=5
|
||||
never-returning-functions=sys.exit,argparse.parse_error
|
||||
|
||||
[REPORTS]
|
||||
evaluation=max(0, 0 if fatal else 10.0 - ((float(5 * error + warning + refactor + convention) / statement) * 10))
|
||||
msg-template=
|
||||
#output-format=
|
||||
reports=no
|
||||
score=no
|
||||
|
||||
[SIMILARITIES]
|
||||
ignore-comments=yes
|
||||
ignore-docstrings=yes
|
||||
ignore-imports=yes
|
||||
ignore-signatures=yes
|
||||
min-similarity-lines=4
|
||||
|
||||
[SPELLING]
|
||||
max-spelling-suggestions=4
|
||||
spelling-dict=
|
||||
spelling-ignore-comment-directives=fmt: on,fmt: off,noqa:,noqa,nosec,isort:skip,mypy:
|
||||
spelling-ignore-words=
|
||||
spelling-private-dict-file=
|
||||
spelling-store-unknown-words=no
|
||||
|
||||
[STRING]
|
||||
check-quote-consistency=no
|
||||
check-str-concat-over-line-jumps=no
|
||||
|
||||
[TYPECHECK]
|
||||
contextmanager-decorators=contextlib.contextmanager
|
||||
generated-members=numpy.*,torch.*,cv2.*
|
||||
ignore-none=yes
|
||||
ignore-on-opaque-inference=yes
|
||||
ignored-checks-for-mixins=no-member,
|
||||
not-async-context-manager,
|
||||
not-context-manager,
|
||||
attribute-defined-outside-init
|
||||
ignored-classes=optparse.Values,thread._local,_thread._local,argparse.Namespace
|
||||
missing-member-hint=yes
|
||||
missing-member-hint-distance=1
|
||||
missing-member-max-choices=1
|
||||
mixin-class-rgx=.*[Mm]ixin
|
||||
signature-mutators=
|
||||
|
||||
[VARIABLES]
|
||||
additional-builtins=
|
||||
allow-global-unused-variables=yes
|
||||
allowed-redefined-builtins=
|
||||
callbacks=cb_,
|
||||
dummy-variables-rgx=_+$|(_[a-zA-Z0-9_]*[a-zA-Z0-9]+?$)|dummy|^ignored_|^unused_
|
||||
ignored-argument-names=_.*|^ignored_|^unused_
|
||||
init-import=no
|
||||
redefining-builtins-modules=six.moves,past.builtins,future.builtins,builtins,io
|
||||
Vendored
+7
-1
@@ -7,5 +7,11 @@
|
||||
"test*.py"
|
||||
],
|
||||
"python.testing.pytestEnabled": false,
|
||||
"python.testing.unittestEnabled": true
|
||||
"python.testing.unittestEnabled": true,
|
||||
"python.analysis.typeCheckingMode": "basic",
|
||||
"python.linting.pylintEnabled": true,
|
||||
"python.linting.enabled": true,
|
||||
"black-formatter.args": [
|
||||
"--line-length=120"
|
||||
]
|
||||
}
|
||||
+66
-69
@@ -6,21 +6,19 @@ import os
|
||||
|
||||
sys.path.insert(1, os.path.join(sys.path[0], ".."))
|
||||
|
||||
from sendtonegative import SendToNegative
|
||||
import logging
|
||||
|
||||
# pylint: disable=import-error
|
||||
|
||||
from modules import scripts, shared, script_callbacks
|
||||
from modules.processing import StableDiffusionProcessing
|
||||
from modules.shared import opts
|
||||
from sendtonegative import SendToNegative
|
||||
|
||||
|
||||
class SendToNegativeScript(scripts.Script):
|
||||
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)
|
||||
script_callbacks.on_ui_settings(self.__on_ui_settings)
|
||||
self.callbacks_added = True
|
||||
|
||||
def title(self):
|
||||
@@ -30,68 +28,67 @@ class SendToNegativeScript(scripts.Script):
|
||||
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.processPrompt(
|
||||
stn = SendToNegative(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]
|
||||
)
|
||||
|
||||
|
||||
def _on_ui_settings():
|
||||
section = ("send-to-negative", SendToNegative.NAME)
|
||||
shared.opts.add_option(
|
||||
key="stn_tagstart",
|
||||
info=shared.OptionInfo(
|
||||
SendToNegative.DEFAULT_tagStart,
|
||||
label="Tag start",
|
||||
section=section,
|
||||
),
|
||||
)
|
||||
shared.opts.add_option(
|
||||
key="stn_tagend",
|
||||
info=shared.OptionInfo(
|
||||
SendToNegative.DEFAULT_tagEnd,
|
||||
label="Tag end",
|
||||
section=section,
|
||||
),
|
||||
)
|
||||
shared.opts.add_option(
|
||||
key="stn_tagparamstart",
|
||||
info=shared.OptionInfo(
|
||||
SendToNegative.DEFAULT_tagParamStart,
|
||||
label="Tag parameter start",
|
||||
section=section,
|
||||
),
|
||||
)
|
||||
shared.opts.add_option(
|
||||
key="stn_tagparamend",
|
||||
info=shared.OptionInfo(
|
||||
SendToNegative.DEFAULT_tagParamEnd,
|
||||
label="Tag parameter end",
|
||||
section=section,
|
||||
),
|
||||
)
|
||||
shared.opts.add_option(
|
||||
key="stn_separator",
|
||||
info=shared.OptionInfo(
|
||||
SendToNegative.DEFAULT_separator,
|
||||
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. Removes extra spaces or separators (the configured separator).",
|
||||
section=section,
|
||||
),
|
||||
)
|
||||
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(
|
||||
SendToNegative.DEFAULT_SEPARATOR,
|
||||
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 (removes extra spaces or the configured separator)",
|
||||
section=section,
|
||||
),
|
||||
)
|
||||
|
||||
+139
-149
@@ -2,27 +2,26 @@ import logging
|
||||
import re
|
||||
|
||||
|
||||
class SendToNegative:
|
||||
class SendToNegative: # pylint: disable=too-few-public-methods
|
||||
NAME = "Send to Negative"
|
||||
VERSION = "1.0"
|
||||
VERSION = "1.1"
|
||||
|
||||
DEFAULT_tagStart = "<!"
|
||||
DEFAULT_tagEnd = "!>"
|
||||
DEFAULT_tagParamStart = "!"
|
||||
DEFAULT_tagParamEnd = "!"
|
||||
DEFAULT_separator = ", "
|
||||
DEFAULT_TAG_START = "<!"
|
||||
DEFAULT_TAG_END = "!>"
|
||||
DEFAULT_TAG_PARAM_START = "!"
|
||||
DEFAULT_TAG_PARAM_END = "!"
|
||||
DEFAULT_SEPARATOR = ", "
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tagStart=None,
|
||||
tagEnd=None,
|
||||
tagParamStart=None,
|
||||
tagParamEnd=None,
|
||||
tag_start=None,
|
||||
tag_end=None,
|
||||
tag_param_start=None,
|
||||
tag_param_end=None,
|
||||
separator=None,
|
||||
ignoreRepeats=None,
|
||||
ignore_repeats=None,
|
||||
cleanup=None,
|
||||
opts=None,
|
||||
logger=None,
|
||||
):
|
||||
"""
|
||||
Default format for the tag:
|
||||
@@ -39,171 +38,162 @@ class SendToNegative:
|
||||
|
||||
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.
|
||||
"""
|
||||
if logger is None:
|
||||
self.logger = logging.getLogger(__name__)
|
||||
self.logger.setLevel(logging.INFO)
|
||||
else:
|
||||
self.logger = logger
|
||||
self.__logger = logging.getLogger(__name__)
|
||||
|
||||
strStart = (
|
||||
tagStart
|
||||
if tagStart is not None
|
||||
else getattr(opts, "stn_tagstart", self.DEFAULT_tagStart)
|
||||
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_tagStart
|
||||
else self.DEFAULT_TAG_START
|
||||
)
|
||||
strEnd = (
|
||||
tagEnd
|
||||
if tagEnd is not None
|
||||
else getattr(opts, "stn_tagend", self.DEFAULT_tagEnd)
|
||||
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_tagEnd
|
||||
else self.DEFAULT_TAG_END
|
||||
)
|
||||
strParamStart = (
|
||||
tagParamStart
|
||||
if tagParamStart is not None
|
||||
else getattr(opts, "stn_tagparamstart", self.DEFAULT_tagParamStart)
|
||||
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_tagParamStart
|
||||
else self.DEFAULT_TAG_PARAM_START
|
||||
)
|
||||
strParamEnd = (
|
||||
tagParamEnd
|
||||
if tagParamEnd is not None
|
||||
else getattr(opts, "stn_tagparamend", self.DEFAULT_tagParamEnd)
|
||||
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_tagParamEnd
|
||||
else self.DEFAULT_TAG_PARAM_END
|
||||
)
|
||||
escapeSequence = r"(?<!\\)"
|
||||
self.ignoreRepeats = (
|
||||
ignoreRepeats
|
||||
if ignoreRepeats is not None
|
||||
else getattr(opts, "stn_ignorerepeats", True)
|
||||
escape_sequence = r"(?<!\\)"
|
||||
self.__ignore_repeats = (
|
||||
ignore_repeats if ignore_repeats is not None else getattr(opts, "stn_ignorerepeats", True)
|
||||
)
|
||||
self.cleanup = (
|
||||
cleanup
|
||||
if cleanup is not None
|
||||
else getattr(opts, "stn_cleanup", 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
|
||||
)
|
||||
self.separator = (
|
||||
self.__separator = (
|
||||
separator
|
||||
if separator is not None
|
||||
else getattr(opts, "stn_separator", self.DEFAULT_separator)
|
||||
else getattr(opts, "stn_separator", self.DEFAULT_SEPARATOR)
|
||||
if opts is not None
|
||||
else self.DEFAULT_separator
|
||||
else self.DEFAULT_SEPARATOR
|
||||
)
|
||||
self.insertionPointTags = [
|
||||
(strStart + strParamStart + "i" + str(x) + strParamEnd + strEnd)
|
||||
for x in range(10)
|
||||
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(
|
||||
self.__regex = re.compile(
|
||||
"("
|
||||
+ escapeSequence
|
||||
+ re.escape(strStart)
|
||||
+ escape_sequence
|
||||
+ re.escape(str_start)
|
||||
+ "(?:"
|
||||
+ re.escape(strParamStart)
|
||||
+ re.escape(str_param_start)
|
||||
+ "([se]|(?:[pi][0-9]))"
|
||||
+ re.escape(strParamEnd)
|
||||
+ re.escape(str_param_end)
|
||||
+ ")?(.*?)"
|
||||
+ escapeSequence
|
||||
+ re.escape(strEnd)
|
||||
+ escape_sequence
|
||||
+ re.escape(str_end)
|
||||
+ ")",
|
||||
re.S,
|
||||
)
|
||||
|
||||
def processPrompt(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
|
||||
"""
|
||||
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:
|
||||
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 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 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)
|
||||
|
||||
self.__logger.debug(f"Input prompt: {prompt}")
|
||||
self.__logger.debug(f"Input negative_prompt: {negative_prompt}")
|
||||
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}")
|
||||
return prompt, negative_prompt
|
||||
except Exception as e:
|
||||
self.logger.exception(e)
|
||||
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()
|
||||
)
|
||||
return prompt, add_at
|
||||
|
||||
def __add_to_insertion_points(self, negative_prompt, add_at_insertion_point):
|
||||
for n in range(10):
|
||||
ipp = negative_prompt.find(self.__insertion_point_tags[n])
|
||||
if ipp >= 0:
|
||||
ipl = len(self.__insertion_point_tags[n])
|
||||
if negative_prompt[ipp - len(self.__separator) : ipp] == self.__separator:
|
||||
ipp -= len(self.__separator) # adjust for existing start separator
|
||||
ipl += len(self.__separator)
|
||||
add_at_insertion_point[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:
|
||||
add_at_insertion_point[n].append(endPart)
|
||||
negative_prompt = self.__separator.join(add_at_insertion_point[n])
|
||||
else:
|
||||
ipp = 0
|
||||
if negative_prompt.startswith(self.__separator):
|
||||
ipp = len(self.__separator)
|
||||
add_at_insertion_point[n].append(negative_prompt[ipp:])
|
||||
negative_prompt = self.__separator.join(add_at_insertion_point[n])
|
||||
return negative_prompt
|
||||
|
||||
def __add_to_start(self, negative_prompt, add_at_start):
|
||||
if len(negative_prompt) > 0:
|
||||
ipp = 0
|
||||
if negative_prompt.startswith(self.__separator):
|
||||
ipp = len(self.__separator) # adjust for existing end separator
|
||||
add_at_start.append(negative_prompt[ipp:])
|
||||
negative_prompt = self.__separator.join(add_at_start)
|
||||
return negative_prompt
|
||||
|
||||
def __add_to_end(self, negative_prompt, add_at_end):
|
||||
if len(negative_prompt) > 0:
|
||||
ipl = len(negative_prompt)
|
||||
if negative_prompt.endswith(self.__separator):
|
||||
ipl -= len(self.__separator) # adjust for existing start separator
|
||||
add_at_end.insert(0, negative_prompt[:ipl])
|
||||
negative_prompt = self.__separator.join(add_at_end)
|
||||
return negative_prompt
|
||||
|
||||
+29
-33
@@ -4,18 +4,18 @@ import os
|
||||
|
||||
sys.path.insert(1, os.path.join(sys.path[0], ".."))
|
||||
|
||||
from sendtonegative import SendToNegative
|
||||
from sendtonegative import SendToNegative # pylint: disable=import-error
|
||||
|
||||
|
||||
class TestSendToNegative(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.defstn = SendToNegative(
|
||||
tagStart="<!",
|
||||
tagEnd="!>",
|
||||
tagParamStart="!",
|
||||
tagParamEnd="!",
|
||||
tag_start="<!",
|
||||
tag_end="!>",
|
||||
tag_param_start="!",
|
||||
tag_param_end="!",
|
||||
separator=", ",
|
||||
ignoreRepeats=True,
|
||||
ignore_repeats=True,
|
||||
cleanup=True,
|
||||
)
|
||||
|
||||
@@ -27,20 +27,16 @@ class TestSendToNegative(unittest.TestCase):
|
||||
expected_negative_prompt,
|
||||
stn=None,
|
||||
):
|
||||
theObj = self.defstn if stn is None else stn
|
||||
result_prompt, result_negative_prompt = theObj.processPrompt(
|
||||
prompt, negative_prompt
|
||||
)
|
||||
self.assertEqual(
|
||||
result_prompt, expected_prompt, f"Prompt should be '{expected_prompt}'"
|
||||
)
|
||||
the_obj = self.defstn if stn is None else stn
|
||||
result_prompt, result_negative_prompt = the_obj.process_prompt(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):
|
||||
def test_tag_default(self):
|
||||
self.process(
|
||||
"flowers<!red!>",
|
||||
"normal quality, worse quality",
|
||||
@@ -48,7 +44,7 @@ class TestSendToNegative(unittest.TestCase):
|
||||
"red, normal quality, worse quality",
|
||||
)
|
||||
|
||||
def test_tagStart(self):
|
||||
def test_tag_start(self):
|
||||
self.process(
|
||||
"flowers<!!s!red!>",
|
||||
"normal quality, worse quality",
|
||||
@@ -56,7 +52,7 @@ class TestSendToNegative(unittest.TestCase):
|
||||
"red, normal quality, worse quality",
|
||||
)
|
||||
|
||||
def test_tagEnd(self):
|
||||
def test_tag_end(self):
|
||||
self.process(
|
||||
"flowers<!!e!red!>",
|
||||
"normal quality, worse quality",
|
||||
@@ -64,7 +60,7 @@ class TestSendToNegative(unittest.TestCase):
|
||||
"normal quality, worse quality, red",
|
||||
)
|
||||
|
||||
def test_tagInsertion_midSep(self):
|
||||
def test_tag_insertion_mid_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"normal quality, <!!i0!!>, worse quality",
|
||||
@@ -72,7 +68,7 @@ class TestSendToNegative(unittest.TestCase):
|
||||
"normal quality, red, worse quality",
|
||||
)
|
||||
|
||||
def test_tagInsertion_midNoSep(self):
|
||||
def test_tag_insertion_mid_no_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"normal quality<!!i0!!>worse quality",
|
||||
@@ -80,7 +76,7 @@ class TestSendToNegative(unittest.TestCase):
|
||||
"normal quality, red, worse quality",
|
||||
)
|
||||
|
||||
def test_tagInsertion_startSep(self):
|
||||
def test_tag_insertion_start_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"<!!i0!!>, normal quality, worse quality",
|
||||
@@ -88,7 +84,7 @@ class TestSendToNegative(unittest.TestCase):
|
||||
"red, normal quality, worse quality",
|
||||
)
|
||||
|
||||
def test_tagInsertion_startNoSep(self):
|
||||
def test_tag_insertion_start_no_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"<!!i0!!>normal quality, worse quality",
|
||||
@@ -96,7 +92,7 @@ class TestSendToNegative(unittest.TestCase):
|
||||
"red, normal quality, worse quality",
|
||||
)
|
||||
|
||||
def test_tagInsertion_endSep(self):
|
||||
def test_tag_insertion_end_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"normal quality, worse quality, <!!i0!!>",
|
||||
@@ -104,7 +100,7 @@ class TestSendToNegative(unittest.TestCase):
|
||||
"normal quality, worse quality, red",
|
||||
)
|
||||
|
||||
def test_tagInsertion_endNoSep(self):
|
||||
def test_tag_insertion_end_no_sep(self):
|
||||
self.process(
|
||||
"flowers<!!p0!red!>",
|
||||
"normal quality, worse quality<!!i0!!>",
|
||||
@@ -114,25 +110,25 @@ class TestSendToNegative(unittest.TestCase):
|
||||
|
||||
def test_complex(self):
|
||||
self.process(
|
||||
"<!red!> <!!s!pink!>, flowers, <!!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, blue",
|
||||
"red, pink, normal quality, yellow, bad quality, green, worse quality, purple, blue",
|
||||
)
|
||||
|
||||
def test_complexNoCleanUp(self):
|
||||
def test_complex_no_cleanup(self):
|
||||
self.process(
|
||||
"<!red!> <!!s!pink!>, flowers, <!!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, blue",
|
||||
" , flowers , , ",
|
||||
"red, pink, normal quality, yellow, bad quality, green, worse quality, purple, blue",
|
||||
SendToNegative(
|
||||
tagStart="<!",
|
||||
tagEnd="!>",
|
||||
tagParamStart="!",
|
||||
tagParamEnd="!",
|
||||
tag_start="<!",
|
||||
tag_end="!>",
|
||||
tag_param_start="!",
|
||||
tag_param_end="!",
|
||||
separator=", ",
|
||||
ignoreRepeats=True,
|
||||
ignore_repeats=True,
|
||||
cleanup=False,
|
||||
),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user