Compare commits
276
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
136932de40 | ||
|
|
06a3f43a93 | ||
|
|
203a7ad45c | ||
|
|
7a566e6e9e | ||
|
|
fa288c226c | ||
|
|
e648f3bdc3 | ||
|
|
1bedc6ad53 | ||
|
|
de79f3c6af | ||
|
|
eb51fd9289 | ||
|
|
7cca9438f2 | ||
|
|
7d786dfe83 | ||
|
|
85c15ff6bd | ||
|
|
1ab1c87f74 | ||
|
|
2ea0b622b4 | ||
|
|
c4c9561a13 | ||
|
|
ff111a9f7c | ||
|
|
ab1cf5949e | ||
|
|
105562a5bb | ||
|
|
e337ecc019 | ||
|
|
b2bb7b9960 | ||
|
|
f1ade19345 | ||
|
|
15c115c7f3 | ||
|
|
97f050c264 | ||
|
|
ca5848a8aa | ||
|
|
24240aac2d | ||
|
|
883ea6a96d | ||
|
|
49b761ae5b | ||
|
|
bd0dd00d8c | ||
|
|
3f303cb33a | ||
|
|
d009bb782b | ||
|
|
eaf5fa20c4 | ||
|
|
2becb5cdc6 | ||
|
|
e1b12acbdf | ||
|
|
3ea605fbe0 | ||
|
|
f4d5477d4c | ||
|
|
5cb18ab5a9 | ||
|
|
fc5c0a1857 | ||
|
|
f9769fe90c | ||
|
|
09438a35f4 | ||
|
|
d0d34f8b2f | ||
|
|
dcc9379d40 | ||
|
|
9231ec4936 | ||
|
|
0473d90446 | ||
|
|
138df090a6 | ||
|
|
c89dcb135d | ||
|
|
c53717101c | ||
|
|
d185c72a16 | ||
|
|
805ffbf463 | ||
|
|
d14f5b789a | ||
|
|
deba0a0642 | ||
|
|
f98bf25a83 | ||
|
|
2c8727a75a | ||
|
|
86ec30c028 | ||
|
|
e0efac1ebc | ||
|
|
68766215f2 | ||
|
|
287f554a68 | ||
|
|
71465f914c | ||
|
|
d7be7bc29e | ||
|
|
329d4cf95f | ||
|
|
c52ace71aa | ||
|
|
4806cf5959 | ||
|
|
a0ab709f50 | ||
|
|
66fef1ffe8 | ||
|
|
3341e9f81e | ||
|
|
aa6b4608f0 | ||
|
|
d1cc60b00a | ||
|
|
efe8939250 | ||
|
|
39b353b916 | ||
|
|
485bc7f2ab | ||
|
|
1d03ded9dd | ||
|
|
94a4d076e0 | ||
|
|
228dc4b22b | ||
|
|
db523e1f16 | ||
|
|
6b08c7a90e | ||
|
|
76142c4b7e | ||
|
|
1d84fdaf9e | ||
|
|
1c50ae5297 | ||
|
|
51618289e7 | ||
|
|
cea1e5b30f | ||
|
|
55f0574ac7 | ||
|
|
167689cb8b | ||
|
|
2be3abed44 | ||
|
|
f86abb0816 | ||
|
|
a3537a5b2f | ||
|
|
af7e4542d1 | ||
|
|
f37f14b2a2 | ||
|
|
7b8231d36b | ||
|
|
9b5e15fde3 | ||
|
|
d97d30074f | ||
|
|
b2eb9b88ba | ||
|
|
c9459e39f9 | ||
|
|
7507b2b55f | ||
|
|
cf1efecf4c | ||
|
|
ffcf94bcaa | ||
|
|
ec8c40355c | ||
|
|
6e538e0abc | ||
|
|
25c44a1fbb | ||
|
|
72d5490498 | ||
|
|
3de4538326 | ||
|
|
f61af15d52 | ||
|
|
8e59f140ff | ||
|
|
44044e962c | ||
|
|
04f36687c5 | ||
|
|
3a8a360d03 | ||
|
|
d76331315a | ||
|
|
e3a6050536 | ||
|
|
7b001ace7b | ||
|
|
f1de65f257 | ||
|
|
11aaa0ac7b | ||
|
|
85de3ef0d3 | ||
|
|
a73260ff34 | ||
|
|
50a2e0abbf | ||
|
|
5cf45ca264 | ||
|
|
bd4a787400 | ||
|
|
c9e5bc25c3 | ||
|
|
d11ffa6e25 | ||
|
|
8cc73a2e49 | ||
|
|
27ae5f683e | ||
|
|
68cda3663e | ||
|
|
3d46f705b6 | ||
|
|
eec4bc4da9 | ||
|
|
3d3218e831 | ||
|
|
f4e57ec514 | ||
|
|
74f65c1b31 | ||
|
|
4d94cca88f | ||
|
|
4569ecccf9 | ||
|
|
192e6d30d4 | ||
|
|
67112f11e0 | ||
|
|
f761dfac86 | ||
|
|
278e733835 | ||
|
|
2bf65720eb | ||
|
|
99f3af92b7 | ||
|
|
2437cd4daf | ||
|
|
d262a7dc7a | ||
|
|
6c319ad5b4 | ||
|
|
34056cac19 | ||
|
|
8ae436abf1 | ||
|
|
c2ce2ce023 | ||
|
|
7a9e69ec31 | ||
|
|
aaff8dc7da | ||
|
|
3e4722278a | ||
|
|
11cb430396 | ||
|
|
9f1cbfd11c | ||
|
|
5b3a914f1d | ||
|
|
a5ffa1acd7 | ||
|
|
7f8783147b | ||
|
|
3c9b806e5f | ||
|
|
a485c2655a | ||
|
|
f888e69b00 | ||
|
|
b4b0858214 | ||
|
|
7a1cb2cf51 | ||
|
|
fbb6b5c8fa | ||
|
|
6115c095cb | ||
|
|
04d3d2e959 | ||
|
|
e4f64837ef | ||
|
|
4a785b294b | ||
|
|
b3195a6297 | ||
|
|
5c52bffc9d | ||
|
|
ac2d275dfb | ||
|
|
399f992a26 | ||
|
|
672a2a09bb | ||
|
|
e3629961ce | ||
|
|
ddac624ad1 | ||
|
|
99ddfe357e | ||
|
|
110d5248a0 | ||
|
|
61b1ecc88e | ||
|
|
05b0b2ad26 | ||
|
|
5831608c4e | ||
|
|
c815bb44f1 | ||
|
|
e55c50e9d7 | ||
|
|
1b0ff62d10 | ||
|
|
cf93093d59 | ||
|
|
fc15a89a2f | ||
|
|
57c092bccf | ||
|
|
88f77a8124 | ||
|
|
a9c2487c0c | ||
|
|
75bced7d2b | ||
|
|
4b285be07e | ||
|
|
200d9f9daf | ||
|
|
d33208b1c3 | ||
|
|
ffa64816c0 | ||
|
|
6ffbf05d7d | ||
|
|
892a70d53b | ||
|
|
20711358a2 | ||
|
|
453580545c | ||
|
|
98d78df7ba | ||
|
|
bdd56410dc | ||
|
|
0289564e55 | ||
|
|
1e05d1a8cc | ||
|
|
dc6fd0fc63 | ||
|
|
a356bddcc7 | ||
|
|
b8081e5736 | ||
|
|
aa00c26365 | ||
|
|
e913bad73c | ||
|
|
5c1b739b82 | ||
|
|
5e3ab1f51a | ||
|
|
b21de76cd5 | ||
|
|
e4a27d01ee | ||
|
|
42cdfa0f5a | ||
|
|
98292e2bc8 | ||
|
|
f0c8e2e873 | ||
|
|
c4ac37333d | ||
|
|
2534e002ad | ||
|
|
fd4823fd75 | ||
|
|
a6f230ff8b | ||
|
|
633b2f05e0 | ||
|
|
a15135ddc5 | ||
|
|
0d7e2a4e60 | ||
|
|
c5495832c5 | ||
|
|
63d2cb3e0c | ||
|
|
95832e801b | ||
|
|
01bd5568d0 | ||
|
|
36b3638f4e | ||
|
|
d46000ef78 | ||
|
|
aba246a33c | ||
|
|
42ae22db83 | ||
|
|
cb6de285cb | ||
|
|
49a073bb12 | ||
|
|
e9afe779ae | ||
|
|
fa3b4f7da3 | ||
|
|
7a76cc8c72 | ||
|
|
6b1e2a5a8a | ||
|
|
9aee531c09 | ||
|
|
c08bf395a6 | ||
|
|
1964708997 | ||
|
|
1c4b5ce0c4 | ||
|
|
5eabbb419c | ||
|
|
53400a029b | ||
|
|
c39605eec4 | ||
|
|
dc62e638ed | ||
|
|
109cac16ef | ||
|
|
cf6c2b3e6a | ||
|
|
d113d4ba78 | ||
|
|
5bd1d04dcd | ||
|
|
7c10770e07 | ||
|
|
69ea298174 | ||
|
|
4ba4b28bb2 | ||
|
|
f728866b90 | ||
|
|
306c02f57b | ||
|
|
2c519310ac | ||
|
|
3f23d1b14a | ||
|
|
148776fe5d | ||
|
|
cef4a80440 | ||
|
|
cd642b5d42 | ||
|
|
2b323da9a9 | ||
|
|
79b3675c4f | ||
|
|
b8d5b7a7c4 | ||
|
|
a5da586dc5 | ||
|
|
01aa061bef | ||
|
|
2fab4be810 | ||
|
|
127acb7018 | ||
|
|
b952b2f186 | ||
|
|
2732a795fb | ||
|
|
4f78fff892 | ||
|
|
e290cb57ac | ||
|
|
d577d439e7 | ||
|
|
e59d46c8d1 | ||
|
|
bd1c69a517 | ||
|
|
cee19aea67 | ||
|
|
b93bb66aed | ||
|
|
c75d1a6651 | ||
|
|
a4c7f99cc1 | ||
|
|
b7d544c05c | ||
|
|
04c4bd0846 | ||
|
|
5365679a60 | ||
|
|
fa77c158ac | ||
|
|
525cb157ce | ||
|
|
e10950e4da | ||
|
|
91ba4c881f | ||
|
|
106ebe49aa | ||
|
|
549b4347fd | ||
|
|
4cbce5df06 | ||
|
|
21208bd733 | ||
|
|
a4065415e7 | ||
|
|
0d546f1a08 | ||
|
|
390e1ec6b8 |
@@ -7,15 +7,23 @@ on:
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
tests:
|
||||
uses: ./.github/workflows/tests.yml
|
||||
tests_with_comfy:
|
||||
uses: ./.github/workflows/tests_with_comfy.yml
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'asagi4' }}
|
||||
needs: [tests, tests_with_comfy]
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
name: Run parser tests
|
||||
on:
|
||||
- workflow_call
|
||||
- workflow_dispatch
|
||||
- push
|
||||
|
||||
jobs:
|
||||
run-parser-tests:
|
||||
name: Run parser tests
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Check out ComfyUI
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
repository: comfyanonymous/ComfyUI
|
||||
path: ComfyUI
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.11'
|
||||
- run: pip install pytest typing-extensions -r requirements.txt
|
||||
- run: PYTHONPATH=ComfyUI pytest tests/test_parser.py
|
||||
@@ -0,0 +1,40 @@
|
||||
name: Run tests requiring ComfyUI
|
||||
on:
|
||||
workflow_call:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
paths:
|
||||
- prompt_control/adv_encode.py
|
||||
- prompt_control/attention_couple_ppm.py
|
||||
- prompt_control/nodes_lazy.py
|
||||
- prompt_control/prompts.py
|
||||
- prompt_control/parser.py
|
||||
- prompt_control/utils.py
|
||||
|
||||
|
||||
jobs:
|
||||
run-graph-tests:
|
||||
name: Run tests requiring ComfyUI
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Check out ComfyUI
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
repository: comfyanonymous/ComfyUI
|
||||
path: ComfyUI
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.11'
|
||||
cache: pip
|
||||
- name: install-torch
|
||||
run: pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu
|
||||
- name: install ComfyUI
|
||||
run: pip install pytest typing-extensions -r requirements.txt -r ComfyUI/requirements.txt
|
||||
- name: Download clip_l.safetensors
|
||||
run: curl -LO https://huggingface.co/comfyanonymous/flux_text_encoders/resolve/main/clip_l.safetensors
|
||||
- name: Force Comfy to use the CPU
|
||||
run: sed -i "s/^cpu_state = CPUState.GPU/cpu_state = CPUState.CPU/g" ComfyUI/comfy/model_management.py
|
||||
- name: Run graph tests
|
||||
run: PYTHONPATH=ComfyUI pytest tests/test_graph.py tests/test_encode.py
|
||||
@@ -1 +1,2 @@
|
||||
__pycache__
|
||||
.pyre
|
||||
|
||||
@@ -1,8 +1,31 @@
|
||||
all: format check
|
||||
ARGS=
|
||||
all: format check test
|
||||
@echo "Done"
|
||||
|
||||
check:
|
||||
find . -name "*.py" | xargs pyflakes
|
||||
ty check && ruff check
|
||||
|
||||
fix:
|
||||
ruff check --fix
|
||||
|
||||
format:
|
||||
find . -name "*.py" | xargs black -l 120
|
||||
ruff format
|
||||
|
||||
test:
|
||||
PYTHONPATH=../../ pytest tests/test_parser.py tests/test_cutout.py $(ARGS)
|
||||
|
||||
test_graph:
|
||||
PYTHONPATH=../../ pytest tests/test_graph.py $(ARGS)
|
||||
|
||||
test_encode:
|
||||
PYTHONPATH=../../ pytest tests/test_encode.py $(ARGS)
|
||||
|
||||
test_encode_both:
|
||||
TEST_TE="clip_l t5" PYTHONPATH=../../ pytest tests/test_encode.py $(ARGS)
|
||||
|
||||
test_heavy: test_graph test_encode_both
|
||||
|
||||
manual_test:
|
||||
PYTHONPATH=../../ python -im prompt_control.manual_test
|
||||
|
||||
.PHONY: check format all
|
||||
|
||||
@@ -2,63 +2,47 @@
|
||||
|
||||
Control LoRA and prompt scheduling, advanced text encoding, regional prompting, and much more, through your text prompt. Generates dynamic graphs that are literally identical to handcrafted noodle soup.
|
||||
|
||||
## Prompt Control v2
|
||||
Prompt Control comes with `PCTextEncode`, which provides advanced text encoding with many additional features compared to ComfyUI's base `CLIPTextEncode`.
|
||||
|
||||
Prompt control has been almost completely rewritten. It now uses ComfyUI's lazy execution to build graphs from the text prompt at runtime. The generated graph is often exactly equivalent to a manually built workflow using native ComfyUI nodes. There are no more weird sampling hooks that could cause problems with other nodes
|
||||
A `Basic Text to Image` template is included with the extension, and can be loaded from ComfyUI's template library.
|
||||
|
||||
Prompt Control also comes with `PCTextEncode`, which provides advanced text encoding with many additional features compared to ComfyUI's base `CLIPTextEncode`.
|
||||
> [!NOTE]
|
||||
> v3.0.0 is backwards compatible with existing workflows, but requires at least ComfyUI v0.8.0
|
||||
> The parser was rewritten using parsy. It is intended to have the same behaviour as the old parser, but is **significantly** faster.
|
||||
> Please report any bugs or incompatibilities you find.
|
||||
|
||||
### Removed features
|
||||
|
||||
- Prompt interpolation syntax; it was too cumbersome to maintain
|
||||
- LoRA block weight integration; ditto, for now.
|
||||
|
||||
|
||||
### Is it stable now?
|
||||
|
||||
Unless I run into bugs or significant annoyances that require changing the interface, it probably won't change too much, but until I tag 2.0, everything can change.
|
||||
|
||||
### Everything broke, where are the old nodes?
|
||||
|
||||
If you really need them, you can install the [legacy nodes](https://github.com/asagi4/comfyui-prompt-control-legacy). However, I will not fix bugs in those nodes, and I strongly recommend just migrating your workflows to the new nodes.
|
||||
|
||||
You can have both installed at the same time; none of the nodes conflict.
|
||||
|
||||
## What can it do?
|
||||
|
||||
See [features](#features) below. Things you can control via the prompt:
|
||||
- Prompt editing and filtering without noodle soup
|
||||
- LoRA loading and scheduling via ComfyUI's hook system
|
||||
- Masking, composition and area control (regional prompting)
|
||||
- Prompt operations like `BREAK` and `AND`
|
||||
- Weight interpretation types (comfy, A1111, etc.)
|
||||
- Prompt masking with [cutoff](#cutoff)
|
||||
- And a bunch more
|
||||
You can use text prompts to control the following:
|
||||
|
||||
See the [syntax documentation](doc/syntax.md)
|
||||
- A1111-style prompt scheduling and filtering without noodle soup.
|
||||
- LoRA loading and [scheduling](/doc/schedules.md) via the prompt, using ComfyUI's hook system
|
||||
- Masking, composition and area control ([regional prompting](/doc/regional_prompts.md)) with an implementation of [Attention Couple](/doc/attention_couple.md), also fully schedulable.
|
||||
- [Advanced prompt encoding](/doc/basic.md)
|
||||
- Per-encoder prompts for models with multiple text encoders, such as SDXL and Flux
|
||||
- Prompt combinators like `BREAK`, as well as `CAT`, `AVG()` and `AND` corresponding to ComfyUI's `ConditioningConcat`, `ConditioningAverage` and `ConditioningCombine` nodes.
|
||||
- Different weight interpretation types (ComfyUI, A1111, compel, etc.)
|
||||
- Prompt masking with an implementation of [cutoff](https://github.com/BlenderNeko/ComfyUI_Cutoff)
|
||||
- Simple [prompt macros](/doc/macros.md) with `DEF`
|
||||
|
||||
All features are fully schedulable unless otherwise stated. See the [scheduling syntax documentation](doc/schedules.md) to get started.
|
||||
|
||||
If you find prompt scheduling inconvenient for some reason, `PCTextEncode` can be used as a drop-in replacement for `CLIPTextEncode` to get everything else.
|
||||
|
||||
[This workflow](example_workflows/Workflow%20Comparison.json?raw=1) shows LoRA scheduling and prompt editing and compares it with the same prompt implemented with built-in ComfyUI nodes. You can also find it in the template library.
|
||||
|
||||
[This workflow](workflows/example-lazy.json?raw=1) shows LoRA scheduling and prompt editing and compares it with the same prompt implemented with built-in ComfyUI nodes.
|
||||
## Compatibility
|
||||
|
||||
[Here](workflows/example-2pass.json?raw=1) is a two-pass workflow illustrating more features, including custom masks and filtering.
|
||||
|
||||
The tools in this repository combine well with the macro and wildcard functionality in [comfyui-utility-nodes](https://github.com/asagi4/comfyui-utility-nodes)
|
||||
Prompt Control uses graph generation, and tries to delegate functionality to core ComfyUI wherever possible, implementing any hooks and patches in a way that is maximally compatible. This means that it should just work in most cases, even with models and nodes not explicitly supported.
|
||||
|
||||
If you encounter issues as a user or if you're a node developer and Prompt Control somehow breaks something, feel free to file a bug report.
|
||||
|
||||
## Requirements
|
||||
|
||||
For LoRA scheduling to work, you'll need at least version 0.3.7 of ComfyUI (0.3.36 of ComfyUI desktop).
|
||||
The v3 node schema uses features that require at least ComfyUI v0.8.0
|
||||
|
||||
You need to have `lark` installed in your Python environment for parsing to work (If you reuse A1111's venv, it'll already be there).
|
||||
|
||||
If you use the portable version of ComfyUI on Windows with its embedded Python, you must open a terminal in the ComfyUI installation directory and run the command:
|
||||
```
|
||||
.\python_embeded\python.exe -m pip install lark
|
||||
```
|
||||
|
||||
Then restart ComfyUI afterwards.
|
||||
If you run into problems, update ComfyUI first.
|
||||
|
||||
# Core nodes
|
||||
|
||||
@@ -96,79 +80,8 @@ This node attaches masks to a `CLIP` model so that they can be referred to when
|
||||
|
||||
This node configures `PCTextEncode` default values for some functions by attaching the information to a `CLIP` model.
|
||||
|
||||
# Features
|
||||
## Scheduling and LoRA loading
|
||||
|
||||
Prompt control provides a way to easily schedule different prompts and control LoRA loading.
|
||||
|
||||
See the [syntax documentation](doc/syntax.md)
|
||||
|
||||
### Note on how schedules work
|
||||
|
||||
ComfyUI does not use the step number to determine whether to apply conds; instead, it uses the sampler's timestep value which is affected by the scheduler you're using. This means that when the sampler scheduler isn't linear, the schedules generated by prompt control will not be either.
|
||||
|
||||
## Advanced CLIP encoding
|
||||
|
||||
If you use `PCTextEncode`, advanced encodings are available automatically. Thanks to BlenderNeko for the original code.
|
||||
|
||||
Use the syntax `STYLE(weight_interpretation, normalization)` in a prompt to affect how prompts are interpreted.
|
||||
|
||||
The weight interpretations available are:
|
||||
- comfy (default)
|
||||
- comfy++
|
||||
- compel
|
||||
- down_weight
|
||||
- A1111
|
||||
- perp
|
||||
|
||||
Normalizations are:
|
||||
- none (default)
|
||||
- length
|
||||
- mean
|
||||
|
||||
The normalization calculations are independent operations and you can combine them with `+`, eg `STYLE(A1111, length+mean)` or `STYLE(comfy, mean+length)`, or even something silly like `STYLE(perp, mean+length+mean+length)`
|
||||
|
||||
The style can be specified separately for each AND:ed prompt, but the first prompt is special; later prompts will "inherit" it as default. For example:
|
||||
|
||||
```
|
||||
STYLE(A1111) a (red:1.1) cat with (brown:0.9) spots and a long tail AND an (old:0.5) dog AND a (green:1.4) (balloon:1.1)
|
||||
```
|
||||
will interpret everything as A1111, but
|
||||
```
|
||||
a (red:1.1) cat with (brown:0.9) spots and a long tail AND STYLE(A1111) an (old:0.5) dog AND a (green:1.4) (balloon:1.1)
|
||||
```
|
||||
Will interpret the first one using the default ComfyUI behaviour, the second prompt with A1111 and the last prompt with the default again
|
||||
|
||||
For things (ie. the code imports) to work, the nodes must be cloned in a directory named exactly `ComfyUI_ADV_CLIP_emb`.
|
||||
|
||||
## Cutoff
|
||||
|
||||
NOTE: Cutoff syntax might change at some point; it's pretty clunky.
|
||||
|
||||
`PCTextEncode` reimplements cutoff from [ComfyUI Cutoff](https://github.com/BlenderNeko/ComfyUI_Cutoff).
|
||||
|
||||
The syntax is
|
||||
```
|
||||
a group of animals, [CUT:white cat:white], [CUT:brown dog:brown:0.5:1.0:1.0:_]
|
||||
```
|
||||
You should read the prompt as `a group of animals, white cat, brown dog`, but CUT causes the tokens in `target_tokens` to be masked off from the base prompt in `region_text`, so that their effect can be isolated, and you're less likely to get brown cats or white dogs.
|
||||
|
||||
Target tokens are treated individually, separated by space, for example, `[CUT:green apple, red apple, green leaf:green apple]` will mask *both* greens and the apple, giving you `+ +, red +, + leaf`. To mask out just `green apple`, use `[CUT:green apple, red apple:green_apple]` which will result in a masked prompt of `+ +, red apple`. Escape `_` with a `\`.
|
||||
|
||||
the parameters in the `CUT` section are `region_text:target_tokens:weight;strict_mask:start_from_masked:padding_token` of which only the first two are required. The default values are `weight=1.0`, `strict_mask=1.0` `start_from_masked=1.0`, `padding_token=+`
|
||||
|
||||
If `strict_mask`, `start_from_masked` or `padding_token` are specified in more than one CUT, the *last* one becomes the default for any CUTs afterwards that do not explicitly set the parameters. For example, in:
|
||||
|
||||
`[CUT:white cat:white:0.5] and [CUT:black parrot, flying:black:1.0:0.5] and [CUT:green apple:green]`
|
||||
|
||||
`white cat` will a weight of 0.5, and 1.0 for all parameters, and `black parrot` and `green apple` will *both* have a `strict_mask` parameter of 0.5.
|
||||
|
||||
The parameters affect how the masked and unmasked prompts are combined to produce the final embedding. Just play around with them.
|
||||
|
||||
# Known issues
|
||||
|
||||
- ComfyUI's caching mechanism has an issue that makes it unnecessarily invalidate caches for certain inputs; you'll still get some benefit from the lazy nodes, but changing inputs that shouldn't affect downstream nodes (especially if using filtering) will still cause them to be recomputed because ComfyUI doesn't realize the inputs haven't changed.
|
||||
|
||||
If you want to enable a hack to fix this, set `PROMPTCONTROL_ENABLE_CACHE_HACK=1` in your environment. Unset it to disable.
|
||||
|
||||
It's a purely optional performance optimization that allows Prompt Control nodes to override their cache keys in a way that should not interfere with other nodes. Note that the optimization only works if the text input to the lazy nodes is a constant (so either directly on the node or from a primitive); outputs from other nodes can't be optimized.
|
||||
- Cutoff does not work with models that use non-CLIP text encoders, like Flux. This might be fixable, but it's uncertain if cutoff even makes sense for those models.
|
||||
|
||||
+24
-29
@@ -5,46 +5,41 @@
|
||||
@description: Control LoRA and prompt scheduling, advanced text encoding, regional prompting, and much more, through your text prompt. Generates dynamic graphs that are literally identical to handcrafted noodle soup.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import logging
|
||||
import importlib
|
||||
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
log.propagate = False
|
||||
if not log.handlers:
|
||||
h = logging.StreamHandler(sys.stdout)
|
||||
h.setFormatter(logging.Formatter("[PromptControl] %(levelname)s: %(message)s"))
|
||||
log.addHandler(h)
|
||||
|
||||
if os.environ.get("PROMPTCONTROL_DEBUG"):
|
||||
log.setLevel(logging.DEBUG)
|
||||
else:
|
||||
log.setLevel(logging.INFO)
|
||||
|
||||
cache_hack = importlib.import_module(".prompt_control.cache_hack", package=__name__)
|
||||
cache_hack.init()
|
||||
WEB_DIRECTORY = "web"
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
v1_modules = []
|
||||
v3_modules = []
|
||||
# Importing things here breaks pytest for whatever reason...
|
||||
if "PYTEST_CURRENT_TEST" not in os.environ:
|
||||
import importlib
|
||||
|
||||
nodes = ["base", "lazy", "tools"]
|
||||
optional_nodes = ["attnmask"]
|
||||
if importlib.util.find_spec("comfy.hooks"):
|
||||
nodes.extend(["hooks"])
|
||||
else:
|
||||
log.error("Your ComfyUI version is too old, can't import comfy.hooks. Update your installation.")
|
||||
from comfy_api.latest import ComfyExtension
|
||||
|
||||
for node in nodes:
|
||||
mod = importlib.import_module(f".prompt_control.nodes_{node}", package=__name__)
|
||||
NODE_CLASS_MAPPINGS.update(mod.NODE_CLASS_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(mod.NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
for node in optional_nodes:
|
||||
try:
|
||||
if not log.handlers:
|
||||
h = logging.StreamHandler(sys.stdout)
|
||||
h.setFormatter(logging.Formatter("[PromptControl] %(levelname)s: %(message)s"))
|
||||
log.addHandler(h)
|
||||
for node in ["base", "hooks", "tools", "lazy"]:
|
||||
mod = importlib.import_module(f".prompt_control.nodes_{node}", package=__name__)
|
||||
NODE_CLASS_MAPPINGS.update(mod.NODE_CLASS_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(mod.NODE_DISPLAY_NAME_MAPPINGS)
|
||||
except ImportError:
|
||||
log.info(f"Could not import optional nodes: {node}; continuing anyway")
|
||||
v3_modules.append(mod)
|
||||
|
||||
class PromptControlExtension(ComfyExtension):
|
||||
async def get_node_list(self):
|
||||
r = []
|
||||
for m in v3_modules:
|
||||
r.extend(m.NODES)
|
||||
return r
|
||||
|
||||
async def comfy_entrypoint():
|
||||
return PromptControlExtension()
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
# Attention Couple
|
||||
|
||||
NOTE: This is still considered an experimental feature, so the syntax may change.
|
||||
|
||||
Attention Couple is an attention-based implementation of regional prompting. it is faster and often more flexible than latent-based masking.
|
||||
|
||||
The implementation is based on the one by [pamparamm](https://github.com/pamparamm/ComfyUI-ppm.git), modified to use ComfyUI's hook system. This enables it to work with prompt scheduling.
|
||||
|
||||
By default, the implementation produces slightly different results from Pamparamm's implementation because ComfyUI will only run the hook for conds that have it attached and can't batch negative conditionings.
|
||||
|
||||
As a consequence of this, however, you can also use `COUPLE` in your negative prompt, and it will work correctly.
|
||||
|
||||
To enable batching negative prompts, run your positive and negative prompt through the `PPCAttentionCoupleBatchNegative` node. This will make the outputs identical to pamparamm's implementation and will also improve performance. It will fall back to the default behaviour in cases where batching can't be done, so it should always be safe to use.
|
||||
|
||||
|
||||
## Syntax
|
||||
|
||||
See also the [regional prompting documentation](/doc/regional_prompts.md) for information about `MASK` etc.
|
||||
|
||||
### COUPLE: Trigger Attention Couple
|
||||
|
||||
You can use `COUPLE` to attach attention-coupled prompts to a base prompt:
|
||||
|
||||
`base_prompt COUPLE MASK(0 0.5) coupled prompt 1 with mask COUPLE IMASK(0) coupled prompt 2 with custom mask`
|
||||
|
||||
as a shortcut, `COUPLE(maskparams)` is expanded to `COUPLE MASK(maskparams)`, so the above prompt can also be written as:
|
||||
|
||||
`base_prompt COUPLE(0 0.5) coupled prompt 1 with mask COUPLE IMASK(0) coupled prompt 2 with custom mask`
|
||||
|
||||
Behaviour:
|
||||
- If no mask is specified, an implicit `MASK()` is assumed.
|
||||
|
||||
- For the base prompt, you can also use `FILL()` to automatically mask all parts not masked by coupled prompts
|
||||
|
||||
- If the base prompt has weight set to zero (ie. ´:0` at the end), then the first coupled prompt with non-zero weight becomes the base prompt.
|
||||
|
||||
For example:
|
||||
```
|
||||
dog FILL() COUPLE(0.5 1) cat
|
||||
```
|
||||
|
||||
Note that because the generation still sees and diffuses the full latent, attention coupling is not guaranteed to perfectly limit the effect of your prompt to the masked area.
|
||||
+210
@@ -0,0 +1,210 @@
|
||||
# Basic Prompt Syntax
|
||||
|
||||
The syntax below documents the features of `PCTextEncode`
|
||||
|
||||
## Combining prompts
|
||||
|
||||
### AND
|
||||
|
||||
`AND` can be used to create "prompt segments". By default, it works as if you had combined the different prompts with `ConditioningCombine`.
|
||||
|
||||
It is also used with regional prompting, see `MASK` and `COUPLE` below.
|
||||
|
||||
Prompts can have a weight at the end:
|
||||
```
|
||||
cat :1 AND dog :2
|
||||
```
|
||||
`AND` is processed after schedule parsing, so you can change the weight mid-prompt: `cat:[1:2:0.5] AND dog`
|
||||
|
||||
The weight defaults to 1. If a prompt's weight is set to 0, it's **skipped entirely.** This can be useful when scheduling to completely disable a prompt:
|
||||
|
||||
```
|
||||
cat [\:0::0.5] AND dog
|
||||
```
|
||||
Note that the `:` needs to be escaped with a `\` or it will be interpreted as scheduling syntax.
|
||||
|
||||
If `AND` is placed inside quotes (eg. `Text saying "CAT AND DOG"`) it will be treated as regular text.
|
||||
|
||||
## Note about processing order
|
||||
|
||||
Prompt operators are processed in the following order, meaning that all features "below" another can be affected by the feature above it. That is, `BREAK` can go inside a `TE()` call, but not `AND` or `CAT`.
|
||||
|
||||
- DEF macros are expanded
|
||||
- Scheduling is expanded, and for each scheduled prompt:
|
||||
- The prompt is split by AND, and for each:
|
||||
- Prompts are split by COUPLE. and for each:
|
||||
- Most functions (like MASK) and cutoffs are evaluated
|
||||
- prompts are split by `AVG()` or CAT
|
||||
- the TE() function is evaluated to set per-encoder prompts
|
||||
- BREAK is evaluated
|
||||
- Everything else
|
||||
- Prompts are combined with `ConditioningAverage` (for `AVG`) or `ConditioningConcat` (for `CAT`)
|
||||
- If coupled prompts exist, the base cond is set up for attention coupling and returned
|
||||
- Prompts split with `AND` are combined with `ConditioningCombine`
|
||||
- Each scheduled prompt is restricted to its effective range with `ConditioningSetTimestepRange`
|
||||
|
||||
## Functions
|
||||
|
||||
There are some "functions" that can be included in a prompt to affect how it is interpreted.
|
||||
|
||||
Functions have the form `FUNCNAME(param1, param2, ...)`. How parameters are interpreted is up to the function.
|
||||
|
||||
In general, function parameters will have default values that are used if the parameter is left empty.
|
||||
|
||||
Note: Whitespace is usually *not* stripped from string parameters by default. Commas can be escaped with `\,`
|
||||
|
||||
Like `AND`, functions are parsed after regular scheduling syntax has been expanded, allowing things like `[AREA:MASK:0.3](...)`, in case that's somehow useful.
|
||||
|
||||
like AND, if any function is placed inside quotes, it will *not* activate and is instead treated as regular text.
|
||||
|
||||
### BREAK
|
||||
The keyword `BREAK` causes the prompt to be tokenized in separate chunks, padding each chunk to the text encoder's maximum size before encoding.
|
||||
|
||||
For some text encoders (like t5), this operation doesn't really make sense and BREAKs are simply ignored.
|
||||
|
||||
### CAT
|
||||
|
||||
`CAT` encodes each prompt separately before concatenating the resulting tensors into a single conditioning. It behaves identically to ComfyUI's `ConditioningConcat`.
|
||||
|
||||
### AVG()
|
||||
|
||||
`prompt1 AVG(weight) prompt2` encodes prompt1 and prompt2 separately, and then combines them using `ConditioningAverage`. The default for `weight` is `0.5`.
|
||||
|
||||
`AVG` is processed before `BREAK` but after `AND`
|
||||
|
||||
`p1 AVG() p2 AVG() p3` combines `p1` and `p2` first, then combines the result with `p3`.
|
||||
|
||||
## Prompt weighting (also known as "Advanced CLIP Encode")
|
||||
|
||||
### STYLE
|
||||
|
||||
Use the syntax `STYLE(weight_interpretation, normalization)` in a prompt to affect how prompts are interpreted.
|
||||
|
||||
The weight interpretations available are:
|
||||
- comfy (default)
|
||||
- comfy++
|
||||
- compel
|
||||
- down_weight
|
||||
- A1111
|
||||
- perp
|
||||
|
||||
Normalizations are:
|
||||
- none (default)
|
||||
- length
|
||||
- mean
|
||||
|
||||
The normalization calculations are independent operations and you can combine them with `+`, eg `STYLE(A1111, length+mean)` or `STYLE(comfy, mean+length)`, or even something silly like `STYLE(perp, mean+length+mean+length)`
|
||||
|
||||
The style can be specified separately for each AND:ed prompt, but the first prompt is special; later prompts will "inherit" it as default. For example:
|
||||
|
||||
```
|
||||
STYLE(A1111) a (red:1.1) cat with (brown:0.9) spots and a long tail AND an (old:0.5) dog AND a (green:1.4) (balloon:1.1)
|
||||
```
|
||||
will interpret everything as A1111, but
|
||||
```
|
||||
a (red:1.1) cat with (brown:0.9) spots and a long tail AND STYLE(A1111) an (old:0.5) dog AND a (green:1.4) (balloon:1.1)
|
||||
```
|
||||
Will interpret the first one using the default ComfyUI behaviour, the second prompt with A1111 and the last prompt with the default again
|
||||
|
||||
### SDXL: Configure SDXL prompting parameters
|
||||
|
||||
The nodes do not treat SDXL models specially, but there are some utilities that enable SDXL specific functionality.
|
||||
|
||||
You can use the function `SDXL(width height, target_width target_height, crop_w crop_h)` to set SDXL prompt parameters. `SDXL()` is equivalent to `SDXL(1024 1024, 1024 1024, 0 0)` unless the default values have been overridden by `PCScheduleSettings`.
|
||||
|
||||
### TE: Per-encoder prompts for multi-encoder models
|
||||
|
||||
You can specify per-encoder prompts using the `TE` function. The syntax is as follows:
|
||||
`TE(encoder_name=prompt)`. Whitespace surrounding the prompt and encoder name are ignored.
|
||||
|
||||
For example:
|
||||
```
|
||||
TE(l=cat) TE(g = (dog:1.1)) TE(t5xxl=tiger)
|
||||
```
|
||||
The keys to use depend on what key ComfyUI uses for the encoder; for example `l` for CLIP L, `g` for CLIP G, and `t5xxl` for T5 XXL (Flux text encoder).
|
||||
|
||||
Use `TE(help)` to print a help text listing available keys.
|
||||
|
||||
Things to note:
|
||||
- If you set a prompt with `TE`, it will override the prompt outside the function for the specified text encoder.
|
||||
- Multiple instances of `TE` are joined with a space. That is, `TE(l=foo)TE(l=bar)` is the same as `TE(l=foo bar)`
|
||||
- `AND` and `BREAK` are processed before `TE`, so they do not do anything sensible; `TE(l=foo AND bar)` will parse as two prompts `TE(foo` and `bar)`. `SHIFT`, `SHUFFLE` and `OLDBREAK` do work, however.
|
||||
|
||||
### SHUFFLE and SHIFT: Create prompt permutations
|
||||
|
||||
Default parameters: `SHUFFLE(seed=0, separator=,, joiner=,)`, `SHIFT(steps=0, separator=,, joiner=,)`
|
||||
|
||||
`SHIFT` moves elements to the left by `steps`. The default is 0 so `SHIFT()` does nothing
|
||||
`SHUFFLE` generates a random permutation with `seed` as its seed.
|
||||
|
||||
These functions are applied to each prompt chunk **after** `BREAK`, `AND` etc. have been parsed. The prompt is split by `separator`, the operation is applied, and it's then joined back by `joiner`.
|
||||
|
||||
Multiple instances of these functions are applied in the order they appear in the prompt.
|
||||
|
||||
**NOTE** To avoid breaking emphasis syntax, the functions ignore any separators inside parentheses
|
||||
|
||||
For example:
|
||||
- `SHIFT(1) cat, dog, tiger, mouse` does a shift and results in `dog, tiger, mouse, cat`. (whitespace may vary)
|
||||
- `SHIFT(1,;) cat, dog ; tiger, mouse` results in `tiger, mouse, cat, dog`
|
||||
- `SHUFFLE() cat, dog, tiger, mouse` results in `cat, dog, mouse, tiger`
|
||||
- `SHUFFLE() SHIFT(1) cat, dog, tiger, mouse` results in `dog, mouse, tiger, cat`
|
||||
|
||||
- `SHIFT(1) cat,dog BREAK tiger,mouse` results in `dog,cat BREAK tiger,mouse`
|
||||
- `SHIFT(1) cat, dog AND SHIFT(1) tiger, mouse` results in `dog, cat BREAK mouse, tiger`
|
||||
|
||||
Whitespace is *not* stripped and may also be used as a joiner or separator
|
||||
- `SHIFT(1,, ) cat,dog` results in `dog cat`
|
||||
|
||||
### NOISE: Add noise to a prompt
|
||||
|
||||
The function `NOISE(weight, seed)` adds some random noise into the cond tensor. The seed is optional, and if not specified, the global RNG is used. `weight` should be between 0 and 1.
|
||||
|
||||
The usefulness of this is questionable, but it wasn't difficult to implement, so here it is.
|
||||
|
||||
## Regional prompting
|
||||
|
||||
See [Regional prompting](/doc/regional_prompting.md)
|
||||
|
||||
## Cutoff
|
||||
|
||||
NOTE: Cutoff syntax might change at some point; it's pretty clunky.
|
||||
|
||||
`PCTextEncode` reimplements cutoff from [ComfyUI Cutoff](https://github.com/BlenderNeko/ComfyUI_Cutoff).
|
||||
|
||||
The syntax is
|
||||
```
|
||||
a group of animals, [CUT:white cat:white], [CUT:brown dog:brown:0.5:1.0:1.0:_]
|
||||
```
|
||||
You should read the prompt as `a group of animals, white cat, brown dog`, but CUT causes the tokens in `target_tokens` to be masked off from the base prompt in `region_text`, so that their effect can be isolated, and you're less likely to get brown cats or white dogs.
|
||||
|
||||
Target tokens are treated individually, separated by space, for example, `[CUT:green apple, red apple, green leaf:green apple]` will mask *both* greens and the apple, giving you `+ +, red +, + leaf`. To mask out just `green apple`, use `[CUT:green apple, red apple:green_apple]` which will result in a masked prompt of `+ +, red apple`. Escape `_` with a `\`.
|
||||
|
||||
the parameters in the `CUT` section are `region_text:target_tokens:weight;strict_mask:start_from_masked:padding_token` of which only the first two are required. The default values are `weight=1.0`, `strict_mask=1.0` `start_from_masked=1.0`, `padding_token=+`
|
||||
|
||||
If `strict_mask`, `start_from_masked` or `padding_token` are specified in more than one CUT, the *last* one becomes the default for any CUTs afterwards that do not explicitly set the parameters. For example, in:
|
||||
|
||||
`[CUT:white cat:white:0.5] and [CUT:black parrot, flying:black:1.0:0.5] and [CUT:green apple:green]`
|
||||
|
||||
`white cat` will a weight of 0.5, and 1.0 for all parameters, and `black parrot` and `green apple` will *both* have a `strict_mask` parameter of 0.5.
|
||||
|
||||
The parameters affect how the masked and unmasked prompts are combined to produce the final embedding. Just play around with them.
|
||||
|
||||
## Miscellaneous
|
||||
- `<emb:xyz>` is alternative syntax for `embedding:xyz` to work around a syntax conflict with `[embedding:xyz:0.5]` which is parsed as a schedule that switches from `embedding` to `xyz`.
|
||||
|
||||
# Experimental features
|
||||
|
||||
> [!WARN]
|
||||
> These features are may change or disappear without warning
|
||||
|
||||
## COUPLE: Attention couple
|
||||
|
||||
See [here](/doc/attention_couple.md)
|
||||
|
||||
## TE_WEIGHT
|
||||
|
||||
For models using multiple text encoders, you can set weights per TE using the syntax `TE_WEIGHT(clipname=weight, clipname2=weight2, ...)` where `clipname` is one of the encoder names printed by `TE(help)`. For example with SDXL, try `TE_WEIGHT(g=0.25, l=0.75)`.
|
||||
|
||||
The weights are applied as a multiplier to the TE output. You can also override pooled output multipliers using eg. `l_pooled`.
|
||||
|
||||
To set a default value for all encoders, use `TE_WEIGHT(all=weight)`
|
||||
@@ -0,0 +1,60 @@
|
||||
## DEF: Lightweight prompt macros
|
||||
|
||||
You can define "prompt macros" by using `DEF`. Macros are expanded before any other parsing takes place. The expansion continues until no further changes occur. Recursion will raise an error.
|
||||
|
||||
`PCLazyTextEncode` and `PCLazyLoraLoader` expand macros, but `PCTextEncode` **does not**. If you need to expand macros for a single prompt, use `PCMacroExpand`
|
||||
|
||||
```
|
||||
DEF(MYMACRO=this is a prompt)
|
||||
[(MYMACRO:0.6):(MYMACRO:1.1):0.5]
|
||||
```
|
||||
is equivalent to
|
||||
```
|
||||
[(this is a prompt:0.5):(this is a prompt:1.1):0.5]
|
||||
```
|
||||
### Macro parameters
|
||||
It's also possible to give parameters to a macro:
|
||||
```
|
||||
DEF(MYMACRO=[(prompt $1:$2):(prompt $1:$3):$4])
|
||||
MYMACRO(test; 1.1; 0.7; 0.2)
|
||||
```
|
||||
gives
|
||||
```
|
||||
[(prompt test:1.1):(prompt test:0.7):0.2]
|
||||
```
|
||||
in this form, the variables $N (where N is any number corresponding to a positional parameter) will be replaced with the given parameter. The parameters must be separated with a semicolon, and can be empty.
|
||||
|
||||
You can also optionally specify default values:
|
||||
|
||||
```
|
||||
DEF(MACRO(example; 0; 1)=[$1:$2,$3])
|
||||
MACRO MACRO(test; 0.2)
|
||||
```
|
||||
gives
|
||||
```
|
||||
[example:0,1] [test:0.2,1]
|
||||
```
|
||||
|
||||
```
|
||||
DEF(MACRO() = [a:$1:0.5])
|
||||
```
|
||||
sets the default value of `$1` to an empty string.
|
||||
|
||||
### Unspecified parameters in macros
|
||||
|
||||
Unspecified parameters (either via defaults or explicitly given) will not be substituted. Compare:
|
||||
|
||||
```
|
||||
DEF(mything=a "$1" b "$2")
|
||||
mything
|
||||
mything()
|
||||
mything(A)
|
||||
```
|
||||
|
||||
gives
|
||||
|
||||
```
|
||||
a "$1" b "$2"
|
||||
a "" b "$2"
|
||||
a "A" b "$2"
|
||||
```
|
||||
@@ -0,0 +1,63 @@
|
||||
# Regional prompting
|
||||
|
||||
This section documents the masking functionality of `PCTextEncode`
|
||||
|
||||
See also [Attention Couple](/doc/attention_couple.md)
|
||||
|
||||
Remember that when using the lazy nodes, prompt scheduling applies to masks as well, so you can change or enable/disable regional prompts at any point during sampling.
|
||||
|
||||
## Behaviour
|
||||
|
||||
For each prompt separated by `AND`, you can specify either latent masks or an area.
|
||||
|
||||
- When masked, ComfyUI generates the model output using the **full latent** as the input, and then applies the mask to the output before adding it to your latent for the next step.
|
||||
- When an area is specified, ComfyUI generates a separate model output using the **part of the latent specified by the area** and then composites it into the full latent afterwards.
|
||||
- You can have *both* an AREA and a MASK specified, in which case the mask is applied to the latent specified by the AREA.
|
||||
|
||||
For example, consider a 1024 by 1024 (width x height) generation:
|
||||
|
||||
- `cat MASK(0 0.5, 0 1) AND dog MASK(0.5 1, 0 1)` generates two outputs at 1024x1024 for "dog" and "cat", then masks half of them off and adds the results together. The following step still see both the dog and the cat from the previous step, so they may blend slightly.
|
||||
|
||||
- `cat AREA(0 0.5, 0 1) AND dog AREA(0.5 1, 0 1)` generates two completely separate outputs at **512**x1024 and then composites them together into the 1024x1024 latent. Because the areas do not overlap, the generation for `cat` will not see the output of `dog` and vice versa in subsequent steps as long as the area restriction is in effect.
|
||||
|
||||
## MASK, IMASK and AREA
|
||||
|
||||
You can use `MASK(x1 x2, y1 y2, weight, op)` to specify a region mask for a prompt. The values are specified as a percentage with a float between `0` and `1`, or as absolute pixel values (these can't be mixed). `1` will be interpreted as a percentage instead of a pixel value.
|
||||
|
||||
Multiple `MASK` or `IMASK` calls will be composited together using ComfyUI's `MaskComposite` node, using `op` as the `operation` parameter (defaulting to `multiply`).
|
||||
|
||||
Similarly, you can use `AREA(x1 x2, y1 y2, weight)` to specify an area for the prompt (see ComfyUI's area composition examples). The area is calculated by ComfyUI relative to your latent size.
|
||||
|
||||
### Custom masks: IMASK and `PCAddMaskToCLIP`
|
||||
|
||||
You can attach custom masks to a `CLIP` with the `PC: Attach Mask` nodes and then refer to those masks in the prompt using `IMASK(index, weight, op)`. Indexing starts from zero, so 0 is the first attached mask etc. `PCSCheduleAddMasks` ignores empty inputs, so if you only add a mask to the `mask4` input, it will still have index 0.
|
||||
|
||||
Applying the nodes multiple times *appends* masks rather than overriding existing ones, so if you need more than 4, you can just use it more than once.
|
||||
|
||||
### Behaviour of multiple masks
|
||||
If multiple `MASK`s are specified, they are combined together with ComfyUI's `MaskComposite` node, with `op` specifying the operation to use (default `multiply`). In this case, the combined mask weight can be set with `MASKW(weight)` (defaults to 1.0).
|
||||
|
||||
Masks assume a size of `(512, 512)`, unless overridden with `PC: Configure PCTextEncode` and pixel values will be relative to that. ComfyUI will scale the mask to match the image resolution. You can change it manually by using `MASK_SIZE(width, height)` anywhere in the prompt,
|
||||
|
||||
These are handled per `AND`-ed prompt, so in `prompt1 AND MASK(...) prompt2`, the mask will only affect prompt2.
|
||||
|
||||
The default values are `MASK(0 1, 0 1, 1)` and you can omit unnecessary ones, that is, `MASK(0 0.5, 0.3)` is `MASK(0 0.5, 0.3 1, 1)`
|
||||
|
||||
Note that because the default values are percentages, `MASK(0 256, 64 512)` is valid, but `MASK(0 200)` will raise an error.
|
||||
|
||||
Masking does not affect LoRA scheduling unless you set unet weights to 0 for a LoRA.
|
||||
|
||||
## FEATHER: Mask operations
|
||||
|
||||
When you use `MASK` or `IMASK`, you can also call `FEATHER(left top right bottom)` to apply feathering using ComfyUI's `FeatherMask` node. The values are in pixels and default to `0`.
|
||||
|
||||
If multiple masks are used, `FEATHER` is applied *before compositing* in the order they appear in the prompt, and any leftovers are applied to the combined mask. If you want to skip feathering a mask while compositing, just use `FEATHER()` with no arguments.
|
||||
|
||||
For example:
|
||||
```
|
||||
MASK(1) MASK(2) MASK(3) FEATHER(1) FEATHER() FEATHER(3) weirdmask FEATHER(4)
|
||||
```
|
||||
|
||||
gives you a mask that is a combination of 1, 2 and 3, where 1 and 3 are feathered before compositing and then `FEATHER(4)` is applied to the composite.
|
||||
|
||||
The order of the `FEATHER` and `MASK` calls doesn't matter; you can have `FEATHER` before `MASK` or even interleave them.
|
||||
@@ -0,0 +1,117 @@
|
||||
# Prompt Schedule Syntax
|
||||
|
||||
> [!TIP]
|
||||
> If you're viewing this on GitHub, I recommend opening the outline by clicking the button in the top right corner of the text view (it is annoyingly easy to miss).
|
||||
|
||||
> [!NOTE]
|
||||
> The syntax documented in this section is only available with the `PC: Schedule Prompt` and `PC: Schedule LoRAs` nodes and their advanced variants.
|
||||
|
||||
Scheduling syntax is available with is similar to A1111, but only fractions are supported for steps. LoRAs are scheduled by including them in a scheduling expression.
|
||||
|
||||
Besides the syntax documented below, the [basic syntax](/doc/basic.md) and [prompt macro](/doc/macros.md) features are also automatically available.
|
||||
|
||||
```
|
||||
a [large::0.1] [cat|dog:0.05] [<lora:somelora:0.5:0.6>::0.5]
|
||||
[in a park:in space:0.4]
|
||||
```
|
||||
## Comments and escaping
|
||||
|
||||
In schedules, any text on a line following a `#` is considered a comment and removed, including the `#` character.
|
||||
You can escape the following characters in places where they would otherwise conflict with syntax:
|
||||
|
||||
- `#` with `\#`
|
||||
- `:` with `\:`
|
||||
- `\` with `\\`
|
||||
|
||||
Escaping is only required if it would otherwise be considered syntax, that is `\o/` will be interpreted literally and the `\` does not need to be escaped, but in `[embedding:a:0.5]` you would need to escape the `:`.
|
||||
|
||||
## Scheduled prompts
|
||||
|
||||
There are two forms of scheduled prompts.
|
||||
|
||||
### Basic scheduling expressions
|
||||
Basic expressions take the form `[before:after:X]` where `X` is the switch point, a decimal number between 0.0 and 1.0 inclusive, representing 0 to 100% of timesteps. Either prompt can also be empty.
|
||||
For example:
|
||||
```
|
||||
a [red:blue:0.5] cat
|
||||
```
|
||||
switches from `a red cat` to `a blue cat` at 0.5. `before` and `after` can be arbitrary prompts (`after` can also be empty), including other scheduling expressions, allowing nesting:
|
||||
```
|
||||
a [red:[blue::0.7]:0.5] cat
|
||||
```
|
||||
|
||||
switches from `a red cat` to `a blue cat` at 0.5 and to `a cat` at 0.7
|
||||
|
||||
For convenience `[cat:0.5]` is equivalent to `[:cat:0.5]` meaning it switches from empty to `cat` at 0.5.
|
||||
|
||||
### Range expressions
|
||||
|
||||
The most general form of a schedule is a range expression: For example, in `prompt [before:during:after:0.3,0.7]`, The prompt be `prompt before` until 0.3, `prompt during` until 0.7, and then `prompt after`. This form is equivalent to `prompt [before:[during:after:0.7]:0.3]`
|
||||
|
||||
For convenience, `[during:0.1,0.4]` is equivalent to `[:during::0.1,0.4]` and `[during:after:0.1,0.4]` is equivalent to `[:during:after:0.1,0.4]`.
|
||||
|
||||
`[before:during:after:0.1]` is the same as `[before:during:after:0.1,1.0]` which is same as `[before:during:0.1]`
|
||||
|
||||
|
||||
### Using step numbers with the Advanced nodes
|
||||
|
||||
If you provide a non-zero value to `num_steps` to the `Advanced` versions of the scheduling nodes, you will be able to use step numbers in prompts.
|
||||
|
||||
For now, a value between 0 and 1.0 will be interpreted as a percentage if it contains a ., and as an absolute step otherwise.
|
||||
|
||||
This is just syntactic sugar. Behind the scenes, the values are converted to percentages and have normal ComfyUI scheduling behaviour.
|
||||
|
||||
## Tag selection
|
||||
Using the `FilterSchedule` node, in addition to step percentages, you can use a *tag* to select part of an input:
|
||||
```
|
||||
a large [dog:cat<lora:catlora:0.5>:SECOND_PASS]
|
||||
```
|
||||
Set the `tags` parameter in the `FilterSchedule` node to filter the prompt. If the tag matches any tag `tags` (comma-separated), the second option is returned (`cat`, in this case, with the LoRA). Otherwise, the first option is chosen (`dog`, without LoRA).
|
||||
|
||||
the values in `tags` are case-insensitive, but the tags in the input **must** be uppercase A-Z and underscores only, or they won't be recognized. That is, `[dog:cat:hr]` will not work.
|
||||
|
||||
For example, a prompt
|
||||
```
|
||||
a [black:blue:X] [cat:dog:Y] [walking:running:Z] in space
|
||||
```
|
||||
with `tags` `x,z` would result in the prompt `a blue cat running in space`
|
||||
|
||||
The three prompt form `[a:b:c:TAG]` is parsed, but ignores `b` and is equivalent to `[a:c:TAG]`.
|
||||
|
||||
## LoRA Scheduling
|
||||
When using the lazy graph building nodes, LoRAs can be scheduled by referring to them in a scheduling expression, like so:
|
||||
|
||||
`<lora:fulllora:1> [<lora:partialora:1>::0.5]`
|
||||
|
||||
This will schedule `fulllora` for the entire duration of the prompt and `partiallora` until half of sampling is complete.
|
||||
|
||||
You can refer to LoRAs by using the filename without extension and subdirectories will also be searched. For example, `<lora:cats:1>`. will match both `cats.safetensors` and `sd15/animals/cats.safetensors`. If there are multiple LoRAs with the same name, the first match will be loaded.
|
||||
|
||||
Alternatively, the name can include the full directory path relative to ComfyUI's search paths, without extension: `<lora:XL/sdxllora:0.5>`. In this case, the *full* path must match.
|
||||
|
||||
You can also give the exact path (including the extension) as shown in `LoRALoader`.
|
||||
|
||||
If no match is found, the node will try to replace spaces with underscores and search again. That is, `<lora:cats and dogs:1>` will find `cats_and_dogs.safetensors`. This helps with some autocompletion scripts that replace underscores with spaces.
|
||||
|
||||
Finally, if none of the above produce a match, the search term will be split by whitespace and files that contain all of the parts in any order will be considered. If this returns only a single match, it will be loaded. For example, consider LoRAs:
|
||||
|
||||
- `xl/red_cats.safetensors`
|
||||
- `flux/blue_cats.safetensors`
|
||||
- `flux/red_cats.safetensors`
|
||||
|
||||
Then `<lora:cats xl:1>` would match the red cats LoRA, but `cats flux` would be ambiguous and not match.
|
||||
|
||||
## Alternating
|
||||
|
||||
Alternating syntax is `[a|b:pct_steps]`, causing the prompt to alternate every `pct_steps`. `pct_steps` defaults to 0.1 if not specified. You can also have more than two options.
|
||||
|
||||
|
||||
## Sequences
|
||||
|
||||
The syntax `[SEQ:a:N1:b:N2:c:N3]` is shorthand for `[a:[b:[c::N3]:N2]:N1]` ie. it switches from `a` to `b` to `c` to nothing at the specified points in sequence.
|
||||
|
||||
Might be useful with Jinja templating (see https://github.com/asagi4/comfyui-utility-nodes). For example:
|
||||
```
|
||||
[SEQ<% for x in steps(0.1, 0.9, 0.1) %>:<lora:test:<= sin(x*pi) + 0.1 =>>:<= x =><% endfor %>]
|
||||
```
|
||||
generates a LoRA schedule based on a sinewave
|
||||
-209
@@ -1,209 +0,0 @@
|
||||
# Scheduling syntax
|
||||
|
||||
Syntax is like A1111 for now, but only fractions are supported for steps. LoRAs are scheduled by including them in a scheduling expression.
|
||||
|
||||
```
|
||||
a [large::0.1] [cat|dog:0.05] [<lora:somelora:0.5:0.6>::0.5]
|
||||
[in a park:in space:0.4]
|
||||
```
|
||||
|
||||
## Scheduled prompts
|
||||
|
||||
There are two forms of scheduled prompts.
|
||||
|
||||
### Basic scheduling expressions
|
||||
Basic expressions take the form `[before:after:X]` where `X` is the switch point, a decimal number between 0.0 and 1.0 inclusive, representing 0 to 100% of timesteps.
|
||||
For example:
|
||||
```
|
||||
a [red:blue:0.5] cat
|
||||
```
|
||||
switches from `a red cat` to `a blue cat` at 0.5. `before` and `after` can be arbitrary prompts (`after` can also be empty), including other scheduling expressions, allowing nesting:
|
||||
```
|
||||
a [red:[blue::0.7]:0.5] cat
|
||||
```
|
||||
|
||||
switches from `a red cat` to `a blue cat` at 0.5 and to `a cat` at 0.7
|
||||
|
||||
|
||||
**Note:** As a special case, `[cat:0.5]` is like `[:cat:0.5]` meaning it switches from empty to `cat` at 0.5. Currently, `[:cat:0.5]` doesn't actually parse correctly, so you **must** use the shortcut form
|
||||
|
||||
### Range expressions
|
||||
|
||||
You can also use `a [during:after:0.3,0.7]` as a shortcut. The prompt be `a` until 0.3, `a during` until 0.7, and then `a after`. This form is equivalent to `[[during:after:0.7]:0.3]`
|
||||
For convenience, `[during:0.1,0.4]` is equivalent to `[during::0.1,0.4]`
|
||||
|
||||
## Tag selection
|
||||
Using the `FilterSchedule` node, in addition to step percentages, you can use a *tag* to select part of an input:
|
||||
```
|
||||
a large [dog:cat<lora:catlora:0.5>:SECOND_PASS]
|
||||
```
|
||||
Set the `tags` parameter in the `FilterSchedule` node to filter the prompt. If the tag matches any tag `tags` (comma-separated), the second option is returned (`cat`, in this case, with the LoRA). Otherwise, the first option is chosen (`dog`, without LoRA).
|
||||
|
||||
the values in `tags` are case-insensitive, but the tags in the input **must** be uppercase A-Z and underscores only, or they won't be recognized. That is, `[dog:cat:hr]` will not work.
|
||||
|
||||
For example, a prompt
|
||||
```
|
||||
a [black:blue:X] [cat:dog:Y] [walking:running:Z] in space
|
||||
```
|
||||
with `tags` `x,z` would result in the prompt `a blue cat running in space`
|
||||
|
||||
|
||||
## LoRA Scheduling
|
||||
When using the lazy graph building nodes, LoRAs can be scheduled by referring to them in a scheduling expression, like so:
|
||||
|
||||
`<lora:fulllora:1> [<lora:partialora:1>::0.5]`
|
||||
|
||||
This will schedule `fulllora` for the entire duration of the prompt and `partiallora` until half of sampling is complete.
|
||||
|
||||
You can refer to LoRAs by using the filename without extension and subdirectories will also be searched. For example, `<lora:cats:1>`. will match both `cats.safetensors` and `sd15/animals/cats.safetensors`. If there are multiple LoRAs with the same name, the first match will be loaded.
|
||||
|
||||
Alternatively, the name can include the full directory path relative to ComfyUI's search paths, without extension: `<lora:XL/sdxllora:0.5>`. In this case, the *full* path must match.
|
||||
|
||||
If no match is found, the node will try to replace spaces with underscores and search again. That is, `<lora:cats and dogs:1>` will find `cats_and_dogs.safetensors`. This helps with some autocompletion scripts that replace underscores with spaces.
|
||||
|
||||
Finally, you can give the exact path (including the extension) as shown in `LoRALoader`.
|
||||
|
||||
|
||||
## Alternating
|
||||
|
||||
Alternating syntax is `[a|b:pct_steps]`, causing the prompt to alternate every `pct_steps`. `pct_steps` defaults to 0.1 if not specified. You can also have more than two options.
|
||||
|
||||
|
||||
## Sequences
|
||||
|
||||
The syntax `[SEQ:a:N1:b:N2:c:N3]` is shorthand for `[a:[b:[c::N3]:N2]:N1]` ie. it switches from `a` to `b` to `c` to nothing at the specified points in sequence.
|
||||
|
||||
Might be useful with Jinja templating (see https://github.com/asagi4/comfyui-utility-nodes). For example:
|
||||
```
|
||||
[SEQ<% for x in steps(0.1, 0.9, 0.1) %>:<lora:test:<= sin(x*pi) + 0.1 =>>:<= x =><% endfor %>]
|
||||
```
|
||||
generates a LoRA schedule based on a sinewave
|
||||
|
||||
|
||||
# Basic prompt syntax
|
||||
|
||||
This syntax is also available in outside scheduled prompts, where applicable.
|
||||
|
||||
## LoRA loading
|
||||
|
||||
The A111-style syntax `<lora:loraname:weight>` can be used to load LoRAs via the prompt. See LoRA scheduling above.
|
||||
|
||||
## Combining prompts, A1111-style
|
||||
|
||||
- The keyword `BREAK` causes the prompt to be tokenized in separate chunks, which results in each chunk being individually padded to the text encoder's maximum token length. This is mostly equivalent to the `ConditioningConcat` node.
|
||||
|
||||
`AND` can be used to combine prompts. You can also use a weight at the end. It does a weighted sum of each prompt,
|
||||
```
|
||||
cat :1 AND dog :2
|
||||
```
|
||||
The weight defaults to 1 and are normalized so that `a:2 AND b:2` is equal to `a AND b`. `AND` is processed after schedule parsing, so you can change the weight mid-prompt: `cat:[1:2:0.5] AND dog`
|
||||
|
||||
|
||||
## Functions
|
||||
|
||||
There are some "functions" that can be included in a prompt to do various things.
|
||||
|
||||
Functions have the form `FUNCNAME(param1, param2, ...)`. How parameters are interpreted is up to the function.
|
||||
Note: Whitespace is *not* stripped from string parameters by default. Commas can be escaped with `\,`
|
||||
|
||||
Like `AND`, these functions are parsed after regular scheduling syntax has been expanded, allowing things like `[AREA:MASK:0.3](...)`, in case that's somehow useful.
|
||||
|
||||
### SDXL
|
||||
|
||||
The nodes do not treat SDXL models specially, but there are some utilities that enable SDXL specific functionality.
|
||||
|
||||
You can use the function `SDXL(width height, target_width target_height, crop_w crop_h)` to set SDXL prompt parameters. `SDXL()` is equivalent to `SDXL(1024 1024, 1024 1024, 0 0)` unless the default values have been overridden by `PCScheduleSettings`.
|
||||
|
||||
To set the `clip_l` prompt, as with `CLIPTextEncodeSDXL`, use the function `CLIP_L(prompt text goes here)`.
|
||||
|
||||
Things to note:
|
||||
- Multiple instances of `CLIP_L` are joined with a space. That is, `CLIP_L(foo)CLIP_L(bar)` is the same as `CLIP_L(foo bar)`
|
||||
- Using `BREAK` isn't supported in it; it'll just parse as the plain word BREAK.
|
||||
- similarly, `AND` inside `CLIP_L` does not do anything sensible; `CLIP_L(foo AND bar)` will parse as two prompts `CLIP_L(foo` and `bar)`
|
||||
- `CLIP_L` and `SDXL` have no effect on SD 1.5.
|
||||
- The rest of the prompt becomes the `clip_g` prompt.
|
||||
- If there is no `CLIP_L` or `SDXL`, the prompts will work as with `CLIPTextEncode`.
|
||||
|
||||
### SHUFFLE and SHIFT
|
||||
|
||||
Default parameters: `SHUFFLE(seed=0, separator=,, joiner=,)`, `SHIFT(steps=0, separator=,, joiner=,)`
|
||||
|
||||
`SHIFT` moves elements to the left by `steps`. The default is 0 so `SHIFT()` does nothing
|
||||
`SHUFFLE` generates a random permutation with `seed` as its seed.
|
||||
|
||||
These functions are applied to each prompt chunk **after** `BREAK`, `AND` etc. have been parsed. The prompt is split by `separator`, the operation is applied, and it's then joined back by `joiner`.
|
||||
|
||||
Multiple instances of these functions are applied in the order they appear in the prompt.
|
||||
|
||||
**NOTE** To avoid breaking emphasis syntax, the functions ignore any separators inside parentheses
|
||||
|
||||
For example:
|
||||
- `SHIFT(1) cat, dog, tiger, mouse` does a shift and results in `dog, tiger, mouse, cat`. (whitespace may vary)
|
||||
- `SHIFT(1,;) cat, dog ; tiger, mouse` results in `tiger, mouse, cat, dog`
|
||||
- `SHUFFLE() cat, dog, tiger, mouse` results in `cat, dog, mouse, tiger`
|
||||
- `SHUFFLE() SHIFT(1) cat, dog, tiger, mouse` results in `dog, mouse, tiger, cat`
|
||||
|
||||
- `SHIFT(1) cat,dog BREAK tiger,mouse` results in `dog,cat BREAK tiger,mouse`
|
||||
- `SHIFT(1) cat, dog AND SHIFT(1) tiger, mouse` results in `dog, cat BREAK mouse, tiger`
|
||||
|
||||
Whitespace is *not* stripped and may also be used as a joiner or separator
|
||||
- `SHIFT(1,, ) cat,dog` results in `dog cat`
|
||||
|
||||
### NOISE
|
||||
|
||||
The function `NOISE(weight, seed)` adds some random noise into the prompt. The seed is optional, and if not specified, the global RNG is used. `weight` should be between 0 and 1.
|
||||
|
||||
### MASK, IMASK and AREA
|
||||
|
||||
You can use `MASK(x1 x2, y1 y2, weight, op)` to specify a region mask for a prompt. The values are specified as a percentage with a float between `0` and `1`, or as absolute pixel values (these can't be mixed). `1` will be interpreted as a percentage instead of a pixel value.
|
||||
|
||||
Similarly, you can use `AREA(x1 x2, y1 y2, weight)` to specify an area for the prompt (see ComfyUI's area composition examples). The area is calculated by ComfyUI relative to your latent size.
|
||||
|
||||
#### Custom masks: IMASK and `PCAddMaskToCLIP`
|
||||
|
||||
You can attach custom masks to a `CLIP` with the `PC: Attach Mask` nodes and then refer to those masks in the prompt using `IMASK(index, weight, op)`. Indexing starts from zero, so 0 is the first attached mask etc. `PCSCheduleAddMasks` ignores empty inputs, so if you only add a mask to the `mask4` input, it will still have index 0.
|
||||
|
||||
Applying the nodes multiple times *appends* masks rather than overriding existing ones, so if you need more than 4, you can just use it more than once.
|
||||
|
||||
#### Behaviour of masks
|
||||
If multiple `MASK`s are specified, they are combined together with ComfyUI's `MaskComposite` node, with `op` specifying the operation to use (default `multiply`). In this case, the combined mask weight can be set with `MASKW(weight)` (defaults to 1.0).
|
||||
|
||||
Masks assume a size of `(512, 512)`, unless overridden with `PC: Configure PCTextEncode` and pixel values will be relative to that. ComfyUI will scale the mask to match the image resolution. You can change it manually by using `MASK_SIZE(width, height)` anywhere in the prompt,
|
||||
|
||||
These are handled per `AND`-ed prompt, so in `prompt1 AND MASK(...) prompt2`, the mask will only affect prompt2.
|
||||
|
||||
The default values are `MASK(0 1, 0 1, 1)` and you can omit unnecessary ones, that is, `MASK(0 0.5, 0.3)` is `MASK(0 0.5, 0.3 1, 1)`
|
||||
|
||||
Note that because the default values are percentages, `MASK(0 256, 64 512)` is valid, but `MASK(0 200)` will raise an error.
|
||||
|
||||
Masking does not affect LoRA scheduling unless you set unet weights to 0 for a LoRA.
|
||||
|
||||
### FEATHER
|
||||
|
||||
When you use `MASK` or `IMASK`, you can also call `FEATHER(left top right bottom)` to apply feathering using ComfyUI's `FeatherMask` node. The values are in pixels and default to `0`.
|
||||
|
||||
If multiple masks are used, `FEATHER` is applied *before compositing* in the order they appear in the prompt, and any leftovers are applied to the combined mask. If you want to skip feathering a mask while compositing, just use `FEATHER()` with no arguments.
|
||||
|
||||
For example:
|
||||
```
|
||||
MASK(1) MASK(2) MASK(3) FEATHER(1) FEATHER() FEATHER(3) weirdmask FEATHER(4)
|
||||
```
|
||||
|
||||
gives you a mask that is a combination of 1, 2 and 3, where 1 and 3 are feathered before compositing and then `FEATHER(4)` is applied to the composite.
|
||||
|
||||
The order of the `FEATHER` and `MASK` calls doesn't matter; you can have `FEATHER` before `MASK` or even interleave them.
|
||||
|
||||
## Miscellaneous
|
||||
- `<emb:xyz>` is alternative syntax for `embedding:xyz` to work around a syntax conflict with `[embedding:xyz:0.5]` which is parsed as a schedule that switches from `embedding` to `xyz`.
|
||||
|
||||
# Experimental features
|
||||
|
||||
Experimental features are unstable and may disappear or break without warning.
|
||||
|
||||
## Attention masking
|
||||
|
||||
Use `ATTN()` in combination with `MASK()` or `IMASK()` to enable attention masking. Currently, it's pretty slow and only works with SDXL. You need to have a recent enough version of ComfyUI for this to work.
|
||||
|
||||
## TE_WEIGHT
|
||||
|
||||
For models using multiple text encoders, you can set weights per TE using the syntax `TE_WEIGHT(clipname=weight, clipname2=weight2, ...)` where `clipname` is one of `g`, `l`, or `t5xxl`. For example with SDXL, try `TE_WEIGHT(g=0.25, l=0.75`). The weights are applied as a multiplier to the TE output.
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 102 KiB |
@@ -0,0 +1,672 @@
|
||||
{
|
||||
"id": "e820c2fb-9502-45b7-a864-684757dddcdf",
|
||||
"revision": 0,
|
||||
"last_node_id": 18,
|
||||
"last_link_id": 20,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 1,
|
||||
"type": "CheckpointLoaderSimple",
|
||||
"pos": [
|
||||
-135,
|
||||
-930
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
98
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "MODEL",
|
||||
"type": "MODEL",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
2
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "CLIP",
|
||||
"type": "CLIP",
|
||||
"slot_index": 1,
|
||||
"links": [
|
||||
3
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "VAE",
|
||||
"type": "VAE",
|
||||
"slot_index": 2,
|
||||
"links": [
|
||||
18
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.18",
|
||||
"Node name for S&R": "CheckpointLoaderSimple"
|
||||
},
|
||||
"widgets_values": [
|
||||
"NoobAI-XL-Vpred-v1.0.safetensors"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"type": "PCLazyTextEncode",
|
||||
"pos": [
|
||||
555,
|
||||
-720
|
||||
],
|
||||
"size": [
|
||||
252,
|
||||
78
|
||||
],
|
||||
"flags": {
|
||||
"collapsed": true
|
||||
},
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "clip",
|
||||
"type": "CLIP",
|
||||
"link": 5
|
||||
},
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "text"
|
||||
},
|
||||
"link": 7
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "CONDITIONING",
|
||||
"type": "CONDITIONING",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
12
|
||||
]
|
||||
}
|
||||
],
|
||||
"title": "PC: Schedule Prompt (positive)",
|
||||
"properties": {
|
||||
"cnr_id": "comfyui-prompt-control",
|
||||
"ver": "2.0.0-beta.7",
|
||||
"Node name for S&R": "PCLazyTextEncode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"STYLE(A1111) 1girl, [painting \\(medium\\), realistic,::0.2] fennec fox girl, animal ear fluff, [[purple:white pupils, purple:0.2] eyes:sparkling eyes:0.85], cargo pants, long sleeves, cardigan, winter, snow, steaming cup, coffee mug, [thermos,:0.1] [long hair,:0.25] [BREAK:0.3]\n[(masterpiece, best quality, newest, very awa,):0.1], night sky, full moon, star \\(sky\\),"
|
||||
],
|
||||
"color": "#232",
|
||||
"bgcolor": "#353"
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"type": "PCLazyLoraLoader",
|
||||
"pos": [
|
||||
257.5,
|
||||
-745
|
||||
],
|
||||
"size": [
|
||||
210,
|
||||
78
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"shape": 7,
|
||||
"type": "MODEL",
|
||||
"link": 2
|
||||
},
|
||||
{
|
||||
"name": "clip",
|
||||
"shape": 7,
|
||||
"type": "CLIP",
|
||||
"link": 3
|
||||
},
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "text"
|
||||
},
|
||||
"link": 6
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "MODEL",
|
||||
"type": "MODEL",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
17
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "CLIP",
|
||||
"type": "CLIP",
|
||||
"slot_index": 1,
|
||||
"links": [
|
||||
5,
|
||||
9
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfyui-prompt-control",
|
||||
"ver": "2.0.0-beta.7",
|
||||
"Node name for S&R": "PCLazyLoraLoader"
|
||||
},
|
||||
"widgets_values": [
|
||||
"STYLE(A1111) 1girl, [painting \\(medium\\), realistic,::0.2] fennec fox girl, animal ear fluff, [[purple:white pupils, purple:0.2] eyes:sparkling eyes:0.85], cargo pants, long sleeves, cardigan, winter, snow, steaming cup, coffee mug, [thermos,:0.1] [long hair,:0.25] [BREAK:0.3]\n[(masterpiece, best quality, newest, very awa,):0.1], night sky, full moon, star \\(sky\\),"
|
||||
],
|
||||
"color": "#223",
|
||||
"bgcolor": "#335"
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"type": "KSampler",
|
||||
"pos": [
|
||||
930,
|
||||
-780
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
474
|
||||
],
|
||||
"flags": {},
|
||||
"order": 9,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "MODEL",
|
||||
"link": 17
|
||||
},
|
||||
{
|
||||
"name": "positive",
|
||||
"type": "CONDITIONING",
|
||||
"link": 12
|
||||
},
|
||||
{
|
||||
"name": "negative",
|
||||
"type": "CONDITIONING",
|
||||
"link": 13
|
||||
},
|
||||
{
|
||||
"name": "latent_image",
|
||||
"type": "LATENT",
|
||||
"link": 14
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "LATENT",
|
||||
"type": "LATENT",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
15
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.18",
|
||||
"Node name for S&R": "KSampler"
|
||||
},
|
||||
"widgets_values": [
|
||||
2,
|
||||
"fixed",
|
||||
25,
|
||||
1.4000000000000001,
|
||||
"euler_cfg_pp",
|
||||
"simple",
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "PrimitiveNode",
|
||||
"pos": [
|
||||
-270,
|
||||
-780
|
||||
],
|
||||
"size": [
|
||||
495,
|
||||
225
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "text"
|
||||
},
|
||||
"links": [
|
||||
6,
|
||||
7
|
||||
]
|
||||
}
|
||||
],
|
||||
"title": "Positive prompt (with LoRAs)",
|
||||
"properties": {
|
||||
"Run widget replace on values": false
|
||||
},
|
||||
"widgets_values": [
|
||||
"STYLE(A1111) 1girl, [painting \\(medium\\), realistic,::0.2] fennec fox girl, animal ear fluff, [[purple:white pupils, purple:0.2] eyes:sparkling eyes:0.85], cargo pants, long sleeves, cardigan, winter, snow, steaming cup, coffee mug, [thermos,:0.1] [long hair,:0.25] [BREAK:0.3]\n[(masterpiece, best quality, newest, very awa,):0.1], night sky, full moon, star \\(sky\\),"
|
||||
],
|
||||
"color": "#232",
|
||||
"bgcolor": "#353"
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "PrimitiveNode",
|
||||
"pos": [
|
||||
-270,
|
||||
-510
|
||||
],
|
||||
"size": [
|
||||
480,
|
||||
225
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "text"
|
||||
},
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
8
|
||||
]
|
||||
}
|
||||
],
|
||||
"title": "Negative prompt",
|
||||
"properties": {
|
||||
"Run widget replace on values": false
|
||||
},
|
||||
"widgets_values": [
|
||||
"chibi, [bad hands,low quality, worst quality,:0.05], simple background, blurry, sketch, unfinished, [holding two cups,no pupils,:0.1]"
|
||||
],
|
||||
"color": "#322",
|
||||
"bgcolor": "#533"
|
||||
},
|
||||
{
|
||||
"id": 7,
|
||||
"type": "PCLazyTextEncode",
|
||||
"pos": [
|
||||
555,
|
||||
-675
|
||||
],
|
||||
"size": [
|
||||
252,
|
||||
78
|
||||
],
|
||||
"flags": {
|
||||
"collapsed": true
|
||||
},
|
||||
"order": 8,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "clip",
|
||||
"type": "CLIP",
|
||||
"link": 9
|
||||
},
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "text"
|
||||
},
|
||||
"link": 8
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "CONDITIONING",
|
||||
"type": "CONDITIONING",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
13
|
||||
]
|
||||
}
|
||||
],
|
||||
"title": "PC: Schedule Prompt (negative)",
|
||||
"properties": {
|
||||
"cnr_id": "comfyui-prompt-control",
|
||||
"ver": "2.0.0-beta.7",
|
||||
"Node name for S&R": "PCLazyTextEncode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"chibi, [bad hands,low quality, worst quality,:0.05], simple background, blurry, sketch, unfinished, [holding two cups,no pupils,:0.1]"
|
||||
],
|
||||
"color": "#322",
|
||||
"bgcolor": "#533"
|
||||
},
|
||||
{
|
||||
"id": 9,
|
||||
"type": "EmptyLatentImage",
|
||||
"pos": [
|
||||
525,
|
||||
-615
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
106
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "LATENT",
|
||||
"type": "LATENT",
|
||||
"links": [
|
||||
14
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.18",
|
||||
"Node name for S&R": "EmptyLatentImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
896,
|
||||
1152,
|
||||
1
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 10,
|
||||
"type": "VAEDecode",
|
||||
"pos": [
|
||||
1290,
|
||||
-780
|
||||
],
|
||||
"size": [
|
||||
210,
|
||||
46
|
||||
],
|
||||
"flags": {},
|
||||
"order": 10,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "samples",
|
||||
"type": "LATENT",
|
||||
"link": 15
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"type": "VAE",
|
||||
"link": 18
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
20
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.18",
|
||||
"Node name for S&R": "VAEDecode"
|
||||
},
|
||||
"widgets_values": []
|
||||
},
|
||||
{
|
||||
"id": 13,
|
||||
"type": "MarkdownNote",
|
||||
"pos": [
|
||||
240,
|
||||
-615
|
||||
],
|
||||
"size": [
|
||||
240,
|
||||
105
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"If you do not need LoRA scheduling, you can simply skip this node."
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 15,
|
||||
"type": "MarkdownNote",
|
||||
"pos": [
|
||||
240,
|
||||
-450
|
||||
],
|
||||
"size": [
|
||||
600,
|
||||
210
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"`PC: Schedule prompt` will expand into instances of `PCTextEncode`. `PC: Schedule LoRAs` will expand into the required `LoRALoader`s and `CLIP` hooks required to schedule LoRAs in the prompt.\n\nYou can pass the same prompt to both nodes; `PC: Schedule Prompt` will simply ignore any `<lora:xyz:1>` elements, so they will not affect the prompt.\nSee the [full syntax available in the prompts](https://github.com/asagi4/comfyui-prompt-control/blob/master/doc/syntax.md) on GitHub"
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 18,
|
||||
"type": "SaveImage",
|
||||
"pos": [
|
||||
1290,
|
||||
-690
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
405
|
||||
],
|
||||
"flags": {},
|
||||
"order": 11,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 20
|
||||
}
|
||||
],
|
||||
"outputs": [],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.18"
|
||||
},
|
||||
"widgets_values": [
|
||||
"PromptControl"
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
2,
|
||||
1,
|
||||
0,
|
||||
3,
|
||||
0,
|
||||
"MODEL"
|
||||
],
|
||||
[
|
||||
3,
|
||||
1,
|
||||
1,
|
||||
3,
|
||||
1,
|
||||
"CLIP"
|
||||
],
|
||||
[
|
||||
5,
|
||||
3,
|
||||
1,
|
||||
2,
|
||||
0,
|
||||
"CLIP"
|
||||
],
|
||||
[
|
||||
6,
|
||||
5,
|
||||
0,
|
||||
3,
|
||||
2,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
7,
|
||||
5,
|
||||
0,
|
||||
2,
|
||||
1,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
8,
|
||||
6,
|
||||
0,
|
||||
7,
|
||||
1,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
9,
|
||||
3,
|
||||
1,
|
||||
7,
|
||||
0,
|
||||
"CLIP"
|
||||
],
|
||||
[
|
||||
12,
|
||||
2,
|
||||
0,
|
||||
4,
|
||||
1,
|
||||
"CONDITIONING"
|
||||
],
|
||||
[
|
||||
13,
|
||||
7,
|
||||
0,
|
||||
4,
|
||||
2,
|
||||
"CONDITIONING"
|
||||
],
|
||||
[
|
||||
14,
|
||||
9,
|
||||
0,
|
||||
4,
|
||||
3,
|
||||
"LATENT"
|
||||
],
|
||||
[
|
||||
15,
|
||||
4,
|
||||
0,
|
||||
10,
|
||||
0,
|
||||
"LATENT"
|
||||
],
|
||||
[
|
||||
17,
|
||||
3,
|
||||
0,
|
||||
4,
|
||||
0,
|
||||
"MODEL"
|
||||
],
|
||||
[
|
||||
20,
|
||||
10,
|
||||
0,
|
||||
18,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
18,
|
||||
1,
|
||||
2,
|
||||
10,
|
||||
1,
|
||||
"VAE"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.8,
|
||||
"offset": [
|
||||
591.75,
|
||||
1235
|
||||
]
|
||||
},
|
||||
"linkExtensions": [
|
||||
{
|
||||
"id": 18,
|
||||
"parentId": 1
|
||||
}
|
||||
],
|
||||
"reroutes": [
|
||||
{
|
||||
"id": 1,
|
||||
"pos": [
|
||||
1273.75,
|
||||
-879.5
|
||||
],
|
||||
"linkIds": [
|
||||
18
|
||||
]
|
||||
}
|
||||
]
|
||||
},
|
||||
"version": 0.4,
|
||||
"models": [{
|
||||
"name": "NoobAI-XL-Vpred-v1.0.safetensors",
|
||||
"url": "https://huggingface.co/Laxhar/noobai-XL-Vpred-1.0/resolve/main/NoobAI-XL-Vpred-v1.0.safetensors",
|
||||
"directory": "checkpoints"
|
||||
}]
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
+316
-167
@@ -1,6 +1,17 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
import itertools
|
||||
import logging
|
||||
from math import copysign
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def _norm_mag(w, n):
|
||||
d = w - 1
|
||||
return 1 + np.sign(d) * np.sqrt(np.abs(d) ** 2 / n)
|
||||
# return np.sign(w) * np.sqrt(np.abs(w)**2 / n)
|
||||
|
||||
|
||||
def _grouper(n, iterable):
|
||||
@@ -12,97 +23,41 @@ def _grouper(n, iterable):
|
||||
yield chunk
|
||||
|
||||
|
||||
def _norm_mag(w, n):
|
||||
d = w - 1
|
||||
return 1 + np.sign(d) * np.sqrt(np.abs(d) ** 2 / n)
|
||||
# return np.sign(w) * np.sqrt(np.abs(w)**2 / n)
|
||||
def batched_clip_encode(tokens, length, encode_func, num_chunks):
|
||||
embs = []
|
||||
for e in _grouper(32, tokens):
|
||||
enc, pooled, *_ = encode_func(e)
|
||||
enc = enc.reshape((len(e), length, -1))
|
||||
embs.append(enc)
|
||||
|
||||
embs = torch.cat(embs)
|
||||
embs = embs.reshape((len(tokens) // num_chunks, length * num_chunks, -1))
|
||||
return embs
|
||||
|
||||
|
||||
def weights_like(weights, emb):
|
||||
return torch.tensor(weights, dtype=emb.dtype, device=emb.device).reshape(1, -1, 1).expand(emb.shape)
|
||||
|
||||
|
||||
def divide_length(word_ids, weights):
|
||||
sums = dict(zip(*np.unique(word_ids, return_counts=True)))
|
||||
sums[0] = 1
|
||||
weights = [[_norm_mag(w, sums[id]) if id != 0 else 1.0 for w, id in zip(x, y)] for x, y in zip(weights, word_ids)]
|
||||
return weights
|
||||
|
||||
|
||||
def shift_mean_weight(word_ids, weights):
|
||||
delta = 1 - np.mean([w for x, y in zip(weights, word_ids) for w, id in zip(x, y) if id != 0])
|
||||
weights = [[w if id == 0 else w + delta for w, id in zip(x, y)] for x, y in zip(weights, word_ids)]
|
||||
return weights
|
||||
|
||||
|
||||
def scale_to_norm(weights, word_ids, w_max):
|
||||
top = np.max(weights)
|
||||
w_max = min(top, w_max)
|
||||
weights = [[w_max if id == 0 else (w / top) * w_max for w, id in zip(x, y)] for x, y in zip(weights, word_ids)]
|
||||
weights = [
|
||||
[w_max if id == 0 else (w / top) * w_max for w, id in zip(x, y, strict=False)]
|
||||
for x, y in zip(weights, word_ids, strict=False)
|
||||
]
|
||||
return weights
|
||||
|
||||
|
||||
def mask_word_id(tokens, word_ids, target_id, mask_token):
|
||||
new_tokens = [[mask_token if wid == target_id else t for t, wid in zip(x, y)] for x, y in zip(tokens, word_ids)]
|
||||
new_tokens = [
|
||||
[mask_token if wid == target_id else t for t, wid in zip(x, y, strict=False)]
|
||||
for x, y in zip(tokens, word_ids, strict=False)
|
||||
]
|
||||
mask = np.array(word_ids) == target_id
|
||||
return (new_tokens, mask)
|
||||
|
||||
|
||||
def batched_clip_encode(tokens, length, encode_func, num_chunks):
|
||||
embs = []
|
||||
for e in _grouper(32, tokens):
|
||||
enc, pooled = encode_func(e)
|
||||
enc = enc.reshape((len(e), length, -1))
|
||||
embs.append(enc)
|
||||
embs = torch.cat(embs)
|
||||
embs = embs.reshape((len(tokens) // num_chunks, length * num_chunks, -1))
|
||||
return embs
|
||||
|
||||
|
||||
def from_masked(tokens, weights, word_ids, base_emb, length, encode_func, m_token=266):
|
||||
pooled_base = base_emb[0, length - 1 : length, :]
|
||||
wids, inds = np.unique(np.array(word_ids).reshape(-1), return_index=True)
|
||||
weight_dict = dict((id, w) for id, w in zip(wids, np.array(weights).reshape(-1)[inds]) if w != 1.0)
|
||||
|
||||
if len(weight_dict) == 0:
|
||||
return torch.zeros_like(base_emb), base_emb[0, length - 1 : length, :]
|
||||
|
||||
weight_tensor = weights_like(weights, base_emb)
|
||||
|
||||
# m_token = (clip.tokenizer.end_token, 1.0) if clip.tokenizer.pad_with_end else (0,1.0)
|
||||
# TODO: find most suitable masking token here
|
||||
m_token = (m_token, 1.0)
|
||||
|
||||
ws = []
|
||||
masked_tokens = []
|
||||
masks = []
|
||||
|
||||
# create prompts
|
||||
for id, w in weight_dict.items():
|
||||
masked, m = mask_word_id(tokens, word_ids, id, m_token)
|
||||
masked_tokens.extend(masked)
|
||||
masks.append(weights_like(m, base_emb))
|
||||
|
||||
ws.append(w)
|
||||
|
||||
# batch process prompts
|
||||
embs = batched_clip_encode(masked_tokens, length, encode_func, len(tokens))
|
||||
masks = torch.cat(masks)
|
||||
|
||||
embs = base_emb.expand(embs.shape) - embs
|
||||
pooled = embs[0, length - 1 : length, :]
|
||||
|
||||
embs *= masks
|
||||
embs = embs.sum(axis=0, keepdim=True)
|
||||
|
||||
pooled_start = pooled_base.expand(len(ws), -1)
|
||||
ws = torch.tensor(ws).reshape(-1, 1).expand(pooled_start.shape)
|
||||
pooled = (pooled - pooled_start) * (ws - 1)
|
||||
pooled = pooled.mean(axis=0, keepdim=True)
|
||||
|
||||
return ((weight_tensor - 1) * embs), pooled_base + pooled
|
||||
|
||||
|
||||
def mask_inds(tokens, inds, mask_token):
|
||||
clip_len = len(tokens[0])
|
||||
inds_set = set(inds)
|
||||
@@ -112,34 +67,6 @@ def mask_inds(tokens, inds, mask_token):
|
||||
return new_tokens
|
||||
|
||||
|
||||
def down_weight(tokens, weights, word_ids, base_emb, length, encode_func, m_token=266):
|
||||
w, w_inv = np.unique(weights, return_inverse=True)
|
||||
|
||||
if np.sum(w < 1) == 0:
|
||||
return base_emb, tokens, base_emb[0, length - 1 : length, :]
|
||||
# m_token = (clip.tokenizer.end_token, 1.0) if clip.tokenizer.pad_with_end else (0,1.0)
|
||||
# using the comma token as a masking token seems to work better than aos tokens for SD 1.x
|
||||
m_token = (m_token, 1.0)
|
||||
|
||||
masked_tokens = []
|
||||
|
||||
masked_current = tokens
|
||||
for i in range(len(w)):
|
||||
if w[i] >= 1:
|
||||
continue
|
||||
masked_current = mask_inds(masked_current, np.where(w_inv == i)[0], m_token)
|
||||
masked_tokens.extend(masked_current)
|
||||
|
||||
embs = batched_clip_encode(masked_tokens, length, encode_func, len(tokens))
|
||||
embs = torch.cat([base_emb, embs])
|
||||
w = w[w <= 1.0]
|
||||
w_mix = np.diff([0] + w.tolist())
|
||||
w_mix = torch.tensor(w_mix, dtype=embs.dtype, device=embs.device).reshape((-1, 1, 1))
|
||||
|
||||
weighted_emb = (w_mix * embs).sum(axis=0, keepdim=True)
|
||||
return weighted_emb, masked_current, weighted_emb[0, length - 1 : length, :]
|
||||
|
||||
|
||||
def scale_emb_to_mag(base_emb, weighted_emb):
|
||||
norm_base = torch.linalg.norm(base_emb)
|
||||
norm_weighted = torch.linalg.norm(weighted_emb)
|
||||
@@ -147,12 +74,6 @@ def scale_emb_to_mag(base_emb, weighted_emb):
|
||||
return embeddings_final
|
||||
|
||||
|
||||
def recover_dist(base_emb, weighted_emb):
|
||||
fixed_std = (base_emb.std() / weighted_emb.std()) * (weighted_emb - weighted_emb.mean())
|
||||
embeddings_final = fixed_std + (base_emb.mean() - fixed_std.mean())
|
||||
return embeddings_final
|
||||
|
||||
|
||||
def perp_weight(weights, unweighted_embs, empty_embs):
|
||||
unweighted, unweighted_pooled = unweighted_embs
|
||||
zero, zero_pooled = empty_embs
|
||||
@@ -171,72 +92,300 @@ def perp_weight(weights, unweighted_embs, empty_embs):
|
||||
result[~over1] = (unweighted - (1 - weights) * perp)[~over1]
|
||||
result[weights == 0.0] = zero[weights == 0.0]
|
||||
|
||||
# Not sure if this is an implementation bug or if this just doesn't make sense with T5
|
||||
nans = result.isnan()
|
||||
if nans.any():
|
||||
log.warning("perp weight returned NaNs (known to happen with T5), replacing with 0")
|
||||
result[nans] = 0.0
|
||||
|
||||
return result, unweighted_pooled
|
||||
|
||||
|
||||
def style_comfy(encoder, tokens, **kwargs):
|
||||
tokens = encoder.without_word_ids(tokens)
|
||||
return encoder.encode_fn(tokens)
|
||||
|
||||
|
||||
def style_a1111(encoder, tokens, **kwargs):
|
||||
base_emb, pooled, *extra = encoder.base_emb(tokens)
|
||||
weighted_emb = base_emb * weights_like(encoder.weights(tokens), base_emb)
|
||||
weighted_emb = (base_emb.mean() / weighted_emb.mean()) * weighted_emb # renormalize
|
||||
return (weighted_emb, pooled) + tuple(extra)
|
||||
|
||||
|
||||
def style_compel(encoder, tokens, **kwargs):
|
||||
pos_tokens = encoder.weighted_with(tokens, lambda w: w if w > 1.0 else 1.0)
|
||||
weighted_emb, pooled, *extra = encoder.encode_fn(pos_tokens)
|
||||
weighted_emb, _, pooled = encoder.down_weight(
|
||||
pos_tokens, encoder.weights(tokens), encoder.word_ids(tokens), weighted_emb, pooled
|
||||
)
|
||||
return (weighted_emb, pooled) + tuple(extra)
|
||||
|
||||
|
||||
def style_comfypp(encoder, tokens, **kwargs):
|
||||
unweighted_tokens = encoder.unweighted(tokens)
|
||||
base_emb, pooled_base, *extra = encoder.base_emb(tokens)
|
||||
weighted_emb, tokens_down, _ = encoder.down_weight(
|
||||
unweighted_tokens, encoder.weights(tokens), encoder.word_ids(tokens), base_emb, pooled_base
|
||||
)
|
||||
weights = encoder.weights(encoder.weighted_with(tokens, lambda w: w if w > 1.0 else 1.0))
|
||||
embs, pooled = encoder.from_masked(
|
||||
unweighted_tokens,
|
||||
weights,
|
||||
encoder.word_ids(tokens),
|
||||
base_emb,
|
||||
pooled_base,
|
||||
)
|
||||
weighted_emb += embs
|
||||
|
||||
return (weighted_emb, pooled) + tuple(extra)
|
||||
|
||||
|
||||
def style_downweight(encoder, tokens, **kwargs):
|
||||
weights = scale_to_norm(encoder.weights(tokens), encoder.word_ids(tokens), encoder.w_max)
|
||||
base_emb, pooled_base, *extra = encoder.base_emb(tokens)
|
||||
weighted_emb, _, pooled = encoder.down_weight(
|
||||
encoder.unweighted(tokens), weights, encoder.word_ids(tokens), base_emb, pooled_base
|
||||
)
|
||||
|
||||
return (weighted_emb, pooled) + tuple(extra)
|
||||
|
||||
|
||||
def style_perp(encoder, tokens, **kwargs):
|
||||
zero_emb, zero_pooled, *_ = encoder.encode_fn(encoder.tokenizer.tokenize_with_weights(""))
|
||||
base_emb, pooled, *extra = encoder.base_emb(tokens)
|
||||
return perp_weight(encoder.weights(tokens), (base_emb, pooled), (zero_emb, zero_pooled)) + tuple(extra)
|
||||
|
||||
|
||||
def apply_negpip(encoder, emb, pooled, **kwargs):
|
||||
original_tokens = kwargs["original_tokens"]
|
||||
emb_negpip = torch.empty_like(emb).repeat(1, 2, 1)
|
||||
emb_negpip[:, 0::2, :] = emb
|
||||
emb_negpip[:, 1::2, :] = emb * weights_like(encoder.signs(original_tokens), emb)
|
||||
return emb_negpip, pooled
|
||||
|
||||
|
||||
def norm_length(encoder, tokens, **kwargs):
|
||||
word_ids = encoder.word_ids(tokens)
|
||||
sums = dict(zip(*np.unique(word_ids, return_counts=True), strict=False))
|
||||
sums[0] = 1
|
||||
tokens = [[(t, _norm_mag(w, sums[id]) if id != 0 else 1.0, id) for (t, w, id) in x] for x in tokens]
|
||||
return tokens
|
||||
|
||||
|
||||
def norm_mean(encoder, tokens, **kwargs):
|
||||
weights = encoder.weights(tokens)
|
||||
word_ids = encoder.word_ids(tokens)
|
||||
delta = 1 - np.mean(
|
||||
[w for x, y in zip(weights, word_ids, strict=False) for w, id in zip(x, y, strict=False) if id != 0]
|
||||
)
|
||||
tokens = [[(t, w if id == 0 else w + delta, id) for (t, w, id) in x] for x in tokens]
|
||||
return tokens
|
||||
|
||||
|
||||
def norm_none(encoder, tokens, **kwargs):
|
||||
return tokens
|
||||
|
||||
|
||||
class AdvancedEncoder:
|
||||
STYLES = {
|
||||
"A1111": style_a1111,
|
||||
"comfy": style_comfy,
|
||||
"comfy++": style_comfypp,
|
||||
"compel": style_compel,
|
||||
"down_weight": style_downweight,
|
||||
"perp": style_perp,
|
||||
}
|
||||
NORMALIZATION_OPS = {
|
||||
"none": norm_none,
|
||||
"length": norm_length,
|
||||
"mean": norm_mean,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def add_encoder(cls, name, fn):
|
||||
cls.STYLES[name] = fn
|
||||
|
||||
@classmethod
|
||||
def add_normalization_op(cls, name, fn):
|
||||
cls.NORMALIZATION_OPS[name] = fn
|
||||
|
||||
@classmethod
|
||||
def weighted_with(cls, tokens, fn=id, word_ids=True):
|
||||
w = ([(t, fn(w), id) for t, w, id in x] for x in tokens)
|
||||
if not word_ids:
|
||||
w = cls.without_word_ids(w)
|
||||
return list(w)
|
||||
|
||||
@classmethod
|
||||
def unweighted(cls, tokens, word_ids=False):
|
||||
return cls.weighted_with(tokens, fn=lambda w: 1.0, word_ids=word_ids)
|
||||
|
||||
@classmethod
|
||||
def tokens_only(cls, tokens):
|
||||
return list([t[0] for t in x] for x in tokens)
|
||||
|
||||
@classmethod
|
||||
def weights(cls, tokens):
|
||||
return list([t[1] for t in x] for x in tokens)
|
||||
|
||||
@classmethod
|
||||
def word_ids(cls, tokens):
|
||||
return list([t[2] for t in x] for x in tokens)
|
||||
|
||||
@classmethod
|
||||
def signs(cls, tokens):
|
||||
return list([copysign(1, t[1]) for t in x] for x in tokens)
|
||||
|
||||
@classmethod
|
||||
def without_word_ids(cls, tokens):
|
||||
return list([(t, w) for t, w, _ in x] for x in tokens)
|
||||
|
||||
def __init__(self, encode_fn, style, normalization, tokenizer, m_token="+", w_max=1.0, **extra_args):
|
||||
self.encode_fn = encode_fn
|
||||
self.preprocessors = []
|
||||
self.postprocessors = []
|
||||
self.tokenizer = tokenizer
|
||||
self.extra_args = extra_args
|
||||
self.m_token = tokenizer.tokenize_with_weights(m_token)[0][tokenizer.tokens_start]
|
||||
self.max_length = tokenizer.max_length if tokenizer.pad_to_max_length else None
|
||||
self.w_max = w_max
|
||||
|
||||
if style == "comfy++" and not self.max_length:
|
||||
log.warning("comfy++ does not work with tokenizer %s, using default weighting", tokenizer)
|
||||
style = "comfy"
|
||||
|
||||
norms = normalization.split("+")
|
||||
assert style in self.STYLES, f"Invalid weight interpretation: {style}"
|
||||
self.weight_fn = self.STYLES[style]
|
||||
for n in norms:
|
||||
n = n.strip()
|
||||
assert n in self.NORMALIZATION_OPS, f"Invalid normalization: {normalization}"
|
||||
self.preprocessors.append(self.NORMALIZATION_OPS[n])
|
||||
|
||||
negpip = extra_args.get("has_negpip")
|
||||
if negpip:
|
||||
|
||||
def _encode(t):
|
||||
emb, pooled, *extra = encode_fn(t)
|
||||
return (emb[:, 0::2, :], pooled) + tuple(extra)
|
||||
|
||||
self.encode_fn = _encode
|
||||
self.preprocessors.insert(0, lambda encoder, tokens, **kwargs: encoder.weighted_with(tokens, abs))
|
||||
self.postprocessors.insert(0, apply_negpip)
|
||||
|
||||
def base_emb(self, tokens):
|
||||
unweighted = self.unweighted(tokens)
|
||||
return self.encode_fn(unweighted)
|
||||
|
||||
def down_weight(self, tokens, weights, word_ids, base_emb, pooled_base):
|
||||
w, w_inv = np.unique(weights, return_inverse=True)
|
||||
|
||||
if np.sum(w < 1) == 0:
|
||||
return (
|
||||
base_emb,
|
||||
tokens,
|
||||
(
|
||||
base_emb[0, self.max_length - 1 : self.max_length, :]
|
||||
if (pooled_base is not None and self.max_length)
|
||||
else None
|
||||
),
|
||||
)
|
||||
|
||||
masked_current = tokens
|
||||
emblist = [base_emb]
|
||||
for i in range(len(w)):
|
||||
if w[i] >= 1:
|
||||
continue
|
||||
masked_current = mask_inds(masked_current, np.where(w_inv == i)[0], self.m_token)
|
||||
masked, _, *extra = self.encode_fn(masked_current)
|
||||
emblist.append(masked)
|
||||
|
||||
embs = torch.cat(emblist)
|
||||
w = w[w <= 1.0]
|
||||
w_mix = np.diff([0] + w.tolist())
|
||||
w_mix = torch.tensor(w_mix, dtype=embs.dtype, device=embs.device).reshape((-1, 1, 1))
|
||||
|
||||
weighted_emb = (w_mix * embs).sum(dim=0, keepdim=True)
|
||||
pooled = pooled_base
|
||||
if pooled is not None and self.max_length:
|
||||
pooled = weighted_emb[0, self.max_length - 1 : self.max_length, :]
|
||||
return weighted_emb, masked_current, pooled
|
||||
|
||||
def from_masked(self, tokens, weights, word_ids, base_emb, pooled_base):
|
||||
wids, inds = np.unique(np.array(word_ids).reshape(-1), return_index=True)
|
||||
weight_dict = dict(
|
||||
(id, w) for id, w in zip(wids, np.array(weights).reshape(-1)[inds], strict=False) if w != 1.0
|
||||
)
|
||||
|
||||
if len(weight_dict) == 0:
|
||||
return torch.zeros_like(base_emb), torch.zeros_like(pooled_base) if pooled_base is not None else None
|
||||
|
||||
weight_tensor = weights_like(weights, base_emb)
|
||||
|
||||
ws = []
|
||||
masked_tokens = []
|
||||
masks = []
|
||||
|
||||
# create prompts
|
||||
for id, w in weight_dict.items():
|
||||
masked, m = mask_word_id(tokens, word_ids, id, self.m_token)
|
||||
masks.append(weights_like(m, base_emb))
|
||||
masked_tokens.extend(masked)
|
||||
|
||||
ws.append(w)
|
||||
|
||||
# TODO: figure out how to get rid of this
|
||||
embs = batched_clip_encode(masked_tokens, self.max_length, self.encode_fn, len(tokens))
|
||||
masks = torch.cat(masks)
|
||||
|
||||
embs = base_emb.expand(embs.shape) - embs
|
||||
pooled = None
|
||||
if pooled_base is not None and self.max_length:
|
||||
pooled = embs[0, self.max_length - 1 : self.max_length, :]
|
||||
pooled_start = pooled_base.expand(len(ws), -1)
|
||||
ws = torch.tensor(ws).reshape(-1, 1).expand(pooled_start.shape)
|
||||
pooled = (pooled - pooled_start) * (ws - 1)
|
||||
pooled = pooled.mean(dim=0, keepdim=True)
|
||||
pooled = pooled_base + pooled
|
||||
|
||||
if embs.shape[0] != masks.shape[0]:
|
||||
embs = embs.repeat(masks.shape[0], 1, 1)
|
||||
embs *= masks
|
||||
embs = embs.sum(axis=0, keepdim=True)
|
||||
|
||||
return ((weight_tensor - 1) * embs), pooled
|
||||
|
||||
def __call__(self, tokens, apply_to_pooled=False, return_pooled=False):
|
||||
normalized_tokens = tokens
|
||||
for op in self.preprocessors:
|
||||
normalized_tokens = op(self, normalized_tokens)
|
||||
|
||||
emb, pooled, *extra = self.weight_fn(self, normalized_tokens, original_tokens=tokens)
|
||||
|
||||
for fn in self.postprocessors:
|
||||
emb, pooled = fn(self, emb, pooled, tokens=tokens, original_tokens=tokens)
|
||||
|
||||
if not return_pooled:
|
||||
pooled = None
|
||||
elif not apply_to_pooled:
|
||||
_, pooled, *_ = self.base_emb(tokens)
|
||||
return (emb, pooled) + tuple(extra)
|
||||
|
||||
|
||||
def advanced_encode_from_tokens(
|
||||
tokenized,
|
||||
token_normalization,
|
||||
weight_interpretation,
|
||||
encode_func,
|
||||
m_token=266,
|
||||
length=77,
|
||||
m_token="+",
|
||||
w_max=1.0,
|
||||
return_pooled=False,
|
||||
apply_to_pooled=False,
|
||||
**extra_args
|
||||
tokenizer=None,
|
||||
**extra_args,
|
||||
):
|
||||
tokens = [[t for t, _, _ in x] for x in tokenized]
|
||||
weights = [[w for _, w, _ in x] for x in tokenized]
|
||||
word_ids = [[wid for _, _, wid in x] for x in tokenized]
|
||||
|
||||
for op in token_normalization.split("+"):
|
||||
op = op.strip()
|
||||
if op == "length":
|
||||
# distribute down/up weights over word lengths
|
||||
weights = divide_length(word_ids, weights)
|
||||
if op == "mean":
|
||||
weights = shift_mean_weight(word_ids, weights)
|
||||
|
||||
pooled = None
|
||||
|
||||
if weight_interpretation == "comfy":
|
||||
weighted_tokens = [[(t, w) for t, w in zip(x, y)] for x, y in zip(tokens, weights)]
|
||||
weighted_emb, pooled_base = encode_func(weighted_tokens)
|
||||
pooled = pooled_base
|
||||
else:
|
||||
unweighted_tokens = [[(t, 1.0) for t, _, _ in x] for x in tokenized]
|
||||
base_emb, pooled_base = encode_func(unweighted_tokens)
|
||||
|
||||
if weight_interpretation == "A1111":
|
||||
weighted_emb = base_emb * weights_like(weights, base_emb) # from_zero
|
||||
weighted_emb = (base_emb.mean() / weighted_emb.mean()) * weighted_emb # renormalize
|
||||
pooled = pooled_base
|
||||
|
||||
if weight_interpretation == "compel":
|
||||
pos_tokens = [[(t, w) if w >= 1.0 else (t, 1.0) for t, w in zip(x, y)] for x, y in zip(tokens, weights)]
|
||||
weighted_emb, _ = encode_func(pos_tokens)
|
||||
weighted_emb, _, pooled = down_weight(pos_tokens, weights, word_ids, weighted_emb, length, encode_func)
|
||||
|
||||
if weight_interpretation == "comfy++":
|
||||
weighted_emb, tokens_down, _ = down_weight(unweighted_tokens, weights, word_ids, base_emb, length, encode_func)
|
||||
weights = [[w if w > 1.0 else 1.0 for w in x] for x in weights]
|
||||
# unweighted_tokens = [[(t,1.0) for t, _,_ in x] for x in tokens_down]
|
||||
embs, pooled = from_masked(unweighted_tokens, weights, word_ids, base_emb, length, encode_func)
|
||||
weighted_emb += embs
|
||||
|
||||
if weight_interpretation == "down_weight":
|
||||
weights = scale_to_norm(weights, word_ids, w_max)
|
||||
weighted_emb, _, pooled = down_weight(unweighted_tokens, weights, word_ids, base_emb, length, encode_func)
|
||||
|
||||
if weight_interpretation == "perp":
|
||||
weighted_emb, pooled = perp_weight(
|
||||
weights, (base_emb, pooled_base), encode_func(extra_args["tokenizer"].tokenize_with_weights(""))
|
||||
)
|
||||
|
||||
if return_pooled:
|
||||
if apply_to_pooled:
|
||||
return weighted_emb, pooled
|
||||
else:
|
||||
return weighted_emb, pooled_base
|
||||
return weighted_emb, None
|
||||
enc = AdvancedEncoder(
|
||||
encode_func, weight_interpretation, token_normalization, tokenizer, m_token, w_max, **extra_args
|
||||
)
|
||||
return enc(tokenized, return_pooled=return_pooled, apply_to_pooled=apply_to_pooled)
|
||||
|
||||
@@ -0,0 +1,247 @@
|
||||
# Lifted from https://github.com/pamparamm/ComfyUI-ppm/blob/c3e6b673ee2d424405dcb99aeed89f21943c89ac/nodes_ppm/attention_couple_ppm.py
|
||||
# Original implementation by laksjdjf, hako-mikan, Haoming02 licensed under GPL-3.0
|
||||
# https://github.com/laksjdjf/cgem156-ComfyUI/blob/1f5533f7f31345bafe4b833cbee15a3c4ad74167/scripts/attention_couple/node.py
|
||||
# https://github.com/Haoming02/sd-forge-couple/blob/e8e258e982a8d149ba59a4bc43b945467604311c/scripts/attention_couple.py
|
||||
import itertools
|
||||
import logging
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from comfy.hooks import EnumHookScope, HookGroup, TransformerOptionsHook, set_hooks_for_conditioning
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def set_cond_attnmask(base_cond, extra_conds, fill=False):
|
||||
hook = AttentionCoupleHook()
|
||||
c = [base_cond[0][0], base_cond[0][1].copy()]
|
||||
# hook uses these, remove them to avoid doing latent masking
|
||||
c[1].pop("mask", None)
|
||||
c[1].pop("strength", None)
|
||||
c[1].pop("mask_strength", None)
|
||||
c = [c]
|
||||
c.extend(base_cond[1:])
|
||||
|
||||
hook.initialize_regions(base_cond[0], extra_conds, fill=fill)
|
||||
group = HookGroup()
|
||||
group.add(hook)
|
||||
|
||||
return set_hooks_for_conditioning(c, hooks=group, append_hooks=True)
|
||||
|
||||
|
||||
def get_mask(mask, batch_size, num_tokens, extra_options):
|
||||
activations_shape = extra_options["activations_shape"]
|
||||
size = activations_shape[-2:]
|
||||
|
||||
num_conds = mask.shape[0]
|
||||
mask_downsample = F.interpolate(mask, size=size, mode="nearest")
|
||||
mask_downsample_reshaped = mask_downsample.view(num_conds, num_tokens, 1).repeat_interleave(batch_size, dim=0)
|
||||
|
||||
return mask_downsample_reshaped
|
||||
|
||||
|
||||
class Proxy:
|
||||
def __init__(self, function):
|
||||
self.function = function
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
self.function.__self__.to(*args, **kwargs)
|
||||
return self
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
return self.function(*args, *kwargs)
|
||||
|
||||
|
||||
class AttentionCoupleHook(TransformerOptionsHook):
|
||||
COND_UNCOND_COUPLE_OPTION = "cond_or_uncond_hook_couple"
|
||||
COND = 0
|
||||
UNCOND = 1
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(hook_scope=EnumHookScope.HookedOnly)
|
||||
|
||||
self.transformers_dict = {
|
||||
"patches": {
|
||||
"attn2_output_patch": [Proxy(self.attn2_output_patch)],
|
||||
"attn2_patch": [Proxy(self.attn2_patch)],
|
||||
}
|
||||
}
|
||||
self.has_negpip = False
|
||||
# calculate later. All clones must refer to the same kv dict
|
||||
self.kv = {"k": [], "v": []}
|
||||
|
||||
def initialize_regions(self, base_cond, conds, fill):
|
||||
self.num_conds = len(conds) + 1
|
||||
self.base_strength = base_cond[1].get("strength", 1.0)
|
||||
self.strengths: list[float] = [cond[1].get("strength", 1.0) for cond in conds]
|
||||
self.conds: list[torch.Tensor] = [base_cond[0]] + [cond[0] for cond in conds]
|
||||
base_mask = base_cond[1].get("mask", None)
|
||||
masks = [cond[1].get("mask") * cond[1].get("mask_strength") for cond in conds]
|
||||
if len(masks) < 1:
|
||||
raise ValueError("Attention Couple hook makes no sense without masked conds")
|
||||
|
||||
if any(m is None for m in masks):
|
||||
raise ValueError("All conds given to Attention Couple must have masks")
|
||||
|
||||
if any(m.shape != masks[0].shape for m in masks) or (
|
||||
base_mask is not None and base_mask.shape != masks[0].shape
|
||||
):
|
||||
largest_shape = max(m.shape for m in masks)
|
||||
if base_mask is not None:
|
||||
largest_shape = max(largest_shape, base_mask.shape)
|
||||
log.warning("Attention Couple: Masks are irregularly shaped, resizing them all to match the largest")
|
||||
for i in range(len(masks)):
|
||||
masks[i] = F.interpolate(masks[i].unsqueeze(1), size=largest_shape[1:], mode="nearest-exact").squeeze(1)
|
||||
|
||||
if base_mask is not None:
|
||||
base_mask = F.interpolate(base_mask.unsqueeze(1), size=largest_shape[1:], mode="nearest-exact").squeeze(
|
||||
1
|
||||
)
|
||||
|
||||
if base_mask is None:
|
||||
if not fill:
|
||||
raise ValueError("You must specify a base mask when fill=False")
|
||||
sum = torch.stack(masks, dim=0).sum(dim=0)
|
||||
base_mask = torch.zeros_like(sum)
|
||||
base_mask[sum <= 0] = 1.0
|
||||
|
||||
mask = [base_mask] + masks
|
||||
mask = torch.stack(mask, dim=0)
|
||||
if mask.sum(dim=0).min() <= 0 and not fill:
|
||||
raise ValueError("Masks contain non-filled areas")
|
||||
|
||||
self.mask = mask / mask.sum(dim=0, keepdim=True)
|
||||
|
||||
def on_apply_hooks(self, model: ModelPatcher, transformer_options: dict[str, Any]):
|
||||
if self.kv["k"] is None:
|
||||
self.has_negpip = model.model_options.get("ppm_negpip", False)
|
||||
log.debug("AttentionCouple has_negpip=%s", self.has_negpip)
|
||||
|
||||
# Skip the base cond here, which is always first
|
||||
if self.has_negpip:
|
||||
self.kv["k"] = [cond[:, 0::2] for cond in self.conds[1:]]
|
||||
self.kv["v"] = [cond[:, 1::2] for cond in self.conds[1:]]
|
||||
else:
|
||||
self.kv["k"] = self.kv["v"] = self.conds[1:]
|
||||
|
||||
return super().on_apply_hooks(model, transformer_options)
|
||||
|
||||
def clone(self):
|
||||
c: AttentionCoupleHook = super().clone()
|
||||
c.mask = self.mask
|
||||
c.conds = self.conds
|
||||
c.kv = self.kv
|
||||
c.has_negpip = self.has_negpip
|
||||
c.base_strength = self.base_strength
|
||||
c.strengths = self.strengths
|
||||
c.num_conds = self.num_conds
|
||||
return c
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
self.conds = [c.to(*args, **kwargs) for c in self.conds]
|
||||
self.mask = self.mask.to(*args, **kwargs)
|
||||
if self.kv["k"] is not None:
|
||||
self.kv["k"] = [c.to(*args, **kwargs) for c in self.kv["k"]]
|
||||
self.kv["v"] = [c.to(*args, **kwargs) for c in self.kv["v"]]
|
||||
return self
|
||||
|
||||
def attn2_patch(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, extra_options):
|
||||
cond_or_uncond = extra_options["cond_or_uncond"]
|
||||
cond_or_uncond_couple = extra_options[self.COND_UNCOND_COUPLE_OPTION] = list(cond_or_uncond)
|
||||
num_chunks = len(cond_or_uncond)
|
||||
|
||||
# Cloning messes up the device sometimes
|
||||
if self.kv["k"][0].device != k.device:
|
||||
self.to(k)
|
||||
|
||||
conds_k = self.kv["k"]
|
||||
conds_v = self.kv["v"]
|
||||
|
||||
lcm_tokens_k = math.lcm(k.shape[1], *(cond.shape[1] for cond in conds_k))
|
||||
lcm_tokens_v = math.lcm(v.shape[1], *(cond.shape[1] for cond in conds_v))
|
||||
q_chunks = q.chunk(num_chunks, dim=0)
|
||||
k_chunks = k.chunk(num_chunks, dim=0)
|
||||
v_chunks = v.chunk(num_chunks, dim=0)
|
||||
|
||||
bs = q.shape[0] // num_chunks
|
||||
|
||||
conds_k_tensor = conds_v_tensor = torch.cat(
|
||||
[cond.repeat(bs, lcm_tokens_k // cond.shape[1], 1) * self.strengths[i] for i, cond in enumerate(conds_k)],
|
||||
dim=0,
|
||||
)
|
||||
if self.has_negpip:
|
||||
conds_v_tensor = torch.cat(
|
||||
[
|
||||
cond.repeat(bs, lcm_tokens_v // cond.shape[1], 1) * self.strengths[i]
|
||||
for i, cond in enumerate(conds_v)
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
|
||||
qs, ks, vs = [], [], []
|
||||
cond_or_uncond_couple.clear()
|
||||
|
||||
for i, cond_type in enumerate(cond_or_uncond):
|
||||
q_target = q_chunks[i]
|
||||
k_target = k_chunks[i].repeat(1, lcm_tokens_k // k.shape[1], 1)
|
||||
v_target = v_chunks[i].repeat(1, lcm_tokens_v // v.shape[1], 1)
|
||||
if cond_type == self.UNCOND:
|
||||
qs.append(q_target)
|
||||
ks.append(k_target)
|
||||
vs.append(v_target)
|
||||
cond_or_uncond_couple.append(self.UNCOND)
|
||||
else:
|
||||
qs.append(q_target.repeat(self.num_conds, 1, 1))
|
||||
ks.append(
|
||||
torch.cat(
|
||||
[
|
||||
k_target * self.base_strength,
|
||||
conds_k_tensor,
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
vs.append(
|
||||
torch.cat(
|
||||
[
|
||||
v_target * self.base_strength,
|
||||
conds_v_tensor,
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
assert self.num_conds is not None, "this is a bug"
|
||||
cond_or_uncond_couple.extend(itertools.repeat(self.COND, self.num_conds))
|
||||
|
||||
q = torch.cat(qs, dim=0)
|
||||
k = torch.cat(ks, dim=0)
|
||||
v = torch.cat(vs, dim=0)
|
||||
|
||||
return q, k, v
|
||||
|
||||
def attn2_output_patch(self, out, extra_options):
|
||||
cond_or_uncond = extra_options[self.COND_UNCOND_COUPLE_OPTION]
|
||||
bs = out.shape[0] // len(cond_or_uncond)
|
||||
mask_downsample = get_mask(self.mask, bs, out.shape[1], extra_options)
|
||||
outputs = []
|
||||
cond_outputs = []
|
||||
i_cond = 0
|
||||
for i, cond_type in enumerate(cond_or_uncond):
|
||||
pos, next_pos = i * bs, (i + 1) * bs
|
||||
|
||||
if cond_type == self.UNCOND:
|
||||
outputs.append(out[pos:next_pos])
|
||||
else:
|
||||
pos_cond, next_pos_cond = i_cond * bs, (i_cond + 1) * bs
|
||||
masked_output = out[pos:next_pos] * mask_downsample[pos_cond:next_pos_cond]
|
||||
cond_outputs.append(masked_output)
|
||||
i_cond += 1
|
||||
|
||||
if len(cond_outputs) > 0:
|
||||
cond_output = torch.stack(cond_outputs).sum(0)
|
||||
outputs.append(cond_output)
|
||||
|
||||
return torch.cat(outputs, dim=0)
|
||||
@@ -1,47 +0,0 @@
|
||||
import comfy_execution.caching
|
||||
from comfy_execution.graph_utils import is_link
|
||||
import nodes
|
||||
from os import environ
|
||||
|
||||
import logging
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
include_unique_id_in_input = comfy_execution.caching.include_unique_id_in_input
|
||||
|
||||
|
||||
def promptcontrol_get_immediate_node_signature(self, dynprompt, node_id, ancestor_order_mapping):
|
||||
if not dynprompt.has_node(node_id):
|
||||
# This node doesn't exist -- we can't cache it.
|
||||
return [float("NaN")]
|
||||
node = dynprompt.get_node(node_id)
|
||||
class_type = node["class_type"]
|
||||
class_def = nodes.NODE_CLASS_MAPPINGS[class_type]
|
||||
inputs = node["inputs"]
|
||||
if hasattr(class_def, "CACHE_KEY"):
|
||||
inputs = getattr(class_def, "CACHE_KEY")(inputs)
|
||||
signature = [class_type, self.is_changed_cache.get(node_id)]
|
||||
if (
|
||||
self.include_node_id_in_input()
|
||||
or (hasattr(class_def, "NOT_IDEMPOTENT") and class_def.NOT_IDEMPOTENT)
|
||||
or include_unique_id_in_input(class_type)
|
||||
):
|
||||
signature.append(node_id)
|
||||
for key in sorted(inputs.keys()):
|
||||
if is_link(inputs[key]):
|
||||
(ancestor_id, ancestor_socket) = inputs[key]
|
||||
ancestor_index = ancestor_order_mapping[ancestor_id]
|
||||
signature.append((key, ("ANCESTOR", ancestor_index, ancestor_socket)))
|
||||
else:
|
||||
signature.append((key, inputs[key]))
|
||||
return signature
|
||||
|
||||
|
||||
def init():
|
||||
if environ.get("PROMPTCONTROL_ENABLE_CACHE_HACK") != "1":
|
||||
return
|
||||
log.warning("Enabling Prompt Control cache hack")
|
||||
comfy_execution.caching.CacheKeySetInputSignature.get_immediate_node_signature = (
|
||||
promptcontrol_get_immediate_node_signature
|
||||
)
|
||||
@@ -1,9 +1,9 @@
|
||||
import torch
|
||||
import copy
|
||||
import logging
|
||||
import re
|
||||
|
||||
import numpy as np
|
||||
import logging
|
||||
import torch
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
@@ -69,10 +69,7 @@ def cutoff_add_region(
|
||||
clip_regions["start_from_masked"] = float(start_from_masked)
|
||||
if mask_token is not None:
|
||||
clip_regions["mask_token"] = tokenizer.tokenizer(mask_token)["input_ids"][1]
|
||||
if weight is None:
|
||||
weight = 1.0
|
||||
else:
|
||||
weight = float(weight)
|
||||
weight = 1.0 if weight is None else float(weight)
|
||||
|
||||
region_text = region_text.strip()
|
||||
target_text = target_text.strip()
|
||||
@@ -139,7 +136,7 @@ def cutoff_add_region(
|
||||
|
||||
|
||||
def create_masked_prompt(weighted_tokens, mask, mask_token):
|
||||
mask_ids = list(zip(*np.nonzero(mask.reshape((len(weighted_tokens), -1)))))
|
||||
mask_ids = list(zip(*np.nonzero(mask.reshape((len(weighted_tokens), -1))), strict=False))
|
||||
new_prompt = copy.deepcopy(weighted_tokens)
|
||||
for x, y in mask_ids:
|
||||
new_prompt[x][y] = (mask_token,) + new_prompt[x][y][1:]
|
||||
@@ -200,7 +197,9 @@ def encode_regions(clip_regions, encode, tokenizer):
|
||||
base_embedding_outer = base_embedding_full * (1 - strict_mask) + base_embedding_masked * strict_mask
|
||||
|
||||
region_embeddings = []
|
||||
for region, target, weight in zip(clip_regions["regions"], clip_regions["targets"], clip_regions["weights"]):
|
||||
for region, target, weight in zip(
|
||||
clip_regions["regions"], clip_regions["targets"], clip_regions["weights"], strict=False
|
||||
):
|
||||
region_masking = torch.tensor(
|
||||
regions_normalized * region * weight, dtype=base_embedding_full.dtype, device=base_embedding_full.device
|
||||
).unsqueeze(-1)
|
||||
@@ -209,14 +208,21 @@ def encode_regions(clip_regions, encode, tokenizer):
|
||||
debug_tokens("region", region_prompt, tokenizer)
|
||||
region_emb, _ = encode(region_prompt)
|
||||
region_emb -= base_embedding_start
|
||||
# NegPiP support:
|
||||
if region_emb.shape[1] == 2 * region_masking.shape[1]:
|
||||
region_masking = torch.repeat_interleave(region_masking, 2, dim=1)
|
||||
region_emb *= region_masking
|
||||
|
||||
region_embeddings.append(region_emb)
|
||||
region_embeddings = torch.stack(region_embeddings).sum(axis=0)
|
||||
region_embeddings = torch.stack(region_embeddings).sum(dim=0)
|
||||
|
||||
embeddings_final_mask = torch.tensor(
|
||||
global_region_mask, dtype=base_embedding_full.dtype, device=base_embedding_full.device
|
||||
).unsqueeze(-1)
|
||||
# NegPiP support:
|
||||
if region_embeddings.shape[1] == 2 * embeddings_final_mask.shape[1]:
|
||||
embeddings_final_mask = torch.repeat_interleave(embeddings_final_mask, 2, dim=1)
|
||||
|
||||
embeddings_final = base_embedding_start * embeddings_final_mask + base_embedding_outer * (1 - embeddings_final_mask)
|
||||
embeddings_final += region_embeddings
|
||||
return embeddings_final, pool
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
import re
|
||||
|
||||
from .utils import parse_args
|
||||
|
||||
CUTOFF_RE = re.compile(r"\[CUT:((.*?):(.*?))\]")
|
||||
|
||||
|
||||
def noop(x):
|
||||
return x
|
||||
|
||||
|
||||
def parse_cuts(string):
|
||||
text = CUTOFF_RE.sub(r"\2", string)
|
||||
cutoffs = CUTOFF_RE.findall(string)
|
||||
cs = []
|
||||
for x, *_ in cutoffs:
|
||||
p = x.split(":")
|
||||
args = parse_args(
|
||||
p, [(str, ""), (str, ""), (float, 0), (float, None), (float, None), (noop, None)], strip=False
|
||||
)
|
||||
args = tuple(args)
|
||||
if not args[0] or not args[1] or (args[5] is not None and not args[5].strip()):
|
||||
raise ValueError(f"Invalid CUT spec: [CUT:{x}]")
|
||||
cs.append(args)
|
||||
return text, cs
|
||||
@@ -0,0 +1,80 @@
|
||||
# vim: sw=4 ts=4
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
|
||||
from .utils import find_closing_paren, get_function
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def parse_search(search):
|
||||
arg_start = search.find("(")
|
||||
args = ""
|
||||
name = search.strip()
|
||||
if arg_start > 0:
|
||||
arg_end = find_closing_paren(search, arg_start + 1)
|
||||
if arg_end < 0:
|
||||
arg_end = len(search)
|
||||
name = search[:arg_start].strip()
|
||||
args = search[arg_start + 1 : arg_end]
|
||||
|
||||
if not name:
|
||||
return None
|
||||
args = args.strip()
|
||||
# If using the form DEF(F()=$1) then the default value of $1 is the empty string
|
||||
args = [a.strip() for a in args.split(";")] if arg_start > 0 else []
|
||||
return name, args
|
||||
|
||||
|
||||
def expand_macros(text):
|
||||
text, defs = get_function(text, "DEF", defaults=None)
|
||||
res = text
|
||||
prevres = text
|
||||
replacements = []
|
||||
for d in defs:
|
||||
if not d.args:
|
||||
continue
|
||||
r = d.args[0].split("=", 1)
|
||||
search = parse_search(r[0].strip())
|
||||
if not search or len(r) != 2:
|
||||
log.warning("Ignoring invalid DEF(%s)", d)
|
||||
continue
|
||||
replacements.append((search, r[1].strip()))
|
||||
iterations = 0
|
||||
while True:
|
||||
iterations += 1
|
||||
if iterations > 10:
|
||||
raise ValueError("Unable to resolve DEFs, make sure there are no cycles!")
|
||||
return text
|
||||
for search, replace in replacements:
|
||||
res = substitute_defcall(res, search, replace)
|
||||
if res == prevres:
|
||||
break
|
||||
prevres = res
|
||||
if res.strip() != text.strip():
|
||||
res = res.strip()
|
||||
log.info("DEFs expanded to: %s", res)
|
||||
return res
|
||||
|
||||
|
||||
def substitute_defcall(text, search, replace):
|
||||
name, default_args = search
|
||||
text, defns = get_function(text, name, defaults=None, placeholder=f"DEFNCALL{name}", require_args=False)
|
||||
for i, d in enumerate(defns):
|
||||
ph = d.placeholder
|
||||
assert ph is not None, "This is a bug"
|
||||
parameters = d.args
|
||||
paramvals = []
|
||||
if parameters:
|
||||
paramvals = [x.strip() for x in parameters[0].split(";")]
|
||||
r = replace
|
||||
for i, v in enumerate(paramvals):
|
||||
r = re.sub(rf"\${i + 1}\b", v, r)
|
||||
|
||||
for i, v in enumerate(default_args):
|
||||
r = re.sub(rf"\${i + 1}\b", v, r)
|
||||
|
||||
text = text.replace(ph, r)
|
||||
return text
|
||||
@@ -1,79 +0,0 @@
|
||||
import logging
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
from comfy.hooks import TransformerOptionsHook, HookGroup, EnumHookScope
|
||||
from comfy.ldm.modules.attention import optimized_attention
|
||||
import torch.nn.functional as F
|
||||
import torch
|
||||
from math import sqrt
|
||||
|
||||
|
||||
class MaskedAttn2:
|
||||
def __init__(self, mask):
|
||||
self.mask = mask
|
||||
|
||||
def __call__(self, q, k, v, extra_options):
|
||||
mask = self.mask
|
||||
orig_shape = extra_options["original_shape"]
|
||||
_, _, oh, ow = orig_shape
|
||||
seq_len = q.shape[1]
|
||||
mask_h = oh / sqrt(oh * ow / seq_len)
|
||||
mask_h = int(mask_h) + int((seq_len % int(mask_h)) != 0)
|
||||
mask_w = seq_len // mask_h
|
||||
r = optimized_attention(q, k, v, extra_options["n_heads"])
|
||||
mask = F.interpolate(mask.unsqueeze(1), size=(mask_h, mask_w), mode="nearest").squeeze(1)
|
||||
mask = mask.view(mask.shape[0], -1, 1).repeat(1, 1, r.shape[2])
|
||||
|
||||
return mask * r
|
||||
|
||||
|
||||
def create_attention_hook(mask):
|
||||
attn_replacements = {}
|
||||
mask = mask.detach().to(device="cuda", dtype=torch.float16)
|
||||
|
||||
masked_attention = MaskedAttn2(mask)
|
||||
|
||||
for id in [4, 5, 7, 8]: # id of input_blocks that have cross attention
|
||||
block_indices = range(2) if id in [4, 5] else range(10) # transformer_depth
|
||||
for index in block_indices:
|
||||
k = ("input", id, index)
|
||||
attn_replacements[k] = masked_attention
|
||||
for id in range(6): # id of output_blocks that have cross attention
|
||||
block_indices = range(2) if id in [3, 4, 5] else range(10) # transformer_depth
|
||||
for index in block_indices:
|
||||
k = ("output", id, index)
|
||||
attn_replacements[k] = masked_attention
|
||||
for index in range(10):
|
||||
k = ("middle", 1, index)
|
||||
attn_replacements[k] = masked_attention
|
||||
|
||||
hook = TransformerOptionsHook(
|
||||
transformers_dict={"patches_replace": {"attn2": attn_replacements}}, hook_scope=EnumHookScope.HookedOnly
|
||||
)
|
||||
group = HookGroup()
|
||||
group.add(hook)
|
||||
|
||||
return group
|
||||
|
||||
|
||||
class AttentionMaskHookExperimental:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"mask": ("MASK",)},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("HOOKS",)
|
||||
CATEGORY = "promptcontrol/_testing"
|
||||
FUNCTION = "apply"
|
||||
EXPERIMENTAL = True
|
||||
DESCRIPTION = "Experimental attention masking hook. For testing only"
|
||||
|
||||
def apply(self, mask):
|
||||
return (create_attention_hook(mask),)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"AttentionMaskHookExperimental": AttentionMaskHookExperimental}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
@@ -1,51 +1,60 @@
|
||||
import logging
|
||||
|
||||
from comfy_api.latest import io
|
||||
|
||||
from .prompts import encode_prompt
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
class PCTextEncodeWithRange:
|
||||
class PCTextEncodeWithRange(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP",), "text": ("STRING", {"multiline": True})},
|
||||
"optional": {
|
||||
"start": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 0.0, "step": 0.01}),
|
||||
"end": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 1.0, "step": 0.01}),
|
||||
},
|
||||
}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCTextEncodeWithRange",
|
||||
display_name="PC: Text Encode with Range (no scheduling)",
|
||||
category="promptcontrol/tools",
|
||||
description="Like PCTextEncode, but if you know the range you need for a prompt, can be slightly more efficient when you have LoRAs scheduled on a CLIP model.",
|
||||
inputs=[
|
||||
io.Clip.Input("clip"),
|
||||
io.String.Input("text", multiline=True),
|
||||
io.Float.Input("start", default=0.0, min=0.0, max=1.0, step=0.01, optional=True),
|
||||
io.Float.Input("end", default=1.0, min=0.0, max=1.0, step=0.01, optional=True),
|
||||
],
|
||||
outputs=[io.Conditioning.Output()],
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "promptcontrol/tools"
|
||||
FUNCTION = "apply"
|
||||
DESCRIPTION = "Like PCTextEncode, but if you know the range you need for a prompt, can be slightly more efficient when you have LoRAs scheduled on a CLIP model"
|
||||
|
||||
def apply(self, clip, text, start=0.0, end=1.0):
|
||||
@classmethod
|
||||
def execute(cls, clip, text, start=0.0, end=1.0) -> io.NodeOutput: # ty: ignore[invalid-method-override]
|
||||
log.debug("PCTextEncode: Encoding '%s'", text)
|
||||
defaults = clip.patcher.model_options.get("x-promptcontrol.defaults", {})
|
||||
masks = clip.patcher.model_options.get("x-promptcontrol.masks", None)
|
||||
return (encode_prompt(clip, text, start, end, defaults, masks),)
|
||||
out = encode_prompt(clip, text, start, end, defaults, masks)
|
||||
return io.NodeOutput(out)
|
||||
|
||||
|
||||
class PCTextEncode:
|
||||
class PCTextEncode(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP",), "text": ("STRING", {"multiline": True})},
|
||||
}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCTextEncode",
|
||||
display_name="PC: Text Encode (no scheduling)",
|
||||
category="promptcontrol",
|
||||
description="Encodes a prompt with extra goodies from Prompt Control. This node does *not* support scheduling.",
|
||||
inputs=[
|
||||
io.Clip.Input("clip"),
|
||||
io.String.Input("text", multiline=True),
|
||||
],
|
||||
outputs=[io.Conditioning.Output()],
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
DESCRIPTION = "Encodes a prompt with extra goodies from Prompt Control. This node does *not* support scheduling"
|
||||
|
||||
def apply(self, clip, text):
|
||||
return PCTextEncodeWithRange.apply(self, clip, text, 0.0, 1.0)
|
||||
@classmethod
|
||||
def execute(cls, clip, text) -> io.NodeOutput: # ty: ignore[invalid-method-override]
|
||||
# Use the WithRange node for the range 0.0, 1.0
|
||||
return PCTextEncodeWithRange.execute(clip, text, 0.0, 1.0)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"PCTextEncode": PCTextEncode, "PCTextEncodeWithRange": PCTextEncodeWithRange}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"PCTextEncode": "PC: Text Encode (no scheduling)",
|
||||
"PCTextEncodeWithRange": "PC: Text Encode with Range (no scheduling)",
|
||||
}
|
||||
NODES = [
|
||||
PCTextEncodeWithRange,
|
||||
PCTextEncode,
|
||||
]
|
||||
|
||||
@@ -1,31 +1,39 @@
|
||||
import logging
|
||||
import comfy.utils
|
||||
|
||||
import comfy.hooks
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
from .utils import consolidate_schedule
|
||||
from comfy_api.latest import io
|
||||
from typing_extensions import override
|
||||
|
||||
from .attention_couple_ppm import AttentionCoupleHook
|
||||
from .parser import parse_prompt_schedules
|
||||
from .utils import consolidate_schedule
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
class PCLoraHooksFromText:
|
||||
class PCLoraHooksFromText(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"text": ("STRING",)},
|
||||
}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCLoraHooksFromText",
|
||||
display_name="PC: LoRA Hooks From Text (non-lazy)",
|
||||
category="promptcontrol/v2",
|
||||
description="set of hooks created from the prompt schedule",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.String.Input("text", multiline=True),
|
||||
],
|
||||
outputs=[io.Hooks.Output()],
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("HOOKS",)
|
||||
OUTPUT_TOOLTIPS = ("set of hooks created from the prompt schedule",)
|
||||
CATEGORY = "promptcontrol/v2"
|
||||
FUNCTION = "apply"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def apply(self, text):
|
||||
@classmethod
|
||||
def execute(cls, text) -> io.NodeOutput: # ty: ignore[invalid-method-override]
|
||||
prompt_schedule = parse_prompt_schedules(text)
|
||||
consolidated = consolidate_schedule(prompt_schedule)
|
||||
hooks = lora_hooks_from_schedule(consolidated, {})
|
||||
return (hooks,)
|
||||
return io.NodeOutput(hooks)
|
||||
|
||||
|
||||
def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
@@ -33,8 +41,7 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
lora_cache = {}
|
||||
all_hooks = []
|
||||
|
||||
def create_hook(loraspec, start_pct, end_pct, non_scheduled):
|
||||
nonlocal lora_cache
|
||||
def create_hook(loras, start_pct, end_pct, non_scheduled):
|
||||
hooks = []
|
||||
hook_kf = comfy.hooks.HookKeyframeGroup()
|
||||
for path, info in loras.items():
|
||||
@@ -48,8 +55,8 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
new_hook = comfy.hooks.create_hook_lora(
|
||||
lora_cache[path], strength_model=info["weight"], strength_clip=info["weight_clip"]
|
||||
)
|
||||
# Set hook_ref so that identical hooks compare equal
|
||||
new_hook.hooks[0].hook_ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
|
||||
ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
|
||||
new_hook.hooks[0].hook_ref = ref
|
||||
hooks.append(new_hook)
|
||||
if start_pct > 0.0:
|
||||
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=0.0)
|
||||
@@ -70,19 +77,55 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
all_hooks.append(hook)
|
||||
start_pct = end_pct
|
||||
|
||||
del lora_cache
|
||||
|
||||
all_hooks = [x for x in all_hooks if x]
|
||||
|
||||
if all_hooks:
|
||||
hooks = comfy.hooks.HookGroup.combine_all_hooks(all_hooks)
|
||||
return hooks
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"PCLoraHooksFromText": PCLoraHooksFromText,
|
||||
}
|
||||
class PCAttentionCoupleBatchNegative(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCAttentionCoupleBatchNegative",
|
||||
display_name="PC: Attention Couple (batch negative)",
|
||||
category="promptcontrol/v2",
|
||||
description="Batch negatives, carrying over Attention Couple hooks",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.Conditioning.Input("positive"),
|
||||
io.Conditioning.Input("negative"),
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output("positive"),
|
||||
io.Conditioning.Output("negative"),
|
||||
],
|
||||
)
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"PCLoraHooksFromText": "PC: LoRA Hooks From Text (non-lazy)",
|
||||
}
|
||||
@classmethod
|
||||
@override
|
||||
def execute(cls, positive, negative) -> io.NodeOutput: # ty: ignore[invalid-method-override]
|
||||
if len(negative) != 1:
|
||||
log.warning("Batching scheduled negatives is not supported yet")
|
||||
return io.NodeOutput(positive, negative)
|
||||
|
||||
negative_batch = []
|
||||
for p in positive:
|
||||
n = [negative[0][0], negative[0][1].copy()]
|
||||
n_hook_group = n[1].get("hooks", comfy.hooks.HookGroup()).clone()
|
||||
p_hook_group = p[1].get("hooks", comfy.hooks.HookGroup())
|
||||
attn_couple = [hook for hook in p_hook_group.hooks if isinstance(hook, AttentionCoupleHook)]
|
||||
for hook in attn_couple:
|
||||
n_hook_group.add(hook)
|
||||
n[1]["hooks"] = p_hook_group if n_hook_group.hooks == p_hook_group.hooks else n_hook_group
|
||||
n[1]["start_percent"] = p[1].get("start_percent", 0.0)
|
||||
n[1]["end_percent"] = p[1].get("end_percent", 1.0)
|
||||
negative_batch.append(n)
|
||||
|
||||
return io.NodeOutput(positive, negative_batch)
|
||||
|
||||
|
||||
NODES = [
|
||||
PCLoraHooksFromText,
|
||||
PCAttentionCoupleBatchNegative,
|
||||
]
|
||||
|
||||
+119
-139
@@ -1,30 +1,18 @@
|
||||
import logging
|
||||
from .parser import parse_prompt_schedules
|
||||
from comfy_execution.graph_utils import GraphBuilder, is_link
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
from __future__ import annotations
|
||||
|
||||
from .prompts import get_function
|
||||
import json
|
||||
import logging
|
||||
|
||||
from comfy_api.latest import io
|
||||
from comfy_execution.graph import ExecutionBlocker
|
||||
from comfy_execution.graph_utils import GraphBuilder
|
||||
|
||||
from .parser import parse_prompt_schedules
|
||||
from .utils import consolidate_schedule, find_nonscheduled_loras, get_function
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
from .utils import consolidate_schedule, find_nonscheduled_loras
|
||||
import json
|
||||
|
||||
|
||||
def _cache_key(cachekey, inputs):
|
||||
out = inputs.copy()
|
||||
text = inputs.get("text")
|
||||
if text is not None and not is_link(text):
|
||||
out["text"] = cache_key_from_inputs(cachekey, **inputs)
|
||||
return out
|
||||
|
||||
|
||||
def cache_key_prompt(inputs):
|
||||
return _cache_key("prompt", inputs)
|
||||
|
||||
|
||||
def cache_key_lora(inputs):
|
||||
return _cache_key("loras", inputs)
|
||||
|
||||
|
||||
def create_lora_loader_nodes(graph, model, clip, loras):
|
||||
for path, info in loras.items():
|
||||
@@ -85,11 +73,15 @@ def create_hook_nodes_for_lora(graph, path, info, existing_node, start_pct, end_
|
||||
return hook_node, next_keyframe
|
||||
|
||||
|
||||
def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True, return_hooks=True):
|
||||
def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True):
|
||||
# This gets rid of non-existent LoRAs
|
||||
consolidated = consolidate_schedule(schedule)
|
||||
non_scheduled = find_nonscheduled_loras(consolidated)
|
||||
model, clip = create_lora_loader_nodes(graph, model, clip, non_scheduled)
|
||||
if model is not None:
|
||||
non_scheduled = find_nonscheduled_loras(consolidated)
|
||||
model, clip = create_lora_loader_nodes(graph, model, clip, non_scheduled)
|
||||
else:
|
||||
non_scheduled = {}
|
||||
model = ExecutionBlocker("No model provided to PCLazyLoRALoader or PCLazyLoRALoaderAdvanced")
|
||||
|
||||
hook_nodes = {}
|
||||
start_pct = 0.0
|
||||
@@ -124,7 +116,7 @@ def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True, return_h
|
||||
n.set_input("hooks_B", h.out(0))
|
||||
res = n
|
||||
res = res.out(0)
|
||||
if apply_hooks:
|
||||
if clip is not None and apply_hooks:
|
||||
n = graph.node("SetClipHooks")
|
||||
n.set_input("clip", clip)
|
||||
n.set_input("hooks", res)
|
||||
@@ -132,74 +124,70 @@ def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True, return_h
|
||||
n.set_input("schedule_clip", True)
|
||||
clip = n.out(0)
|
||||
|
||||
if clip is None:
|
||||
clip = ExecutionBlocker("No clip model provided to PCLazyLoRALoader or PCLazyLoRALoaderAdvanced")
|
||||
r = graph.finalize()
|
||||
log.debug("LazyLoraLoader built graph: %s", json.dumps(r))
|
||||
|
||||
if return_hooks:
|
||||
ret = (model, clip, res)
|
||||
else:
|
||||
ret = (model, clip)
|
||||
ret = (model, clip, res)
|
||||
|
||||
return {"result": ret, "expand": r}
|
||||
return io.NodeOutput(*ret, expand=r)
|
||||
|
||||
|
||||
class PCLazyLoraLoaderAdvanced:
|
||||
CACHE_KEY = cache_key_lora
|
||||
class PCLazyLoraLoaderAdvanced(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCLazyLoraLoaderAdvanced",
|
||||
display_name="PC: Schedule LoRAs (Advanced)",
|
||||
enable_expand=True,
|
||||
category="promptcontrol",
|
||||
description="Returns a model and clip with LoRAs scheduled",
|
||||
inputs=[
|
||||
io.Model.Input("model", extra_dict={"rawLink": True}, optional=True),
|
||||
io.Clip.Input("clip", extra_dict={"rawLink": True}, optional=True),
|
||||
io.String.Input("text", multiline=True, default=""),
|
||||
io.Boolean.Input("apply_hooks", default=True),
|
||||
io.String.Input("tags", default=""),
|
||||
io.Float.Input("start", min=0.0, max=1.0, default=0.0, step=0.01),
|
||||
io.Float.Input("end", min=0.0, max=1.0, default=1.0, step=0.01),
|
||||
io.Int.Input("num_steps", min=0, max=10000, default=0, step=1),
|
||||
],
|
||||
outputs=[io.Model.Output("model"), io.Clip.Output("clip"), io.Hooks.Output("hooks")],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
"model": ("MODEL", {"rawLink": True}),
|
||||
"clip": ("CLIP", {"rawLink": True}),
|
||||
},
|
||||
"optional": {
|
||||
"apply_hooks": ("BOOLEAN", {"default": True}),
|
||||
"tags": ("STRING", {"default": ""}),
|
||||
"start": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 0.0, "step": 0.01}),
|
||||
"end": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 1.0, "step": 0.01}),
|
||||
},
|
||||
"hidden": {"unique_id": "UNIQUE_ID"},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CLIP", "HOOKS")
|
||||
OUTPUT_TOOLTIPS = ("Returns a model and clip with LoRAs scheduled",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, model, clip, text, unique_id, apply_hooks=True, tags="", start=0.0, end=1.0):
|
||||
schedule = parse_prompt_schedules(text, filters=tags, start=start, end=end)
|
||||
graph = GraphBuilder(f"PCLazyLoraLoaderAdvanced-{unique_id}")
|
||||
return build_lora_schedule(graph, schedule, model, clip, apply_hooks=apply_hooks, return_hooks=True)
|
||||
def execute(cls, model=None, clip=None, text="", apply_hooks=True, tags="", start=0.0, end=1.0, num_steps=0):
|
||||
schedule = parse_prompt_schedules(text, filters=tags, start=start, end=end, num_steps=num_steps)
|
||||
graph = GraphBuilder()
|
||||
r = build_lora_schedule(graph, schedule, model, clip, apply_hooks=apply_hooks)
|
||||
return r
|
||||
|
||||
|
||||
class PCLazyLoraLoader:
|
||||
CACHE_KEY = cache_key_lora
|
||||
class PCLazyLoraLoader(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCLazyLoraLoader",
|
||||
display_name="PC: Schedule LoRAs",
|
||||
enable_expand=True,
|
||||
category="promptcontrol",
|
||||
description="Returns a model and clip with LoRAs scheduled",
|
||||
inputs=[
|
||||
io.Model.Input("model", extra_dict={"rawLink": True}, optional=True),
|
||||
io.Clip.Input("clip", extra_dict={"rawLink": True}, optional=True),
|
||||
io.String.Input("text", multiline=True, default=""),
|
||||
],
|
||||
outputs=[
|
||||
io.Model.Output("model"),
|
||||
io.Clip.Output("clip"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL", {"rawLink": True}),
|
||||
"clip": ("CLIP", {"rawLink": True}),
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
},
|
||||
"hidden": {"unique_id": "UNIQUE_ID"},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (
|
||||
"MODEL",
|
||||
"CLIP",
|
||||
)
|
||||
OUTPUT_TOOLTIPS = ("Returns a model and clip with LoRAs scheduled",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, model, clip, text, unique_id):
|
||||
graph = GraphBuilder(f"PCLazyLoraLoader-{unique_id}")
|
||||
schedule = parse_prompt_schedules(text)
|
||||
return build_lora_schedule(graph, schedule, model, clip, apply_hooks=True, return_hooks=False)
|
||||
def execute(cls, model, clip, text):
|
||||
no = PCLazyLoraLoaderAdvanced.execute(model, clip, text)
|
||||
return io.NodeOutput(*no.args[:2], expand=no.expand)
|
||||
|
||||
|
||||
def build_scheduled_prompts(graph, schedules, clip):
|
||||
@@ -211,8 +199,7 @@ def build_scheduled_prompts(graph, schedules, clip):
|
||||
classname = "PCTextEncode"
|
||||
paramname = "text"
|
||||
if classnames:
|
||||
classname = classnames[0][0]
|
||||
paramname = classnames[0][1]
|
||||
classname, paramname = classnames[0].args
|
||||
node = graph.node(classname)
|
||||
node.set_input("clip", clip)
|
||||
node.set_input(paramname, p)
|
||||
@@ -232,69 +219,62 @@ def build_scheduled_prompts(graph, schedules, clip):
|
||||
g = graph.finalize()
|
||||
log.debug("Built graph: %s", json.dumps(g))
|
||||
|
||||
return {"result": (node.out(0),), "expand": g}
|
||||
return io.NodeOutput(node.out(0), expand=g)
|
||||
|
||||
|
||||
def cache_key_from_inputs(cachekey, text, tags="", start=0.0, end=1.0, **kwargs):
|
||||
schedules = parse_prompt_schedules(text, filters=tags, start=start, end=end)
|
||||
return [(pct, s[cachekey]) for pct, s in schedules]
|
||||
|
||||
|
||||
class PCLazyTextEncode:
|
||||
CACHE_KEY = cache_key_prompt
|
||||
class PCLazyTextEncodeAdvanced(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCLazyTextEncodeAdvanced",
|
||||
display_name="PC: Schedule prompt (Advanced)",
|
||||
enable_expand=True,
|
||||
category="promptcontrol",
|
||||
inputs=[
|
||||
io.Clip.Input("clip", extra_dict={"rawLink": True}),
|
||||
io.String.Input("text", multiline=True, default=""),
|
||||
io.String.Input("tags", default=""),
|
||||
io.Float.Input("start", min=0.0, max=1.0, default=0.0, step=0.01),
|
||||
io.Float.Input("end", min=0.0, max=1.0, default=1.0, step=0.01),
|
||||
io.Int.Input("num_steps", min=0, max=10000, default=0, step=1),
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output("conditioning"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP", {"rawLink": True}), "text": ("STRING", {"multiline": True})},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
OUTPUT_TOOLTIPS = ("A fully encoded and scheduled conditioning",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, clip, text):
|
||||
schedules = parse_prompt_schedules(text)
|
||||
def execute(cls, clip, text, tags="", start=0.0, end=1.0, num_steps=0):
|
||||
schedules = parse_prompt_schedules(text, filters=tags, start=start, end=end, num_steps=num_steps)
|
||||
graph = GraphBuilder()
|
||||
return build_scheduled_prompts(graph, schedules, clip)
|
||||
|
||||
|
||||
class PCLazyTextEncodeAdvanced:
|
||||
CACHE_KEY = cache_key_prompt
|
||||
class PCLazyTextEncode(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCLazyTextEncode",
|
||||
display_name="PC: Schedule prompt",
|
||||
enable_expand=True,
|
||||
category="promptcontrol",
|
||||
inputs=[
|
||||
io.Clip.Input("clip", extra_dict={"rawLink": True}),
|
||||
io.String.Input("text", multiline=True, default=""),
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output("conditioning"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP", {"rawLink": True}), "text": ("STRING", {"multiline": True})},
|
||||
"optional": {
|
||||
"tags": ("STRING", {"default": ""}),
|
||||
"start": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 0.0, "step": 0.01}),
|
||||
"end": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 1.0, "step": 0.01}),
|
||||
},
|
||||
"hidden": {"unique_id": "UNIQUE_ID"},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, clip, text, unique_id, tags="", start=0.1, end=1.0):
|
||||
schedules = parse_prompt_schedules(text, filters=tags, start=start, end=end)
|
||||
graph = GraphBuilder(f"PCLazyTextEncodeAdvanced-{unique_id}")
|
||||
return build_scheduled_prompts(graph, schedules, clip)
|
||||
def execute(cls, clip, text):
|
||||
return PCLazyTextEncodeAdvanced.execute(clip, text)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"PCLazyTextEncode": PCLazyTextEncode,
|
||||
"PCLazyTextEncodeAdvanced": PCLazyTextEncodeAdvanced,
|
||||
"PCLazyLoraLoader": PCLazyLoraLoader,
|
||||
"PCLazyLoraLoaderAdvanced": PCLazyLoraLoaderAdvanced,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"PCLazyTextEncode": "PC: Schedule Prompt",
|
||||
"PCLazyTextEncodeAdvanced": "PC: Schedule prompt (Advanced)",
|
||||
"PCLazyLoraLoader": "PC: Schedule LoRAs",
|
||||
"PCLazyLoraLoaderAdvanced": "PC: Schedule LoRAs (Advanced)",
|
||||
}
|
||||
NODES = [
|
||||
PCLazyTextEncode,
|
||||
PCLazyTextEncodeAdvanced,
|
||||
PCLazyLoraLoader,
|
||||
PCLazyLoraLoaderAdvanced,
|
||||
]
|
||||
|
||||
+132
-90
@@ -1,103 +1,108 @@
|
||||
import logging
|
||||
|
||||
from comfy_api.latest import io
|
||||
|
||||
from .macros import expand_macros
|
||||
from .parser import parse_prompt_schedules
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
class PCSetLogLevel:
|
||||
class PCSetLogLevel(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"clip": ("CLIP",),
|
||||
},
|
||||
"optional": {
|
||||
"level": (["INFO", "DEBUG", "WARNING", "ERROR"], {"default": "INFO"}),
|
||||
},
|
||||
}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCSetLogLevel",
|
||||
display_name="PC: Configure Logging (for debug)",
|
||||
category="promptcontrol/tools",
|
||||
description="A debug node to configure Prompt Control logging level. Pass a CLIP through it before you run any PC nodes",
|
||||
inputs=[
|
||||
io.Clip.Input("clip"),
|
||||
io.Combo.Input("level", options=["INFO", "DEBUG", "WARNING", "ERROR"], default="INFO", optional=True),
|
||||
],
|
||||
outputs=[io.Clip.Output()],
|
||||
)
|
||||
|
||||
def apply(self, clip, level="INFO"):
|
||||
@classmethod
|
||||
def execute(cls, clip, level="INFO") -> io.NodeOutput:
|
||||
log.setLevel(getattr(logging, level))
|
||||
log.info("Set logging level to %s", level)
|
||||
return (clip,)
|
||||
|
||||
RETURN_TYPES = ("CLIP",)
|
||||
CATEGORY = "promptcontrol/tools"
|
||||
DESCRIPTION = (
|
||||
"A debug node to configure Prompt Control logging level. Pass a CLIP through it before you run any PC nodes"
|
||||
)
|
||||
|
||||
FUNCTION = "apply"
|
||||
return io.NodeOutput(clip)
|
||||
|
||||
|
||||
class PCAddMaskToCLIP:
|
||||
class PCAddMaskToCLIP(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP",)},
|
||||
"optional": {
|
||||
"mask": ("MASK",),
|
||||
},
|
||||
}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCAddMaskToCLIP",
|
||||
display_name="PC: Attach Mask",
|
||||
category="promptcontrol/tools",
|
||||
description="Attaches a mask to a CLIP object so that they can be referred to in a prompt using IMASK(). Using this node multiple times adds more masks rather than replacing existing ones.",
|
||||
inputs=[
|
||||
io.Clip.Input("clip"),
|
||||
io.Mask.Input("mask", optional=True),
|
||||
],
|
||||
outputs=[io.Clip.Output()],
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("CLIP",)
|
||||
CATEGORY = "promptcontrol/tools"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, clip, mask=None):
|
||||
return PCAddMaskToCLIPMany().apply(clip, mask1=mask)
|
||||
|
||||
|
||||
class PCAddMaskToCLIPMany:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP",)},
|
||||
"optional": {
|
||||
"mask1": ("MASK",),
|
||||
"mask2": ("MASK",),
|
||||
"mask3": ("MASK",),
|
||||
"mask4": ("MASK",),
|
||||
},
|
||||
}
|
||||
def execute(cls, clip, mask=None) -> io.NodeOutput:
|
||||
return PCAddMaskToCLIPMany.execute(clip, mask1=mask)
|
||||
|
||||
RETURN_TYPES = ("CLIP",)
|
||||
CATEGORY = "promptcontrol/tools"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, clip, mask1=None, mask2=None, mask3=None, mask4=None):
|
||||
class PCAddMaskToCLIPMany(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCAddMaskToCLIPMany",
|
||||
display_name="PC: Attach Mask (multi)",
|
||||
category="promptcontrol/tools",
|
||||
description="Multi-input version of PCAddMaskToCLIP, for convenience",
|
||||
inputs=[
|
||||
io.Clip.Input("clip"),
|
||||
io.Mask.Input("mask1", optional=True),
|
||||
io.Mask.Input("mask2", optional=True),
|
||||
io.Mask.Input("mask3", optional=True),
|
||||
io.Mask.Input("mask4", optional=True),
|
||||
],
|
||||
outputs=[io.Clip.Output()],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, clip, mask1=None, mask2=None, mask3=None, mask4=None) -> io.NodeOutput:
|
||||
clip = clip.clone()
|
||||
current_masks = clip.patcher.model_options.get("x-promptcontrol.masks", [])
|
||||
current_masks.extend(m for m in (mask1, mask2, mask3, mask4) if m is not None)
|
||||
clip.patcher.model_options["x-promptcontrol.masks"] = current_masks
|
||||
return (clip,)
|
||||
return io.NodeOutput(clip)
|
||||
|
||||
|
||||
class PCSetPCTextEncodeSettings:
|
||||
class PCSetPCTextEncodeSettings(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP",)},
|
||||
"optional": {
|
||||
"steps": ("INT", {"default": 0, "min": 0, "max": 10000}),
|
||||
"mask_width": ("INT", {"default": 512, "min": 64, "max": 4096 * 4}),
|
||||
"mask_height": ("INT", {"default": 512, "min": 64, "max": 4096 * 4}),
|
||||
"sdxl_width": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
|
||||
"sdxl_height": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
|
||||
"sdxl_target_w": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
|
||||
"sdxl_target_h": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
|
||||
"sdxl_crop_w": ("INT", {"default": 0, "min": 0, "max": 4096 * 4}),
|
||||
"sdxl_crop_h": ("INT", {"default": 0, "min": 0, "max": 4096 * 4}),
|
||||
},
|
||||
}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCSetPCTextEncodeSettings",
|
||||
display_name="PC: Configure PCTextEncode",
|
||||
category="promptcontrol/tools",
|
||||
description="Configures default values for PCTextEncode",
|
||||
inputs=[
|
||||
io.Clip.Input("clip"),
|
||||
io.Int.Input("mask_width", default=512, min=64, max=4096 * 4, optional=True),
|
||||
io.Int.Input("mask_height", default=512, min=64, max=4096 * 4, optional=True),
|
||||
io.Int.Input("sdxl_width", default=1024, min=0, max=4096 * 4, optional=True),
|
||||
io.Int.Input("sdxl_height", default=1024, min=0, max=4096 * 4, optional=True),
|
||||
io.Int.Input("sdxl_target_w", default=1024, min=0, max=4096 * 4, optional=True),
|
||||
io.Int.Input("sdxl_target_h", default=1024, min=0, max=4096 * 4, optional=True),
|
||||
io.Int.Input("sdxl_crop_w", default=0, min=0, max=4096 * 4, optional=True),
|
||||
io.Int.Input("sdxl_crop_h", default=0, min=0, max=4096 * 4, optional=True),
|
||||
],
|
||||
outputs=[io.Clip.Output()],
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("CLIP",)
|
||||
CATEGORY = "promptcontrol/tools"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(
|
||||
self,
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
clip,
|
||||
steps=0,
|
||||
mask_width=512,
|
||||
mask_height=512,
|
||||
sdxl_width=1024,
|
||||
@@ -106,9 +111,8 @@ class PCSetPCTextEncodeSettings:
|
||||
sdxl_target_h=1024,
|
||||
sdxl_crop_w=0,
|
||||
sdxl_crop_h=0,
|
||||
):
|
||||
) -> io.NodeOutput:
|
||||
settings = {
|
||||
"steps": steps,
|
||||
"mask_width": mask_width,
|
||||
"mask_height": mask_height,
|
||||
"sdxl_width": sdxl_width,
|
||||
@@ -120,19 +124,57 @@ class PCSetPCTextEncodeSettings:
|
||||
}
|
||||
clip = clip.clone()
|
||||
clip.patcher.model_options["x-promptcontrol.settings"] = settings
|
||||
return (clip,)
|
||||
return io.NodeOutput(clip)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"PCSetPCTextEncodeSettings": PCSetPCTextEncodeSettings,
|
||||
"PCAddMaskToCLIP": PCAddMaskToCLIP,
|
||||
"PCAddMaskToCLIPMany": PCAddMaskToCLIPMany,
|
||||
"PCSetLogLevel": PCSetLogLevel,
|
||||
}
|
||||
class PCExtractScheduledPrompt(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCExtractScheduledPrompt",
|
||||
display_name="PC: Extract Scheduled Prompt",
|
||||
category="promptcontrol/tools",
|
||||
description="Parses the input prompt and returns the prompt scheduled at the specified point",
|
||||
inputs=[
|
||||
io.String.Input("text", multiline=True),
|
||||
io.Float.Input("at", min=0.0, max=1.0, default=1.0, step=0.01),
|
||||
io.String.Input("tags", default="", optional=True),
|
||||
],
|
||||
outputs=[io.String.Output()],
|
||||
)
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"PCSetPCTextEncodeSettings": "PC: Configure PCTextEncode",
|
||||
"PCAddMaskToCLIP": "PC: Attach Mask",
|
||||
"PCAddMaskToCLIPMany": "PC: Attach Mask (multi)",
|
||||
"PCSetLogLevel": "PC: Configure Logging (for debug)",
|
||||
}
|
||||
@classmethod
|
||||
def execute(cls, text, at, tags="") -> io.NodeOutput:
|
||||
schedule = parse_prompt_schedules(text, filters=tags)
|
||||
_, entry = schedule.at_step(at, total_steps=1)
|
||||
prompt_text = entry.get("prompt", "")
|
||||
return io.NodeOutput(prompt_text)
|
||||
|
||||
|
||||
class PCMacroExpand(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="PCMacroExpand",
|
||||
display_name="PC: Expand Macros",
|
||||
category="promptcontrol/tools",
|
||||
description="Expands DEF macros in a string and returns the result",
|
||||
inputs=[
|
||||
io.String.Input("text", multiline=True),
|
||||
],
|
||||
outputs=[io.String.Output()],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, text) -> io.NodeOutput:
|
||||
return io.NodeOutput(expand_macros(text))
|
||||
|
||||
|
||||
NODES = [
|
||||
PCSetPCTextEncodeSettings,
|
||||
PCAddMaskToCLIP,
|
||||
PCAddMaskToCLIPMany,
|
||||
PCSetLogLevel,
|
||||
PCExtractScheduledPrompt,
|
||||
PCMacroExpand,
|
||||
]
|
||||
|
||||
+6
-365
@@ -1,369 +1,10 @@
|
||||
import lark
|
||||
import logging
|
||||
from math import ceil
|
||||
import os
|
||||
|
||||
logging.basicConfig()
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
import re
|
||||
|
||||
from functools import lru_cache
|
||||
from .utils import get_function
|
||||
|
||||
if lark.__version__ == "0.12.0":
|
||||
from sys import executable
|
||||
|
||||
x = "\n".join(
|
||||
[
|
||||
"Your lark package reports an ancient version (0.12.0) and will not work. If you have the 'lark-parser' package in your Python environment, remove that and *reinstall* lark!",
|
||||
f"{executable} -m pip uninstall lark-parser lark",
|
||||
f"{executable} -m pip install lark",
|
||||
]
|
||||
)
|
||||
log.error(x)
|
||||
raise ImportError(x)
|
||||
|
||||
|
||||
prompt_parser = lark.Lark(
|
||||
r"""
|
||||
!start: (prompt | /[][():|]/+)*
|
||||
prompt: (emphasized | embedding | scheduled | alternate | sequence | loraspec | PLAIN | /</ | />/ | WHITESPACE)+
|
||||
!emphasized: "(" prompt? ")"
|
||||
| "(" prompt ":" prompt ")"
|
||||
| "[" prompt "]"
|
||||
scheduled: "[" [prompt ":"] [prompt] ":" _WS? NUMBER ["," NUMBER] "]"
|
||||
| "[" [prompt ":"] [prompt] ":" _WS? TAG "]"
|
||||
sequence: "[SEQ" ":" [prompt] ":" NUMBER (":" [prompt] ":" NUMBER)+ "]"
|
||||
alternate: "[" [prompt] ("|" [prompt])+ [":" NUMBER] "]"
|
||||
loraspec.99: "<lora:" FILENAME lora_weights [lora_block_weights] ">"
|
||||
lora_weights.1: (":" _WS? NUMBER)~1..2
|
||||
lora_block_weights.-1: ":" PLAIN
|
||||
embedding.100: "<emb:" FILENAME ">"
|
||||
WHITESPACE: /\s+/
|
||||
_WS: WHITESPACE
|
||||
PLAIN: /([^<>\\\[\]():|]|\\.)+/
|
||||
FILENAME: /[^<>:]+/
|
||||
TAG: /[A-Z_]+/
|
||||
%import common.SIGNED_NUMBER -> NUMBER
|
||||
""",
|
||||
lexer="dynamic",
|
||||
)
|
||||
|
||||
|
||||
cut_parser = lark.Lark(
|
||||
r"""
|
||||
!start: (prompt | /[][:()]/+)*
|
||||
prompt: (cut | PLAIN | WHITESPACE)+
|
||||
cut: "[CUT:" prompt ":" prompt [":" NUMBER [ ":" NUMBER [":" NUMBER [ ":" PLAIN ] ] ] ]"]"
|
||||
WHITESPACE: /\s+/
|
||||
PLAIN: /([^\[\]:])+/
|
||||
%import common.SIGNED_NUMBER -> NUMBER
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
class CutTransform(lark.Transformer):
|
||||
def __default__(self, data, children, meta):
|
||||
return children
|
||||
|
||||
def cut(self, args):
|
||||
prompt, cutout, weight, strict_mask, start_from_masked, mask_token = args
|
||||
|
||||
return ("".join(flatten(prompt)), "".join(flatten(cutout)), weight, strict_mask, start_from_masked, mask_token)
|
||||
|
||||
def start(self, args):
|
||||
prompt = []
|
||||
cuts = []
|
||||
for a in flatten(args):
|
||||
if isinstance(a, str):
|
||||
prompt.append(a)
|
||||
else:
|
||||
prompt.append(a[0])
|
||||
cuts.append(a)
|
||||
return "".join(prompt), cuts
|
||||
|
||||
def PLAIN(self, args):
|
||||
return args
|
||||
|
||||
|
||||
def parse_cuts(text):
|
||||
return CutTransform().transform(cut_parser.parse(text))
|
||||
|
||||
|
||||
def flatten(x):
|
||||
if type(x) in [str, tuple] or isinstance(x, dict) and "type" in x:
|
||||
yield x
|
||||
else:
|
||||
for g in x:
|
||||
yield from flatten(g)
|
||||
|
||||
|
||||
def clamp(a, b, c):
|
||||
"""clamp b between a and c"""
|
||||
return min(max(a, b), c)
|
||||
|
||||
|
||||
def get_steps(tree):
|
||||
res = [100]
|
||||
|
||||
def tostep(s):
|
||||
w = float(s) * 100
|
||||
w = int(clamp(0, w, 100))
|
||||
return w
|
||||
|
||||
class CollectSteps(lark.Visitor):
|
||||
def scheduled(self, tree):
|
||||
i = tree.children[-1]
|
||||
if i and i.type == "TAG":
|
||||
return
|
||||
for i in [-1, -2]:
|
||||
if tree.children[i] is not None:
|
||||
tree.children[i] = tostep(tree.children[i])
|
||||
res.append(tree.children[i])
|
||||
|
||||
def interp_steps(self, tree):
|
||||
tree.children[-1] = tostep(tree.children[-1] or 0.1)
|
||||
for i, _ in enumerate(tree.children[:-1]):
|
||||
tree.children[i] = tostep(tree.children[i])
|
||||
|
||||
res.extend(tree.children[:-1])
|
||||
|
||||
def sequence(self, tree):
|
||||
steps = tree.children[1::2]
|
||||
for i, steps in enumerate(steps):
|
||||
w = float(tree.children[i * 2 + 1]) * 100
|
||||
tree.children[i * 2 + 1] = clamp(0, w, 100)
|
||||
res.append(w)
|
||||
|
||||
def alternate(self, tree):
|
||||
step_size = int(round(float(tree.children[-1] or 0.1), 2) * 100)
|
||||
step_size = clamp(1, step_size, 100)
|
||||
tree.children[-1] = step_size
|
||||
res.extend([x for x in range(step_size, 100, step_size)])
|
||||
|
||||
CollectSteps().visit(tree)
|
||||
|
||||
return sorted(set(res))
|
||||
|
||||
|
||||
def at_step(step, filters, tree):
|
||||
class AtStep(lark.Transformer):
|
||||
def scheduled(self, args):
|
||||
when_end = None
|
||||
before, after, when, *rest = args
|
||||
if isinstance(when, str):
|
||||
return before or "" if when not in filters else after or ""
|
||||
|
||||
if rest:
|
||||
when_end = rest[0]
|
||||
|
||||
if when_end is not None and step <= when and before is not None:
|
||||
return ""
|
||||
|
||||
if when_end is not None and (step > when and step <= when_end):
|
||||
# handle [a:0,1]
|
||||
if before is None:
|
||||
return after or ""
|
||||
return before or ""
|
||||
|
||||
if when_end is not None and step >= when_end:
|
||||
# handle [a:0,1]
|
||||
if before is None:
|
||||
return ""
|
||||
return after or ""
|
||||
|
||||
if step <= when:
|
||||
return before or ""
|
||||
else:
|
||||
return after or ""
|
||||
|
||||
def sequence(self, args):
|
||||
previous_step = 0.0
|
||||
prompts = args[::2]
|
||||
steps = args[1::2]
|
||||
for s, p in zip(steps, prompts):
|
||||
if s >= step and step >= previous_step:
|
||||
previous_step = step
|
||||
return p or ""
|
||||
else:
|
||||
previous_step = s
|
||||
return ""
|
||||
|
||||
def alternate(self, args):
|
||||
step_size = args[-1]
|
||||
idx = ceil(step / step_size)
|
||||
return args[(idx - 1) % (len(args) - 1)] or ""
|
||||
|
||||
def start(self, args):
|
||||
prompt = []
|
||||
loraspecs = {}
|
||||
args = flatten(args)
|
||||
for a in args:
|
||||
if isinstance(a, str):
|
||||
prompt.append(a)
|
||||
elif isinstance(a, tuple):
|
||||
# sum identical specs together
|
||||
n = a[0]
|
||||
# if clip weight is not provided, use unet weight
|
||||
w, w_clip = a[1][0], a[1][1 % len(a[1])]
|
||||
e = loraspecs.get(n, {})
|
||||
loraspecs[n] = {
|
||||
"weight": round(e.get("weight", 0.0) + w, 2),
|
||||
"weight_clip": round(e.get("weight_clip", 0.0) + w_clip, 2),
|
||||
}
|
||||
lbw = a[2]
|
||||
if lbw:
|
||||
loraspecs[n]["lbw"] = lbw
|
||||
if loraspecs[n]["weight"] == 0 and loraspecs[n]["weight_clip"] == 0 and not lbw:
|
||||
del loraspecs[n]
|
||||
else:
|
||||
pass
|
||||
p = "".join(prompt)
|
||||
return {"prompt": p, "loras": loraspecs}
|
||||
|
||||
def PLAIN(self, args):
|
||||
return args.replace("\\:", ":")
|
||||
|
||||
def FILENAME(self, value):
|
||||
return str(value)
|
||||
|
||||
def embedding(self, args):
|
||||
return "embedding:" + str(args[0])
|
||||
|
||||
def lora_weights(self, args):
|
||||
return [float(str(a)) for a in args]
|
||||
|
||||
def lora_block_weights(self, args):
|
||||
vals = args[0].split(";")
|
||||
r = {}
|
||||
for v in vals:
|
||||
x = v.split("=", 2)
|
||||
if len(x) != 2:
|
||||
continue
|
||||
k, v = x[0].strip().upper(), x[1].strip()
|
||||
r[k] = v
|
||||
return r
|
||||
|
||||
def loraspec(self, args):
|
||||
name = args[0]
|
||||
params = args[1]
|
||||
lbw = args[2]
|
||||
|
||||
return name, params, lbw
|
||||
|
||||
def __default__(self, data, children, meta):
|
||||
for child in children:
|
||||
yield child
|
||||
|
||||
return AtStep().transform(tree)
|
||||
|
||||
|
||||
class PromptSchedule(object):
|
||||
def __init__(self, prompt, filters="", start=0.0, end=1.0):
|
||||
self.filters = filters
|
||||
self.start = start
|
||||
self.end = end
|
||||
self.prompt = prompt.strip()
|
||||
self.defaults = {}
|
||||
self.loaded_loras = {}
|
||||
|
||||
self.parsed_prompt = self._parse()
|
||||
|
||||
def __iter__(self):
|
||||
# Filter out zero, it's only useful for interpolation
|
||||
return (x for x in self.parsed_prompt if x[0] != 0)
|
||||
|
||||
def _parse(self):
|
||||
filters = [x.strip() for x in self.filters.upper().split(",")]
|
||||
try:
|
||||
parsed = []
|
||||
tree = prompt_parser.parse(self.prompt)
|
||||
steps = get_steps(tree)
|
||||
|
||||
def f(x):
|
||||
return round(x / 100, 2)
|
||||
|
||||
for t in steps:
|
||||
p = at_step(t, filters, tree)
|
||||
parsed.append([f(t), p])
|
||||
|
||||
except lark.exceptions.LarkError as e:
|
||||
log.error("Prompt editing parse error: %s", e)
|
||||
parsed = [[1.0, {"prompt": self.prompt, "loras": {}}]]
|
||||
|
||||
# Tag filtering may return redundant prompts, so filter them out here
|
||||
res = []
|
||||
prev_end = -1
|
||||
|
||||
for end_at, p in parsed:
|
||||
if end_at < self.start:
|
||||
continue
|
||||
elif end_at <= self.end:
|
||||
res.append([end_at, p])
|
||||
prev_end = end_at
|
||||
elif end_at > self.end and prev_end < self.end:
|
||||
res.append([end_at, p])
|
||||
break
|
||||
|
||||
# Always use the last prompt if everything was filtered
|
||||
if len(res) == 0:
|
||||
res = [[1.0, parsed[-1][1]]]
|
||||
|
||||
return res
|
||||
|
||||
def clone(self):
|
||||
return self.with_filters()
|
||||
|
||||
def with_filters(self, filters=None, start=None, end=None, defaults=None):
|
||||
def ifspecified(x, defval):
|
||||
return x if x is not None else defval
|
||||
|
||||
p = PromptSchedule(
|
||||
self.prompt,
|
||||
filters=ifspecified(filters, self.filters),
|
||||
start=ifspecified(start, self.start),
|
||||
end=ifspecified(end, self.end),
|
||||
)
|
||||
return p
|
||||
|
||||
def at_step(self, step, total_steps=1):
|
||||
_, x = self.at_step_idx(step, total_steps)
|
||||
return x
|
||||
|
||||
def at_step_idx(self, step, total_steps=1):
|
||||
for i, x in enumerate(self.parsed_prompt):
|
||||
if x[0] * total_steps >= step:
|
||||
return i, x
|
||||
return len(self.parsed_prompt) - 1, self.parsed_prompt[-1]
|
||||
|
||||
|
||||
def replace_defs(text):
|
||||
text, defs = get_function(text, "DEF", defaults=None)
|
||||
res = text
|
||||
prevres = text
|
||||
replacements = []
|
||||
for d in defs:
|
||||
r = d.split("=", 1)
|
||||
if len(r) != 2 or not r[0].strip():
|
||||
log.warning("Ignoring invalid DEF(%s)", d)
|
||||
continue
|
||||
replacements.append(r[0].strip(), r[1].strip())
|
||||
iterations = 0
|
||||
while True:
|
||||
iterations += 1
|
||||
if iterations > 10:
|
||||
log.error("Unable to resolve DEFs, make sure there are no cycles!")
|
||||
return text
|
||||
for search, replace in replacements:
|
||||
res = re.sub(rf"\b{re.escape(search)}\b", replace, res)
|
||||
if res == prevres:
|
||||
break
|
||||
prevres = res
|
||||
if res != text:
|
||||
log.info("DEFs expanded to: %s", res)
|
||||
return res
|
||||
|
||||
|
||||
@lru_cache
|
||||
def parse_prompt_schedules(prompt, **kwargs):
|
||||
prompt = replace_defs(prompt)
|
||||
return PromptSchedule(prompt, **kwargs)
|
||||
if os.environ.get("PC_USE_OLD_PARSER", "0") != "1":
|
||||
log.info("Using new parser implementation. Set PC_USE_OLD_PARSER=1 to use old parser instead")
|
||||
from .parser_parsy import parse_prompt_schedules # noqa
|
||||
else:
|
||||
from .parser_lark import parse_prompt_schedules # noqa
|
||||
|
||||
@@ -0,0 +1,359 @@
|
||||
# vim: sw=4 ts=4
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from functools import lru_cache
|
||||
from math import ceil
|
||||
|
||||
import lark
|
||||
|
||||
from .macros import expand_macros
|
||||
from .utils import flatten
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
if lark.__version__ == "0.12.0":
|
||||
from sys import executable
|
||||
|
||||
x = "\n".join(
|
||||
[
|
||||
"Your lark package reports an ancient version (0.12.0) and will not work.",
|
||||
"If you have the 'lark-parser' package in your Python environment, remove that and *reinstall* lark!",
|
||||
f"{executable} -m pip uninstall lark-parser lark",
|
||||
f"{executable} -m pip install lark",
|
||||
]
|
||||
)
|
||||
log.error(x)
|
||||
raise ImportError(x)
|
||||
|
||||
|
||||
ESCAPES = [
|
||||
("XxPCBackslashESCAPExX", "\\"),
|
||||
("XxPCColonESCAPExX", ":"),
|
||||
("XxPCCommentESCAPExX", "#"),
|
||||
]
|
||||
|
||||
|
||||
def escape_specials(string: str) -> str:
|
||||
for ph, c in ESCAPES:
|
||||
string = string.replace(rf"\{c}", ph)
|
||||
return string
|
||||
|
||||
|
||||
def restore_escaped(string: str) -> str:
|
||||
for ph, c in ESCAPES:
|
||||
string = string.replace(ph, c)
|
||||
return string
|
||||
|
||||
|
||||
def remove_comments(string: str) -> str:
|
||||
r = []
|
||||
for line in string.split("\n"):
|
||||
comment = line.find("#")
|
||||
if comment >= 0:
|
||||
r.append(line[:comment])
|
||||
else:
|
||||
r.append(line)
|
||||
return "\n".join(r)
|
||||
|
||||
|
||||
prompt_parser = lark.Lark(
|
||||
r"""
|
||||
!start: (prompt | /[][():|]/+)*
|
||||
prompt: (emphasized | embedding | scheduled | alternate | sequence | loraspec | PLAIN | | /\\:/ | /</ | />/ | WHITESPACE)+
|
||||
!emphasized: "(" prompt? ")"
|
||||
| "(" prompt ":" prompt ")"
|
||||
| "[" prompt "]"
|
||||
promptlist: ([prompt] ":")~1..3
|
||||
scheduled: "[" promptlist _WS? NUMBER ["," NUMBER] "]"
|
||||
| "[" promptlist _WS? TAG "]"
|
||||
sequence.5: "[SEQ" ":" [prompt] ":" NUMBER (":" [prompt] ":" NUMBER)* "]"
|
||||
alternate: "[" [prompt] ("|" [prompt])+ [":" NUMBER] "]"
|
||||
loraspec.99: "<lora:" FILENAME lora_weights [lora_block_weights] ">"
|
||||
lora_weights.1: (":" _WS? NUMBER)~1..2
|
||||
lora_block_weights.-1: ":" PLAIN
|
||||
embedding.100: "<emb:" FILENAME ">"
|
||||
WHITESPACE: /\s+/
|
||||
_WS: WHITESPACE
|
||||
PLAIN: /([^<>\\\[\]():|]|\\.)+/
|
||||
FILENAME: /[^<>:]+/
|
||||
TAG: /[A-Z_]+/
|
||||
%import common.SIGNED_NUMBER -> NUMBER
|
||||
""",
|
||||
lexer="dynamic",
|
||||
)
|
||||
|
||||
|
||||
def clamp(a, b, c):
|
||||
"""clamp b between a and c"""
|
||||
return min(max(a, b), c)
|
||||
|
||||
|
||||
def get_steps(tree, num_steps):
|
||||
res = [num_steps or 100]
|
||||
|
||||
def tostep(s):
|
||||
steps = num_steps or 100
|
||||
if "." in str(s) or not num_steps:
|
||||
w = float(s)
|
||||
value = w * steps
|
||||
else:
|
||||
w = int(s)
|
||||
value = w
|
||||
|
||||
if w > 1 and not num_steps:
|
||||
log.warning(
|
||||
"You haven't configured the number of steps for Prompt Control to use, %s will be clipped to 1.0", w
|
||||
)
|
||||
value = steps
|
||||
|
||||
return int(clamp(0, value, steps))
|
||||
|
||||
class CollectSteps(lark.Visitor):
|
||||
def scheduled(self, tree):
|
||||
i = tree.children[-1]
|
||||
if i and i.type == "TAG":
|
||||
return
|
||||
for i in [-1, -2]:
|
||||
if tree.children[i] is not None:
|
||||
tree.children[i] = tostep(tree.children[i])
|
||||
res.append(tree.children[i])
|
||||
|
||||
def interp_steps(self, tree):
|
||||
tree.children[-1] = tostep(tree.children[-1] or 0.1)
|
||||
for i, _ in enumerate(tree.children[:-1]):
|
||||
tree.children[i] = tostep(tree.children[i])
|
||||
|
||||
res.extend(tree.children[:-1])
|
||||
|
||||
def sequence(self, tree):
|
||||
steps = tree.children[1::2]
|
||||
for i, _ in enumerate(steps):
|
||||
w = tostep(tree.children[i * 2 + 1])
|
||||
tree.children[i * 2 + 1] = w
|
||||
res.append(w)
|
||||
|
||||
def alternate(self, tree):
|
||||
step_size = tostep(round(float(tree.children[-1] or 0.1), 2))
|
||||
tree.children[-1] = step_size
|
||||
res.extend([x for x in range(step_size, num_steps or 100, step_size)])
|
||||
|
||||
CollectSteps().visit(tree)
|
||||
|
||||
return sorted(set(res))
|
||||
|
||||
|
||||
def at_step(step, filters, tree):
|
||||
class AtStep(lark.Transformer):
|
||||
def scheduled(self, args):
|
||||
before = None
|
||||
during = None
|
||||
after = None
|
||||
when_end = None
|
||||
pl, when, *rest = args
|
||||
if rest:
|
||||
when_end = rest[0]
|
||||
|
||||
pl = list(pl)
|
||||
if len(pl) == 1:
|
||||
(during,) = pl # [after:0.5] == [::after:0.5,0.5]
|
||||
if when_end is None:
|
||||
when_end = when
|
||||
after = during
|
||||
elif len(pl) == 2:
|
||||
during, after = pl # [during:after:0.5] = [before::after:0.5,0.5]
|
||||
if when_end is None:
|
||||
when_end = when
|
||||
before = during
|
||||
else:
|
||||
before, during, after = pl # [before:during:after:0.5,0.8]
|
||||
|
||||
if isinstance(when, str):
|
||||
return before or "" if when not in filters else after or ""
|
||||
|
||||
if when_end is None:
|
||||
when_end = 1000_000
|
||||
|
||||
if step <= when:
|
||||
return before or ""
|
||||
if when < step <= when_end:
|
||||
return during or ""
|
||||
else:
|
||||
return after or ""
|
||||
|
||||
def sequence(self, args):
|
||||
previous_step = 0.0
|
||||
prompts = args[::2]
|
||||
steps = args[1::2]
|
||||
for s, p in zip(steps, prompts, strict=False):
|
||||
if s >= step and step >= previous_step:
|
||||
previous_step = step
|
||||
return p or ""
|
||||
else:
|
||||
previous_step = s
|
||||
return ""
|
||||
|
||||
def alternate(self, args):
|
||||
step_size = args[-1]
|
||||
idx = ceil(step / step_size)
|
||||
return args[(idx - 1) % (len(args) - 1)] or ""
|
||||
|
||||
def start(self, args):
|
||||
prompt = []
|
||||
loraspecs = {}
|
||||
args = flatten(args)
|
||||
for a in args:
|
||||
if isinstance(a, str):
|
||||
prompt.append(a)
|
||||
elif isinstance(a, tuple):
|
||||
# sum identical specs together
|
||||
n = a[0]
|
||||
# if clip weight is not provided, use unet weight
|
||||
w, w_clip = a[1][0], a[1][1 % len(a[1])]
|
||||
e = loraspecs.get(n, {})
|
||||
loraspecs[n] = {
|
||||
"weight": round(e.get("weight", 0.0) + w, 2),
|
||||
"weight_clip": round(e.get("weight_clip", 0.0) + w_clip, 2),
|
||||
}
|
||||
lbw = a[2]
|
||||
if lbw:
|
||||
loraspecs[n]["lbw"] = lbw
|
||||
if loraspecs[n]["weight"] == 0 and loraspecs[n]["weight_clip"] == 0 and not lbw:
|
||||
del loraspecs[n]
|
||||
else:
|
||||
pass
|
||||
p = "".join(prompt)
|
||||
return {"prompt": p, "loras": loraspecs}
|
||||
|
||||
def PLAIN(self, args):
|
||||
return restore_escaped(args)
|
||||
|
||||
def FILENAME(self, value):
|
||||
return str(value)
|
||||
|
||||
def embedding(self, args):
|
||||
return "embedding:" + str(args[0])
|
||||
|
||||
def lora_weights(self, args):
|
||||
return [float(str(a)) for a in args]
|
||||
|
||||
def lora_block_weights(self, args):
|
||||
vals = args[0].split(";")
|
||||
r = {}
|
||||
for v in vals:
|
||||
x = v.split("=", 2)
|
||||
if len(x) != 2:
|
||||
continue
|
||||
k, v = x[0].strip().upper(), x[1].strip()
|
||||
r[k] = v
|
||||
return r
|
||||
|
||||
def loraspec(self, args):
|
||||
name = args[0]
|
||||
params = args[1]
|
||||
lbw = args[2]
|
||||
|
||||
return name, params, lbw
|
||||
|
||||
def __default__(self, data, children, meta):
|
||||
return children
|
||||
|
||||
return AtStep().transform(tree)
|
||||
|
||||
|
||||
class PromptSchedule:
|
||||
# 0 num_steps means unconfigured
|
||||
def __init__(self, prompt, filters="", start=0.0, end=1.0, num_steps=0):
|
||||
self.filters = filters
|
||||
self.start = start
|
||||
self.end = end
|
||||
self.num_steps = num_steps
|
||||
# placeholder is restored on parse
|
||||
self.prompt = remove_comments(escape_specials(prompt.strip()))
|
||||
self.defaults = {}
|
||||
self.loaded_loras = {}
|
||||
|
||||
self.parsed_prompt = self._parse(num_steps)
|
||||
|
||||
def __iter__(self):
|
||||
# Filter out zero, it's only useful for interpolation
|
||||
return (x for x in self.parsed_prompt if x[0] != 0)
|
||||
|
||||
def _parse(self, num_steps):
|
||||
filters = [x.strip() for x in self.filters.upper().split(",")]
|
||||
try:
|
||||
parsed = []
|
||||
tree = prompt_parser.parse(self.prompt)
|
||||
steps = get_steps(tree, num_steps=num_steps)
|
||||
|
||||
def f(x):
|
||||
return round(x / (num_steps or 100), 2)
|
||||
|
||||
for t in steps:
|
||||
p = at_step(t, filters, tree)
|
||||
parsed.append([f(t), p])
|
||||
|
||||
except lark.exceptions.LarkError as e:
|
||||
log.error("Prompt editing parse error: %s", e)
|
||||
parsed = [[1.0, {"prompt": self.prompt, "loras": {}}]]
|
||||
raise
|
||||
|
||||
# Tag filtering may return redundant prompts, so filter them out here
|
||||
res = []
|
||||
prev_end = -1
|
||||
|
||||
for end_at, p in parsed:
|
||||
if end_at < self.start:
|
||||
continue
|
||||
elif end_at <= self.end:
|
||||
res.append([end_at, p])
|
||||
prev_end = end_at
|
||||
elif end_at > self.end and prev_end < self.end:
|
||||
res.append([end_at, p])
|
||||
break
|
||||
|
||||
# Always use the last prompt if everything was filtered
|
||||
if len(res) == 0:
|
||||
res = [[1.0, parsed[-1][1]]]
|
||||
|
||||
final = [res[0]]
|
||||
|
||||
# Clean up duplicates
|
||||
for p in res[1:]:
|
||||
if p[1] != final[-1][1]:
|
||||
final.append(p)
|
||||
else:
|
||||
final[-1][0] = p[0]
|
||||
return final
|
||||
|
||||
def clone(self):
|
||||
return self.with_filters()
|
||||
|
||||
def with_filters(self, filters=None, start=None, end=None, defaults=None):
|
||||
def ifspecified(x, defval):
|
||||
return x if x is not None else defval
|
||||
|
||||
p = PromptSchedule(
|
||||
self.prompt,
|
||||
filters=ifspecified(filters, self.filters),
|
||||
start=ifspecified(start, self.start),
|
||||
end=ifspecified(end, self.end),
|
||||
num_steps=self.num_steps,
|
||||
)
|
||||
return p
|
||||
|
||||
def at_step(self, step, total_steps=1):
|
||||
_, x = self.at_step_idx(step, total_steps)
|
||||
return x
|
||||
|
||||
def at_step_idx(self, step, total_steps=1):
|
||||
for i, x in enumerate(self.parsed_prompt):
|
||||
if x[0] * total_steps >= step:
|
||||
return i, x
|
||||
return len(self.parsed_prompt) - 1, self.parsed_prompt[-1]
|
||||
|
||||
|
||||
@lru_cache
|
||||
def parse_prompt_schedules(prompt, **kwargs):
|
||||
prompt = expand_macros(prompt)
|
||||
return PromptSchedule(prompt, **kwargs)
|
||||
@@ -0,0 +1,389 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import itertools as it
|
||||
from dataclasses import dataclass
|
||||
from math import ceil
|
||||
from typing import Any, TypeAlias
|
||||
|
||||
from typing_extensions import override
|
||||
|
||||
from .macros import expand_macros
|
||||
from .parsy import any_char, char_from, digit, eof, forward_declaration, generate, regex, seq, string, success
|
||||
|
||||
FOREVER = float("inf")
|
||||
|
||||
EvalResult: TypeAlias = tuple[float, str, list["LoRA"]]
|
||||
|
||||
|
||||
def merge_until(i: EvalResult, minimum: float):
|
||||
until, p, loras = i
|
||||
until = min(until, minimum)
|
||||
return until, p, loras
|
||||
|
||||
|
||||
def batched(iterable, n, *, strict=False):
|
||||
# batched('ABCDEFG', 2) → AB CD EF G
|
||||
if n < 1:
|
||||
raise ValueError("n must be at least one")
|
||||
iterator = iter(iterable)
|
||||
while batch := tuple(it.islice(iterator, n)):
|
||||
if strict and len(batch) != n:
|
||||
raise ValueError("batched(): incomplete batch")
|
||||
yield batch
|
||||
|
||||
|
||||
EvalResult: TypeAlias = tuple[float, str, list["LoRA"]]
|
||||
|
||||
|
||||
class Expression:
|
||||
def eval(self, step: float, tags: list[str]) -> EvalResult:
|
||||
return (FOREVER, "", [])
|
||||
|
||||
def required_steps(self, max_steps: float) -> set[float]:
|
||||
return set()
|
||||
|
||||
|
||||
@dataclass
|
||||
class Text(Expression):
|
||||
string: str
|
||||
|
||||
@override
|
||||
def eval(self, step: float, tags: list[str]) -> EvalResult:
|
||||
assert isinstance(self.string, str)
|
||||
return FOREVER, self.string, []
|
||||
|
||||
|
||||
@dataclass
|
||||
class Alternate(Expression):
|
||||
prompts: list[Expression]
|
||||
step: float = 0.1
|
||||
|
||||
@override
|
||||
def eval(self, step: float, tags: list[str]) -> EvalResult:
|
||||
SCALE = 10_000
|
||||
step = max(step, self.step)
|
||||
position = (step * SCALE) / (self.step * SCALE)
|
||||
idx = (ceil(position) - 1) % len(self.prompts)
|
||||
|
||||
r = self.prompts[max(0, idx)].eval(step, tags)
|
||||
r = merge_until(r, max(self.step, ceil(position) * self.step))
|
||||
return r
|
||||
|
||||
@override
|
||||
def required_steps(self, max_steps: float):
|
||||
r = set()
|
||||
for x in self.prompts:
|
||||
r.update(x.required_steps(max_steps))
|
||||
r.update(set(x / 100 for x in range(0, int(max_steps * 100), int(self.step * 100))))
|
||||
return r
|
||||
|
||||
|
||||
@dataclass
|
||||
class Sequence(Expression):
|
||||
prompts: list[tuple[Expression, float]]
|
||||
|
||||
@override
|
||||
def eval(self, step: float, tags: list[str]) -> EvalResult:
|
||||
item = Text("")
|
||||
found_step = FOREVER
|
||||
for prompt, switch_step in self.prompts:
|
||||
if step <= switch_step:
|
||||
found_step = switch_step
|
||||
item = prompt
|
||||
break
|
||||
|
||||
return merge_until(item.eval(step, tags), found_step)
|
||||
|
||||
@override
|
||||
def required_steps(self, max_steps: float):
|
||||
return set(step for _, step in self.prompts if step <= max_steps)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Schedule(Expression):
|
||||
before: Prompt
|
||||
during: Prompt
|
||||
after: Prompt
|
||||
start: float
|
||||
end: float
|
||||
tag: str | None
|
||||
|
||||
def tag_matches(self, tags: list[str]):
|
||||
return self.tag in tags
|
||||
|
||||
@override
|
||||
def eval(self, step: float, tags: list[str]) -> EvalResult:
|
||||
if self.tag is not None and not self.tag_matches(tags):
|
||||
return self.before.eval(step, tags)
|
||||
if self.tag_matches(tags):
|
||||
return self.during.eval(step, tags)
|
||||
|
||||
if step <= self.start:
|
||||
return merge_until(self.before.eval(step, tags), self.start)
|
||||
if self.start < step <= self.end:
|
||||
return merge_until(self.during.eval(step, tags), self.end)
|
||||
if step > self.end:
|
||||
return self.after.eval(step, tags)
|
||||
raise AssertionError("How are you here?")
|
||||
|
||||
@override
|
||||
def required_steps(self, max_steps: float):
|
||||
r = set()
|
||||
if self.tag is not None:
|
||||
return r
|
||||
if self.start < max_steps:
|
||||
r.add(self.start)
|
||||
if self.end < max_steps:
|
||||
r.add(self.end)
|
||||
r.update(self.before.required_steps(max_steps))
|
||||
r.update(self.during.required_steps(max_steps))
|
||||
r.update(self.after.required_steps(max_steps))
|
||||
return r
|
||||
|
||||
|
||||
@dataclass
|
||||
class Prompt(Expression):
|
||||
data: list[Expression]
|
||||
|
||||
@override
|
||||
def eval(self, step: float, tags: list[str]) -> EvalResult:
|
||||
evals = [x.eval(step, tags) for x in self.data]
|
||||
text = "".join(x[1] for x in evals)
|
||||
untils = [x[0] for x in evals]
|
||||
loras = []
|
||||
for x in evals:
|
||||
loras.extend(x[2])
|
||||
until = FOREVER if not untils else min(untils)
|
||||
return until, text, loras
|
||||
|
||||
@override
|
||||
def required_steps(self, max_steps):
|
||||
r = set()
|
||||
for x in self.data:
|
||||
r.update(x.required_steps(max_steps))
|
||||
return r
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoRA(Expression):
|
||||
filename: str
|
||||
w_model: float = 1.0
|
||||
w_te: float = 1.0
|
||||
|
||||
def eval(self, step: float, tags: list[str]) -> EvalResult:
|
||||
return FOREVER, "", [self]
|
||||
|
||||
|
||||
def find_weight_at(weights: list[tuple[float, float]], step: float, until: float):
|
||||
res_w = 0
|
||||
for this, next in zip(weights, it.chain(weights[1:], [(0, FOREVER)]), strict=False):
|
||||
w, start = this
|
||||
_, next_start = next
|
||||
if start > step or next_start < step:
|
||||
until = min(until, start)
|
||||
continue
|
||||
res_w = w
|
||||
return until, res_w
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoRACTL(Expression):
|
||||
filename: str
|
||||
w_model: list[tuple[float, float]]
|
||||
w_te: list[tuple[float, float]]
|
||||
|
||||
def eval(self, step: float, tags: list[str]) -> EvalResult:
|
||||
until, w1 = find_weight_at(self.w_model, step, FOREVER)
|
||||
until, w2 = find_weight_at(self.w_te, step, until)
|
||||
lora = []
|
||||
if w1 != 0 or w1 != 0:
|
||||
lora = [LoRA(self.filename, w1, w2)]
|
||||
|
||||
return until, "", lora
|
||||
|
||||
def required_steps(self, max_steps):
|
||||
r = set(x[1] for x in self.w_model)
|
||||
r.update(set(x[1] for x in self.w_te))
|
||||
return r
|
||||
|
||||
|
||||
def combine_arglist(prompts, start_end) -> Schedule:
|
||||
a, b, c = prompts
|
||||
start_or_tag, end = start_end
|
||||
empty = Prompt([])
|
||||
start = start_or_tag
|
||||
# Handle [a:b:TAG]
|
||||
if isinstance(start_or_tag, str):
|
||||
if b is None:
|
||||
before = empty
|
||||
during = a # [a:TAG] produces a when tag is active
|
||||
else:
|
||||
before, during = a, b # [a:b:TAG] changes from a to b when tag is active
|
||||
return Schedule(before, during, empty, start=0.0, end=FOREVER, tag=start_or_tag)
|
||||
during = before = after = empty
|
||||
if end is not None:
|
||||
if b is None: # [a:0,0.5] == [:a:0,0.5]
|
||||
during = a
|
||||
before = after = empty
|
||||
elif c is None: # [a:b:0,0.5]
|
||||
before = empty
|
||||
during = a
|
||||
after = b
|
||||
else:
|
||||
before, during, after = a, b, c
|
||||
else:
|
||||
end = FOREVER
|
||||
if b is None: # [a:0.5] == [::a:0.5,0.5]
|
||||
before = empty
|
||||
during = a
|
||||
after = a
|
||||
else:
|
||||
before = a
|
||||
during = b
|
||||
after = b
|
||||
# c always gets ignored
|
||||
start = float(start) # for typechecking
|
||||
return Schedule(before, during, after, start, end, tag=None)
|
||||
|
||||
|
||||
def token(s: str):
|
||||
return string(s).map(Text)
|
||||
|
||||
|
||||
def combine_prompt(*prompts):
|
||||
p = prompts
|
||||
if len(p) == 1:
|
||||
p = p[0]
|
||||
if isinstance(p, Prompt):
|
||||
p = p.data[0] if len(p.data) == 1 else combine_prompt(*p.data)
|
||||
if isinstance(p, Expression):
|
||||
return p
|
||||
p = [combine_prompt(x) for x in p]
|
||||
return Prompt(p)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptSchedule:
|
||||
parse_tree: Expression
|
||||
filters: list[str]
|
||||
start: float
|
||||
end: float
|
||||
num_steps: int
|
||||
|
||||
def at_step(self, step: float) -> tuple[float, dict[str, Any]]:
|
||||
max_step = self.num_steps or 1.0
|
||||
if max_step > 1 and step < 1:
|
||||
step = step * max_step
|
||||
until, p, lora_list = self.parse_tree.eval(step, self.filters)
|
||||
loras = {}
|
||||
for lora in lora_list:
|
||||
d = loras.get(lora.filename, {})
|
||||
d["weight"] = d.get("weight", 0) + lora.w_model
|
||||
d["weight_clip"] = d.get("weight_clip", 0) + lora.w_te
|
||||
loras[lora.filename] = d
|
||||
if max_step > 0 and until > 1:
|
||||
# TODO: better logic for this?
|
||||
until = min(until / max_step, 1.0)
|
||||
return (min(max_step, round(until, 2)), {"prompt": p, "loras": loras})
|
||||
|
||||
def with_filters(self, filters: str | None = None, start: float | None = None, end: float | None = None):
|
||||
return PromptSchedule(
|
||||
self.parse_tree,
|
||||
self.filters if filters is None else parse_filters(filters),
|
||||
self.start if start is None else start,
|
||||
self.end if end is None else end,
|
||||
self.num_steps,
|
||||
)
|
||||
|
||||
def clone(self):
|
||||
return self.with_filters()
|
||||
|
||||
def __iter__(self):
|
||||
return (x for x in self.parsed_prompt if x[0] != 0)
|
||||
|
||||
@property
|
||||
def parsed_prompt(self):
|
||||
max_step = self.num_steps or 1.0
|
||||
required_steps = self.parse_tree.required_steps(max_step).union({max_step})
|
||||
|
||||
prompts = list(sorted((self.at_step(step) for step in required_steps), key=lambda x: x[0]))
|
||||
res = []
|
||||
prev_end = -1
|
||||
for end_at, p in prompts:
|
||||
if end_at < self.start:
|
||||
continue
|
||||
elif end_at < self.end and prev_end < end_at:
|
||||
res.append([end_at, p])
|
||||
prev_end = end_at
|
||||
elif end_at >= self.end and prev_end < self.end:
|
||||
res.append([end_at, p])
|
||||
break
|
||||
|
||||
if len(res) == 0:
|
||||
res = [[1.0], prompts[-1][1]]
|
||||
|
||||
return res
|
||||
|
||||
|
||||
def lora_weights(p):
|
||||
@generate
|
||||
def parser():
|
||||
w_model = yield col >> p
|
||||
w_te = yield (col >> p).optional(w_model)
|
||||
return [w_model, w_te]
|
||||
|
||||
return parser.desc("lora_weights")
|
||||
|
||||
|
||||
prompt = forward_declaration()
|
||||
empty = Text("")
|
||||
comma = token(",")
|
||||
col = token(":")
|
||||
lsq = token("[")
|
||||
rsq = token("]")
|
||||
lpar = token("(")
|
||||
rpar = token(")")
|
||||
tag = regex(r"[A-Z_]+")
|
||||
non_special = regex(r"[^:\[\]()|\\<>#]+").map(Text)
|
||||
filename = regex(r"[^:<>]+")
|
||||
|
||||
comment = string("#") >> any_char.until(eof | char_from("\n")) >> success(empty)
|
||||
escape = (string("\\") >> char_from("\\[]:#") | string(r"\(") | string(r"\)")).map(Text)
|
||||
emphasis = seq(lpar, (prompt | col).at_least(0), rpar)
|
||||
sign = string("+") | string("-")
|
||||
number = (
|
||||
(sign.optional("") + (digit.at_least(1) + string(".") * 1 + digit.many() | digit.at_least(1)).concat())
|
||||
.concat()
|
||||
.map(float)
|
||||
)
|
||||
|
||||
opt_prompt = prompt.optional(empty)
|
||||
step_range = seq(number | tag, (comma >> number).optional())
|
||||
arglist = seq((opt_prompt << col).optional() * 3, step_range)
|
||||
schedule = lsq >> arglist.combine(combine_arglist) << rsq
|
||||
|
||||
alternate = (lsq >> seq(prompt.sep_by(string("|"), min=1), (col >> number).optional(0.1)) << rsq).combine(Alternate)
|
||||
sequence = (lsq >> string("SEQ") >> seq(col >> opt_prompt << col, number).at_least(1) << rsq).map(Sequence)
|
||||
bracketed = seq(lsq, prompt.at_least(0), rsq) | sequence | schedule | alternate
|
||||
lora = (string("<lora:") >> filename * 1 + lora_weights(number) << string(">")).combine(LoRA)
|
||||
ctlweight = seq(number, (string("@") >> number).optional(0)).sep_by(comma, min=1)
|
||||
loractl = (string("<loractl:") >> filename * 1 + lora_weights(ctlweight) << string(">")).combine(LoRACTL)
|
||||
emb = (string("<emb:") >> filename << string(">")).map(lambda f: Text(f"embedding:{f}"))
|
||||
|
||||
expr = escape | comment | non_special | bracketed | emphasis.combine(combine_prompt) | lora | loractl | emb
|
||||
prompt_ = expr.at_least(1).combine(combine_prompt)
|
||||
prompt.become(prompt_)
|
||||
# Treat any character that isn't valid prompt syntax as just text
|
||||
all = (prompt | any_char.map(Text)).at_least(0).combine(combine_prompt)
|
||||
|
||||
|
||||
def parse_filters(filters: str):
|
||||
return [x.strip().upper() for x in filters.split(",") if x.strip()]
|
||||
|
||||
|
||||
def parse(text):
|
||||
return combine_prompt(all.parse(text))
|
||||
|
||||
|
||||
def parse_prompt_schedules(text, filters="", start=0, end=1.0, num_steps=0):
|
||||
return PromptSchedule(parse(expand_macros(text.strip())), parse_filters(filters), start, end, num_steps)
|
||||
@@ -0,0 +1,720 @@
|
||||
# Vendored from https://github.com/python-parsy/parsy/blob/master/src/parsy/__init__.py
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import operator
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from functools import wraps
|
||||
from typing import Any, Callable, FrozenSet
|
||||
|
||||
__version__ = "2.2"
|
||||
|
||||
noop = lambda x: x
|
||||
|
||||
|
||||
def line_info_at(stream, index):
|
||||
if index > len(stream):
|
||||
raise ValueError("invalid index")
|
||||
line = stream.count("\n", 0, index)
|
||||
last_nl = stream.rfind("\n", 0, index)
|
||||
col = index - (last_nl + 1)
|
||||
return (line, col)
|
||||
|
||||
|
||||
class ParseError(RuntimeError):
|
||||
def __init__(self, expected, stream, index):
|
||||
self.expected = expected
|
||||
self.stream = stream
|
||||
self.index = index
|
||||
|
||||
def line_info(self) -> str:
|
||||
try:
|
||||
return "{}:{}".format(*line_info_at(self.stream, self.index))
|
||||
except (TypeError, AttributeError): # not a str
|
||||
return str(self.index)
|
||||
|
||||
def __str__(self):
|
||||
expected_list = sorted(repr(e) for e in self.expected)
|
||||
|
||||
if len(expected_list) == 1:
|
||||
return f"expected {expected_list[0]} at {self.line_info()}"
|
||||
else:
|
||||
return f"expected one of {', '.join(expected_list)} at {self.line_info()}"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Result:
|
||||
status: bool
|
||||
index: int
|
||||
value: Any
|
||||
furthest: int
|
||||
expected: FrozenSet[str]
|
||||
|
||||
@staticmethod
|
||||
def success(index, value) -> Result:
|
||||
return Result(True, index, value, -1, frozenset())
|
||||
|
||||
@staticmethod
|
||||
def failure(index, expected) -> Result:
|
||||
return Result(False, -1, None, index, frozenset([expected]))
|
||||
|
||||
# collect the furthest failure from self and other
|
||||
def aggregate(self, other) -> Result:
|
||||
if not other:
|
||||
return self
|
||||
|
||||
if self.furthest > other.furthest:
|
||||
return self
|
||||
elif self.furthest == other.furthest:
|
||||
# if we both have the same failure index, we combine the expected messages.
|
||||
return Result(self.status, self.index, self.value, self.furthest, self.expected | other.expected)
|
||||
else:
|
||||
return Result(self.status, self.index, self.value, other.furthest, other.expected)
|
||||
|
||||
|
||||
# Roughly, a stream is str|bytes|list, but in practice we are duck-typed
|
||||
# and could accept other things.
|
||||
# We should switch to this alias when all supported Python versions allow it:
|
||||
# type Stream = str | bytes | list
|
||||
|
||||
|
||||
class Parser:
|
||||
"""
|
||||
A Parser is an object that wraps a function whose arguments are
|
||||
a string to be parsed and the index on which to begin parsing.
|
||||
The function should return either Result.success(next_index, value),
|
||||
where the next index is where to continue the parse and the value is
|
||||
the yielded value, or Result.failure(index, expected), where expected
|
||||
is a string indicating what was expected, and the index is the index
|
||||
of the failure.
|
||||
"""
|
||||
|
||||
def __init__(self, wrapped_fn: Callable[[str | bytes | list, int], Result]):
|
||||
"""
|
||||
Creates a new Parser from a function that takes a stream
|
||||
and returns a Result.
|
||||
"""
|
||||
self.wrapped_fn = wrapped_fn
|
||||
|
||||
def __call__(self, stream: str | bytes | list, index: int) -> Any:
|
||||
return self.wrapped_fn(stream, index)
|
||||
|
||||
def parse(self, stream: str | bytes | list) -> Any:
|
||||
"""Parses a string or list of tokens and returns the result or raise a ParseError."""
|
||||
(result, _) = (self << eof).parse_partial(stream)
|
||||
return result
|
||||
|
||||
def parse_partial(self, stream: str | bytes | list) -> tuple[Any, str | bytes | list]:
|
||||
"""
|
||||
Parses the longest possible prefix of a given string.
|
||||
Returns a tuple of the result and the unparsed remainder,
|
||||
or raises ParseError
|
||||
"""
|
||||
result = self(stream, 0)
|
||||
|
||||
if result.status:
|
||||
return (result.value, stream[result.index :])
|
||||
else:
|
||||
raise ParseError(result.expected, stream, result.furthest)
|
||||
|
||||
def bind(self, bind_fn: Callable[[Any], Parser]) -> Parser:
|
||||
@Parser
|
||||
def bound_parser(stream: str | bytes | list, index: int) -> Result:
|
||||
result = self(stream, index)
|
||||
|
||||
if result.status:
|
||||
next_parser = bind_fn(result.value)
|
||||
return next_parser(stream, result.index).aggregate(result)
|
||||
else:
|
||||
return result
|
||||
|
||||
return bound_parser
|
||||
|
||||
def map(self, map_function: Callable) -> Parser:
|
||||
"""
|
||||
Returns a parser that transforms the produced value of the initial parser with map_function.
|
||||
"""
|
||||
return self.bind(lambda res: success(map_function(res)))
|
||||
|
||||
def combine(self, combine_fn: Callable) -> Parser:
|
||||
"""
|
||||
Returns a parser that transforms the produced values of the initial parser
|
||||
with ``combine_fn``, passing the arguments using ``*args`` syntax.
|
||||
|
||||
The initial parser should return a list/sequence of parse results.
|
||||
"""
|
||||
return self.bind(lambda res: success(combine_fn(*res)))
|
||||
|
||||
def combine_dict(self, combine_fn: Callable) -> Parser:
|
||||
"""
|
||||
Returns a parser that transforms the value produced by the initial parser
|
||||
using the supplied function/callable, passing the arguments using the
|
||||
``**kwargs`` syntax.
|
||||
|
||||
The value produced by the initial parser must be a mapping/dictionary from
|
||||
names to values, or a list of two-tuples, or something else that can be
|
||||
passed to the ``dict`` constructor.
|
||||
|
||||
If ``None`` is present as a key in the dictionary it will be removed
|
||||
before passing to ``fn``, as will all keys starting with ``_``.
|
||||
"""
|
||||
return self.bind(
|
||||
lambda res: success(
|
||||
combine_fn(
|
||||
**{
|
||||
k: v
|
||||
for k, v in dict(res).items()
|
||||
if k is not None and not (isinstance(k, str) and k.startswith("_"))
|
||||
}
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
def concat(self) -> Parser:
|
||||
"""
|
||||
Returns a parser that concatenates together (as a string) the previously
|
||||
produced values.
|
||||
"""
|
||||
return self.map("".join)
|
||||
|
||||
def then(self, other: Parser) -> Parser:
|
||||
"""
|
||||
Returns a parser which, if the initial parser succeeds, will
|
||||
continue parsing with ``other``. This will produce the
|
||||
value produced by ``other``.
|
||||
|
||||
"""
|
||||
return seq(self, other).combine(lambda left, right: right)
|
||||
|
||||
def skip(self, other: Parser) -> Parser:
|
||||
"""
|
||||
Returns a parser which, if the initial parser succeeds, will
|
||||
continue parsing with ``other``. It will produce the
|
||||
value produced by the initial parser.
|
||||
"""
|
||||
return seq(self, other).combine(lambda left, right: left)
|
||||
|
||||
def result(self, value: Any) -> Parser:
|
||||
"""
|
||||
Returns a parser that, if the initial parser succeeds, always produces
|
||||
the passed in ``value``.
|
||||
"""
|
||||
return self >> success(value)
|
||||
|
||||
def many(self) -> Parser:
|
||||
"""
|
||||
Returns a parser that expects the initial parser 0 or more times, and
|
||||
produces a list of the results.
|
||||
"""
|
||||
return self.times(0, float("inf"))
|
||||
|
||||
def times(self, min: int, max: int = None) -> Parser:
|
||||
"""
|
||||
Returns a parser that expects the initial parser at least ``min`` times,
|
||||
and at most ``max`` times, and produces a list of the results. If only one
|
||||
argument is given, the parser is expected exactly that number of times.
|
||||
"""
|
||||
if max is None:
|
||||
max = min
|
||||
|
||||
@Parser
|
||||
def times_parser(stream: str | bytes | list, index: int) -> Result:
|
||||
values = []
|
||||
times = 0
|
||||
result = None
|
||||
|
||||
while times < max:
|
||||
result = self(stream, index).aggregate(result)
|
||||
if result.status:
|
||||
values.append(result.value)
|
||||
index = result.index
|
||||
times += 1
|
||||
elif times >= min:
|
||||
break
|
||||
else:
|
||||
return result
|
||||
|
||||
return Result.success(index, values).aggregate(result)
|
||||
|
||||
return times_parser
|
||||
|
||||
def at_most(self, n: int) -> Parser:
|
||||
"""
|
||||
Returns a parser that expects the initial parser at most ``n`` times, and
|
||||
produces a list of the results.
|
||||
"""
|
||||
return self.times(0, n)
|
||||
|
||||
def at_least(self, n: int) -> Parser:
|
||||
"""
|
||||
Returns a parser that expects the initial parser at least ``n`` times, and
|
||||
produces a list of the results.
|
||||
"""
|
||||
return self.times(n) + self.many()
|
||||
|
||||
def optional(self, default: Any = None) -> Parser:
|
||||
"""
|
||||
Returns a parser that expects the initial parser zero or once, and maps
|
||||
the result to a given default value in the case of no match. If no default
|
||||
value is given, ``None`` is used.
|
||||
"""
|
||||
return self.times(0, 1).map(lambda v: v[0] if v else default)
|
||||
|
||||
def until(self, other: Parser, min: int = 0, max: int = float("inf"), consume_other: bool = False) -> Parser:
|
||||
"""
|
||||
Returns a parser that expects the initial parser followed by ``other``.
|
||||
The initial parser is expected at least ``min`` times and at most ``max`` times.
|
||||
By default, it does not consume ``other`` and it produces a list of the
|
||||
results excluding ``other``. If ``consume_other`` is ``True`` then
|
||||
``other`` is consumed and its result is included in the list of results.
|
||||
"""
|
||||
|
||||
@Parser
|
||||
def until_parser(stream: str | bytes | list, index: int) -> Result:
|
||||
values = []
|
||||
times = 0
|
||||
while True:
|
||||
|
||||
# try parser first
|
||||
res = other(stream, index)
|
||||
if res.status and times >= min:
|
||||
if consume_other:
|
||||
# consume other
|
||||
values.append(res.value)
|
||||
index = res.index
|
||||
return Result.success(index, values)
|
||||
|
||||
# exceeded max?
|
||||
if times >= max:
|
||||
# return failure, it matched parser more than max times
|
||||
return Result.failure(index, f"at most {max} items")
|
||||
|
||||
# failed, try parser
|
||||
result = self(stream, index)
|
||||
if result.status:
|
||||
# consume
|
||||
values.append(result.value)
|
||||
index = result.index
|
||||
times += 1
|
||||
elif times >= min:
|
||||
# return failure, parser is not followed by other
|
||||
return Result.failure(index, "did not find other parser")
|
||||
else:
|
||||
# return failure, it did not match parser at least min times
|
||||
return Result.failure(index, f"at least {min} items; got {times} item(s)")
|
||||
|
||||
return until_parser
|
||||
|
||||
def sep_by(self, sep: Parser, *, min: int = 0, max: int = float("inf")) -> Parser:
|
||||
"""
|
||||
Returns a new parser that repeats the initial parser and
|
||||
collects the results in a list. Between each item, the ``sep`` parser
|
||||
is run (and its return value is discarded). By default it
|
||||
repeats with no limit, but minimum and maximum values can be supplied.
|
||||
"""
|
||||
zero_times = success([])
|
||||
if max == 0:
|
||||
return zero_times
|
||||
res = self.times(1) + (sep >> self).times(min - 1, max - 1)
|
||||
if min == 0:
|
||||
res |= zero_times
|
||||
return res
|
||||
|
||||
def desc(self, description: str) -> Parser:
|
||||
"""
|
||||
Returns a new parser with a description added, which is used in the error message
|
||||
if parsing fails.
|
||||
"""
|
||||
|
||||
@Parser
|
||||
def desc_parser(stream: str | bytes | list, index: int) -> Result:
|
||||
result = self(stream, index)
|
||||
if result.status:
|
||||
return result
|
||||
else:
|
||||
return Result.failure(index, description)
|
||||
|
||||
return desc_parser
|
||||
|
||||
def mark(self) -> Parser:
|
||||
"""
|
||||
Returns a parser that wraps the initial parser's result in a value
|
||||
containing column and line information of the match, as well as the
|
||||
original value. The new value is a 3-tuple:
|
||||
|
||||
((start_row, start_column),
|
||||
original_value,
|
||||
(end_row, end_column))
|
||||
"""
|
||||
|
||||
@generate
|
||||
def marked():
|
||||
start = yield line_info
|
||||
body = yield self
|
||||
end = yield line_info
|
||||
return (start, body, end)
|
||||
|
||||
return marked
|
||||
|
||||
def tag(self, name: str) -> Parser:
|
||||
"""
|
||||
Returns a parser that wraps the produced value of the initial parser in a
|
||||
2 tuple containing ``(name, value)``. This provides a very simple way to
|
||||
label parsed components
|
||||
"""
|
||||
return self.map(lambda v: (name, v))
|
||||
|
||||
def should_fail(self, description: str) -> Parser:
|
||||
"""
|
||||
Returns a parser that fails when the initial parser succeeds, and succeeds
|
||||
when the initial parser fails (consuming no input). A description must
|
||||
be passed which is used in parse failure messages.
|
||||
|
||||
This is essentially a negative lookahead
|
||||
"""
|
||||
|
||||
@Parser
|
||||
def fail_parser(stream: str | bytes | list, index: int) -> Result:
|
||||
res = self(stream, index)
|
||||
if res.status:
|
||||
return Result.failure(index, description)
|
||||
return Result.success(index, res)
|
||||
|
||||
return fail_parser
|
||||
|
||||
def __add__(self, other: Parser) -> Parser:
|
||||
return seq(self, other).combine(operator.add)
|
||||
|
||||
def __mul__(self, other: int | range) -> Parser:
|
||||
if isinstance(other, range):
|
||||
return self.times(other.start, other.stop - 1)
|
||||
return self.times(other)
|
||||
|
||||
def __or__(self, other: Parser) -> Parser:
|
||||
return alt(self, other)
|
||||
|
||||
# haskelley operators, for fun #
|
||||
|
||||
# >>
|
||||
def __rshift__(self, other: Parser) -> Parser:
|
||||
return self.then(other)
|
||||
|
||||
# <<
|
||||
def __lshift__(self, other: Parser) -> Parser:
|
||||
return self.skip(other)
|
||||
|
||||
|
||||
def alt(*parsers: Parser) -> Parser:
|
||||
"""
|
||||
Creates a parser from the passed in argument list of alternative
|
||||
parsers, which are tried in order, moving to the next one if the
|
||||
current one fails.
|
||||
"""
|
||||
if not parsers:
|
||||
return fail("<empty alt>")
|
||||
|
||||
@Parser
|
||||
def alt_parser(stream: str | bytes | list, index: int) -> Result:
|
||||
result = None
|
||||
for parser in parsers:
|
||||
result = parser(stream, index).aggregate(result)
|
||||
if result.status:
|
||||
return result
|
||||
|
||||
return result
|
||||
|
||||
return alt_parser
|
||||
|
||||
|
||||
def seq(*parsers: Parser, **kw_parsers: Parser) -> Parser:
|
||||
"""
|
||||
Takes a list of parsers, runs them in order,
|
||||
and collects their individuals results in a list,
|
||||
or in a dictionary if you pass them as keyword arguments.
|
||||
"""
|
||||
if not parsers and not kw_parsers:
|
||||
return success([])
|
||||
|
||||
if parsers and kw_parsers:
|
||||
raise ValueError("Use either positional arguments or keyword arguments with seq, not both")
|
||||
|
||||
if parsers:
|
||||
|
||||
@Parser
|
||||
def seq_parser(stream: str | bytes | list, index: int) -> Result:
|
||||
result = None
|
||||
values = []
|
||||
for parser in parsers:
|
||||
result = parser(stream, index).aggregate(result)
|
||||
if not result.status:
|
||||
return result
|
||||
index = result.index
|
||||
values.append(result.value)
|
||||
return Result.success(index, values).aggregate(result)
|
||||
|
||||
return seq_parser
|
||||
else:
|
||||
|
||||
@Parser
|
||||
def seq_kwarg_parser(stream: str | bytes | list, index: int) -> Result:
|
||||
result = None
|
||||
values = {}
|
||||
for name, parser in kw_parsers.items():
|
||||
result = parser(stream, index).aggregate(result)
|
||||
if not result.status:
|
||||
return result
|
||||
index = result.index
|
||||
values[name] = result.value
|
||||
return Result.success(index, values).aggregate(result)
|
||||
|
||||
return seq_kwarg_parser
|
||||
|
||||
|
||||
def generate(fn) -> Parser:
|
||||
"""
|
||||
Creates a parser from a generator function
|
||||
"""
|
||||
if isinstance(fn, str):
|
||||
return lambda f: generate(f).desc(fn)
|
||||
|
||||
@Parser
|
||||
@wraps(fn)
|
||||
def generated(stream: str | bytes | list, index: int) -> Result:
|
||||
# start up the generator
|
||||
iterator = fn()
|
||||
|
||||
result = None
|
||||
value = None
|
||||
try:
|
||||
while True:
|
||||
next_parser = iterator.send(value)
|
||||
result = next_parser(stream, index).aggregate(result)
|
||||
if not result.status:
|
||||
return result
|
||||
value = result.value
|
||||
index = result.index
|
||||
except StopIteration as stop:
|
||||
returnVal = stop.value
|
||||
if isinstance(returnVal, Parser):
|
||||
return returnVal(stream, index).aggregate(result)
|
||||
|
||||
return Result.success(index, returnVal).aggregate(result)
|
||||
|
||||
return generated
|
||||
|
||||
|
||||
index = Parser(lambda _, index: Result.success(index, index))
|
||||
line_info = Parser(lambda stream, index: Result.success(index, line_info_at(stream, index)))
|
||||
|
||||
|
||||
def success(value: Any) -> Parser:
|
||||
"""
|
||||
Returns a parser that does not consume any of the stream, but
|
||||
produces ``value``.
|
||||
"""
|
||||
return Parser(lambda _, index: Result.success(index, value))
|
||||
|
||||
|
||||
def fail(expected: str) -> Parser:
|
||||
"""
|
||||
Returns a parser that always fails with the provided error message.
|
||||
"""
|
||||
return Parser(lambda _, index: Result.failure(index, expected))
|
||||
|
||||
|
||||
def string(expected_string: str, transform: Callable[[str], str] = noop) -> Parser:
|
||||
"""
|
||||
Returns a parser that expects the ``expected_string`` and produces
|
||||
that string value.
|
||||
|
||||
Optionally, a transform function can be passed, which will be used on both
|
||||
the expected string and tested string.
|
||||
"""
|
||||
|
||||
slen = len(expected_string)
|
||||
transformed_s = transform(expected_string)
|
||||
|
||||
@Parser
|
||||
def string_parser(stream: str, index: int) -> Result:
|
||||
if transform(stream[index : index + slen]) == transformed_s:
|
||||
return Result.success(index + slen, expected_string)
|
||||
else:
|
||||
return Result.failure(index, expected_string)
|
||||
|
||||
return string_parser
|
||||
|
||||
|
||||
def regex(exp: str, flags=0, group: int | str | tuple = 0) -> Parser:
|
||||
"""
|
||||
Returns a parser that expects the given ``exp``, and produces the
|
||||
matched string. ``exp`` can be a compiled regular expression, or a
|
||||
string which will be compiled with the given ``flags``.
|
||||
|
||||
Optionally, accepts ``group``, which is passed to re.Match.group
|
||||
https://docs.python.org/3/library/re.html#re.Match.group> to
|
||||
return the text from a capturing group in the regex instead of the
|
||||
entire match.
|
||||
"""
|
||||
|
||||
if isinstance(exp, (str, bytes)):
|
||||
exp = re.compile(exp, flags)
|
||||
if isinstance(group, (str, int)):
|
||||
group = (group,)
|
||||
|
||||
@Parser
|
||||
def regex_parser(stream: str | bytes | list, index: int) -> Result:
|
||||
match = exp.match(stream, index)
|
||||
if match:
|
||||
return Result.success(match.end(), match.group(*group))
|
||||
else:
|
||||
return Result.failure(index, exp.pattern)
|
||||
|
||||
return regex_parser
|
||||
|
||||
|
||||
def test_item(func: Callable[..., bool], description: str) -> Parser:
|
||||
"""
|
||||
Returns a parser that tests a single item from the list of items being
|
||||
consumed, using the callable ``func``. If ``func`` returns ``True``, the
|
||||
parse succeeds, otherwise the parse fails with the description
|
||||
``description``.
|
||||
"""
|
||||
|
||||
@Parser
|
||||
def test_item_parser(stream: str | bytes | list, index: int) -> Result:
|
||||
if index < len(stream):
|
||||
if isinstance(stream, bytes):
|
||||
# Subscripting bytes with `[index]` instead of
|
||||
# `[index:index + 1]` returns an int
|
||||
item = stream[index : index + 1]
|
||||
else:
|
||||
item = stream[index]
|
||||
if func(item):
|
||||
return Result.success(index + 1, item)
|
||||
return Result.failure(index, description)
|
||||
|
||||
return test_item_parser
|
||||
|
||||
|
||||
def test_char(func: Callable[..., bool], description: str) -> Parser:
|
||||
"""
|
||||
Returns a parser that tests a single character with the callable
|
||||
``func``. If ``func`` returns ``True``, the parse succeeds, otherwise
|
||||
the parse fails with the description ``description``.
|
||||
"""
|
||||
# Implementation is identical to test_item
|
||||
return test_item(func, description)
|
||||
|
||||
|
||||
def match_item(item: Any, description: str = None) -> Parser:
|
||||
"""
|
||||
Returns a parser that tests the next item (or character) from the stream (or
|
||||
string) for equality against the provided item. Optionally a string
|
||||
description can be passed.
|
||||
"""
|
||||
|
||||
if description is None:
|
||||
description = str(item)
|
||||
return test_item(lambda i: item == i, description)
|
||||
|
||||
|
||||
def string_from(*strings: str, transform: Callable[[str], str] = noop):
|
||||
"""
|
||||
Accepts a sequence of strings as positional arguments, and returns a parser
|
||||
that matches and returns one string from the list. The list is first sorted
|
||||
in descending length order, so that overlapping strings are handled correctly
|
||||
by checking the longest one first.
|
||||
"""
|
||||
# Sort longest first, so that overlapping options work correctly
|
||||
return alt(*(string(s, transform) for s in sorted(strings, key=len, reverse=True)))
|
||||
|
||||
|
||||
def char_from(string: str | bytes) -> Parser:
|
||||
"""
|
||||
Accepts a string and returns a parser that matches and returns one character
|
||||
from the string.
|
||||
"""
|
||||
if isinstance(string, bytes):
|
||||
return test_char(lambda c: c in string, b"[" + string + b"]")
|
||||
else:
|
||||
return test_char(lambda c: c in string, "[" + string + "]")
|
||||
|
||||
|
||||
def peek(parser: Parser) -> Parser:
|
||||
"""
|
||||
Returns a lookahead parser that parses the input stream without consuming
|
||||
chars.
|
||||
"""
|
||||
|
||||
@Parser
|
||||
def peek_parser(stream: str | bytes | list, index: int) -> Result:
|
||||
result = parser(stream, index)
|
||||
if result.status:
|
||||
return Result.success(index, result.value)
|
||||
else:
|
||||
return result
|
||||
|
||||
return peek_parser
|
||||
|
||||
|
||||
any_char = test_char(lambda c: True, "any character")
|
||||
|
||||
whitespace = regex(r"\s+")
|
||||
|
||||
letter = test_char(lambda c: c.isalpha(), "a letter")
|
||||
|
||||
digit = test_char(lambda c: c.isdigit(), "a digit")
|
||||
|
||||
decimal_digit = char_from("0123456789")
|
||||
|
||||
|
||||
@Parser
|
||||
def eof(stream: str | bytes | list, index: int) -> Result:
|
||||
"""
|
||||
A parser that only succeeds if the end of the stream has been reached.
|
||||
"""
|
||||
|
||||
if index >= len(stream):
|
||||
return Result.success(index, None)
|
||||
else:
|
||||
return Result.failure(index, "EOF")
|
||||
|
||||
|
||||
def from_enum(enum_cls: type[enum.Enum], transform=noop) -> Parser:
|
||||
"""
|
||||
Given a class that is an enum.Enum class
|
||||
https://docs.python.org/3/library/enum.html , returns a parser that
|
||||
will parse the values (or the string representations of the values)
|
||||
and return the corresponding enum item.
|
||||
"""
|
||||
|
||||
items = sorted(
|
||||
((str(enum_item.value), enum_item) for enum_item in enum_cls), key=lambda t: len(t[0]), reverse=True
|
||||
)
|
||||
return alt(*(string(value, transform=transform).result(enum_item) for value, enum_item in items))
|
||||
|
||||
|
||||
class forward_declaration(Parser):
|
||||
"""
|
||||
An empty parser that can be used as a forward declaration,
|
||||
especially for parsers that need to be defined recursively.
|
||||
|
||||
You must use `.become(parser)` before using.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def _raise_error(self, *args, **kwargs):
|
||||
raise ValueError("You must use 'become' before attempting to call `parse` or `parse_partial`")
|
||||
|
||||
parse = _raise_error
|
||||
parse_partial = _raise_error
|
||||
|
||||
def become(self, other: Parser):
|
||||
"""
|
||||
Take on the behavior of the given parser.
|
||||
"""
|
||||
self.__dict__ = other.__dict__
|
||||
self.__class__ = other.__class__
|
||||
+324
-155
@@ -1,28 +1,31 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
import re
|
||||
import torch
|
||||
from collections import defaultdict
|
||||
from functools import partial
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from comfy_extras.nodes_mask import FeatherMask, MaskComposite
|
||||
from nodes import ConditioningAverage
|
||||
|
||||
from .utils import safe_float, get_function, parse_floats, smarter_split
|
||||
from .adv_encode import advanced_encode_from_tokens
|
||||
from .attention_couple_ppm import set_cond_attnmask
|
||||
from .cutoff import process_cuts
|
||||
from .parser import parse_cuts
|
||||
|
||||
try:
|
||||
from .nodes_attnmask import create_attention_hook
|
||||
from comfy.hooks import set_hooks_for_conditioning
|
||||
|
||||
def set_cond_attnmask(cond, mask):
|
||||
hook = create_attention_hook(mask)
|
||||
return set_hooks_for_conditioning(cond, hooks=hook)
|
||||
|
||||
except ImportError:
|
||||
|
||||
def set_cond_attnmask(cond, mask):
|
||||
log.info("Attention masking is not available")
|
||||
return cond
|
||||
|
||||
from .cutoff_parser import parse_cuts
|
||||
from .utils import (
|
||||
ComfyConditioning,
|
||||
FunctionSpec,
|
||||
call_node,
|
||||
get_function,
|
||||
parse_floats,
|
||||
safe_float,
|
||||
smarter_split,
|
||||
split_by_function,
|
||||
split_quotable,
|
||||
)
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
@@ -32,12 +35,12 @@ AVAILABLE_NORMALIZATIONS = ["none", "mean", "length", "length+mean"]
|
||||
SHUFFLE_GEN = torch.Generator(device="cpu")
|
||||
|
||||
|
||||
def get_sdxl(text, defaults):
|
||||
def get_sdxl(text: str, defaults: dict[str, Any]) -> tuple[str, dict[str, int]]:
|
||||
# Defaults fail to parse and get looked up from the defaults dict
|
||||
text, sdxl = get_function(text, "SDXL", ["none", "none", "none"])
|
||||
if not sdxl:
|
||||
return text, {}
|
||||
args = sdxl[0]
|
||||
args = sdxl[0].args
|
||||
d = defaults
|
||||
w, h = parse_floats(args[0], [d.get("sdxl_width", 1024), d.get("sdxl_height", 1024)], split_re="\\s+")
|
||||
tw, th = parse_floats(args[1], [d.get("sdxl_twidth", 1024), d.get("sdxl_theight", 1024)], split_re="\\s+")
|
||||
@@ -54,11 +57,11 @@ def get_sdxl(text, defaults):
|
||||
return text, opts
|
||||
|
||||
|
||||
def get_clipweights(text, existing_spec=None):
|
||||
def get_clipweights(text: str, existing_spec: dict[str, float] | None = None) -> tuple[dict[str, float], str]:
|
||||
text, spec = get_function(text, "TE_WEIGHT", defaults=None)
|
||||
if not spec:
|
||||
return existing_spec or {}, text
|
||||
args = spec[0].strip()
|
||||
args = spec[0].args[0].strip()
|
||||
res = {}
|
||||
for arg in args.split(","):
|
||||
try:
|
||||
@@ -70,26 +73,28 @@ def get_clipweights(text, existing_spec=None):
|
||||
return res, text
|
||||
|
||||
|
||||
def get_style(text, default_style="comfy", default_normalization="none"):
|
||||
def get_style(text: str, default_style="comfy", default_normalization="none") -> tuple[str, str, str]:
|
||||
text, styles = get_function(text, "STYLE", [default_style, default_normalization])
|
||||
if not styles:
|
||||
return default_style, default_normalization, text
|
||||
style, normalization = styles[0]
|
||||
style, normalization = styles[0].args
|
||||
style = style.strip()
|
||||
normalization = normalization.strip()
|
||||
if style not in AVAILABLE_STYLES:
|
||||
if style.replace("old+", "") not in AVAILABLE_STYLES:
|
||||
log.warning("Unrecognized prompt style: %s. Using %s", style, default_style)
|
||||
style = default_style
|
||||
|
||||
if normalization not in AVAILABLE_NORMALIZATIONS:
|
||||
log.warning("Unrecognized prompt normalization: %s. Using %s", normalization, default_normalization)
|
||||
normalization = default_normalization
|
||||
for part in normalization.split("+"):
|
||||
if part not in AVAILABLE_NORMALIZATIONS:
|
||||
log.warning("Unrecognized prompt normalization: %s. Using %s", normalization, default_normalization)
|
||||
normalization = default_normalization
|
||||
break
|
||||
|
||||
return style, normalization, text
|
||||
|
||||
|
||||
def shuffle_chunk(shuffle, c):
|
||||
func, shuffle = shuffle
|
||||
def shuffle_chunk(func_spec: FunctionSpec, c: str) -> str:
|
||||
func = func_spec.name
|
||||
shuffle = func_spec.args
|
||||
shuffle_count = int(safe_float(shuffle[0], 0))
|
||||
_, separator, joiner = shuffle
|
||||
if separator == "default":
|
||||
@@ -123,7 +128,8 @@ def shuffle_chunk(shuffle, c):
|
||||
|
||||
|
||||
def fix_word_ids(tokens):
|
||||
"""Fix word indexes. Tokenizing separately (when BREAKs exist) causes the indexes to restart which causes problems with some weighting algorithms that rely on them"""
|
||||
"""Fix word indexes. Tokenizing separately (when BREAKs exist) causes the indexes
|
||||
to restart which causes problems with some weighting algorithms that rely on them"""
|
||||
for key in tokens:
|
||||
max_idx = 0
|
||||
for group in range(len(tokens[key])):
|
||||
@@ -138,6 +144,82 @@ def fix_word_ids(tokens):
|
||||
return tokens
|
||||
|
||||
|
||||
def tokenize_chunks(clip, text, need_word_ids, can_break):
|
||||
chunks = list(split_quotable(text, r"\bBREAK\b"))
|
||||
token_chunks = []
|
||||
shuffled_chunks = []
|
||||
for c in chunks:
|
||||
c, shuffles = get_function(c.strip(), "(SHIFT|SHUFFLE)", ["0", "default", "default"])
|
||||
r = c
|
||||
for s in shuffles:
|
||||
r = shuffle_chunk(s, r)
|
||||
if r != c:
|
||||
log.info("Shuffled prompt chunk to %s", r)
|
||||
shuffled_chunks.append(r)
|
||||
t = clip.tokenize(c, return_word_ids=need_word_ids)
|
||||
token_chunks.append(t)
|
||||
|
||||
tokens = token_chunks[0]
|
||||
full_prompt = "".join(shuffled_chunks)
|
||||
full_tokenized = tokens
|
||||
if len(chunks) > 1:
|
||||
full_tokenized = clip.tokenize(full_prompt, return_word_ids=need_word_ids)
|
||||
for key in tokens:
|
||||
if not can_break.get(key):
|
||||
log.warning("BREAK does not make sense for %s, tokenizing as one chunk. Use CAT instead.", key)
|
||||
tokens[key] = full_tokenized[key]
|
||||
continue
|
||||
for c in token_chunks[1:]:
|
||||
tokens[key].extend(c[key])
|
||||
|
||||
return tokens
|
||||
|
||||
|
||||
def tokenize(clip, text, can_break, empty_tokens):
|
||||
# defaults=None means there is no argument parsing at all
|
||||
text, l_prompts = get_function(text, "CLIP_L", defaults=None)
|
||||
text, te_prompts = get_function(text, "TE", defaults=None)
|
||||
need_word_ids = True
|
||||
tokens = tokenize_chunks(clip, text, need_word_ids, can_break)
|
||||
|
||||
per_te_prompts = defaultdict(list)
|
||||
if l_prompts:
|
||||
log.warning("Note: CLIP_L is deprecated. Use TE(l=prompt) instead")
|
||||
per_te_prompts["l"] = [x.args for x in l_prompts]
|
||||
|
||||
for prompt in te_prompts:
|
||||
prompt = prompt.args[0]
|
||||
if prompt.strip() == "help":
|
||||
log.info("Encoders available for TE: %s", ", ".join(tokens.keys()))
|
||||
continue
|
||||
params = prompt.split("=", 1)
|
||||
if len(params) != 2:
|
||||
log.warning("Invalid TE call, ignoring: %s", prompt)
|
||||
continue
|
||||
te = params[0].strip()
|
||||
prompt = params[1].strip()
|
||||
if te not in tokens:
|
||||
log.warning("Invalid TE call, no TE with key '%s', ignoring: %s", te)
|
||||
log.info("Encoders available for TE: %s", ", ".join(tokens.keys()))
|
||||
continue
|
||||
per_te_prompts[te].append(prompt)
|
||||
|
||||
if per_te_prompts:
|
||||
for key in per_te_prompts:
|
||||
prompt = " ".join(per_te_prompts[key])
|
||||
tokens[key] = tokenize_chunks(clip, prompt, need_word_ids, can_break)[key]
|
||||
log.info("Encoded prompt with TE '%s': %s", key, prompt)
|
||||
|
||||
maxlen = max([0] + [len(tokens[k]) for k in tokens if can_break[k]])
|
||||
for k in tokens:
|
||||
if not can_break[k]:
|
||||
continue
|
||||
while len(tokens[k]) < maxlen:
|
||||
tokens[k] += empty_tokens[k]
|
||||
|
||||
return fix_word_ids(tokens)
|
||||
|
||||
|
||||
def encode_prompt_segment(
|
||||
clip,
|
||||
text,
|
||||
@@ -145,7 +227,7 @@ def encode_prompt_segment(
|
||||
default_style="comfy",
|
||||
default_normalization="none",
|
||||
clip_weights=None,
|
||||
) -> list[tuple[torch.Tensor, dict[str]]]:
|
||||
) -> list[ComfyConditioning]:
|
||||
style, normalization, text = get_style(text, default_style, default_normalization)
|
||||
clip_weights, text = get_clipweights(text, clip_weights)
|
||||
text, cuts = parse_cuts(text)
|
||||
@@ -155,52 +237,62 @@ def encode_prompt_segment(
|
||||
if cuts:
|
||||
extra["cuts"] = cuts
|
||||
|
||||
# defaults=None means there is no argument parsing at all
|
||||
text, l_prompts = get_function(text, "CLIP_L", defaults=None)
|
||||
chunks = re.split(r"\bBREAK\b", text)
|
||||
token_chunks = []
|
||||
need_word_ids = True
|
||||
for c in chunks:
|
||||
c, shuffles = get_function(c.strip(), "(SHIFT|SHUFFLE)", ["0", "default", "default"], return_func_name=True)
|
||||
r = c
|
||||
for s in shuffles:
|
||||
r = shuffle_chunk(s, r)
|
||||
if r != c:
|
||||
log.info("Shuffled prompt chunk to %s", r)
|
||||
c = r
|
||||
t = clip.tokenize(c, return_word_ids=need_word_ids)
|
||||
token_chunks.append(t)
|
||||
tokens = token_chunks[0]
|
||||
empty = clip.tokenize("", return_word_ids=True)
|
||||
can_break = {}
|
||||
for k in empty:
|
||||
tokenizer = getattr(clip.tokenizer, f"clip_{k}", getattr(clip.tokenizer, k, None))
|
||||
can_break[k] = tokenizer and getattr(tokenizer, "pad_to_max_length", False)
|
||||
|
||||
for key in tokens:
|
||||
for c in token_chunks[1:]:
|
||||
tokens[key].extend(c[key])
|
||||
clip = hook_te(clip, empty.keys(), style, normalization, extra)
|
||||
|
||||
# Non-SDXL has only "l"
|
||||
if "g" in tokens and l_prompts:
|
||||
text_l = " ".join(l_prompts)
|
||||
log.info("Encoded SDXL CLIP_L prompt: %s", text_l)
|
||||
tokens["l"] = clip.tokenize(text_l, return_word_ids=need_word_ids)["l"]
|
||||
# Chunks to ConditioningAverage:
|
||||
|
||||
if "g" in tokens and "l" in tokens and len(tokens["l"]) != len(tokens["g"]):
|
||||
empty = clip.tokenize("", return_word_ids=need_word_ids)
|
||||
while len(tokens["l"]) < len(tokens["g"]):
|
||||
tokens["l"] += empty["l"]
|
||||
while len(tokens["l"]) > len(tokens["g"]):
|
||||
tokens["g"] += empty["g"]
|
||||
text, averages = split_by_function(text, "AVG", ["0.5"], require_args=False)
|
||||
prompts_to_avg = []
|
||||
for chunk, avg in averages:
|
||||
w = safe_float(avg.args[0], 0.5)
|
||||
prompts_to_avg.append((text, w))
|
||||
text = chunk
|
||||
prompts_to_avg.append((text, 1.0))
|
||||
|
||||
tokens = fix_word_ids(tokens)
|
||||
conds_to_avg = []
|
||||
for prompt, weight in prompts_to_avg:
|
||||
conds_to_cat = []
|
||||
for c in split_quotable(prompt, r"\bCAT\b"):
|
||||
tokens = tokenize(clip, c, can_break, empty)
|
||||
conds_to_cat.append(clip.encode_from_tokens_scheduled(tokens, add_dict=settings))
|
||||
|
||||
tes = []
|
||||
for k in tokens:
|
||||
if k in ["g", "l"]:
|
||||
tes.append(f"clip_{k}")
|
||||
else:
|
||||
tes.append(k)
|
||||
base = conds_to_cat[0]
|
||||
for cond in conds_to_cat[1:]:
|
||||
assert len(cond) == len(base), "Conditioning length mismatch"
|
||||
# Pooled gets ignored
|
||||
for i in range(len(base)):
|
||||
c1 = base[i][0]
|
||||
c2 = cond[i][0]
|
||||
base[i][0] = torch.cat((c1, c2), 1)
|
||||
conds_to_avg.append((base, weight))
|
||||
|
||||
clip = hook_te(clip, tes, style, normalization, extra)
|
||||
base, w = conds_to_avg[0]
|
||||
for cond, next_w in conds_to_avg[1:]:
|
||||
assert len(base) == len(cond), "Conditioning length mismatch"
|
||||
if w == 1.0:
|
||||
w = next_w
|
||||
continue
|
||||
for i in range(len(base)):
|
||||
(cond,) = call_node(ConditioningAverage, [base[i]], [cond[i]], w)
|
||||
base[i] = cond[0]
|
||||
w = next_w
|
||||
|
||||
return clip.encode_from_tokens_scheduled(tokens, add_dict=settings)
|
||||
return base
|
||||
|
||||
|
||||
def calc_w(tensor, w):
|
||||
if math.isclose(w, 0):
|
||||
return torch.zeros_like(tensor)
|
||||
elif math.isclose(w, 1.0):
|
||||
return tensor
|
||||
else:
|
||||
return tensor * w
|
||||
|
||||
|
||||
def apply_weights(output, te_name, spec):
|
||||
@@ -211,21 +303,29 @@ def apply_weights(output, te_name, spec):
|
||||
if te_name.startswith("clip_"):
|
||||
te_name = te_name[5:]
|
||||
|
||||
if isinstance(output, tuple):
|
||||
out, pooled = output
|
||||
if te_name in spec:
|
||||
log.info("Weighting %s output by %s", te_name, spec[te_name])
|
||||
out = out * spec[te_name]
|
||||
pkey = te_name + "_pooled"
|
||||
if pkey in spec:
|
||||
log.info("Weighting %s pooled output by %s", te_name, spec[pkey])
|
||||
pooled = pooled * spec[pkey]
|
||||
default = spec.get("all", None)
|
||||
|
||||
return out, pooled
|
||||
if isinstance(output, tuple):
|
||||
out, pooled, *extra = output
|
||||
pkey = te_name + "_pooled"
|
||||
if te_name in spec or pkey in spec or default is not None:
|
||||
w = spec.get(te_name, default)
|
||||
pooled_w = spec.get(pkey, w)
|
||||
if w is None:
|
||||
w = 1.0
|
||||
if pooled_w is None:
|
||||
pooled_w = 1.0
|
||||
log.info("Weighting %s output by %s, pooled by %s", te_name, w, pooled_w)
|
||||
out = calc_w(out, w)
|
||||
if pooled is not None:
|
||||
pooled = calc_w(pooled, pooled_w)
|
||||
|
||||
return (out, pooled) + tuple(extra)
|
||||
else:
|
||||
if te_name in spec:
|
||||
log.info("Weighting %s output by %s", te_name, spec[te_name])
|
||||
output = output * spec[te_name]
|
||||
if te_name in spec or default is not None:
|
||||
w = spec.get(te_name, default)
|
||||
log.info("Weighting %s output by %s", te_name, w)
|
||||
output = calc_w(output, w)
|
||||
return output
|
||||
|
||||
|
||||
@@ -246,15 +346,24 @@ def hook_te(clip, te_names, style, normalization, extra):
|
||||
return clip
|
||||
newclip = clip.clone()
|
||||
for te_name in te_names:
|
||||
if hasattr(clip.patcher.model, te_name):
|
||||
tokenizer = getattr(clip.tokenizer, f"clip_{te_name}", getattr(clip.tokenizer, te_name, None))
|
||||
if tokenizer:
|
||||
x = extra.copy()
|
||||
x["tokenizer"] = getattr(clip.tokenizer, te_name)
|
||||
log.debug("Hooked into %s with style=%s, normalization=%s", te_name, style, normalization)
|
||||
x["tokenizer"] = tokenizer
|
||||
if not hasattr(clip.patcher.model, te_name):
|
||||
te_name = "clip_" + te_name
|
||||
if not hasattr(clip.patcher.model, te_name):
|
||||
log.warning("TE model %s not found on model patcher. Skipping...", te_name)
|
||||
continue
|
||||
|
||||
log.debug("Hooked into te=%s with style=%s, normalization=%s", te_name, style, normalization)
|
||||
encode = clip.patcher.get_model_object(f"{te_name}.encode_token_weights")
|
||||
x["has_negpip"] = clip.patcher.model_options.get("ppm_negpip", False)
|
||||
newclip.patcher.add_object_patch(
|
||||
f"{te_name}.encode_token_weights",
|
||||
make_patch(
|
||||
te_name,
|
||||
clip.patcher.get_model_object(f"{te_name}.encode_token_weights"),
|
||||
encode,
|
||||
normalization,
|
||||
style,
|
||||
x,
|
||||
@@ -262,7 +371,7 @@ def hook_te(clip, te_names, style, normalization, extra):
|
||||
)
|
||||
# 'g' and 'l' exist in these are clip_g and clip_l
|
||||
else:
|
||||
log.debug("Tokens contain items with key %s but no TE found on object with that name.", te_name)
|
||||
log.warning("Tokens contain items with key %s but no tokenizer found on object with that name.", te_name)
|
||||
return newclip
|
||||
|
||||
|
||||
@@ -271,7 +380,7 @@ def get_area(text):
|
||||
if not areas:
|
||||
return text, None
|
||||
|
||||
args = areas[0]
|
||||
args = areas[0].args
|
||||
x, w = parse_floats(args[0], [0.0, 1.0], split_re="\\s+")
|
||||
y, h = parse_floats(args[1], [0.0, 1.0], split_re="\\s+")
|
||||
weight = safe_float(args[2], 1.0)
|
||||
@@ -288,7 +397,8 @@ def get_area(text):
|
||||
area = (int(h) // 8, int(w) // 8, int(y) // 8, int(x) // 8)
|
||||
else:
|
||||
raise Exception(
|
||||
f"AREA specified with invalid size {x} {w}, {h} {y}. They must either all be percentages between 0 and 1 or positive integer pixel values excluding 1"
|
||||
f"AREA specified with invalid size {x} {w}, {h} {y}. They must either all"
|
||||
" be percentages between 0 and 1 or positive integer pixel values excluding 1"
|
||||
)
|
||||
|
||||
return text, (area, weight)
|
||||
@@ -298,7 +408,7 @@ def get_mask_size(text, defaults):
|
||||
text, sizes = get_function(text, "MASK_SIZE", ["512", "512"])
|
||||
if not sizes:
|
||||
return text, (defaults.get("mask_width", 512), defaults.get("mask_height", 512))
|
||||
w, h = sizes[0]
|
||||
w, h = sizes[0].args
|
||||
return text, (int(w), int(h))
|
||||
|
||||
|
||||
@@ -322,13 +432,14 @@ def make_mask(args, size, weight):
|
||||
ys = int(y1), int(y2)
|
||||
else:
|
||||
raise Exception(
|
||||
f"MASK specified with invalid size {x1} {x2}, {y1} {y2}. They must either all be percentages between 0 and 1 or positive integer pixel values excluding 1"
|
||||
f"MASK specified with invalid size {x1} {x2}, {y1} {y2}. They must either all"
|
||||
" be percentages between 0 and 1 or positive integer pixel values excluding 1"
|
||||
)
|
||||
|
||||
mask = torch.full((h, w), 0, dtype=torch.float32, device="cpu")
|
||||
mask[ys[0] : ys[1], xs[0] : xs[1]] = weight
|
||||
mask = mask.unsqueeze(0)
|
||||
log.info("Mask xs=%s, ys=%s, shape=%s, weight=%s", xs, ys, mask.shape, weight)
|
||||
log.debug("Mask xs=%s, ys=%s, shape=%s, weight=%s", xs, ys, mask.shape, weight)
|
||||
return mask
|
||||
|
||||
|
||||
@@ -343,47 +454,51 @@ def get_mask(text, size, input_masks):
|
||||
return text, None, None
|
||||
|
||||
def feather(f, mask):
|
||||
l, t, r, b, *_ = [int(x) for x in parse_floats(f[0], [0, 0, 0, 0], split_re="\\s+")]
|
||||
mask = FeatherMask().feather(mask, l, t, r, b)[0]
|
||||
log.info("FeatherMask l=%s, t=%s, r=%s, b=%s", l, t, r, b)
|
||||
left, top, right, bottom, *_ = [int(x) for x in parse_floats(f[0], [0, 0, 0, 0], split_re="\\s+")]
|
||||
mask = call_node(FeatherMask, mask, left, top, right, bottom)[0]
|
||||
log.info("FeatherMask l=%s, t=%s, r=%s, b=%s", left, top, right, bottom)
|
||||
return mask
|
||||
|
||||
mask = None
|
||||
totalweight = 1.0
|
||||
if maskw:
|
||||
totalweight = safe_float(maskw[0][0], 1.0)
|
||||
totalweight = safe_float(maskw[0].args[0], 1.0)
|
||||
i = 0
|
||||
for m in masks:
|
||||
weight = safe_float(m[2], 1.0)
|
||||
op = m[3]
|
||||
nextmask = make_mask(m, size, weight)
|
||||
weight = safe_float(m.args[2], 1.0)
|
||||
op = m.args[3]
|
||||
nextmask = make_mask(m.args, size, weight)
|
||||
if i < len(feathers):
|
||||
nextmask = feather(feathers[i], nextmask)
|
||||
nextmask = feather(feathers[i].args, nextmask)
|
||||
i += 1
|
||||
if mask is not None:
|
||||
log.info("MaskComposite op=%s", op)
|
||||
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
|
||||
mask = call_node(MaskComposite, mask, nextmask, 0, 0, op)[0]
|
||||
else:
|
||||
mask = nextmask
|
||||
|
||||
for idx, w, op in imasks:
|
||||
for im in imasks:
|
||||
idx, w, op = im.args
|
||||
idx = int(safe_float(idx, 0.0))
|
||||
w = safe_float(w, 1.0)
|
||||
if input_masks is None:
|
||||
log.warning(
|
||||
"IMASK requires you to attach custom masks to the CLIP object using PCAddMasksToClIP before using it"
|
||||
)
|
||||
input_masks = []
|
||||
|
||||
if len(input_masks) < idx + 1:
|
||||
log.warn("IMASK index %s not found, ignoring...", idx)
|
||||
log.warning("IMASK index %s not found, ignoring...", idx)
|
||||
continue
|
||||
nextmask = input_masks[idx] * w
|
||||
if i < len(feathers):
|
||||
nextmask = feather(feathers[i], nextmask)
|
||||
nextmask = feather(feathers[i].args, nextmask)
|
||||
i += 1
|
||||
if mask is not None:
|
||||
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
|
||||
else:
|
||||
mask = nextmask
|
||||
mask = call_node(MaskComposite, mask, nextmask, 0, 0, op)[0] if mask is not None else nextmask
|
||||
|
||||
# apply leftover FEATHER() specs to the whole
|
||||
for f in feathers[i:]:
|
||||
mask = feather(f, mask)
|
||||
mask = feather(f.args, mask)
|
||||
|
||||
return text, mask, totalweight
|
||||
|
||||
@@ -398,14 +513,15 @@ def get_noise(text):
|
||||
return text, None, None
|
||||
w = 0
|
||||
# Only take seed from first noise spec, for simplicity
|
||||
seed = safe_float(noises[0][1], "none")
|
||||
seed = noises[0].args[0].strip()
|
||||
if seed == "none":
|
||||
gen = None
|
||||
else:
|
||||
seed = safe_float(seed, 0)
|
||||
gen = torch.Generator()
|
||||
gen.manual_seed(int(seed))
|
||||
for n in noises:
|
||||
w += safe_float(n[0], 0.0)
|
||||
w += safe_float(n.args[0], 0.0)
|
||||
return text, max(min(w, 1.0), 0.0), gen
|
||||
|
||||
|
||||
@@ -418,21 +534,15 @@ def apply_noise(cond, weight, gen):
|
||||
return cond * (1 - weight) + n * weight
|
||||
|
||||
|
||||
def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
|
||||
# First style modifier applies to ANDed prompts too unless overridden
|
||||
style, normalization, text = get_style(text)
|
||||
text, mask_size = get_mask_size(text, defaults)
|
||||
|
||||
prompts = [p.strip() for p in re.split(r"\bAND\b", text)]
|
||||
|
||||
p, sdxl_opts = get_sdxl(prompts[0], defaults)
|
||||
prompts[0] = p
|
||||
def process_settings(prompt, defaults, masks, mask_size, sdxl_opts):
|
||||
if "ATTN()" in prompt:
|
||||
raise ValueError("ATTN() no longer works and has been replaced by COUPLE()")
|
||||
|
||||
def weight(t):
|
||||
opts = {}
|
||||
m = re.search(r":(-?\d\.?\d*)(![A-Za-z]+)?$", t)
|
||||
m = re.search(r":(-?\d\.?\d*)(![A-Za-z]+)?$", t.strip())
|
||||
if not m:
|
||||
return (1.0, opts, t)
|
||||
return (None, opts, t)
|
||||
w = float(m[1])
|
||||
tag = m[2]
|
||||
t = t[: m.span()[0]]
|
||||
@@ -441,42 +551,101 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
|
||||
|
||||
return w, opts, t
|
||||
|
||||
settings = {"prompt": prompt}
|
||||
|
||||
if "FILL()" in prompt:
|
||||
prompt = prompt.replace("FILL()", "")
|
||||
settings["x-promptcontrol.fill"] = True
|
||||
prompt, mask, mask_weight = get_mask(prompt, mask_size, masks)
|
||||
prompt, area = get_area(prompt)
|
||||
prompt, local_sdxl_opts = get_sdxl(prompt, defaults)
|
||||
# Get weight last so other syntax doesn't interfere with it
|
||||
w, opts, prompt = weight(prompt)
|
||||
if w is not None:
|
||||
settings["strength"] = w
|
||||
settings.update(sdxl_opts)
|
||||
settings.update(local_sdxl_opts)
|
||||
if area:
|
||||
settings["area"] = area[0]
|
||||
settings["strength"] = area[1]
|
||||
settings["set_area_to_bounds"] = False
|
||||
if mask is not None:
|
||||
settings["mask"] = mask
|
||||
settings["mask_strength"] = mask_weight
|
||||
|
||||
return prompt, settings
|
||||
|
||||
|
||||
def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
|
||||
# First style modifier applies to ANDed prompts too unless overridden
|
||||
style, normalization, text = get_style(text)
|
||||
text, mask_size = get_mask_size(text, defaults)
|
||||
|
||||
prompts = list(split_quotable(text, r"\bAND\b"))
|
||||
|
||||
p, sdxl_opts = get_sdxl(prompts[0], defaults)
|
||||
prompts[0] = p
|
||||
|
||||
conds = []
|
||||
# TODO: is this still needed?
|
||||
# scale = sum(abs(weight(p)[0]) for p in prompts if not ("AREA(" in p or "MASK(" in p))
|
||||
|
||||
def ensure_mask(c):
|
||||
if "mask" not in c[1]:
|
||||
_, mask, _ = get_mask("MASK()", mask_size, masks)
|
||||
c[1]["mask"] = mask
|
||||
c[1]["mask_strength"] = 1.0
|
||||
return c
|
||||
|
||||
def couple_mask(args):
|
||||
assert len(args) <= 1, "Argument parsing failure. This is a bug in Prompt Control"
|
||||
if not args:
|
||||
return ""
|
||||
return f"MASK({args[0]})"
|
||||
|
||||
for prompt in prompts:
|
||||
attn = False
|
||||
if "ATTN()" in prompt:
|
||||
prompt = prompt.replace("ATTN()", "")
|
||||
attn = True
|
||||
log.info("Using attention masking for prompt segment")
|
||||
prompt, mask, mask_weight = get_mask(prompt, mask_size, masks)
|
||||
w, opts, prompt = weight(prompt)
|
||||
text, noise_w, generator = get_noise(text)
|
||||
if not w:
|
||||
continue
|
||||
prompt, area = get_area(prompt)
|
||||
prompt, local_sdxl_opts = get_sdxl(prompt, defaults)
|
||||
settings = {"prompt": prompt}
|
||||
settings["strength"] = w
|
||||
settings.update(sdxl_opts)
|
||||
settings.update(local_sdxl_opts)
|
||||
if area:
|
||||
settings["area"] = area[0]
|
||||
settings["strength"] = area[1]
|
||||
settings["set_area_to_bounds"] = False
|
||||
if mask is not None:
|
||||
settings["mask"] = mask
|
||||
settings["mask_strength"] = mask_weight
|
||||
base_prompt, attn_couple_prompts = split_by_function(prompt, "COUPLE", defaults=None, require_args=False)
|
||||
|
||||
settings["start_percent"] = start_pct
|
||||
settings["end_percent"] = end_pct
|
||||
x = encode_prompt_segment(clip, prompt, settings, style, normalization)
|
||||
if attn and mask is not None:
|
||||
mask = settings.pop("mask")
|
||||
strength = settings.pop("mask_strength")
|
||||
x = set_cond_attnmask(x, mask * strength)
|
||||
prompts = [base_prompt] + [couple_mask(f.args) + chunk for (chunk, f) in attn_couple_prompts]
|
||||
encoded = []
|
||||
for p in prompts:
|
||||
p, settings = process_settings(p, defaults, masks, mask_size, sdxl_opts)
|
||||
if settings.get("strength") == 0: # weight is explicitly set to 0, skip
|
||||
continue
|
||||
settings["start_percent"] = start_pct
|
||||
settings["end_percent"] = end_pct
|
||||
x = encode_prompt_segment(clip, p, settings, style, normalization)
|
||||
encoded.append(x)
|
||||
|
||||
conds.extend(x)
|
||||
assert all(len(c) == len(encoded[0]) for c in encoded), (
|
||||
"All encoded prompts didn't produce the same number of conds, I don't know what to do in this situation."
|
||||
)
|
||||
|
||||
# each call to encode_prompt_segment can produce a number of conds based on any
|
||||
# scheduled LoRA hooks on the clip model. Zip them together with coupled prompts
|
||||
base_cond = []
|
||||
for base_cond, *attention_couple in zip(*encoded, strict=False):
|
||||
s = base_cond[1]
|
||||
# If there are LoRAs on the CLIP, we need to fix start_percent and
|
||||
# end_percent on the new conds for things to work properly.
|
||||
s["start_percent"] = s.get("clip_start_percent", s["start_percent"])
|
||||
s["end_percent"] = s.get("clip_end_percent", s["end_percent"])
|
||||
s.pop("clip_start_percent", None)
|
||||
s.pop("clip_end_percent", None)
|
||||
base_cond = [base_cond]
|
||||
if attention_couple:
|
||||
fill = base_cond[0][1].get("x-promptcontrol.fill")
|
||||
if not fill:
|
||||
ensure_mask(base_cond[0])
|
||||
# else, set_cond_attnmask will have the base mask fill any unspecified areas
|
||||
base_cond = set_cond_attnmask(
|
||||
base_cond,
|
||||
[ensure_mask(c) for c in attention_couple],
|
||||
fill=fill,
|
||||
)
|
||||
|
||||
base_cond = [[apply_noise(c[0], noise_w, generator), c[1]] for c in base_cond]
|
||||
conds.extend(base_cond)
|
||||
|
||||
return conds
|
||||
|
||||
+183
-29
@@ -1,19 +1,57 @@
|
||||
from pathlib import Path
|
||||
import re
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import logging
|
||||
import re
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, TypeAlias, TypeVar
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import torch # flakes8: noqa
|
||||
|
||||
FunctionArgs: TypeAlias = list[str]
|
||||
ComfyConditioning: TypeAlias = tuple["torch.Tensor", dict[str, Any]]
|
||||
|
||||
|
||||
@dataclass
|
||||
class FunctionSpec:
|
||||
name: str
|
||||
args: FunctionArgs
|
||||
position: int
|
||||
placeholder: str | None
|
||||
|
||||
|
||||
# Allow testing
|
||||
try:
|
||||
from folder_paths import get_filename_list
|
||||
except ImportError:
|
||||
|
||||
def get_filename_list(x):
|
||||
raise NotImplementedError("How did you get here?")
|
||||
def get_filename_list(folder_name) -> list[str]:
|
||||
return []
|
||||
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def flatten(x):
|
||||
if type(x) in [str, tuple, int, type(None)] or isinstance(x, dict) and "type" in x:
|
||||
yield x
|
||||
else:
|
||||
for g in x:
|
||||
yield from flatten(g)
|
||||
|
||||
|
||||
def call_node(cls, *args, **kwargs):
|
||||
if hasattr(cls, "execute"):
|
||||
# v3 node
|
||||
return cls.execute(*args, **kwargs)
|
||||
else:
|
||||
func = getattr(cls(), cls.FUNCTION)
|
||||
return func(*args, **kwargs)
|
||||
|
||||
|
||||
def consolidate_schedule(prompt_schedule):
|
||||
prev_loras = {}
|
||||
not_found = []
|
||||
@@ -54,10 +92,11 @@ def find_nonscheduled_loras(consolidated_schedule):
|
||||
return {k: v for (k, v) in candidate_loras.items() if k not in to_remove}
|
||||
|
||||
|
||||
def smarter_split(separator, string):
|
||||
def smarter_split(separator: str, string: str) -> list[str]:
|
||||
"""Does not break () when splitting"""
|
||||
splits = []
|
||||
prev = 0
|
||||
idx = 0
|
||||
stack = 0
|
||||
escape = False
|
||||
for idx, x in enumerate(string):
|
||||
@@ -74,7 +113,7 @@ def smarter_split(separator, string):
|
||||
return splits
|
||||
|
||||
|
||||
def find_closing_paren(text, start):
|
||||
def find_closing_paren(text: str, start: int) -> int:
|
||||
stack = 1
|
||||
for i, char in enumerate(text[start:]):
|
||||
if char == ")":
|
||||
@@ -83,51 +122,125 @@ def find_closing_paren(text, start):
|
||||
stack += 1
|
||||
if stack == 0:
|
||||
return start + i
|
||||
# Implicit closing paren after end
|
||||
return len(text)
|
||||
return -1
|
||||
|
||||
|
||||
def get_function(text, func, defaults, return_func_name=False):
|
||||
rex = re.compile(rf"\b{func}\(", re.MULTILINE)
|
||||
instances = []
|
||||
def find_function_spans(
|
||||
text: str, func: str, require_args: bool, defaults: FunctionArgs | None
|
||||
) -> Iterator[tuple[int, int, str, FunctionArgs]]:
|
||||
e = r"\(" if require_args else r"\b"
|
||||
rex = re.compile(rf"\b{func}{e}", re.MULTILINE)
|
||||
|
||||
idx = 0
|
||||
match = rex.search(text)
|
||||
while match:
|
||||
# Match start, content start
|
||||
start, after_first_paren = match.span()
|
||||
funcname = text[start : after_first_paren - 1]
|
||||
end = find_closing_paren(text, after_first_paren)
|
||||
args = parse_strings(text[after_first_paren:end], defaults)
|
||||
if return_func_name:
|
||||
instances.append((funcname, args))
|
||||
start, at_paren = match.span()
|
||||
if require_args:
|
||||
at_paren = at_paren - 1
|
||||
funcname = text[start:at_paren]
|
||||
after_first_paren = at_paren + 1
|
||||
if text[at_paren:after_first_paren] == "(":
|
||||
end = find_closing_paren(text, after_first_paren)
|
||||
if end < 0:
|
||||
continue
|
||||
args = parse_strings(text[after_first_paren:end], defaults)
|
||||
end += 1
|
||||
else:
|
||||
instances.append(args)
|
||||
|
||||
text = text[:start] + text[end + 1 :]
|
||||
end = at_paren
|
||||
args = defaults or []
|
||||
yield idx + start, idx + end, funcname, args
|
||||
idx = idx + end
|
||||
text = text[end:]
|
||||
match = rex.search(text)
|
||||
|
||||
|
||||
def get_function(
|
||||
text: str, func: str, defaults: list[str] | None, placeholder: str = "", require_args: bool = True
|
||||
) -> tuple[str, list[FunctionSpec]]:
|
||||
spans = [x.span() for x in re.finditer(r'".+?"', text)]
|
||||
instances = []
|
||||
count = 0
|
||||
chunks = []
|
||||
current = 0
|
||||
skipped = 0
|
||||
for start, end, funcname, args in find_function_spans(text, func, require_args, defaults):
|
||||
ph = None
|
||||
if spans_include(spans, start, end):
|
||||
continue
|
||||
if placeholder:
|
||||
ph = f"\0{placeholder}{count}\0"
|
||||
instances.append(FunctionSpec(funcname, args, start - skipped, ph))
|
||||
skipped += end - start
|
||||
chunks.append(text[current:start] + (ph or ""))
|
||||
current = end
|
||||
count += 1
|
||||
chunks.append(text[current:])
|
||||
text = "".join(chunks)
|
||||
return text, instances
|
||||
|
||||
|
||||
def parse_args(strings, arg_spec, strip=True):
|
||||
def spans_include(spans: list[tuple[int, int]], s: int, e: int) -> bool:
|
||||
return any((s > a and e < b) for a, b in spans)
|
||||
|
||||
|
||||
def split_quotable(text: str, regexp: str) -> Iterator[str]:
|
||||
start_from = 0
|
||||
spans = [x.span() for x in re.finditer(r'".+?"', text)]
|
||||
for x in re.finditer(regexp, text):
|
||||
s, e = x.span()
|
||||
if not spans_include(spans, s, e):
|
||||
yield text[start_from:s].strip()
|
||||
start_from = e
|
||||
yield text[start_from:].strip()
|
||||
|
||||
|
||||
def split_by_function(
|
||||
text: str, func: str, defaults: list[str] | None = None, require_args: bool = True
|
||||
) -> tuple[str, list[tuple[str, FunctionSpec]]]:
|
||||
"""
|
||||
Splits a string by function calls, returning the leftover text
|
||||
along with a list of functions with their associated text chunk.
|
||||
"""
|
||||
text, functions = get_function(text, func, defaults, require_args=require_args)
|
||||
chunks = []
|
||||
prev = 0
|
||||
for f in functions:
|
||||
chunks.append(text[prev : f.position])
|
||||
prev = f.position
|
||||
chunks.append(text[prev:])
|
||||
r = []
|
||||
for i, f in enumerate(functions):
|
||||
r.append((chunks[i + 1], f))
|
||||
return chunks[0], r
|
||||
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def parse_args(strings: list[str], arg_spec: list[tuple[Any, T]], strip: bool = True) -> list[T]:
|
||||
args = [s[1] for s in arg_spec]
|
||||
for i, spec in list(enumerate(arg_spec))[: len(strings)]:
|
||||
try:
|
||||
if strip:
|
||||
strings[i] = strings[i].strip()
|
||||
args[i] = spec[0](strings[i])
|
||||
f = spec[0]
|
||||
args[i] = f(strings[i])
|
||||
except ValueError:
|
||||
pass
|
||||
return args
|
||||
|
||||
|
||||
def parse_floats(string, defaults, split_re=","):
|
||||
def parse_floats(string: str, defaults: list[float], split_re: str = ",") -> list[float]:
|
||||
spec = [(float, d) for d in defaults]
|
||||
return parse_args(re.split(split_re, string.strip()), spec)
|
||||
|
||||
|
||||
def parse_strings(string, defaults, split_re=r"(?<!\\),", replace=(r"\,", ",")):
|
||||
def parse_strings(
|
||||
string: str, defaults: FunctionArgs | None, split_re: str = r"(?<!\\),", replace: tuple[str, str] = (r"\,", ",")
|
||||
) -> FunctionArgs:
|
||||
if defaults is None:
|
||||
return string
|
||||
spec = [(lambda x: x, d) for d in defaults]
|
||||
return [string]
|
||||
spec = [(str, d) for d in defaults]
|
||||
splits = re.split(split_re, string)
|
||||
if replace:
|
||||
f, t = replace
|
||||
@@ -135,7 +248,7 @@ def parse_strings(string, defaults, split_re=r"(?<!\\),", replace=(r"\,", ",")):
|
||||
return parse_args(splits, spec, strip=False)
|
||||
|
||||
|
||||
def safe_float(f, default):
|
||||
def safe_float(f: Any, default: float) -> float:
|
||||
if f is None:
|
||||
return default
|
||||
try:
|
||||
@@ -144,7 +257,7 @@ def safe_float(f, default):
|
||||
return default
|
||||
|
||||
|
||||
def lora_name_to_file(name):
|
||||
def lora_name_to_file(name: str) -> str | None:
|
||||
filenames = get_filename_list("loras")
|
||||
# Return exact matches as is
|
||||
if name in filenames:
|
||||
@@ -155,4 +268,45 @@ def lora_name_to_file(name):
|
||||
p = Path(f).with_suffix("")
|
||||
if p.name == n or str(p) == n:
|
||||
return f
|
||||
# Finally, try to find unique match from parts
|
||||
parts = name.split()
|
||||
search = [f for f in filenames if all(p in f for p in parts)]
|
||||
if len(search) == 1:
|
||||
return search[0]
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def map_inputs(input_map, inputs):
|
||||
new_inputs = {}
|
||||
for k in inputs:
|
||||
key = inputs[k]
|
||||
new_inputs[k] = key
|
||||
if isinstance(key, list):
|
||||
key = tuple(key)
|
||||
x = input_map.get(key, inputs[k])
|
||||
new_inputs[k] = x
|
||||
return new_inputs
|
||||
|
||||
|
||||
def expand_graph(node_mappings, graph):
|
||||
input_map = {}
|
||||
new_graph = copy.deepcopy(graph)
|
||||
for k in graph:
|
||||
data = graph[k]
|
||||
if not isinstance(data, dict) or "class_type" not in data or data["class_type"] not in node_mappings:
|
||||
continue
|
||||
node = node_mappings[data["class_type"]]()
|
||||
inputs = map_inputs(input_map, data["inputs"].copy())
|
||||
inputs["unique_id"] = k
|
||||
fn = getattr(node, node.FUNCTION)
|
||||
expansion = fn(**inputs)
|
||||
for i, v in enumerate(expansion["result"]):
|
||||
input_map[(k, i)] = v
|
||||
del new_graph[k]
|
||||
new_graph.update(expansion["expand"])
|
||||
|
||||
for k in new_graph:
|
||||
data = new_graph[k]
|
||||
data["inputs"] = map_inputs(input_map, data["inputs"])
|
||||
return new_graph
|
||||
|
||||
+37
-5
@@ -1,16 +1,48 @@
|
||||
[project]
|
||||
name = "comfyui-prompt-control"
|
||||
description = "Nodes for convenient prompt editing, making many common operations prompt-controllable"
|
||||
version = "2.0.0-beta.4"
|
||||
description = "Nodes for prompt editing and LoRA scheduling, advanced regional prompting (including attention masking) and advanced prompt encoding, all controlled through your text prompt. Feature keywords: comfyui-prompt-control, schedule, macros, attention couple, loractl, A1111"
|
||||
version = "3.0.0-beta.3"
|
||||
license = { file = "LICENSE" }
|
||||
# some lark versions older than 1.1.9 apparently have a bug that breaks things, see https://github.com/asagi4/comfyui-prompt-control/issues/35
|
||||
dependencies = ["lark >= 1.1.9"]
|
||||
|
||||
requires-python = ">= 3.10"
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/asagi4/comfyui-prompt-control"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "asagi4"
|
||||
DisplayName = "ComfyUI Prompt Control"
|
||||
Icon = ""
|
||||
|
||||
[tool.pyright]
|
||||
extraPaths = ["../../"]
|
||||
exclude = ["prompt_control/*test*"]
|
||||
|
||||
[tool.ty.src]
|
||||
exclude = ["tests/*.py", "prompt_control/*test*.py", "prompt_control/parsy.py"]
|
||||
|
||||
[tool.ty.environment]
|
||||
extra-paths = ["../.."]
|
||||
|
||||
[tool.ty.rules]
|
||||
# ComfyUI executes give this...
|
||||
invalid-method-override = "ignore"
|
||||
|
||||
[tool.ruff]
|
||||
exclude = ["prompt_control/parsy.py"]
|
||||
line-length = 120
|
||||
|
||||
[tool.ruff.lint]
|
||||
# Ignore line length and let the formatter handle it
|
||||
ignore = ["E501"]
|
||||
select = [
|
||||
"E",
|
||||
"F",
|
||||
"UP",
|
||||
"B",
|
||||
"SIM",
|
||||
"I",
|
||||
]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = ["tests"]
|
||||
|
||||
+1
-2
@@ -1,2 +1 @@
|
||||
# some lark versions older than 1.1.9 apparently have a bug that breaks things, see https://github.com/asagi4/comfyui-prompt-control/issues/35
|
||||
lark >= 1.1.9
|
||||
# Nothing for now
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
import logging
|
||||
|
||||
|
||||
def pytest_runtest_setup(item):
|
||||
logging.getLogger("comfyui-prompt-control").setLevel(logging.CRITICAL)
|
||||
@@ -0,0 +1,26 @@
|
||||
import pytest
|
||||
|
||||
from prompt_control.cutoff_parser import parse_cuts
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
def parser():
|
||||
return parse_cuts
|
||||
|
||||
|
||||
def test_parse_no_cuts(parser):
|
||||
prompt, cutouts = parse_cuts("a b c")
|
||||
assert prompt == "a b c"
|
||||
assert cutouts == []
|
||||
|
||||
|
||||
def test_parse_cuts(parser):
|
||||
prompt, cutouts = parse_cuts("a [CUT:b:d:0] c")
|
||||
assert prompt == "a b c"
|
||||
assert cutouts == [("b", "d", 0.0, None, None, None)]
|
||||
|
||||
|
||||
def test_parse_cuts_multiple(parser):
|
||||
prompt, cutouts = parse_cuts("a [CUT:b:d:0] [CUT:c:e:1.0:0.5:0.9:-]")
|
||||
assert prompt == "a b c"
|
||||
assert cutouts == [("b", "d", 0.0, None, None, None), ("c", "e", 1.0, 0.5, 0.9, "-")]
|
||||
@@ -0,0 +1,243 @@
|
||||
import numpy.testing as npt
|
||||
import pytest
|
||||
|
||||
|
||||
def run(f, *args):
|
||||
if hasattr(f, "execute"):
|
||||
return f.execute(*args)
|
||||
else:
|
||||
return getattr(f, f.FUNCTION)(*args)
|
||||
|
||||
|
||||
def compare_hookgroup_mask(h1, h2):
|
||||
assert len(h1.hooks) == len(h2.hooks)
|
||||
for a, b in zip(h1.hooks, h2.hooks, strict=True):
|
||||
assert (a.mask == b.mask).all()
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def text_encoder_clips():
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from comfy.sd import load_clip
|
||||
|
||||
clips = []
|
||||
to_test = os.environ.get("TEST_TE", "clip_l").split()
|
||||
model_dir = os.environ.get("COMFYUI_TE_DIR", ".")
|
||||
te_root = Path(model_dir).resolve()
|
||||
|
||||
if "clip_l" in to_test:
|
||||
clip_l = load_clip(
|
||||
ckpt_paths=[str(te_root / "clip_l.safetensors")], clip_type="stable_diffusion", model_options={}
|
||||
)
|
||||
clips.append(("clip_l", clip_l))
|
||||
|
||||
if "t5" in to_test:
|
||||
dual = load_clip(
|
||||
[str(te_root / "clip_l.safetensors"), str(te_root / "t5xxl_fp16.safetensors")],
|
||||
clip_type="flux",
|
||||
model_options={},
|
||||
)
|
||||
clips.append(("clip_l+t5", dual))
|
||||
|
||||
return clips
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def pc_text_encode():
|
||||
from prompt_control.nodes_base import PCTextEncode
|
||||
|
||||
return PCTextEncode()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def node_class_objs():
|
||||
import comfy_extras.nodes_mask
|
||||
import nodes
|
||||
|
||||
# Return all used node class objects
|
||||
return {
|
||||
"comfy": nodes.CLIPTextEncode(),
|
||||
"combine": nodes.ConditioningCombine(),
|
||||
"average": nodes.ConditioningAverage(),
|
||||
"concat": nodes.ConditioningConcat(),
|
||||
"zeroout": nodes.ConditioningZeroOut(),
|
||||
"strength": nodes.ConditioningSetAreaStrength(),
|
||||
"solidmask": comfy_extras.nodes_mask.SolidMask(),
|
||||
"setmask": nodes.ConditioningSetMask(),
|
||||
}
|
||||
|
||||
|
||||
def tensors_equal(t1, t2):
|
||||
npt.assert_equal(t1.detach().numpy(), t2.detach().numpy())
|
||||
|
||||
|
||||
def cond_equal(c1, c2, key=None, key_assert=None):
|
||||
assert len(c1) == len(c2)
|
||||
for i in range(len(c1)):
|
||||
a, b = c1[i], c2[i]
|
||||
if key:
|
||||
(key_assert or assert_equal)(a[1].get(key), b[1].get(key))
|
||||
else:
|
||||
tensors_equal(a[0], b[0])
|
||||
|
||||
|
||||
def assert_equal(a, b):
|
||||
assert a == b
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("text_encoder_clips", "pc_text_encode", "node_class_objs")
|
||||
class TestPCTextEncode:
|
||||
def test_basic_encode(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
comfy = node_class_objs["comfy"]
|
||||
combine = node_class_objs["combine"]
|
||||
average = node_class_objs["average"]
|
||||
concat = node_class_objs["concat"]
|
||||
zeroout = node_class_objs["zeroout"]
|
||||
|
||||
for _k, clip in text_encoder_clips:
|
||||
# No exceptions
|
||||
run(
|
||||
pc_text_encode,
|
||||
clip,
|
||||
"test AND test (test:1.2) BREAK test AND TE_WEIGHT(all=0) SDXL() AND AREA(,,) test CAT test",
|
||||
)
|
||||
|
||||
# Basic
|
||||
(c1,) = run(pc_text_encode, clip, "test")
|
||||
(c2,) = run(comfy, clip, "test")
|
||||
c = c2 # Used in later tests
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Quotes
|
||||
(c1,) = run(pc_text_encode, clip, 'Text saying "DOG MASK AND CAT COUPLE MASK(X)"')
|
||||
(c2,) = run(comfy, clip, 'Text saying "DOG MASK AND CAT COUPLE MASK(X)"')
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Function cornercase
|
||||
(c1,) = run(pc_text_encode, clip, "test SDXL function")
|
||||
(c2,) = run(comfy, clip, "test SDXL function")
|
||||
(c3,) = run(pc_text_encode, clip, "test SDXL() function")
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Weights
|
||||
(c1,) = run(pc_text_encode, clip, "(test:1.2) (test:0.6)")
|
||||
(c2,) = run(comfy, clip, "(test:1.2) (test:0.6)")
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Concat
|
||||
(c1,) = run(pc_text_encode, clip, "test CAT test")
|
||||
(c2,) = run(concat, c, c)
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Combine
|
||||
(c1,) = run(pc_text_encode, clip, "test AND test")
|
||||
(c2,) = run(combine, c, c)
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Zero out
|
||||
(c1,) = run(pc_text_encode, clip, "test TE_WEIGHT(all=0)")
|
||||
(c2,) = run(zeroout, c)
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Average
|
||||
(c1,) = run(comfy, clip, "test1")
|
||||
(c2,) = run(comfy, clip, "test2")
|
||||
(c3,) = run(pc_text_encode, clip, "test1 AVG() test2")
|
||||
(c4,) = run(pc_text_encode, clip, "test1 AVG test2")
|
||||
(avg,) = run(average, c1, c2, 0.5)
|
||||
cond_equal(avg, c3)
|
||||
cond_equal(avg, c4)
|
||||
|
||||
def test_avg(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
comfy = node_class_objs["comfy"]
|
||||
average = node_class_objs["average"]
|
||||
for _k, clip in text_encoder_clips:
|
||||
(c1,) = run(comfy, clip, "test1")
|
||||
(c2,) = run(comfy, clip, "test2")
|
||||
(c3,) = run(comfy, clip, "test3")
|
||||
(c4,) = run(pc_text_encode, clip, "test1 AVG() test2 AVG() test3")
|
||||
(c5,) = run(pc_text_encode, clip, "test1 AVG test2 AVG test3")
|
||||
(avg1,) = run(average, c1, c2, 0.5)
|
||||
(avg,) = run(average, avg1, c3, 0.5)
|
||||
cond_equal(avg, c4)
|
||||
cond_equal(avg, c5)
|
||||
|
||||
@pytest.mark.xfail
|
||||
def test_failure(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
comfy = node_class_objs["comfy"]
|
||||
for _k, clip in text_encoder_clips:
|
||||
(c1,) = run(comfy, clip, "test SDXL function")
|
||||
(c2,) = run(pc_text_encode, clip, "test SDXL() function")
|
||||
cond_equal(c1, c2)
|
||||
|
||||
def test_weight(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
comfy = node_class_objs["comfy"]
|
||||
combine = node_class_objs["combine"]
|
||||
strength = node_class_objs["strength"]
|
||||
|
||||
for _k, clip in text_encoder_clips:
|
||||
(c,) = run(comfy, clip, "test")
|
||||
(c2,) = run(strength, c, 0.5)
|
||||
# Conditioning weights
|
||||
(a,) = run(pc_text_encode, clip, "test :0.5 AND test :0.5")
|
||||
(b,) = run(combine, c2, c2)
|
||||
cond_equal(a, b)
|
||||
cond_equal(a, b, "strength")
|
||||
# Weight == 0
|
||||
(a,) = run(pc_text_encode, clip, "test :0.5 AND test :0 AND test")
|
||||
(b,) = run(combine, c2, c)
|
||||
cond_equal(a, b)
|
||||
cond_equal(a, b, "strength")
|
||||
|
||||
def test_attn_couple(self, text_encoder_clips, pc_text_encode):
|
||||
for _k, clip in text_encoder_clips:
|
||||
(c,) = run(pc_text_encode, clip, "test COUPLE prompt1 AND test2 COUPLE prompt2")
|
||||
(c2,) = run(pc_text_encode, clip, "test COUPLE prompt1 COUPLE test2 COUPLE prompt2")
|
||||
assert len(c) == 2
|
||||
assert len(c2) == 1
|
||||
|
||||
def test_styles(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
comfy = node_class_objs["comfy"]
|
||||
for _k, clip in text_encoder_clips:
|
||||
(no_weights,) = run(comfy, clip, "this prompt has no weights")
|
||||
for style in ["comfy", "A1111", "comfy++", "compel", "down_weight", "perp"]:
|
||||
# no weights equal comfy
|
||||
(c,) = run(pc_text_encode, clip, "this prompt has no weights")
|
||||
cond_equal(no_weights, c)
|
||||
# does not fail when encoding weights
|
||||
for normalization in ["none", "mean", "length", "mean+length", "length+mean"]:
|
||||
run(
|
||||
pc_text_encode,
|
||||
clip,
|
||||
f"STYLE({style}, {normalization}) (this prompt) (has weights:0.9), (a:1.2) (b:1.2)",
|
||||
)
|
||||
# Just checking for exceptions
|
||||
|
||||
def test_masks(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
comfy = node_class_objs["comfy"]
|
||||
solidmask = node_class_objs["solidmask"]
|
||||
setmask = node_class_objs["setmask"]
|
||||
for _k, clip in text_encoder_clips:
|
||||
(c1,) = run(pc_text_encode, clip, "test MASK()")
|
||||
(c2,) = run(comfy, clip, "test")
|
||||
(c2,) = run(setmask, c2, run(solidmask, 1.0, 512, 512)[0], "default", 1.0)
|
||||
cond_equal(c1, c2)
|
||||
cond_equal(c1, c2, "mask", tensors_equal)
|
||||
|
||||
def test_cutoff_nofail(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
for _k, clip in text_encoder_clips:
|
||||
(c1,) = run(pc_text_encode, clip, "test [CUT:a:b:0.5]")
|
||||
|
||||
def test_couple_mask_shortcut(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
for _k, clip in text_encoder_clips:
|
||||
(c,) = run(pc_text_encode, clip, "test COUPLE() prompt1")
|
||||
(c2,) = run(pc_text_encode, clip, "test COUPLE MASK() prompt1")
|
||||
cond_equal(c, c2)
|
||||
cond_equal(c, c2, "hooks", compare_hookgroup_mask)
|
||||
|
||||
(c,) = run(pc_text_encode, clip, "test COUPLE(0 0.2, 0.5) prompt1")
|
||||
(c2,) = run(pc_text_encode, clip, "test COUPLE MASK(0 0.2, 0.5) prompt1")
|
||||
cond_equal(c, c2)
|
||||
cond_equal(c, c2, "hooks", compare_hookgroup_mask)
|
||||
@@ -0,0 +1,584 @@
|
||||
import logging
|
||||
|
||||
import pytest
|
||||
from comfy_execution.graph_utils import GraphBuilder
|
||||
|
||||
from prompt_control.nodes_lazy import (
|
||||
PCLazyLoraLoader,
|
||||
PCLazyLoraLoaderAdvanced,
|
||||
PCLazyTextEncode,
|
||||
PCLazyTextEncodeAdvanced,
|
||||
)
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def reset_graphbuilder_state():
|
||||
GraphBuilder.set_default_prefix("UID", 0, 0)
|
||||
|
||||
|
||||
def find_file(name):
|
||||
names = {"test": "test.safetensors", "other": "some/other.safetensors"}
|
||||
return names.get(name)
|
||||
|
||||
|
||||
def as_dict(out):
|
||||
return {"result": out.result, "expand": out.expand}
|
||||
|
||||
|
||||
def loraloader(text, adv=False, **kwargs):
|
||||
reset_graphbuilder_state()
|
||||
cls = PCLazyLoraLoaderAdvanced if adv else PCLazyLoraLoader
|
||||
model = [0, 1]
|
||||
clip = [0, 0]
|
||||
return as_dict(cls.execute(model=model, clip=clip, text=text, **kwargs))
|
||||
|
||||
|
||||
def te(text, adv=False, **kwargs):
|
||||
cls = PCLazyTextEncode if adv else PCLazyTextEncodeAdvanced
|
||||
reset_graphbuilder_state()
|
||||
clip = [0, 0]
|
||||
return as_dict(cls.execute(clip=clip, text=text, **kwargs))
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def patch_lora_name_to_file(monkeypatch):
|
||||
import prompt_control.utils
|
||||
|
||||
monkeypatch.setattr(prompt_control.utils, "lora_name_to_file", find_file)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def patch_torch_cuda_current_device(monkeypatch):
|
||||
import torch.cuda
|
||||
|
||||
monkeypatch.setattr(torch.cuda, "current_device", lambda: "cpu")
|
||||
|
||||
|
||||
def test_textencode_expansion():
|
||||
for p in ["test", "[test:0.2] test", "[test[test::0.5]]<lora:test:1>"]:
|
||||
r1 = te(p)
|
||||
r2 = te(p, adv=True)
|
||||
assert r1 == r2
|
||||
|
||||
|
||||
def test_textencode_alternating():
|
||||
r = te("[a|b]")
|
||||
expected_result = {
|
||||
"expand": {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {
|
||||
"clip": [
|
||||
0,
|
||||
0,
|
||||
],
|
||||
"text": "a",
|
||||
},
|
||||
},
|
||||
"UID.0.0.10": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {
|
||||
"conditioning": [
|
||||
"UID.0.0.9",
|
||||
0,
|
||||
],
|
||||
"end": 0.5,
|
||||
"start": 0.4,
|
||||
},
|
||||
},
|
||||
"UID.0.0.11": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {
|
||||
"clip": [
|
||||
0,
|
||||
0,
|
||||
],
|
||||
"text": "b",
|
||||
},
|
||||
},
|
||||
"UID.0.0.12": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {
|
||||
"conditioning": [
|
||||
"UID.0.0.11",
|
||||
0,
|
||||
],
|
||||
"end": 0.6,
|
||||
"start": 0.5,
|
||||
},
|
||||
},
|
||||
"UID.0.0.13": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {
|
||||
"clip": [
|
||||
0,
|
||||
0,
|
||||
],
|
||||
"text": "a",
|
||||
},
|
||||
},
|
||||
"UID.0.0.14": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {
|
||||
"conditioning": [
|
||||
"UID.0.0.13",
|
||||
0,
|
||||
],
|
||||
"end": 0.7,
|
||||
"start": 0.6,
|
||||
},
|
||||
},
|
||||
"UID.0.0.15": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {
|
||||
"clip": [
|
||||
0,
|
||||
0,
|
||||
],
|
||||
"text": "b",
|
||||
},
|
||||
},
|
||||
"UID.0.0.16": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {
|
||||
"conditioning": [
|
||||
"UID.0.0.15",
|
||||
0,
|
||||
],
|
||||
"end": 0.8,
|
||||
"start": 0.7,
|
||||
},
|
||||
},
|
||||
"UID.0.0.17": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {
|
||||
"clip": [
|
||||
0,
|
||||
0,
|
||||
],
|
||||
"text": "a",
|
||||
},
|
||||
},
|
||||
"UID.0.0.18": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {
|
||||
"conditioning": [
|
||||
"UID.0.0.17",
|
||||
0,
|
||||
],
|
||||
"end": 0.9,
|
||||
"start": 0.8,
|
||||
},
|
||||
},
|
||||
"UID.0.0.19": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {
|
||||
"clip": [
|
||||
0,
|
||||
0,
|
||||
],
|
||||
"text": "b",
|
||||
},
|
||||
},
|
||||
"UID.0.0.2": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {
|
||||
"conditioning": [
|
||||
"UID.0.0.1",
|
||||
0,
|
||||
],
|
||||
"end": 0.1,
|
||||
"start": 0.0,
|
||||
},
|
||||
},
|
||||
"UID.0.0.20": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {
|
||||
"conditioning": [
|
||||
"UID.0.0.19",
|
||||
0,
|
||||
],
|
||||
"end": 1.0,
|
||||
"start": 0.9,
|
||||
},
|
||||
},
|
||||
"UID.0.0.21": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {
|
||||
"conditioning_1": [
|
||||
"UID.0.0.2",
|
||||
0,
|
||||
],
|
||||
"conditioning_2": [
|
||||
"UID.0.0.4",
|
||||
0,
|
||||
],
|
||||
},
|
||||
},
|
||||
"UID.0.0.22": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {
|
||||
"conditioning_1": [
|
||||
"UID.0.0.21",
|
||||
0,
|
||||
],
|
||||
"conditioning_2": [
|
||||
"UID.0.0.6",
|
||||
0,
|
||||
],
|
||||
},
|
||||
},
|
||||
"UID.0.0.23": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {
|
||||
"conditioning_1": [
|
||||
"UID.0.0.22",
|
||||
0,
|
||||
],
|
||||
"conditioning_2": [
|
||||
"UID.0.0.8",
|
||||
0,
|
||||
],
|
||||
},
|
||||
},
|
||||
"UID.0.0.24": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {
|
||||
"conditioning_1": [
|
||||
"UID.0.0.23",
|
||||
0,
|
||||
],
|
||||
"conditioning_2": [
|
||||
"UID.0.0.10",
|
||||
0,
|
||||
],
|
||||
},
|
||||
},
|
||||
"UID.0.0.25": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {
|
||||
"conditioning_1": [
|
||||
"UID.0.0.24",
|
||||
0,
|
||||
],
|
||||
"conditioning_2": [
|
||||
"UID.0.0.12",
|
||||
0,
|
||||
],
|
||||
},
|
||||
},
|
||||
"UID.0.0.26": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {
|
||||
"conditioning_1": [
|
||||
"UID.0.0.25",
|
||||
0,
|
||||
],
|
||||
"conditioning_2": [
|
||||
"UID.0.0.14",
|
||||
0,
|
||||
],
|
||||
},
|
||||
},
|
||||
"UID.0.0.27": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {
|
||||
"conditioning_1": [
|
||||
"UID.0.0.26",
|
||||
0,
|
||||
],
|
||||
"conditioning_2": [
|
||||
"UID.0.0.16",
|
||||
0,
|
||||
],
|
||||
},
|
||||
},
|
||||
"UID.0.0.28": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {
|
||||
"conditioning_1": [
|
||||
"UID.0.0.27",
|
||||
0,
|
||||
],
|
||||
"conditioning_2": [
|
||||
"UID.0.0.18",
|
||||
0,
|
||||
],
|
||||
},
|
||||
},
|
||||
"UID.0.0.29": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {
|
||||
"conditioning_1": [
|
||||
"UID.0.0.28",
|
||||
0,
|
||||
],
|
||||
"conditioning_2": [
|
||||
"UID.0.0.20",
|
||||
0,
|
||||
],
|
||||
},
|
||||
},
|
||||
"UID.0.0.3": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {
|
||||
"clip": [
|
||||
0,
|
||||
0,
|
||||
],
|
||||
"text": "b",
|
||||
},
|
||||
},
|
||||
"UID.0.0.4": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {
|
||||
"conditioning": [
|
||||
"UID.0.0.3",
|
||||
0,
|
||||
],
|
||||
"end": 0.2,
|
||||
"start": 0.1,
|
||||
},
|
||||
},
|
||||
"UID.0.0.5": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {
|
||||
"clip": [
|
||||
0,
|
||||
0,
|
||||
],
|
||||
"text": "a",
|
||||
},
|
||||
},
|
||||
"UID.0.0.6": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {
|
||||
"conditioning": [
|
||||
"UID.0.0.5",
|
||||
0,
|
||||
],
|
||||
"end": 0.3,
|
||||
"start": 0.2,
|
||||
},
|
||||
},
|
||||
"UID.0.0.7": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {
|
||||
"clip": [
|
||||
0,
|
||||
0,
|
||||
],
|
||||
"text": "b",
|
||||
},
|
||||
},
|
||||
"UID.0.0.8": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {
|
||||
"conditioning": [
|
||||
"UID.0.0.7",
|
||||
0,
|
||||
],
|
||||
"end": 0.4,
|
||||
"start": 0.3,
|
||||
},
|
||||
},
|
||||
"UID.0.0.9": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {
|
||||
"clip": [
|
||||
0,
|
||||
0,
|
||||
],
|
||||
"text": "a",
|
||||
},
|
||||
},
|
||||
},
|
||||
"result": (
|
||||
[
|
||||
"UID.0.0.29",
|
||||
0,
|
||||
],
|
||||
),
|
||||
}
|
||||
assert r == expected_result
|
||||
|
||||
|
||||
def test_textencode_lora():
|
||||
reset_graphbuilder_state()
|
||||
r = te("test<lora:test:1>")
|
||||
assert r == {
|
||||
"result": (["UID.0.0.2", 0],),
|
||||
"expand": {
|
||||
"UID.0.0.1": {"class_type": "PCTextEncode", "inputs": {"clip": [0, 0], "text": "test"}},
|
||||
"UID.0.0.2": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {"conditioning": ["UID.0.0.1", 0], "start": 0.0, "end": 1.0},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_textencode_lora_with_schedule():
|
||||
r = te("simple [test:0.1,0.5] prompt<lora:test:1>")
|
||||
assert r == {
|
||||
"result": (["UID.0.0.8", 0],),
|
||||
"expand": {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {"clip": [0, 0], "text": "simple prompt"},
|
||||
},
|
||||
"UID.0.0.2": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {"conditioning": ["UID.0.0.1", 0], "start": 0.0, "end": 0.1},
|
||||
},
|
||||
"UID.0.0.3": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {"clip": [0, 0], "text": "simple test prompt"},
|
||||
},
|
||||
"UID.0.0.4": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {"conditioning": ["UID.0.0.3", 0], "start": 0.1, "end": 0.5},
|
||||
},
|
||||
"UID.0.0.5": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {"clip": [0, 0], "text": "simple prompt"},
|
||||
},
|
||||
"UID.0.0.6": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {"conditioning": ["UID.0.0.5", 0], "start": 0.5, "end": 1.0},
|
||||
},
|
||||
"UID.0.0.7": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {"conditioning_1": ["UID.0.0.2", 0], "conditioning_2": ["UID.0.0.4", 0]},
|
||||
},
|
||||
"UID.0.0.8": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {"conditioning_1": ["UID.0.0.7", 0], "conditioning_2": ["UID.0.0.6", 0]},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_loraloader_empty(monkeypatch, caplog):
|
||||
result = loraloader("prompt here <lora:nonexistent:1.0:0.5>")["expand"]
|
||||
result_adv = loraloader("prompt here <lora:nonexistent:1.0:0.5>", adv=True)["expand"]
|
||||
assert result == {}
|
||||
assert result_adv == {}
|
||||
|
||||
|
||||
def test_loraloader_duplicate_results():
|
||||
result = loraloader("<lora:test:1>")["expand"]
|
||||
result2 = loraloader("prompt here <lora:test:1.0:0.5><lora:test:0:0.5>")["expand"]
|
||||
result3 = loraloader("prompt here <lora:test:1.0:0.5><lora:test:0:0.5>", adv=True)["expand"]
|
||||
assert result == result2
|
||||
assert result2 == result3
|
||||
assert result == {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "LoraLoader",
|
||||
"inputs": {
|
||||
"model": [0, 1],
|
||||
"clip": [0, 0],
|
||||
"strength_model": 1.0,
|
||||
"strength_clip": 1.0,
|
||||
"lora_name": "test.safetensors",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def test_loraloader_multiple_loras():
|
||||
result = loraloader("<lora:test:1><lora:other:0.5>")["expand"]
|
||||
assert result == {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "LoraLoader",
|
||||
"inputs": {
|
||||
"model": [0, 1],
|
||||
"clip": [0, 0],
|
||||
"strength_model": 1.0,
|
||||
"strength_clip": 1.0,
|
||||
"lora_name": "test.safetensors",
|
||||
},
|
||||
},
|
||||
"UID.0.0.2": {
|
||||
"class_type": "LoraLoader",
|
||||
"inputs": {
|
||||
"model": ["UID.0.0.1", 0],
|
||||
"clip": ["UID.0.0.1", 1],
|
||||
"strength_model": 0.5,
|
||||
"strength_clip": 0.5,
|
||||
"lora_name": "some/other.safetensors",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_loraloader_strength_clip():
|
||||
result = loraloader("prompt here <lora:test:1.0:0.5>")["expand"]
|
||||
assert result == {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "LoraLoader",
|
||||
"inputs": {
|
||||
"model": [0, 1],
|
||||
"clip": [0, 0],
|
||||
"strength_model": 1.0,
|
||||
"strength_clip": 0.5,
|
||||
"lora_name": "test.safetensors",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def test_loraloader_scheduled_compare():
|
||||
result = loraloader("prompt [<lora:test:0.5>:0.5]")["expand"]
|
||||
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True)["expand"]
|
||||
assert result == result2
|
||||
expected = {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "CreateHookLora",
|
||||
"inputs": {"lora_name": "test.safetensors", "strength_model": 0.5, "strength_clip": 0.5},
|
||||
},
|
||||
"UID.0.0.2": {
|
||||
"class_type": "CreateHookKeyframe",
|
||||
"inputs": {"strength_mult": 0.0, "start_percent": 0.0},
|
||||
},
|
||||
"UID.0.0.3": {
|
||||
"class_type": "CreateHookKeyframe",
|
||||
"inputs": {"start_percent": 0.5, "prev_hook_kf": ["UID.0.0.2", 0], "strength_mult": 1.0},
|
||||
},
|
||||
"UID.0.0.4": {
|
||||
"class_type": "SetHookKeyframes",
|
||||
"inputs": {"hooks": ["UID.0.0.1", 0], "hook_kf": ["UID.0.0.3", 0]},
|
||||
},
|
||||
"UID.0.0.5": {
|
||||
"class_type": "SetClipHooks",
|
||||
"inputs": {
|
||||
"clip": [0, 0],
|
||||
"hooks": ["UID.0.0.4", 0],
|
||||
"apply_to_conds": True,
|
||||
"schedule_clip": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
assert result == expected
|
||||
|
||||
|
||||
def test_loraloader_adv_start():
|
||||
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True, start=0.6)["expand"]
|
||||
assert result2 == {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "LoraLoader",
|
||||
"inputs": {
|
||||
"model": [0, 1],
|
||||
"clip": [0, 0],
|
||||
"strength_model": 0.5,
|
||||
"strength_clip": 0.5,
|
||||
"lora_name": "test.safetensors",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def test_loraloader_end_zero():
|
||||
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True, end=0.5)["expand"]
|
||||
assert result2 == {}
|
||||
@@ -0,0 +1,379 @@
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
from prompt_control.macros import expand_macros
|
||||
|
||||
|
||||
def lora_dict(*loras):
|
||||
return {lora: {"weight": unet, "weight_clip": te} for lora, unet, te in loras}
|
||||
|
||||
|
||||
def prompt(until, text, *loras):
|
||||
return (until, {"prompt": text, "loras": lora_dict(*loras)})
|
||||
|
||||
|
||||
def prompts_match(a, b):
|
||||
return list(a) == list(b)
|
||||
|
||||
|
||||
def assert_prompt(p, at, until, text, *loras):
|
||||
assert prompts_match(p.at_step(at), prompt(until, text, *loras))
|
||||
|
||||
|
||||
parsers_to_test = os.environ.get("PC_PARSERS_TO_TEST", "new").split()
|
||||
|
||||
params = []
|
||||
if "old" in parsers_to_test:
|
||||
from prompt_control.parser_lark import parse_prompt_schedules as old_parse # noqa
|
||||
|
||||
params.append(old_parse)
|
||||
|
||||
if "new" in parsers_to_test:
|
||||
from prompt_control.parser_parsy import parse_prompt_schedules as new_parse # noqa
|
||||
|
||||
params.append(new_parse)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True, params=params)
|
||||
def parse(request):
|
||||
return request.param
|
||||
|
||||
|
||||
@pytest.mark.parametrize("step", [0, 0.5, 1])
|
||||
def test_no_scheduling(step, parse):
|
||||
p = parse(r"This is a (basic:0.6) (prompt) with [no scheduling] features and \(escaped parens\)")
|
||||
expected = prompt(1.0, r"This is a (basic:0.6) (prompt) with [no scheduling] features and \(escaped parens\)")
|
||||
assert prompts_match(p.at_step(step), expected)
|
||||
|
||||
|
||||
def test_integer_steps(parse):
|
||||
p = parse("[a:b:25]", num_steps=50)
|
||||
assert prompts_match(p.at_step(0), prompt(0.5, "a"))
|
||||
assert prompts_match(p.at_step(25), prompt(0.5, "a"))
|
||||
assert prompts_match(p.at_step(0.5), prompt(0.5, "a"))
|
||||
assert prompts_match(p.at_step(0.51), prompt(1.0, "b"))
|
||||
assert prompts_match(p.at_step(30), prompt(1.0, "b"))
|
||||
|
||||
|
||||
def test_mixed_steps(parse):
|
||||
p = parse("[a:b:25] [c:d:0.25]", num_steps=50)
|
||||
assert prompts_match(p.at_step(0), prompt(0.25, "a c"))
|
||||
assert prompts_match(p.at_step(25), prompt(0.5, "a d"))
|
||||
assert prompts_match(p.at_step(0.5), prompt(0.5, "a d"))
|
||||
assert prompts_match(p.at_step(0.51), prompt(1.0, "b d"))
|
||||
assert prompts_match(p.at_step(30), prompt(1.0, "b d"))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("step", [0, 0.5, 1])
|
||||
def test_quote(step, parse):
|
||||
p = parse('This is a text with a "QUOTED DEF(X=Y)"')
|
||||
expected = prompt(1.0, 'This is a text with a "QUOTED DEF(X=Y)"')
|
||||
assert prompts_match(p.at_step(step), expected)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"group",
|
||||
[
|
||||
["[a:0.1]", "[:a:0.1]", "[:a::0.1,1.0]", "[:a::0.1,1.0]", "[:a::0.1]"],
|
||||
["[before:during:after:0.1]", "[before:during:after:0.1,1.0]", "[before:during:0.1]"],
|
||||
["[a:0.1,0.5]", "[[a:0.1]::0.5]", "[:a::0.1,0.5]", "[a::0.1,0.5]"],
|
||||
["[a:b:0.5]", "[a::b:0.5,0.5]"],
|
||||
["[a::0.5]", "[a:::0.5,0.5]"],
|
||||
],
|
||||
)
|
||||
def test_equivalences(group, parse):
|
||||
objects = [parse(g) for g in group]
|
||||
first = objects[0].parsed_prompt
|
||||
for obj in objects[1:]:
|
||||
assert obj.parsed_prompt == first
|
||||
|
||||
|
||||
def test_basic(parse):
|
||||
p = parse(
|
||||
"This is a (basic:0.6) (prompt) with (very [[simple]:(basic:0.6):0.5]:1.1) [features::0.8][ and this is ignored:1]"
|
||||
)
|
||||
assert_prompt(p, 0.5, 0.5, "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("step", [0, 0.5, 1])
|
||||
def test_basic_cornercase(parse, step):
|
||||
p = parse("This contains[ an ignored segment in:1] the prompt")
|
||||
assert_prompt(p, step, 1.0, "This contains the prompt")
|
||||
|
||||
|
||||
def test_basic_ok(parse):
|
||||
p = parse(
|
||||
"This is a (basic:0.6) (prompt) with (very [[simple]:(basic:0.6):0.5]:1.1) [features::0.8][ and this is ignored:1]"
|
||||
)
|
||||
assert_prompt(p, 0, 0.5, "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features")
|
||||
assert_prompt(p, 0.7, 0.8, "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) features")
|
||||
assert_prompt(p, 1.0, 1.0, "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) ")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("step", [0, 0.5, 1])
|
||||
def test_lora(step, parse):
|
||||
p = parse("This is a (lora:0.6) (prompt) with [no scheduling] features <lora:foo:0.5> <lora:bar:0.5:-1.0>")
|
||||
expected = prompt(
|
||||
1.0, "This is a (lora:0.6) (prompt) with [no scheduling] features ", ("foo", 0.5, 0.5), ("bar", 0.5, -1.0)
|
||||
)
|
||||
assert prompts_match(p.at_step(step), expected)
|
||||
|
||||
|
||||
def test_scheduled_lora(parse):
|
||||
p = parse(
|
||||
"This is a (lora:0.6) (prompt) with [scheduling] features [<lora:foo:0.5>:<lora:bar:0.5:0.2>:0.3] <lora:bar:0.5:1.0>"
|
||||
)
|
||||
assert_prompt(
|
||||
p, 0.1, 0.3, "This is a (lora:0.6) (prompt) with [scheduling] features ", ("foo", 0.5, 0.5), ("bar", 0.5, 1.0)
|
||||
)
|
||||
assert_prompt(p, 0.5, 1.0, "This is a (lora:0.6) (prompt) with [scheduling] features ", ("bar", 1.0, 1.2))
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"text",
|
||||
[
|
||||
"This is a sequence of [SEQ:a:0.2::0.5:c:0.8][SEQ: and x:0.8]",
|
||||
"This is a sequence of [[a:[c:0.5]:0.2]::0.8][ and x::0.8]",
|
||||
],
|
||||
)
|
||||
def test_seq(parse, text):
|
||||
p = parse(text)
|
||||
prompts = {
|
||||
0.2: "This is a sequence of a and x",
|
||||
0.5: "This is a sequence of and x",
|
||||
0.8: "This is a sequence of c and x",
|
||||
1.0: "This is a sequence of ",
|
||||
}
|
||||
for k, v in prompts.items():
|
||||
assert_prompt(p, k, k, v)
|
||||
|
||||
|
||||
def test_shortcuts_scheduling(parse):
|
||||
p = parse("A schedule [a:0.1,0.7] b")
|
||||
p2 = parse("A schedule [[a:0.1]::0.7] b")
|
||||
p3 = parse("A schedule [a:b:0.5,0.8]")
|
||||
p4 = parse("A schedule [[a:0.5]:b:0.8]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
assert p3.parsed_prompt == p4.parsed_prompt
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"step,until,text",
|
||||
[
|
||||
(0, 0.1, "test excluded test"),
|
||||
(0.2, 0.4, "test test"),
|
||||
(0.45, 1.0, "test excluded2 test"),
|
||||
],
|
||||
)
|
||||
def test_range_1(step, until, text, parse):
|
||||
p = parse("test [excluded::excluded2:0.1,0.4] test")
|
||||
assert_prompt(p, step, until, text)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"step,until,text",
|
||||
[
|
||||
(0, 0.1, "test test"),
|
||||
(0.25, 0.3, "test included test"),
|
||||
(0.15, 0.2, "test excluded test"),
|
||||
(0.55, 0.6, "test test"),
|
||||
(0.95, 1.0, "test excluded2 test"),
|
||||
],
|
||||
)
|
||||
def test_range_2(step, until, text, parse):
|
||||
p = parse("test [[:included::0.2,0.8]|[excluded::excluded2:0.4,0.9]:0.1] test")
|
||||
assert_prompt(p, step, until, text)
|
||||
|
||||
|
||||
def test_nested(parse):
|
||||
p = parse(
|
||||
"This [prompt is [SEQ:[crazy:weird:0.2] stuff:0.5:<lora:cool:1>:0.7:nesting:1.0]:completely ignored with tags:HR]"
|
||||
)
|
||||
prompts = {
|
||||
0.2: (0.2, "This prompt is crazy stuff"),
|
||||
0.3: (0.5, "This prompt is weird stuff"),
|
||||
0.5: (0.5, "This prompt is weird stuff"),
|
||||
0.8: (1.0, "This prompt is nesting"),
|
||||
}
|
||||
for k in prompts:
|
||||
exp = [prompts[k][0], {"prompt": prompts[k][1], "loras": {}}]
|
||||
assert prompts_match(p.at_step(k), exp)
|
||||
|
||||
assert_prompt(p, 0.6, 0.7, "This prompt is ", ("cool", 1.0, 1.0))
|
||||
assert_prompt(p, 0.7, 0.7, "This prompt is ", ("cool", 1.0, 1.0))
|
||||
p2 = p.with_filters(filters="hr, xyz")
|
||||
assert prompts_match(p2.at_step(0), p2.at_step(1))
|
||||
|
||||
|
||||
def test_def(parse):
|
||||
p = parse("DEF(X=0.5) [a:b:X] DEF(test = [c:X]) test test")
|
||||
cases = [
|
||||
(0.2, 0.5, "a "),
|
||||
(0.6, 1.0, "b c c"),
|
||||
]
|
||||
for k, until, text in cases:
|
||||
assert_prompt(p, k, until, text)
|
||||
|
||||
p = parse("DEF(X=[($1):($1:$2):$2])X(test;0.7)")
|
||||
p2 = parse("[(test):(test:0.7):0.7]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
p = parse("DEF(X=[($1):($1:$2):$2])DEF(Y=X(test;$1))Y(0.7) Y(0.5)")
|
||||
p2 = parse("[(test):(test:0.7):0.7] [(test):(test:0.5):0.5]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
p = expand_macros("DEF(X(a;b)=$1 $2 $3 d)X(A) X(A;B;C)")
|
||||
assert p == "A b $3 d A B C d"
|
||||
|
||||
p = expand_macros("DEF(MACRO()=[empty:$1:$2])MACRO MACRO(;) MACRO(;0.5) MACRO(a;0.5)")
|
||||
assert p == "[empty::$2] [empty::] [empty::0.5] [empty:a:0.5]"
|
||||
|
||||
p = expand_macros("DEF(X=$1)DEF(Y()=$1)[X Y][X() Y()][X(1) Y(1)]")
|
||||
assert p == "[$1 ][ ][1 1]"
|
||||
|
||||
p = parse("DEF(test(1)=prompt $1)DEF(test2((a); (test))=[$1:$2:0.5])test test2")
|
||||
p2 = parse("prompt 1 [(a):(prompt 1):0.5]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
with pytest.raises(ValueError) as c:
|
||||
expand_macros("DEF(X=recurse Y) DEF(Y=recurse X) X")
|
||||
assert "Unable to resolve DEFs" in str(c.value)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"text, cases",
|
||||
[
|
||||
(r"[embedding\:a:embedding\:b:0.1,0.5]", [(0.15, 0.5, r"embedding:a"), (0.55, 1, r"embedding:b")]),
|
||||
(
|
||||
r"[embedding\:a:embedding\:b:embedding\:c:0.1,0.5]",
|
||||
[(0.0, 0.1, r"embedding:a"), (0.15, 0.5, r"embedding:b"), (0.55, 1, r"embedding:c")],
|
||||
),
|
||||
(r"[a\:b\\:c:0.5]", [(0.0, 0.5, "a:b\\"), (0.55, 1, r"c")]),
|
||||
(r"[a:\#b:0.5]", [(0.0, 0.5, "a"), (0.55, 1, "#b")]),
|
||||
(r"[a:b \(test\):0.2]", [(0, 0.2, r"a"), (0.25, 1, r"b \(test\)")]),
|
||||
],
|
||||
)
|
||||
def test_escapes(text, cases, parse):
|
||||
p = parse(text)
|
||||
for step, until, val in cases:
|
||||
assert_prompt(p, step, until, val)
|
||||
|
||||
|
||||
# I think these were wrong in the old parser too
|
||||
@pytest.mark.xfail(reason="Old parser behaviour, possibly buggy")
|
||||
@pytest.mark.parametrize(
|
||||
"text, cases",
|
||||
[
|
||||
(r"[a:\:a:0.5] :\[a:b:0.5]", [(0, 0.5, r"a :\[a:b:0.5]"), (0.55, 1, r":a :\[a:b:0.5]")]),
|
||||
],
|
||||
)
|
||||
def test_escapes_fail(text, cases, parse):
|
||||
p = parse(text)
|
||||
for step, until, val in cases:
|
||||
assert_prompt(p, step, until, val)
|
||||
|
||||
|
||||
def test_comments(parse):
|
||||
p = parse("this is a # comment")
|
||||
assert_prompt(p, 0, 1.0, "this is a ")
|
||||
p = parse("this is a [comment#:scheduled:0.6]")
|
||||
assert_prompt(p, 0, 1.0, "this is a [comment")
|
||||
p = parse(r"this is a [comment\#:scheduled:0.6]")
|
||||
assert_prompt(p, 0, 0.6, "this is a comment#")
|
||||
assert_prompt(p, 0.65, 1.0, "this is a scheduled")
|
||||
p = parse("#this is a comment\nthis is a prompt")
|
||||
assert_prompt(p, 0, 1.0, "\nthis is a prompt")
|
||||
|
||||
|
||||
def test_misc(parse):
|
||||
p = parse("[[a:c:0.5]:0.7]")
|
||||
p2 = parse("[:[a:c:0.5]:0.7]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
p = parse("test [[a:[b<lora:test:0.5>:0.6]:0.5]:HR]")
|
||||
p2 = parse("test [:[a:[:b<lora:test:0.5>:0.6]:0.5]:HR]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
|
||||
def test_filters(parse):
|
||||
p = parse("test [[a:[b<lora:test:0.5>:0.6]:0.5]:HR]")
|
||||
p2 = parse("test [:[a:[:b<lora:test:0.5>:0.6]:0.5]:HR]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
pf = p.with_filters(filters="hr")
|
||||
assert pf.parsed_prompt == p2.with_filters(filters="hr").parsed_prompt
|
||||
assert_prompt(pf, 0, 0.5, "test a")
|
||||
assert_prompt(pf, 0.55, 0.6, "test ")
|
||||
assert_prompt(pf, 0.8, 1.0, "test b", ("test", 0.5, 0.5))
|
||||
|
||||
p = parse("[:[<lora:test:1>:c:0.5]:0.3]")
|
||||
assert_prompt(p, 0, 0.3, "")
|
||||
assert_prompt(p, 0.4, 0.5, "", ("test", 1.0, 1.0))
|
||||
assert_prompt(p, 1.0, 1.0, "c")
|
||||
|
||||
|
||||
def test_emb(parse):
|
||||
p = parse("an [<emb:foo>:<emb:bar>:0.5]")
|
||||
prompts = {
|
||||
0.2: (0.5, "an embedding:foo"),
|
||||
0.8: (1.0, "an embedding:bar"),
|
||||
}
|
||||
for k, (until, val) in prompts.items():
|
||||
assert_prompt(p, k, until, val)
|
||||
|
||||
|
||||
def test_alternating_defaultstep(parse):
|
||||
p = parse("[cat|dog|tiger]")
|
||||
p2 = parse("[cat|dog|tiger:0.1]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
|
||||
def test_alternating_basic(parse):
|
||||
p = parse("[cat|dog|tiger]")
|
||||
p2 = parse("[cat|dog|tiger:0.1]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"equivalent",
|
||||
[
|
||||
"[cat::0.1][dog:0.1,0.2][tiger:0.2,0.3][cat:0.3,0.4][dog:0.4,0.5][tiger:0.5,0.6][cat:0.6,0.7][dog:0.7,0.8][tiger:0.8,0.9][cat:0.9,1.0]"
|
||||
],
|
||||
)
|
||||
def test_alternating_equivalences(parse, equivalent):
|
||||
p = parse("[cat|dog|tiger]")
|
||||
p2 = parse(equivalent)
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
|
||||
@pytest.mark.xfail(reason="Old parser behaviour")
|
||||
def test_cornercase_failure(parse):
|
||||
"""p1 returns a prompt entry until 0 at the start"""
|
||||
p = parse("[cat:0,0.1]")
|
||||
p2 = parse("[cat::0.1]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
|
||||
def test_cornercase_corrected(parse):
|
||||
p = parse("[cat:0,0.1]")
|
||||
p2 = parse("[cat::0.1]")
|
||||
assert p.parsed_prompt[0][0] == 0.0
|
||||
assert p.parsed_prompt[1:] == p2.parsed_prompt
|
||||
|
||||
|
||||
def test_alternating_lora(parse):
|
||||
p4 = parse("[cat|[dog:wolf<lora:canine:1>:0.5]:0.2]")
|
||||
for i, (text, *_loras) in enumerate(
|
||||
[(["cat"],), (["dog"],), (["cat"],), (["wolf", ("canine", 1.0, 1.0)],), (["cat"],)]
|
||||
):
|
||||
step = round((i * 0.2) + 0.2, 2)
|
||||
assert_prompt(p4, step, step, *text)
|
||||
assert_prompt(p4, 0.7, 0.8, "wolf", ("canine", 1.0, 1.0))
|
||||
|
||||
|
||||
def test_alternating_nested(parse):
|
||||
p3 = parse("[cat|[dog|wolf]|tiger]")
|
||||
catdogtigers = ["cat", "wolf", "tiger", "cat", "dog", "tiger", "cat", "wolf", "tiger", "cat"]
|
||||
for i, x in enumerate(catdogtigers):
|
||||
step = round((i * 0.1) + 0.1, 2)
|
||||
assert_prompt(p3, step, step, x)
|
||||
@@ -0,0 +1,3 @@
|
||||
# PC: Attach Mask
|
||||
|
||||
Attaches custom masks to a CLIP object so that they can be referred to in prompts using `PCTextEncode` or `PC: Schedule prompt`.
|
||||
Symlink
+1
@@ -0,0 +1 @@
|
||||
PCAddMaskToCLIP.md
|
||||
@@ -0,0 +1,7 @@
|
||||
# PC: Attention Couple (batch negative)
|
||||
|
||||
This node applies an optimization that re-enables negative cond batching when Attention Couple is in use.
|
||||
|
||||
It improves performance when negative prompts are not scheduled, but slightly affects outputs and is not required for Attention Couple to work.
|
||||
|
||||
Simply add it to your workflow and pass in your positive and negative prompts. It is always safe to use, as it will not do anything when it detects that the optimization can't be applied (eg. when negative prompts contain schedules)
|
||||
@@ -0,0 +1,7 @@
|
||||
# PC: Schedule LoRAs
|
||||
|
||||
This node is the core of Prompt Control. It evaluates a prompt schedule and dynamically expands into a scheduled workflow consisting of necessary calls to `LoRALoader` and `Create Hook LoRA` (for scheduled LoRAs).
|
||||
|
||||
You can use it in place or in addition to your usual `LoRA Loader` nodes; just pass in a text prompt containing your LoRA schedule (it can be shared with `PC: Schedule Prompt`). Then connect your MODEL output as usual and the CLIP output to your `PC: Schedule Prompt` nodes.
|
||||
|
||||
For documentation on syntax, for now see the [documentation on GitHub](https://github.com/asagi4/comfyui-prompt-control/blob/master/doc/schedules.md)
|
||||
+1
@@ -0,0 +1 @@
|
||||
PCLazyLoraLoader.md
|
||||
@@ -0,0 +1,7 @@
|
||||
# PC: Schedule Prompt
|
||||
|
||||
This node is the core of Prompt Control. It evaluates a prompt schedule and dynamically expands into a scheduled workflow consisting of calls to `PCTextEncode`, `SetConditioningTimesteps` and other necessary nodes.
|
||||
|
||||
To use it, simply replace your usual `CLIP Text Encode` nodes with `PC: Schedule Prompt` nodes. For LoRA Loading, you should use `PC: Schedule LoRAs` in place (or in addition to) of your usual LoRA Loader node.
|
||||
|
||||
For documentation on syntax, for now see the [documentation on GitHub](https://github.com/asagi4/comfyui-prompt-control/blob/master/doc/schedules.md)
|
||||
+1
@@ -0,0 +1 @@
|
||||
PCLazyTextEncode.md
|
||||
@@ -0,0 +1,5 @@
|
||||
# PC: LoRA Hooks from Text (non-lazy)
|
||||
|
||||
Creates cond hooks from a LoRA schedule, if you want to apply them manually for some reason.
|
||||
|
||||
You should not need to use this. Use `PC: Schedule LoRAs`.
|
||||
@@ -0,0 +1,5 @@
|
||||
# PC: Expand Macros
|
||||
|
||||
Expands [prompt macros](https://github.com/asagi4/comfyui-prompt-control/blob/master/doc/macros.md)
|
||||
|
||||
You should not need to use this directly. Use `PC: Schedule Prompt` instead.
|
||||
@@ -0,0 +1,7 @@
|
||||
# PC: Configure PCTextEncode
|
||||
|
||||
Configures a CLIP object with new default values used by `PCTextEncode`. Apply it before everything else.
|
||||
|
||||
This is needed if you want to do scheduling with steps instead of denoising percentages, but otherwise it's completely optional.
|
||||
|
||||
Note that steps are simply syntactic sugar for percentages and may not correspond to actual steps depending on the scheduler used.
|
||||
@@ -0,0 +1,5 @@
|
||||
# PC: Text Encode (no scheduling)
|
||||
|
||||
This node encodes text using some special syntax for advanced features. You should rarely need to use this node directly, and instead use `PC: Schedule Prompt` which uses this node under the hood.
|
||||
|
||||
For documentation on syntax, see the [documentation on GitHub](https://github.com/asagi4/comfyui-prompt-control/blob/master/doc/basic.md)
|
||||
Reference in New Issue
Block a user