Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c8a7b6c815 | ||
|
|
7f36587bfd | ||
|
|
46d5100e7e | ||
|
|
d31736187b | ||
|
|
7ccae56ce0 | ||
|
|
782c8f76c2 | ||
|
|
6ddfc184c0 | ||
|
|
31e7822499 | ||
|
|
d86b997cb6 |
@@ -0,0 +1,46 @@
|
|||||||
|
---
|
||||||
|
name: Issue report
|
||||||
|
about: Create a report to help us improve
|
||||||
|
title: "[Issue] "
|
||||||
|
assignees: acorderob
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Prerequisites
|
||||||
|
|
||||||
|
Please answer the following questions for yourself before submitting an issue. **YOU MAY DELETE THE PREREQUISITES SECTION.**
|
||||||
|
|
||||||
|
- [ ] 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)
|
||||||
|
|
||||||
|
## Current Behavior
|
||||||
|
|
||||||
|
What is the current behavior?
|
||||||
|
|
||||||
|
## 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.
|
||||||
@@ -0,0 +1,24 @@
|
|||||||
|
---
|
||||||
|
name: Feature request
|
||||||
|
about: Suggest an idea for this project
|
||||||
|
title: "[Request]"
|
||||||
|
assignees: acorderob
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Prerequisites
|
||||||
|
|
||||||
|
Please answer the following questions for yourself before submitting an issue. **YOU MAY DELETE THE PREREQUISITES SECTION.**
|
||||||
|
|
||||||
|
- [ ] 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
|
||||||
|
|
||||||
|
Add any other context or screenshots about the feature request here.
|
||||||
@@ -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
|
||||||
@@ -1 +1,5 @@
|
|||||||
**/__pycache__
|
**/__pycache__
|
||||||
|
|
||||||
|
.vscode/**/*
|
||||||
|
!.vscode/settings.json
|
||||||
|
!.vscode/launch.json
|
||||||
|
|||||||
@@ -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
+17
@@ -0,0 +1,17 @@
|
|||||||
|
{
|
||||||
|
"python.testing.unittestArgs": [
|
||||||
|
"-v",
|
||||||
|
"-s",
|
||||||
|
"./tests",
|
||||||
|
"-p",
|
||||||
|
"test*.py"
|
||||||
|
],
|
||||||
|
"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"
|
||||||
|
]
|
||||||
|
}
|
||||||
@@ -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.
|
Extension for the [AUTOMATIC1111 Stable Diffusion WebUI](https://github.com/AUTOMATIC1111/stable-diffusion-webui) or compatible UIs.
|
||||||
|
|
||||||
@@ -8,13 +8,26 @@ 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
|
negative prompt. This allows useful tricks when using a wildcard extension
|
||||||
since you can add negative content from choices made in the positive prompt.
|
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)
|
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 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
|
this is not the case, you can just rename the extension folder so the ordering
|
||||||
works out.
|
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.
|
||||||
|
|
||||||
|
## 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
|
## Usage
|
||||||
|
|
||||||
The format of the tags is like this:
|
The format of the tags is like this:
|
||||||
@@ -54,11 +67,39 @@ that part to the negative prompt.
|
|||||||
|
|
||||||
## Configuration
|
## Configuration
|
||||||
|
|
||||||
The extension settings allow you to change the format of the tag in case there
|
Separator used when adding to the negative prompt: You can specify the separator used when adding to the negative prompt (by default it's ", ").
|
||||||
is some incompatibility with another extension.
|
|
||||||
|
|
||||||
You can also specify the separator added to the negative prompt which by
|
Ignore tags with repeated content: by default it ignores repeated content to avoid repetitions in the negative prompt.
|
||||||
default is ", ".
|
|
||||||
|
|
||||||
By default it ignores repeated content and also tries to clean up the prompt
|
Join attention modifiers (weights) when possible: by default it joins attention modifiers when possible (joins into one, multipliying their values).
|
||||||
after removing the tags, but these can also be changed in the settings.
|
|
||||||
|
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.
|
||||||
|
|||||||
+42
-69
@@ -6,21 +6,19 @@ import os
|
|||||||
|
|
||||||
sys.path.insert(1, os.path.join(sys.path[0], ".."))
|
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 import scripts, shared, script_callbacks
|
||||||
from modules.processing import StableDiffusionProcessing
|
from modules.processing import StableDiffusionProcessing
|
||||||
from modules.shared import opts
|
from modules.shared import opts
|
||||||
|
from sendtonegative import SendToNegative
|
||||||
|
|
||||||
|
|
||||||
class SendToNegativeScript(scripts.Script):
|
class SendToNegativeScript(scripts.Script):
|
||||||
def __init__(self):
|
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"):
|
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
|
self.callbacks_added = True
|
||||||
|
|
||||||
def title(self):
|
def title(self):
|
||||||
@@ -30,68 +28,43 @@ class SendToNegativeScript(scripts.Script):
|
|||||||
return scripts.AlwaysVisible
|
return scripts.AlwaysVisible
|
||||||
|
|
||||||
def process(self, p: StableDiffusionProcessing, *args, **kwargs):
|
def process(self, p: StableDiffusionProcessing, *args, **kwargs):
|
||||||
stn = SendToNegative(opts=opts, logger=self.logger)
|
stn = SendToNegative(opts=opts)
|
||||||
for i in range(len(p.all_prompts)):
|
for i in range(len(p.all_prompts)): # pylint: disable=consider-using-enumerate
|
||||||
p.all_prompts[i], p.all_negative_prompts[i] = stn.processPrompt(
|
p.all_prompts[i], p.all_negative_prompts[i] = stn.process_prompt(
|
||||||
p.all_prompts[i], p.all_negative_prompts[i]
|
p.all_prompts[i], p.all_negative_prompts[i]
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def __on_ui_settings(self):
|
||||||
def _on_ui_settings():
|
section = ("send-to-negative", SendToNegative.NAME)
|
||||||
section = ("send-to-negative", SendToNegative.NAME)
|
shared.opts.add_option(
|
||||||
shared.opts.add_option(
|
key="stn_separator",
|
||||||
key="stn_tagstart",
|
info=shared.OptionInfo(
|
||||||
info=shared.OptionInfo(
|
SendToNegative.DEFAULT_SEPARATOR,
|
||||||
SendToNegative.DEFAULT_tagStart,
|
label="Separator used when adding to the negative prompt",
|
||||||
label="Tag start",
|
section=section,
|
||||||
section=section,
|
),
|
||||||
),
|
)
|
||||||
)
|
shared.opts.add_option(
|
||||||
shared.opts.add_option(
|
key="stn_ignorerepeats",
|
||||||
key="stn_tagend",
|
info=shared.OptionInfo(
|
||||||
info=shared.OptionInfo(
|
True,
|
||||||
SendToNegative.DEFAULT_tagEnd,
|
label="Ignore tags with repeated content",
|
||||||
label="Tag end",
|
section=section,
|
||||||
section=section,
|
),
|
||||||
),
|
)
|
||||||
)
|
shared.opts.add_option(
|
||||||
shared.opts.add_option(
|
key="stn_joinattention",
|
||||||
key="stn_tagparamstart",
|
info=shared.OptionInfo(
|
||||||
info=shared.OptionInfo(
|
True,
|
||||||
SendToNegative.DEFAULT_tagParamStart,
|
label="Join attention modifiers (weights) when possible",
|
||||||
label="Tag parameter start",
|
section=section,
|
||||||
section=section,
|
),
|
||||||
),
|
)
|
||||||
)
|
shared.opts.add_option(
|
||||||
shared.opts.add_option(
|
key="stn_cleanup",
|
||||||
key="stn_tagparamend",
|
info=shared.OptionInfo(
|
||||||
info=shared.OptionInfo(
|
True,
|
||||||
SendToNegative.DEFAULT_tagParamEnd,
|
label="Try to clean-up the prompt after processing (removes extra spaces or the configured separator)",
|
||||||
label="Tag parameter end",
|
section=section,
|
||||||
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,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|||||||
+276
-164
@@ -1,31 +1,25 @@
|
|||||||
|
from collections import namedtuple
|
||||||
import logging
|
import logging
|
||||||
|
import lark
|
||||||
import re
|
import re
|
||||||
|
|
||||||
|
|
||||||
class SendToNegative:
|
class SendToNegative: # pylint: disable=too-few-public-methods
|
||||||
NAME = "Send to Negative"
|
NAME = "Send to Negative"
|
||||||
VERSION = "1.0"
|
VERSION = "2.0"
|
||||||
|
|
||||||
DEFAULT_tagStart = "<!"
|
DEFAULT_SEPARATOR = ", "
|
||||||
DEFAULT_tagEnd = "!>"
|
|
||||||
DEFAULT_tagParamStart = "!"
|
|
||||||
DEFAULT_tagParamEnd = "!"
|
|
||||||
DEFAULT_separator = ", "
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
tagStart=None,
|
|
||||||
tagEnd=None,
|
|
||||||
tagParamStart=None,
|
|
||||||
tagParamEnd=None,
|
|
||||||
separator=None,
|
separator=None,
|
||||||
ignoreRepeats=None,
|
ignore_repeats=None,
|
||||||
|
join_attention=None,
|
||||||
cleanup=None,
|
cleanup=None,
|
||||||
opts=None,
|
opts=None,
|
||||||
logger=None,
|
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Default format for the tag:
|
Format for the tag:
|
||||||
<!content!>
|
<!content!>
|
||||||
|
|
||||||
<!!x!content!>
|
<!!x!content!>
|
||||||
@@ -39,171 +33,289 @@ 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.
|
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 = logging.getLogger(__name__)
|
self.__ignore_repeats = (
|
||||||
self.logger.setLevel(logging.INFO)
|
ignore_repeats if ignore_repeats is not None else getattr(opts, "stn_ignorerepeats", True)
|
||||||
else:
|
|
||||||
self.logger = logger
|
|
||||||
|
|
||||||
strStart = (
|
|
||||||
tagStart
|
|
||||||
if tagStart is not None
|
|
||||||
else getattr(opts, "stn_tagstart", self.DEFAULT_tagStart)
|
|
||||||
if opts is not None
|
|
||||||
else self.DEFAULT_tagStart
|
|
||||||
)
|
)
|
||||||
strEnd = (
|
self.__join_attention = (
|
||||||
tagEnd
|
join_attention
|
||||||
if tagEnd is not None
|
if join_attention is not None
|
||||||
else getattr(opts, "stn_tagend", self.DEFAULT_tagEnd)
|
else getattr(opts, "stn_joinattention", True)
|
||||||
if opts is not None
|
|
||||||
else self.DEFAULT_tagEnd
|
|
||||||
)
|
|
||||||
strParamStart = (
|
|
||||||
tagParamStart
|
|
||||||
if tagParamStart is not None
|
|
||||||
else getattr(opts, "stn_tagparamstart", self.DEFAULT_tagParamStart)
|
|
||||||
if opts is not None
|
|
||||||
else self.DEFAULT_tagParamStart
|
|
||||||
)
|
|
||||||
strParamEnd = (
|
|
||||||
tagParamEnd
|
|
||||||
if tagParamEnd is not None
|
|
||||||
else getattr(opts, "stn_tagparamend", self.DEFAULT_tagParamEnd)
|
|
||||||
if opts is not None
|
|
||||||
else self.DEFAULT_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)
|
|
||||||
if opts is not None
|
if opts is not None
|
||||||
else True
|
else True
|
||||||
)
|
)
|
||||||
self.separator = (
|
self.__cleanup = (
|
||||||
|
cleanup if cleanup is not None else getattr(opts, "stn_cleanup", True) if opts is not None else True
|
||||||
|
)
|
||||||
|
self.__separator = (
|
||||||
separator
|
separator
|
||||||
if separator is not None
|
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
|
if opts is not None
|
||||||
else self.DEFAULT_separator
|
else self.DEFAULT_SEPARATOR
|
||||||
)
|
)
|
||||||
self.insertionPointTags = [
|
self.__insertion_point_tags = [f"<!!i{x}!!>" for x in range(10)]
|
||||||
(strStart + strParamStart + "i" + str(x) + strParamEnd + strEnd)
|
# Process with lark (debug with https://www.lark-parser.org/ide/)
|
||||||
for x in range(10)
|
self.__schedule_parser = lark.Lark(
|
||||||
]
|
r"""
|
||||||
self.regex = re.compile(
|
start: (prompt | /[\][():|<>!]/+)*
|
||||||
"("
|
?prompt: (emphasized | deemphasized | scheduled | alternate | modeltag | negtag | plain)*
|
||||||
+ escapeSequence
|
?nonegprompt: (emphasized | deemphasized | scheduled | alternate | modeltag | plain)*
|
||||||
+ re.escape(strStart)
|
emphasized: "(" prompt [":" numpar] ")"
|
||||||
+ "(?:"
|
deemphasized: "[" prompt "]"
|
||||||
+ re.escape(strParamStart)
|
scheduled: "[" [prompt ":"] prompt ":" numpar "]"
|
||||||
+ "([se]|(?:[pi][0-9]))"
|
alternate: "[" alternateoption ("|" alternateoption)+ "]"
|
||||||
+ re.escape(strParamEnd)
|
alternateoption: prompt
|
||||||
+ ")?(.*?)"
|
negtag: "<!" [negtagparameters] nonegprompt "!>"
|
||||||
+ escapeSequence
|
negtagparameters: "!" /s|e|[ip]\d/ "!"
|
||||||
+ re.escape(strEnd)
|
modeltag: "<" /(?!!)[^>]+/ ">"
|
||||||
+ ")",
|
numpar: WHITESPACE* NUMBER WHITESPACE*
|
||||||
re.S,
|
WHITESPACE: /\s+/
|
||||||
|
?plain: /([^\\[\]():|<>!]|\\.)+/s
|
||||||
|
%import common.SIGNED_NUMBER -> NUMBER
|
||||||
|
""",
|
||||||
|
propagate_positions=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
def processPrompt(self, original_prompt, original_negative_prompt):
|
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
|
Extract from the prompt the tagged parts and add them to the negative prompt
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
prompt = original_prompt
|
prompt = original_prompt
|
||||||
negative_prompt = original_negative_prompt
|
negative_prompt = original_negative_prompt
|
||||||
alreadyProcessed = []
|
self.__logger.debug(f"Input prompt: {prompt.encode('unicode_escape').decode('utf-8')}")
|
||||||
addAtStart = []
|
self.__logger.debug(f"Input negative_prompt: {negative_prompt.encode('unicode_escape').decode('utf-8')}")
|
||||||
addAtEnd = []
|
prompt, add_at = self.__find_tags(prompt)
|
||||||
addAtInsertionPoint = [[] for x in range(10)]
|
negative_prompt = self.__add_to_insertion_points(negative_prompt, add_at["insertion_point"])
|
||||||
# process tags in prompt
|
if len(add_at["start"]) > 0:
|
||||||
matches = self.regex.findall(prompt)
|
negative_prompt = self.__add_to_start(negative_prompt, add_at["start"])
|
||||||
for match in matches:
|
if len(add_at["end"]) > 0:
|
||||||
position = match[1] or "s"
|
negative_prompt = self.__add_to_end(negative_prompt, add_at["end"])
|
||||||
content = match[2]
|
self.__logger.debug(f"Output prompt: {prompt.encode('unicode_escape').decode('utf-8')}")
|
||||||
if len(content) > 0:
|
self.__logger.debug(f"Output negative_prompt: {negative_prompt.encode('unicode_escape').decode('utf-8')}")
|
||||||
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)
|
|
||||||
|
|
||||||
return prompt, negative_prompt
|
return prompt, negative_prompt
|
||||||
except Exception as e:
|
except Exception as e: # pylint: disable=broad-exception-caught
|
||||||
self.logger.exception(e)
|
self.__logger.exception(e)
|
||||||
return original_prompt, original_negative_prompt
|
return original_prompt, original_negative_prompt
|
||||||
|
|
||||||
|
def __find_tags(self, prompt):
|
||||||
|
add_at = {"start": [], "insertion_point": [[] for x in range(10)], "end": []}
|
||||||
|
tree = self.__schedule_parser.parse(prompt)
|
||||||
|
self.__logger.debug(f"Initial tree: {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):
|
||||||
|
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
|
||||||
|
|||||||
+81
-33
@@ -1,23 +1,22 @@
|
|||||||
|
import logging
|
||||||
import unittest
|
import unittest
|
||||||
import sys
|
import sys
|
||||||
import os
|
import os
|
||||||
|
|
||||||
sys.path.insert(1, os.path.join(sys.path[0], ".."))
|
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):
|
class TestSendToNegative(unittest.TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
self.defstn = SendToNegative(
|
self.defstn = SendToNegative(
|
||||||
tagStart="<!",
|
|
||||||
tagEnd="!>",
|
|
||||||
tagParamStart="!",
|
|
||||||
tagParamEnd="!",
|
|
||||||
separator=", ",
|
separator=", ",
|
||||||
ignoreRepeats=True,
|
ignore_repeats=True,
|
||||||
|
join_attention=True,
|
||||||
cleanup=True,
|
cleanup=True,
|
||||||
)
|
)
|
||||||
|
logging.basicConfig(level=logging.DEBUG)
|
||||||
|
|
||||||
def process(
|
def process(
|
||||||
self,
|
self,
|
||||||
@@ -27,20 +26,16 @@ class TestSendToNegative(unittest.TestCase):
|
|||||||
expected_negative_prompt,
|
expected_negative_prompt,
|
||||||
stn=None,
|
stn=None,
|
||||||
):
|
):
|
||||||
theObj = self.defstn if stn is None else stn
|
the_obj = self.defstn if stn is None else stn
|
||||||
result_prompt, result_negative_prompt = theObj.processPrompt(
|
result_prompt, result_negative_prompt = the_obj.process_prompt(prompt, negative_prompt)
|
||||||
prompt, negative_prompt
|
self.assertEqual(result_prompt, expected_prompt, f"Prompt should be '{expected_prompt}'")
|
||||||
)
|
|
||||||
self.assertEqual(
|
|
||||||
result_prompt, expected_prompt, f"Prompt should be '{expected_prompt}'"
|
|
||||||
)
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
result_negative_prompt,
|
result_negative_prompt,
|
||||||
expected_negative_prompt,
|
expected_negative_prompt,
|
||||||
f"Negative Prompt should be '{expected_negative_prompt}'",
|
f"Negative Prompt should be '{expected_negative_prompt}'",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_tagDefault(self):
|
def test_tag_default(self):
|
||||||
self.process(
|
self.process(
|
||||||
"flowers<!red!>",
|
"flowers<!red!>",
|
||||||
"normal quality, worse quality",
|
"normal quality, worse quality",
|
||||||
@@ -48,7 +43,7 @@ class TestSendToNegative(unittest.TestCase):
|
|||||||
"red, normal quality, worse quality",
|
"red, normal quality, worse quality",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_tagStart(self):
|
def test_tag_start(self):
|
||||||
self.process(
|
self.process(
|
||||||
"flowers<!!s!red!>",
|
"flowers<!!s!red!>",
|
||||||
"normal quality, worse quality",
|
"normal quality, worse quality",
|
||||||
@@ -56,7 +51,7 @@ class TestSendToNegative(unittest.TestCase):
|
|||||||
"red, normal quality, worse quality",
|
"red, normal quality, worse quality",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_tagEnd(self):
|
def test_tag_end(self):
|
||||||
self.process(
|
self.process(
|
||||||
"flowers<!!e!red!>",
|
"flowers<!!e!red!>",
|
||||||
"normal quality, worse quality",
|
"normal quality, worse quality",
|
||||||
@@ -64,7 +59,7 @@ class TestSendToNegative(unittest.TestCase):
|
|||||||
"normal quality, worse quality, red",
|
"normal quality, worse quality, red",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_tagInsertion_midSep(self):
|
def test_tag_insertion_mid_sep(self):
|
||||||
self.process(
|
self.process(
|
||||||
"flowers<!!p0!red!>",
|
"flowers<!!p0!red!>",
|
||||||
"normal quality, <!!i0!!>, worse quality",
|
"normal quality, <!!i0!!>, worse quality",
|
||||||
@@ -72,7 +67,7 @@ class TestSendToNegative(unittest.TestCase):
|
|||||||
"normal quality, red, worse quality",
|
"normal quality, red, worse quality",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_tagInsertion_midNoSep(self):
|
def test_tag_insertion_mid_no_sep(self):
|
||||||
self.process(
|
self.process(
|
||||||
"flowers<!!p0!red!>",
|
"flowers<!!p0!red!>",
|
||||||
"normal quality<!!i0!!>worse quality",
|
"normal quality<!!i0!!>worse quality",
|
||||||
@@ -80,7 +75,7 @@ class TestSendToNegative(unittest.TestCase):
|
|||||||
"normal quality, red, worse quality",
|
"normal quality, red, worse quality",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_tagInsertion_startSep(self):
|
def test_tag_insertion_start_sep(self):
|
||||||
self.process(
|
self.process(
|
||||||
"flowers<!!p0!red!>",
|
"flowers<!!p0!red!>",
|
||||||
"<!!i0!!>, normal quality, worse quality",
|
"<!!i0!!>, normal quality, worse quality",
|
||||||
@@ -88,7 +83,7 @@ class TestSendToNegative(unittest.TestCase):
|
|||||||
"red, normal quality, worse quality",
|
"red, normal quality, worse quality",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_tagInsertion_startNoSep(self):
|
def test_tag_insertion_start_no_sep(self):
|
||||||
self.process(
|
self.process(
|
||||||
"flowers<!!p0!red!>",
|
"flowers<!!p0!red!>",
|
||||||
"<!!i0!!>normal quality, worse quality",
|
"<!!i0!!>normal quality, worse quality",
|
||||||
@@ -96,7 +91,7 @@ class TestSendToNegative(unittest.TestCase):
|
|||||||
"red, normal quality, worse quality",
|
"red, normal quality, worse quality",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_tagInsertion_endSep(self):
|
def test_tag_insertion_end_sep(self):
|
||||||
self.process(
|
self.process(
|
||||||
"flowers<!!p0!red!>",
|
"flowers<!!p0!red!>",
|
||||||
"normal quality, worse quality, <!!i0!!>",
|
"normal quality, worse quality, <!!i0!!>",
|
||||||
@@ -104,7 +99,7 @@ class TestSendToNegative(unittest.TestCase):
|
|||||||
"normal quality, worse quality, red",
|
"normal quality, worse quality, red",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_tagInsertion_endNoSep(self):
|
def test_tag_insertion_end_no_sep(self):
|
||||||
self.process(
|
self.process(
|
||||||
"flowers<!!p0!red!>",
|
"flowers<!!p0!red!>",
|
||||||
"normal quality, worse quality<!!i0!!>",
|
"normal quality, worse quality<!!i0!!>",
|
||||||
@@ -114,29 +109,82 @@ class TestSendToNegative(unittest.TestCase):
|
|||||||
|
|
||||||
def test_complex(self):
|
def test_complex(self):
|
||||||
self.process(
|
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",
|
"normal quality, <!!i0!!>, bad quality<!!i1!!>, worse quality",
|
||||||
"flowers",
|
"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(
|
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",
|
"normal quality, <!!i0!!>, bad quality<!!i1!!>, worse quality",
|
||||||
" , flowers, , ",
|
" (), flowers , , ",
|
||||||
"red, pink, normal quality, yellow, bad quality, green, worse quality, blue",
|
"red, (pink), normal quality, yellow, bad quality, green, worse quality, purple, blue",
|
||||||
SendToNegative(
|
SendToNegative(
|
||||||
tagStart="<!",
|
|
||||||
tagEnd="!>",
|
|
||||||
tagParamStart="!",
|
|
||||||
tagParamEnd="!",
|
|
||||||
separator=", ",
|
separator=", ",
|
||||||
ignoreRepeats=True,
|
ignore_repeats=True,
|
||||||
|
join_attention=True,
|
||||||
cleanup=False,
|
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]",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user