Compare commits

...
Author SHA1 Message Date
asagi4 6e3fce9dcb v3.0.0-beta.8 2026-08-19 13:42:56 +03:00
asagi4 b581cf7f24 NODE: Fix SEG expansion
See: #150
2026-08-19 13:33:46 +03:00
asagi4 d7a992b96d NODE: Test for SEG expansion 2026-08-19 13:32:04 +03:00
asagi4 19123570e1 Fix parsing of < and > in schedules, see #149 2026-08-16 17:17:38 +03:00
asagi4 e6bb57cd25 Formatting 2026-08-14 18:13:49 +03:00
asagi4 8864201be5 v3.0.0-beta.7 2026-08-14 18:09:05 +03:00
asagi4 4c9cc9e44c Support empty inputs in NODE helper 2026-08-14 18:09:05 +03:00
asagi4 71a7ae5ca6 Add tests for NODE 2026-08-14 17:54:13 +03:00
asagi4 3ce15e45a5 Doc cleanup 2026-08-09 18:12:32 +03:00
asagi4 ec30db6208 Document NODE better 2026-08-09 17:22:21 +03:00
asagi4 e6b617a7cb Note source of custom concat node in H3 workflow, take 2 2026-08-08 20:29:47 +03:00
asagi4 e069747170 Note source of custom concat node in H3 workflow 2026-08-08 20:24:13 +03:00
asagi4 b30af843c7 Add MiniMax prompt control example workflow 2026-08-08 20:14:05 +03:00
asagi4 6bb80c563a Add a link helper for NODE 2026-08-08 13:12:05 +03:00
asagi4 24d4596c4f Fix breakage 2026-08-08 10:08:35 +03:00
asagi4 ebc53ae6e9 Experiment: Arbitrary extra parameters to NODE
This is extremely cursed, but you can do prompt scheduling with video models using this.

For example:

NODE(UC_AdvancedMiniMaxH3ImageToVideo, prompt, vae 161:152 0; reference_images.reference_image_1 [114:201:0.3] 0; width 120 0; height 120 1; length 176 1)

This changes the reference image parameter mid-sampling.

Syntax for the last parameter is a ;-separated list of paramname node_id output_slot tuples. The given slot from the node will then be plugged in as the named parameter)

Export your workflow in API format to find the node IDs.

This will probably change later.
2026-08-08 00:08:39 +03:00
asagi4 ea7a61a52b Don't spam DEF expansions at INFO log level 2026-07-07 18:12:23 +03:00
asagi4 139808033b Add a test for LoRALoader with SEGs 2026-07-04 21:41:10 +03:00
asagi4 42db26f04e Typo 2026-07-04 21:28:45 +03:00
asagi4 f0fec7ea94 v3.0.0-beta.6 2026-07-04 21:27:00 +03:00
asagi4 e71047fd2f Documentation improvements 2026-07-04 21:26:07 +03:00
asagi4 d4e3078af4 Expand SEGs in LazyLoRALoader before prompt parsing 2026-07-04 21:26:02 +03:00
asagi4 0a698eb7ab Apply Anima attention wrappers just-in-time before execution
This avoids using stale block references and seems to fix LoRAs
2026-07-01 19:00:15 +03:00
asagi4 6547749a0f v3.0.0-beta.5 2026-06-26 11:24:28 +03:00
asagi4 e4e77c2f89 Properly read NOISE() from the prompt. Fixes #148 2026-06-26 11:23:40 +03:00
asagi4 1a16b6b811 Code formatting 2026-06-26 11:22:28 +03:00
asagi4 7b815f1edf Fix test failure 2026-06-24 00:04:19 +03:00
asagi4 0cfc50678e Avoid bad performance with many SEGs 2026-06-23 23:56:43 +03:00
asagi4 8180b423e3 Anima Couple: cond weighting 2026-06-23 23:56:43 +03:00
asagi4 1eb836a575 Rename PC: Extract Scheduled Prompt to PC: Show Prompt 2026-06-23 23:56:43 +03:00
asagi4 673a02391d Add expansion feature to PC: Extract Scheduled Prompt 2026-06-23 23:56:43 +03:00
asagi4 4ee459858f Experimental feature: SUB
This one might go away or change
2026-06-23 23:56:03 +03:00
asagi4 74fdb6791f Anima Couple: brute force workaround for device issue 2026-05-25 00:40:20 +03:00
asagi4 054134b5d5 Anima Attention Couple (VERY EXPERIMENTAL, SEE README)
Code stolen and adapted from ppm's node pack again.
I probably introduced bugs.
2026-05-12 20:07:37 +03:00
asagi4 6a1dd77fe9 Silence type checker 2026-05-10 09:12:28 +03:00
asagi4 ad67d0f3ad Remove extra ty ignores 2026-05-10 09:05:46 +03:00
asagi4 a5dfd55613 Relax macro parameter expansion boundary
See #146
2026-05-10 09:05:20 +03:00
asagi4 981dbed245 Remove old 2-pass example 2026-04-30 22:39:32 +03:00
asagi4 45ebc687d1 Warn if using old parser 2026-04-30 22:38:42 +03:00
asagi4 f9c5da7210 v3.0.0-beta.4 2026-04-30 20:57:08 +03:00
asagi4 244ef49230 README wording 2026-04-30 20:56:03 +03:00
asagi4 3a563e3ceb README tweaks 2026-04-30 20:47:16 +03:00
asagi4 812ad90d17 New feature: SEG 2026-04-30 20:41:17 +03:00
asagi4 7d9e8aa6ac Split macro tests 2026-04-30 20:38:31 +03:00
asagi4 f47428ea8c Remove requirements.txt 2026-04-30 20:37:49 +03:00
asagi4 0aeeb50331 Add a test that actually runs a workflow 2026-03-25 21:34:48 +02:00
asagi4 9931c6fa75 These should be initialized to None, see #143 2026-03-25 20:19:59 +02:00
asagi4 655f6ac4a1 Allow floats without leading zero, fixes #141 2026-03-12 20:55:12 +02:00
asagi4 136932de40 v3.0.0-beta.3 2026-03-10 21:11:17 +02:00
asagi4 06a3f43a93 Don't fail in case TE is missing pad_to_max_length attribute 2026-03-10 21:11:07 +02:00
asagi4 203a7ad45c Fix negative LoRA weights and test for them 2026-03-10 02:28:31 +02:00
asagi4 7a566e6e9e Remove calls to logging.basicConfig 2026-02-28 18:44:10 +02:00
asagi4 fa288c226c Add some keywords to the description to make the extension easier to find, see #139 2026-02-28 11:59:13 +02:00
asagi4 e648f3bdc3 Fix nitpick 2026-02-18 20:44:09 +02:00
asagi4 1bedc6ad53 Properly keep \( in escaped parens 2026-02-18 20:38:22 +02:00
asagi4 de79f3c6af Properly default to new parser 2026-02-18 20:17:46 +02:00
asagi4 eb51fd9289 Now things should actually work 2026-02-18 20:08:08 +02:00
asagi4 7cca9438f2 simplify cutoff parser 2026-02-18 20:06:24 +02:00
asagi4 7d786dfe83 Actually run cutoff tests 2026-02-18 18:25:48 +02:00
asagi4 85c15ff6bd Add lark back as a dep for now, I forgot about the cutoff parser 2026-02-18 18:20:40 +02:00
asagi4 1ab1c87f74 Move flatten to utils 2026-02-18 18:18:35 +02:00
asagi4 2ea0b622b4 Fix tests again 2026-02-18 18:15:27 +02:00
asagi4 c4c9561a13 Fix tests 2026-02-18 18:05:51 +02:00
asagi4 ff111a9f7c v3.0.0-beta.1
Not a complete apocalypse this time
2026-02-18 18:00:11 +02:00
asagi4 ab1cf5949e Make new parser the default 2026-02-04 16:39:50 +02:00
asagi4 105562a5bb Remove requirements.txt, not needed with new parser 2026-02-04 16:39:50 +02:00
asagi4 e337ecc019 tests: COUPLE mask shortcut 2026-02-04 16:39:50 +02:00
asagi4 b2bb7b9960 Tests for step count 2026-02-04 16:39:50 +02:00
asagi4 f1ade19345 test: Verify that basic paren escapes don't get stripped 2026-02-04 16:39:50 +02:00
asagi4 15c115c7f3 Fix type complaint 2026-02-04 16:39:50 +02:00
asagi4 97f050c264 test only new parser 2026-02-04 16:39:50 +02:00
asagi4 ca5848a8aa Test multiple averages 2026-02-04 16:39:50 +02:00
asagi4 24240aac2d Make github tests work again 2026-02-04 16:39:50 +02:00
asagi4 883ea6a96d Disable method override complaint 2026-02-04 16:39:50 +02:00
asagi4 49b761ae5b Working importing for v3 nodes 2026-02-04 16:39:50 +02:00
asagi4 bd0dd00d8c v3: nodes_lazy.py 2026-02-04 16:39:50 +02:00
asagi4 3f303cb33a v3: nodes_tools 2026-02-04 16:39:50 +02:00
asagi4 d009bb782b v3: nodes_hooks.py 2026-02-04 16:39:50 +02:00
asagi4 eaf5fa20c4 v3: nodes_base.py 2026-02-04 16:39:50 +02:00
asagi4 2becb5cdc6 Initial v3 migration 2026-02-04 16:39:50 +02:00
asagi4 e1b12acbdf handle setting steps 2026-02-04 16:39:50 +02:00
asagi4 3ea605fbe0 parsy ruff fixes 2026-02-04 16:39:50 +02:00
asagi4 f4d5477d4c parsy loractl 2026-02-04 16:39:50 +02:00
asagi4 5cb18ab5a9 simplify LoRA weight parsing 2026-02-04 16:39:50 +02:00
asagi4 fc5c0a1857 Rewrite parser to use parsy 2026-02-04 16:39:50 +02:00
asagi4 f9769fe90c Remove broken expand_graph.py 2026-02-04 16:39:50 +02:00
asagi4 09438a35f4 Split cutoff parser to its own file 2026-02-04 16:39:50 +02:00
asagi4 d0d34f8b2f split macros out of parser 2026-02-04 16:39:50 +02:00
asagi4 dcc9379d40 Fix NOISE doing nothing 2026-02-04 16:39:50 +02:00
asagi4 9231ec4936 Fix github workflows 2026-02-04 16:39:50 +02:00
asagi4 0473d90446 Remove old tests 2026-02-04 16:39:50 +02:00
asagi4 138df090a6 Fix LazyLoraLoader tests 2026-02-04 16:39:50 +02:00
asagi4 c89dcb135d Test for discovered corner case behaviour 2026-02-04 16:39:50 +02:00
asagi4 c53717101c Add a graph test for alternating 2026-02-04 16:39:50 +02:00
asagi4 d185c72a16 Refactor tests for new parser 2026-02-04 16:39:50 +02:00
asagi4 805ffbf463 Use pytest tests 2026-02-04 16:39:50 +02:00
asagi4 d14f5b789a Convert tests to pytest
Not 100% sure these fully work yet
2026-02-04 16:39:50 +02:00
asagi4 deba0a0642 Ruff fixes etc. 2026-02-04 16:39:50 +02:00
asagi4 f98bf25a83 Typing fixes etc 2026-02-04 16:39:50 +02:00
asagi4 2c8727a75a Remove cache hack, it's broken anyway 2026-02-04 16:39:50 +02:00
asagi4 86ec30c028 Add ruff and ty 2026-02-04 16:39:50 +02:00
asagi4 e0efac1ebc Add pyright 2026-02-04 16:39:50 +02:00
asagi4 68766215f2 v2.1.3 2026-02-04 16:32:26 +02:00
asagi4 287f554a68 Test COUPLE mask shortcut 2026-01-16 18:20:46 +02:00
asagi4 71465f914c Properly parse COUPLE(), see #134 2026-01-16 18:05:30 +02:00
asagi4 d7be7bc29e v2.1.2 2026-01-13 21:50:09 +02:00
asagi4 329d4cf95f Add a test for #133 2026-01-13 21:48:34 +02:00
asagi4 c52ace71aa #133 properly set function locations in get_function 2026-01-13 20:45:10 +02:00
asagi4 4806cf5959 refactor get_function to make it more consistent 2026-01-13 20:45:10 +02:00
asagi4 a0ab709f50 v2.1.1 2025-12-14 17:57:53 +02:00
asagi4 66fef1ffe8 Avoid splitting AND, CAT and others when inside quotes
See #132
2025-12-05 21:25:28 +02:00
asagi4 3341e9f81e Fix TE_WEIGHT failing with some encoders
See #131

Prompt weighting and attention couple will probably not work, but
this should prevent exceptions.
2025-12-03 16:07:31 +02:00
asagi4 aa6b4608f0 Clarify attention couple docs a bit and add a warning if IMASK is used without attaching custom masks
See #108
2025-12-01 13:19:29 +02:00
asagi4 d1cc60b00a v2.1.0 2025-11-20 20:20:22 +02:00
asagi4 efe8939250 Fix issue with T5 encoding sometimes returning NaNs in test 2025-11-20 20:14:40 +02:00
asagi4 39b353b916 Fix tests for #130 2025-11-20 20:14:37 +02:00
asagi4 485bc7f2ab Prepare for v3 conversion of imported nodes, see #12
This should prevent things from breaking, but it needs a bit of testing.
2025-11-20 15:29:33 +02:00
asagi4 1d03ded9dd Find LoRAs with partial match 2025-11-20 15:23:51 +02:00
asagi4 94a4d076e0 Remove unused attributes 2025-08-30 13:46:00 +03:00
asagi4 228dc4b22b Remove the old advanced encoding implementation 2025-08-27 20:27:30 +03:00
asagi4 db523e1f16 Remove dead code 2025-08-27 20:23:41 +03:00
asagi4 6b08c7a90e v2.0.1 2025-08-27 20:20:13 +03:00
asagi4 76142c4b7e Fix #127 2025-08-27 20:15:30 +03:00
asagi4 1d84fdaf9e Release 2.0.0 2025-08-19 19:53:34 +03:00
asagi4 1c50ae5297 Disable the cache hack for now 2025-08-19 19:53:23 +03:00
asagi4 51618289e7 v2.0.0-rc.9 2025-06-21 13:16:57 +03:00
asagi4 cea1e5b30f Initial embedded documentation 2025-06-21 13:15:58 +03:00
asagi4 55f0574ac7 Clarification 2025-06-21 12:22:05 +03:00
asagi4 167689cb8b Add some explanations, see #121 2025-06-21 12:18:01 +03:00
asagi4 2be3abed44 Fix graph expansion node 2025-06-21 11:32:51 +03:00
asagi4 f86abb0816 Add tool for expanding lazy graphs 2025-06-21 11:19:40 +03:00
asagi4 a3537a5b2f Some progress... 2025-06-13 21:55:17 +03:00
asagi4 af7e4542d1 Let's just bruteforce it 2025-06-13 21:49:59 +03:00
asagi4 f37f14b2a2 Does this work? 2025-06-13 21:33:51 +03:00
asagi4 7b8231d36b ... 2025-06-13 21:16:50 +03:00
asagi4 9b5e15fde3 Forgot to import mock 2025-06-13 21:06:14 +03:00
asagi4 d97d30074f Encoder tests need a CPU mock too for CI 2025-06-13 21:03:52 +03:00
asagi4 b2eb9b88ba Try running encoder tests in CI 2025-06-13 20:59:37 +03:00
asagi4 c9459e39f9 Fix #120 and add a test 2025-06-13 09:08:04 +03:00
asagi4 7507b2b55f Also strip comments if they're at the start of a line 2025-06-12 23:30:54 +03:00
asagi4 cf1efecf4c Fix minor mistake in doc 2025-06-12 23:12:13 +03:00
asagi4 ffcf94bcaa Syntax 2025-06-12 23:05:44 +03:00
asagi4 ec8c40355c Split documentation 2025-06-12 23:03:39 +03:00
asagi4 6e538e0abc Fix markdown syntax 2025-06-12 22:42:45 +03:00
asagi4 25c44a1fbb Documentation 2025-06-12 22:41:31 +03:00
asagi4 72d5490498 Add support for commenting out things with #
You can escape it with \#

Fixes #105
2025-06-12 22:26:34 +03:00
asagi4 3de4538326 Test cleanup 2025-06-12 21:38:56 +03:00
asagi4 f61af15d52 Update the description a bit 2025-06-10 19:05:52 +03:00
asagi4 8e59f140ff v2.0.0-rc.8 2025-06-09 20:07:54 +03:00
asagi4 44044e962c Very basic test for COUPLE 2025-06-09 20:06:25 +03:00
asagi4 04f36687c5 Fix skipping prompt segments by setting weight to 0 2025-06-09 20:05:50 +03:00
asagi4 3a8a360d03 Make testing less stupid 2025-06-09 19:53:05 +03:00
asagi4 d76331315a Deduplicate tests 2025-06-09 19:15:55 +03:00
asagi4 e3a6050536 Don't call to() on every iteration 2025-06-09 18:37:53 +03:00
asagi4 7b001ace7b Fix Attention Couple when combined with hooks on the CLIP (eg. LoRAs)
All clones of the AC hook must maintain the same state. This feels
a bit hacky though; there should be a better way

Fixes #119
2025-06-09 17:25:26 +03:00
asagi4 f1de65f257 Don't override existing hooks. Unfortunately, this doesn't make things quite work; hmm. 2025-06-09 16:22:25 +03:00
asagi4 11aaa0ac7b Fix indexing error 2025-06-09 01:06:44 +03:00
asagi4 85de3ef0d3 Fix links 2025-06-08 23:09:16 +03:00
asagi4 a73260ff34 Split t5 tests 2025-06-08 23:05:11 +03:00
asagi4 50a2e0abbf Change ATTN() to COUPLE() and remove need for AND 2025-06-08 23:01:59 +03:00
asagi4 5cf45ca264 Make functions generally callable without argument lists 2025-06-08 18:57:25 +03:00
asagi4 bd4a787400 Use a helper function to parse function splits 2025-06-08 18:56:30 +03:00
asagi4 c9e5bc25c3 Need to do imports after torch mock, otherwise running tests on CPU torch fails 2025-06-08 17:22:54 +03:00
asagi4 d11ffa6e25 Don't pass in a custom prefix to GraphBuilder
It breaks when PCLazyTextEncode etc. are called with list inputs.
Tests needed adjusting after the change.

Fixes #117
2025-06-08 16:49:37 +03:00
61 changed files with 7124 additions and 4486 deletions
+7 -2
View File
@@ -11,8 +11,13 @@ jobs:
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 -r requirements.txt
- run: python -m prompt_control.test_parser
- run: pip install pytest typing-extensions
- run: PYTHONPATH=ComfyUI pytest tests/test_parser.py
+16 -4
View File
@@ -4,13 +4,17 @@ on:
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 graph tests
name: Run tests requiring ComfyUI
runs-on: ubuntu-latest
steps:
- name: Check out code
@@ -23,6 +27,14 @@ jobs:
- uses: actions/setup-python@v5
with:
python-version: '3.11'
- run: pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu
- run: pip install -r requirements.txt -r ComfyUI/requirements.txt
- run: PYTHONPATH=ComfyUI python -m prompt_control.test_graph
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 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
View File
@@ -1 +1,2 @@
__pycache__
.pyre
+19 -5
View File
@@ -1,18 +1,32 @@
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:
python -m prompt_control.test_parser
PYTHONPATH=../../ pytest tests/test_parser.py tests/test_cutout.py tests/test_macros.py $(ARGS)
test_graph:
PYTHONPATH=../../ python -m prompt_control.test_graph
PYTHONPATH=../../ pytest tests/test_graph.py $(ARGS)
test_encode:
PYTHONPATH=../../ python -m prompt_control.test_encode
PYTHONPATH=../../ pytest tests/test_encode.py $(ARGS)
test_workflow:
PYTHONPATH=../../ pytest tests/test_workflow.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
+19 -44
View File
@@ -1,26 +1,31 @@
# ComfyUI prompt control
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.
Control LoRA and prompt scheduling, advanced text encoding, regional prompting, and much more, through your text prompt. Prompt Control generates dynamic graphs that are literally identical to handcrafted noodle soup, condensing complicated workflows with dozens of nodes into simple text prompts.
Prompt Control comes with `PCTextEncode`, which provides advanced text encoding with many additional features compared to ComfyUI's base `CLIPTextEncode`.
A `Basic Text to Image` template is included with the extension, and can be loaded from ComfyUI's template library.
> [!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.
## What can it do?
You can use text prompts to control the following:
- A1111-style prompt scheduling and filtering without noodle soup.
- LoRA loading and scheduling via ComfyUI's hook system
- Masking, composition and area control (regional prompting) with an implementation of [Attention Couple](doc/attention_couple.md), also fully schedulable.
- 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 with `DEF`
- And a bunch more
- LoRA loading and [scheduling](/doc/schedules.md) using ComfyUI's built-in 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).
- Organize complicated prompts with [segments and prompt macros](/doc/macros.md).
- [Schedule your own encoder nodes](/doc/node_function.md), allowing prompt control of eg. video or audio models with non-text inputs.
All features are fully schedulable unless otherwise stated. See the [syntax documentation](doc/syntax.md) for details on how to use each feature.
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.
@@ -32,33 +37,11 @@ Prompt Control uses graph generation, and tries to delegate functionality to co
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.
## Prompt Control v2
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
### Removed features
- Prompt interpolation syntax; it was too cumbersome to maintain
- LoRA block weight integration; ditto, for now.
### 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.
## 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
@@ -70,10 +53,6 @@ Then restart ComfyUI afterwards.
for example, if you first encode `[cat:dog:0.1]` and later change that to `[cat:dog:0.5]`, no re-encoding takes place.
for added fun, put `NODE(NodeClassName, textinputname)` in a prompt to generate a graph using **any other node** that's compatible. The node can't have required parameters besides a single CLIP parameter (which must be named `clip`) and the text prompt, and it must return a `CONDITIONING` as its first return value. The "default" values are `PCTextEncode` and `text`.
For example, if you for some reason do not want the advanced features of `PCTextEncode`, use `NODE(CLIPTextEncode)` in the prompt and you'll still get scheduling with ComfyUI's regular TE node.
The advanced node enables filtering the prompt for multi-pass workflows.
## PCLazyLoraLoader and PCLazyLoraLoaderAdvanced
@@ -100,8 +79,4 @@ This node configures `PCTextEncode` default values for some functions by attachi
- 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.
+25 -17
View File
@@ -5,33 +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", "hooks"]
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)
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", "anima"]:
mod = importlib.import_module(f".prompt_control.nodes_{node}", package=__name__)
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()
+31 -16
View File
@@ -1,39 +1,54 @@
# 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 `ATTN()` in your negative prompt, and it will work correctly.
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.
## Anima
There is a **very experimental** port of pamparamm's Anima support for Attention Couple in Prompt Control. Because ComfyUI lacks the built-in schedulable hooks required, you must first patch your model with `PC: Anima Attention Couple Model Patch` in addition to using `COUPLE` as usual.
The code was hacked together with minimal thought, so expect bugs and misbehaviour. The port is also currently *not* compatible with NegPIP.
## Syntax
See also the main syntax documentation for `MASK` etc.
See also the [regional prompting documentation](/doc/regional_prompts.md) for information about `MASK` etc.
### ATTN: Trigger Attention Couple
### COUPLE: Trigger Attention Couple
Use `ATTN()` to mark a prompt to be used with Attention Couple. `ATTN()` needs to be combined with either `MASK()` or `IMASK()` to work correctly.
If no mask is specified, an implicit `MASK()` is assumed.
For attention masking to take effect, you need at least two prompt segments with the `ATTN()` marker (separated with `AND`). A single prompt with `ATTN()` will simply ignore the marker.
For the first prompt (and the first prompt only) you can also use `FILL()` to automatically mask all parts not masked by other prompt segments.
You can use `COUPLE` to attach attention-coupled prompts to a base prompt:
For example:
```
dog FILL() ATTN() AND cat MASK(0.5 1) ATTN()
dog FILL() COUPLE(0.5 1) cat
```
If typing `ATTN() MASK()` feels bothersome, try the following macro:
The full syntax looks as follows (to use `IMASK` you need to attach a custom mask)
`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, meaning that the prompt affects the entire image.
- For the base prompt, you can use `FILL()` to automatically mask all parts not masked by other 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:
```
DEF(AM=ATTN() MASK($1))
disabled prompt :0 COUPLE new base prompt COUPLE coupled prompt
```
and then use it like `MASK`: `AM(0 1, 0.5 1)`
You can also schedule the weight normally: `prompt :[1:0:0.35]`
> ![NOTE]
> 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.
+211
View File
@@ -0,0 +1,211 @@
# 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:
- SEGs are processed and the template is expanded
- 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)`
+99
View File
@@ -0,0 +1,99 @@
## 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"
```
## SEG: Split your prompt into named segments
Syntax: `SEG(segment_name)`
To help with organizing prompts, you can use the `SEG` function. For example:
```
This is a comic
Top panel: $CAT. $SEG3
Bottom panel: $DOG
SEG(DOG)
A dog chasing its
tail in a living room.
SEG(CAT)
a sleeping cat
SEG
The cat has orange fur with white stripes
```
This produces:
```
This is a comic
Top panel: a sleeping cat. The cat has orange fur with white stripes
Bottom panel: A dog chasing its
tail in a living room.
```
> [!NOTE]
> Unlike macros, SEGs are processed *after* scheduling syntax has been expanded, except in the LoRA loader (this may change later, but requires a bit of refactoring)
In this case, the first section before any `SEG` becomes the *template* and any text after a `SEG` call becomes part of that segment. Whitespace is stripped from the start and end of segments and the template.
In the template, you can refer to segments by either their index (starting from 1) or the given name, prefixed with a `$SEG`, so in this example, `$SEG1` is the same as `$PANEL1`
Segments can also refer to each other. Recursion will terminate, but produces weird outputs.
Naming segments is optional, in which case you will have to refer to it by its index.
+38
View File
@@ -0,0 +1,38 @@
# The NODE function
The `NODE` function allows you to use any other text encoding node within `PC: Schedule Prompt`, replacing the default `PCTextEncode` and allowing for example video model scheduling.
## Basic usage
Use `NODE(NodeClassName, textinputname)` in a prompt to generate a graph using any node that's compatible. The requirements are as follows:
- The node must have a CLIP parameter (which must be named `clip`)
- It must have a text field
- It must return a `CONDITIONING` as its first return value.
For example, if you for some reason do not want the advanced features of `PCTextEncode`, use `NODE(CLIPTextEncode)` in the prompt and you'll still get scheduling with ComfyUI's regular TE node.
The default parameters are `PCTextEncode` and `text`.
## Advanced Usage with arbitrary parameters
Advanced usage of `NODE` can be complicated. For an example, see [The H3 workflow](/example_workflows/Prompt%20Control%20with%20MiniMax%20H3.json?raw=1). You can also find it in the template library.
The full synopsis of the function is `NODE(NodeClassName, textinputname, arg_spec)` where `arg_spec` is a semicolon-separated list of `parameter_name json_value` pairs. In raw form, it looks like this:
```
NODE(MiniMaxH3ImageToVideo, prompt, vae ["1", 0]; width 1024; height 1024; first_frame ["2", 0])
```
The names and inputs must match the ComfyUI API format which **may differ from frontend names**. You can export your workflow in API format and inspect it to see how inputs are passed in to nodes.
The arrays are literal ComfyUI node links, meaning the `vae` parameter is taken from node ID "1" first output and `first_frame` from node ID "2" first output.
The values are arbitrary JSON literals, meaning that you can also pass in constant values. To pass in literal strings for example, you need to use quotes `"like this"`.
This is intended to be used with the helper node `PC: NODE Input Helper`, which can be used to pass arbitrary parameters (named `$a` to `$n`) to the encoder. The recommended pattern is to put something like the following:
```
SEG(node)
NODE(MiniMaxH3ImageToVideo, prompt, vae $a; width $b; height $c; length $d; first_frame $e)
```
to the helper and then concatenate it at the end of your prompt (use whitespace as a separator). You can then trigger the node with `$node` in your prompt. (see [documentation](/doc/macros.md) for `SEG`)
The helper will replace the parameters with the correct ComfyUI link values.
+63
View File
@@ -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.
+117
View File
@@ -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
-400
View File
@@ -1,400 +0,0 @@
# Prompt Control Syntax
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).
Scheduling syntax is similar to A1111, 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. 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 `[before:during:after:0.3,0.7]`, The prompt be `a before` until 0.3, `a during` until 0.7, and then `a after`. This form is equivalent to `[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.
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 with the `PCTextEncode` node, where applicable.
## 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 to separate different prompts; see `MASK` and `ATTN` 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.
## 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
- Prompts are split by AND
- Most functions (like STYLE, MASK) and cutoffs are evaluated
- prompts are split by AVG()
- prompts are split by CAT
- the TE() function is evaluated to set per-encoder prompts
- BREAK is evaluated
- Everything else
## 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.
### 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 also [Attention Couple](#attention-couple) below
### 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 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.
## 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
Experimental features are unstable and may disappear or change without warning.
## 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"
```
## ATTN: 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)`
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+49 -50
View File
@@ -1,9 +1,9 @@
import torch
import numpy as np
from math import copysign
import logging
import itertools
from .adv_encode_old import old_advanced_encode_from_tokens
import logging
from math import copysign
import numpy as np
import torch
log = logging.getLogger("comfyui-prompt-control")
@@ -26,7 +26,7 @@ def _grouper(n, iterable):
def batched_clip_encode(tokens, length, encode_func, num_chunks):
embs = []
for e in _grouper(32, tokens):
enc, pooled = encode_func(e)
enc, pooled, *_ = encode_func(e)
enc = enc.reshape((len(e), length, -1))
embs.append(enc)
@@ -42,12 +42,18 @@ def weights_like(weights, emb):
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)
@@ -101,24 +107,24 @@ def style_comfy(encoder, tokens, **kwargs):
def style_a1111(encoder, tokens, **kwargs):
base_emb, pooled = encoder.base_emb(tokens)
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
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 = encoder.encode_fn(pos_tokens)
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
return (weighted_emb, pooled) + tuple(extra)
def style_comfypp(encoder, tokens, **kwargs):
unweighted_tokens = encoder.unweighted(tokens)
base_emb, pooled_base = encoder.base_emb(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
)
@@ -132,23 +138,23 @@ def style_comfypp(encoder, tokens, **kwargs):
)
weighted_emb += embs
return weighted_emb, pooled
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 = encoder.base_emb(tokens)
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
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 = encoder.base_emb(tokens)
return perp_weight(encoder.weights(tokens), (base_emb, pooled), (zero_emb, zero_pooled))
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):
@@ -161,7 +167,7 @@ def apply_negpip(encoder, emb, pooled, **kwargs):
def norm_length(encoder, tokens, **kwargs):
word_ids = encoder.word_ids(tokens)
sums = dict(zip(*np.unique(word_ids, return_counts=True)))
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
@@ -170,7 +176,9 @@ def norm_length(encoder, tokens, **kwargs):
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) for w, id in zip(x, y) if id != 0])
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
@@ -198,6 +206,7 @@ class AdvancedEncoder:
def add_encoder(cls, name, fn):
cls.STYLES[name] = fn
@classmethod
def add_normalization_op(cls, name, fn):
cls.NORMALIZATION_OPS[name] = fn
@@ -258,8 +267,8 @@ class AdvancedEncoder:
if negpip:
def _encode(t):
emb, pooled = encode_fn(t)
return emb[:, 0::2, :], pooled
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))
@@ -289,7 +298,7 @@ class AdvancedEncoder:
if w[i] >= 1:
continue
masked_current = mask_inds(masked_current, np.where(w_inv == i)[0], self.m_token)
masked, _ = self.encode_fn(masked_current)
masked, _, *extra = self.encode_fn(masked_current)
emblist.append(masked)
embs = torch.cat(emblist)
@@ -297,7 +306,7 @@ class AdvancedEncoder:
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)
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, :]
@@ -305,7 +314,9 @@ class AdvancedEncoder:
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]) if w != 1.0)
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
@@ -329,12 +340,13 @@ class AdvancedEncoder:
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(axis=0, keepdim=True)
pooled = pooled.mean(dim=0, keepdim=True)
pooled = pooled_base + pooled
if embs.shape[0] != masks.shape[0]:
@@ -349,16 +361,16 @@ class AdvancedEncoder:
for op in self.preprocessors:
normalized_tokens = op(self, normalized_tokens)
emb, pooled = self.weight_fn(self, normalized_tokens, original_tokens=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 return_pooled:
if not apply_to_pooled:
_, pooled = self.base_emb(tokens)
return emb, pooled
return emb, None
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(
@@ -373,20 +385,7 @@ def advanced_encode_from_tokens(
tokenizer=None,
**extra_args,
):
if "old+" not in weight_interpretation:
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)
else:
weight_interpretation = weight_interpretation.replace("old+", "")
log.warning("Using old implementation of %s", weight_interpretation)
return old_advanced_encode_from_tokens(
tokenized,
token_normalization,
weight_interpretation,
encode_func,
266,
return_pooled=return_pooled,
apply_to_pooled=apply_to_pooled,
)
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)
-235
View File
@@ -1,235 +0,0 @@
import torch
import numpy as np
import logging
import itertools
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):
it = iter(iterable)
while True:
chunk = list(itertools.islice(it, n))
if not chunk:
return
yield chunk
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)]
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)]
mask = np.array(word_ids) == target_id
return (new_tokens, mask)
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 = torch.tensor(weights, dtype=base_emb.dtype, device=base_emb.device)
weight_tensor = weight_tensor.reshape(1, -1, 1).expand(base_emb.shape)
# 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)
m = torch.tensor(m, dtype=base_emb.dtype, device=base_emb.device)
m = m.reshape(1, -1, 1).expand(base_emb.shape)
masks.append(m)
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)
new_tokens = [
[mask_token if i * clip_len + j in inds_set else t for j, t in enumerate(x)] for i, x in enumerate(tokens)
]
return new_tokens
def down_weight(tokens, weights, word_ids, base_emb, length, encode_func):
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 = (266, 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)
embeddings_final = (norm_base / norm_weighted) * weighted_emb
return embeddings_final
# For verification
def A1111_renorm(base_emb, weighted_emb):
embeddings_final = (base_emb.mean() / weighted_emb.mean()) * weighted_emb
return embeddings_final
def from_zero(weights, base_emb):
weight_tensor = torch.tensor(weights, dtype=base_emb.dtype, device=base_emb.device)
weight_tensor = weight_tensor.reshape(1, -1, 1).expand(base_emb.shape)
return base_emb * weight_tensor
def old_advanced_encode_from_tokens(
tokenized,
token_normalization,
weight_interpretation,
encode_func,
m_token=266,
w_max=1.0,
return_pooled=False,
apply_to_pooled=False,
**extra_args,
):
length = 77
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]
# weight normalization
# ====================
# distribute down/up weights over word lengths
if token_normalization.startswith("length"):
weights = divide_length(word_ids, weights)
# make mean of word tokens 1
if token_normalization.endswith("mean"):
weights = shift_mean_weight(word_ids, weights)
# weight interpretation
# =====================
pooled = None
if weight_interpretation in ["comfy", "perp"]:
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 = from_zero(weights, base_emb)
weighted_emb = A1111_renorm(base_emb, weighted_emb)
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 return_pooled:
if apply_to_pooled:
return weighted_emb, pooled
else:
return weighted_emb, pooled_base
return weighted_emb, None
+167
View File
@@ -0,0 +1,167 @@
# Adapted from https://github.com/pamparamm/ComfyUI-ppm
import itertools
from collections.abc import Callable
from functools import partial
from math import lcm
import torch
import torch.nn.functional as F
from comfy.ldm.anima.model import Anima as AnimaDIT
from comfy.ldm.cosmos.predict2 import Attention as CosmosAttention
from comfy.patcher_extension import WrapperExecutor
from comfy.sampler_helpers import convert_cond
from comfy.samplers import process_conds
COND = 0
UNCOND = 1
def reshape_mask(mask: torch.Tensor, size: tuple[int, int], bs: int, num_tokens: int) -> torch.Tensor:
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(bs, dim=0)
return mask_downsample_reshaped
def wrap_forwards(anima_model):
backups = {}
for block_name, b in (
(n, b) for n, b in anima_model.named_modules() if "cross_attn" in n and isinstance(b, CosmosAttention)
):
backups[block_name] = b.forward
b.forward = partial(cosmos_attention_forward_couple, b.forward)
return backups
def unwrap_forwards(anima_model, backups):
for block_name, b in (
(n, b) for n, b in anima_model.named_modules() if "cross_attn" in n and isinstance(b, CosmosAttention)
):
b.forward = backups[block_name]
def anima_sample_wrapper(executor, *args, **kwargs):
guider, _, extra_options, _, noise, latent_image, denoise_mask, *_ = args
seed = extra_options["seed"]
device = "cuda" # TODO: fix
def pc_process_conds(pc_conds):
conds = [convert_cond([c])[0] for c in pc_conds]
conds = process_conds(
guider.inner_model,
noise,
{"positive": conds},
device,
latent_image,
denoise_mask,
seed,
latent_shapes=[latent_image.shape],
)
return [
c["model_conds"]["c_crossattn"].cond * pc_conds[i][1].get("strength", 1.0)
for i, c in enumerate(conds["positive"])
]
extra_options["model_options"]["transformer_options"]["pc_process_conds"] = pc_process_conds
return executor(*args, **kwargs)
def anima_forward_wrapper(executor: WrapperExecutor, *args, **kwargs):
"""Model wrapper does something with activation shapes?"""
anima_model: AnimaDIT = executor.class_obj # type: ignore
x: torch.Tensor = args[0]
transformer_options: dict = kwargs.get("transformer_options", {}).copy()
pc = transformer_options.get("pc_couple")
if pc and "processed_conds" not in pc:
pc["processed_conds"] = transformer_options["pc_process_conds"](pc["conds"])
patch_spatial = anima_model.patch_spatial
activations_shape = list(x.shape)
activations_shape[-2] = activations_shape[-2] // patch_spatial
activations_shape[-1] = activations_shape[-1] // patch_spatial
transformer_options["activations_shape"] = activations_shape
kwargs["transformer_options"] = transformer_options
b = {}
if pc:
b = wrap_forwards(anima_model)
r = executor(*args, **kwargs)
if pc:
unwrap_forwards(anima_model, b)
return r
def cosmos_attention_forward_couple(_forward: Callable, x, context, rope_emb, transformer_options):
"""attention block wrapper"""
if "pc_couple" not in transformer_options:
return _forward(x, context, rope_emb, transformer_options)
c: torch.Tensor = context
# FIXME: base cond weight
# c = args["processed_conds"][0]
args = transformer_options["pc_couple"]
mask = args["mask"]
conds = args["processed_conds"][1:]
num_conds = len(conds) + 1
num_tokens_c: list[int] = [c.shape[1] for c in conds]
cond_or_uncond = transformer_options["cond_or_uncond"]
cond_or_uncond_couple = []
num_chunks = len(cond_or_uncond)
bs = x.shape[0] // num_chunks
x_chunks = x.chunk(num_chunks, dim=0)
c_chunks = c.chunk(num_chunks, dim=0)
lcm_tokens_c = lcm(c.shape[1], *num_tokens_c)
conds_c_tensor = torch.cat(
[cond.repeat(bs, lcm_tokens_c // num_tokens_c[i], 1) for i, cond in enumerate(conds)],
dim=0,
)
xs, cs = [], []
for i, cond_type in enumerate(cond_or_uncond):
x_target = x_chunks[i]
c_target = c_chunks[i].repeat(1, lcm_tokens_c // c.shape[1], 1)
if cond_type == UNCOND:
xs.append(x_target)
cs.append(c_target)
cond_or_uncond_couple.append(UNCOND)
else:
xs.append(x_target.repeat(num_conds, 1, 1))
cs.append(torch.cat([c_target, conds_c_tensor], dim=0))
cond_or_uncond_couple.extend(itertools.repeat(COND, num_conds))
xs = torch.cat(xs, dim=0)
cs = torch.cat(cs, dim=0)
out = _forward(xs, cs, rope_emb, transformer_options)
size = tuple(transformer_options["activations_shape"][-2:])
num_tokens = out.shape[1]
mask_downsample = reshape_mask(mask, size, bs, num_tokens)
outputs = []
cond_outputs = []
i_cond = 0
for i, cond_type in enumerate(cond_or_uncond_couple):
pos, next_pos = i * bs, (i + 1) * bs
if cond_type == 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)
+40 -26
View File
@@ -9,7 +9,6 @@ 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
@@ -30,7 +29,7 @@ def set_cond_attnmask(base_cond, extra_conds, fill=False):
group = HookGroup()
group.add(hook)
return set_hooks_for_conditioning(c, hooks=group)
return set_hooks_for_conditioning(c, hooks=group, append_hooks=True)
def get_mask(mask, batch_size, num_tokens, extra_options):
@@ -64,26 +63,23 @@ class AttentionCoupleHook(TransformerOptionsHook):
def __init__(self):
super().__init__(hook_scope=EnumHookScope.HookedOnly)
self.transformers_dict = {
self.transformers_dict: dict[str, Any] = {
"patches": {
"attn2_output_patch": [Proxy(self.attn2_output_patch)],
"attn2_patch": [Proxy(self.attn2_patch)],
}
},
"pc_couple": {},
}
self.has_negpip = False
# calculate later
self.conds_k: list[torch.Tensor] = None
self.conds_v: list[torch.Tensor] = None
self.has_negpip = False
# The list will be calculated later. All clones must refer to the same kv dict
self.kv: dict[str, list] = {"k": None, "v": None} # type: ignore
def initialize_regions(self, base_cond, conds, fill):
self._base_cond = base_cond
self._conds = conds
self._fill = fill
self.num_conds = len(conds) + 1
self.base_strength = base_cond[1].get("strength", 1.0)
self.strengths = [cond[1].get("strength", 1.0) for cond in conds]
self.strengths: list[float] = [cond[1].get("strength", 1.0) for cond in conds]
self.comfy_conds = [base_cond] + 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]
@@ -99,7 +95,6 @@ class AttentionCoupleHook(TransformerOptionsHook):
largest_shape = max(m.shape for m in masks)
if base_mask is not None:
largest_shape = max(largest_shape, base_mask.shape)
print("largest shape x", largest_shape, [m.shape for m in masks], 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)
@@ -124,27 +119,41 @@ class AttentionCoupleHook(TransformerOptionsHook):
self.mask = mask / mask.sum(dim=0, keepdim=True)
def on_apply_hooks(self, model: ModelPatcher, transformer_options: dict[str, Any]):
if self.conds_k is None:
self.transformers_dict["pc_couple"] = {
"conds": self.comfy_conds,
"num_conds": self.num_conds,
"mask": self.mask,
}
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.conds_k = [cond[:, 0::2] for cond in self.conds[1:]]
self.conds_v = [cond[:, 1::2] for cond in self.conds[1:]]
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.conds_k = self.conds_v = self.conds[1:]
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.initialize_regions(self._base_cond, self._conds, self._fill)
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):
@@ -152,8 +161,15 @@ class AttentionCoupleHook(TransformerOptionsHook):
cond_or_uncond_couple = extra_options[self.COND_UNCOND_COUPLE_OPTION] = list(cond_or_uncond)
num_chunks = len(cond_or_uncond)
lcm_tokens_k = math.lcm(k.shape[1], *(cond.shape[1] for cond in self.conds_k))
lcm_tokens_v = math.lcm(v.shape[1], *(cond.shape[1] for cond in self.conds_v))
# 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)
@@ -161,17 +177,14 @@ class AttentionCoupleHook(TransformerOptionsHook):
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(self.conds_k)
],
[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(self.conds_v)
for i, cond in enumerate(conds_v)
],
dim=0,
)
@@ -208,6 +221,7 @@ class AttentionCoupleHook(TransformerOptionsHook):
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)
-47
View File
@@ -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
)
+8 -9
View File
@@ -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)
@@ -215,7 +214,7 @@ def encode_regions(clip_regions, encode, tokenizer):
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
+25
View File
@@ -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
+128
View File
@@ -0,0 +1,128 @@
# vim: sw=4 ts=4
from __future__ import annotations
import logging
import re
from .utils import find_closing_paren, get_function, split_by_function
log = logging.getLogger("comfyui-prompt-control")
def substitute_template(template, segments, do_subs):
def _substitute(template, segments, stack):
name = ""
if "$" in template:
for name, value in sorted(segments):
value = substitute_var(value, name, "")
if name not in stack:
stack.add(name)
value = _substitute(value, segments, stack)
stack.remove(name)
template = substitute_var(template, name, value)
if do_subs and name not in stack:
template = expand_subs(template)
return template
return _substitute(template, segments, set())
def expand_segs(text, do_subs=True):
template, segments = split_by_function(text, "SEG", defaults=[""], require_args=True)
named_segs = [(f.args[0].strip() or f"SEG{i + 1}", c.strip()) for i, (c, f) in enumerate(segments)]
new_text = substitute_template(template, named_segs, do_subs).strip()
if new_text != text.strip():
log.debug("Template expanded to: %s", new_text)
return new_text
def expand_subs(text):
text, subs = get_function(text, "SUB", defaults=None)
subs = [spec.strip() for f in subs for spec in f.args[0].split(";")]
for spec in subs:
if len(spec) <= 3 or spec[0] != "s":
log.warning("Invalid SUB spec ignored: '%s'", spec)
continue
splitchar = spec[1]
search, replace, *_ = spec[2:].split(splitchar)
text = re.sub(search, replace, text)
return text
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!")
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.debug("DEFs expanded to: %s", res)
return res
def substitute_var(text, name, replace, boundary=r"\b"):
if f"${name}" not in text:
return text
name = re.escape(str(name))
return re.sub(rf"\${name}{boundary}", replace, text)
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
end_re = r"(?![0-9])"
for i, v in enumerate(paramvals):
r = substitute_var(r, i + 1, v, boundary=end_re)
for i, v in enumerate(default_args):
r = substitute_var(r, i + 1, v, boundary=end_re)
text = text.replace(ph, r)
return text
-46
View File
@@ -1,46 +0,0 @@
import main
import nodes
import prompt_control.adv_encode
(l,) = nodes.CLIPLoader.load_clip(None, "clip_l.safetensors")
(t5,) = nodes.CLIPLoader.load_clip(None, "t5base.safetensors")
id(main) # get rid of warning
def adv(t, text, style="A1111", norm="none", new=True, **kwargs):
c = t.tokenize(text, return_word_ids=True)
if new:
style = "new+" + style
if t is t5:
te = t.patcher.model.t5base.encode_token_weights
token = t.tokenizer.clip_t5base
tok = c["t5base"]
else:
te = t.patcher.model.clip_l.encode_token_weights
token = t.tokenizer.clip_l
tok = c["l"]
return prompt_control.adv_encode.advanced_encode_from_tokens(tok, norm, style, te, tokenizer=token)
def adv_all(t, text, styles=[], **kwargs):
r = []
for s in styles or prompt_control.adv_encode.AdvancedEncoder.STYLES:
print("Testing", s, kwargs)
r.append([s, adv(t, text, style=s, **kwargs)])
return r
def replacenan(t):
t[t.isnan()] = 42.123321
return t
def adv_equal(t, text, **kwargs):
old = adv_all(t, text, new=False, **kwargs)
new = adv_all(t, text, new=True, **kwargs)
r = {}
for i, o in enumerate(old):
n = new[i]
r[n[0]] = (replacenan(n[1][0]) == replacenan(o[1][0])).all()
return r
+52
View File
@@ -0,0 +1,52 @@
# Adapted from ComfyUI-ppm into hook form
import comfy.model_management
import comfy.patcher_extension
from comfy.model_base import Anima
from comfy.model_patcher import ModelPatcher
from comfy_api.latest import io
from .anima_couple import (
anima_forward_wrapper,
anima_sample_wrapper,
)
class PCAnimaAttnCouplePatch(io.ComfyNode):
@classmethod
def define_schema(cls) -> io.Schema:
return io.Schema(
node_id="PCAnimaAttnCouplePatch",
display_name="PC: Anima attention Couple Model Patch",
category="promptcontrol/experimental",
inputs=[
io.Model.Input("model"),
],
outputs=[
io.Model.Output(),
],
)
@classmethod
def execute(cls, model: ModelPatcher) -> io.NodeOutput:
model_type = type(model.model)
m = model
if issubclass(model_type, Anima):
m = model.clone()
m.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL,
cls.__name__,
anima_forward_wrapper,
)
m.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE,
cls.__name__,
anima_sample_wrapper,
)
return io.NodeOutput(m)
NODES = [PCAnimaAttnCouplePatch]
+45 -34
View File
@@ -1,51 +1,62 @@
import logging
from comfy_api.latest import io
from .macros import expand_segs
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:
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),)
text = expand_segs(text)
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:
# 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,
]
+49 -51
View File
@@ -3,7 +3,8 @@ import logging
import comfy.hooks
import comfy.utils
import folder_paths
from comfy.comfy_types.node_typing import IO, ComfyNodeABC, InputTypeDict
from comfy_api.latest import io
from typing_extensions import override
from .attention_couple_ppm import AttentionCoupleHook
from .parser import parse_prompt_schedules
@@ -12,24 +13,27 @@ 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:
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):
@@ -37,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():
@@ -52,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)
@@ -74,43 +77,43 @@ 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
class PCAttentionCoupleBatchNegative(ComfyNodeABC):
class PCAttentionCoupleBatchNegative(io.ComfyNode):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
return {
"required": {
"positive": (IO.CONDITIONING, {}),
"negative": (IO.CONDITIONING, {}),
},
}
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"),
],
)
RETURN_TYPES = (IO.CONDITIONING, IO.CONDITIONING)
RETURN_NAMES = ("positive", "negative")
CATEGORY = "promptcontrol/v2"
FUNCTION = "batch"
EXPERIMENTAL = True
# May cause side-effects?
# TODO: Support scheduling in negative prompt
def batch(self, positive, negative):
@classmethod
@override
def execute(cls, positive, negative) -> io.NodeOutput:
if len(negative) != 1:
log.warning("Batching scheduled negatives is not supported yet")
return (positive, negative)
return io.NodeOutput(positive, negative)
negative_batch = []
for p in positive:
n = [negative[0][0], negative[0][1].copy()]
n_hook_group: comfy.hooks.HookGroup = n[1].get("hooks", comfy.hooks.HookGroup()).clone()
p_hook_group: comfy.hooks.HookGroup = p[1].get("hooks", comfy.hooks.HookGroup())
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)
@@ -119,15 +122,10 @@ class PCAttentionCoupleBatchNegative(ComfyNodeABC):
n[1]["end_percent"] = p[1].get("end_percent", 1.0)
negative_batch.append(n)
return (positive, negative_batch)
return io.NodeOutput(positive, negative_batch)
NODE_CLASS_MAPPINGS = {
"PCLoraHooksFromText": PCLoraHooksFromText,
"PCAttentionCoupleBatchNegative": PCAttentionCoupleBatchNegative,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"PCLoraHooksFromText": "PC: LoRA Hooks From Text (non-lazy)",
"PCAttentionCoupleBatchNegative": "PC: Attention Couple (batch negative)",
}
NODES = [
PCLoraHooksFromText,
PCAttentionCoupleBatchNegative,
]
+133 -122
View File
@@ -1,32 +1,19 @@
# pyright: reportSelfClsParameterName=false
from __future__ import annotations
import json
import logging
from .parser import parse_prompt_schedules
from comfy_execution.graph_utils import GraphBuilder, is_link
from comfy_api.latest import io
from comfy_execution.graph import ExecutionBlocker
from comfy_execution.graph_utils import GraphBuilder
from .utils import get_function
from .macros import expand_macros, expand_segs
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():
@@ -145,64 +132,64 @@ def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True):
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 {
"optional": {
"model": ("MODEL", {"rawLink": True}),
"clip": ("CLIP", {"rawLink": True}),
"text": ("STRING", {"multiline": True, "default": ""}),
"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}),
"num_steps": ("INT", {"min": 0, "max": 10000, "default": 0, "step": 1}),
},
"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, unique_id, model=None, clip=None, text="", apply_hooks=True, tags="", start=0.0, end=1.0, num_steps=0
):
def execute(cls, model=None, clip=None, text="", apply_hooks=True, tags="", start=0.0, end=1.0, num_steps=0):
text = expand_segs(expand_macros(text))
schedule = parse_prompt_schedules(text, filters=tags, start=start, end=end, num_steps=num_steps)
graph = GraphBuilder(f"{unique_id}-")
graph = GraphBuilder()
r = build_lora_schedule(graph, schedule, model, clip, apply_hooks=apply_hooks)
return r
class PCLazyLoraLoader(PCLazyLoraLoaderAdvanced):
class PCLazyLoraLoader(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {
"optional": {
"model": ("MODEL", {"rawLink": True}),
"clip": ("CLIP", {"rawLink": True}),
"text": ("STRING", {"multiline": True, "default": ""}),
},
"hidden": {"unique_id": "UNIQUE_ID"},
}
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"),
],
)
RETURN_TYPES = (
"MODEL",
"CLIP",
)
CATEGORY = "promptcontrol"
def apply(self, *args, **kwargs):
r = super().apply(*args, **kwargs)
r["result"] = r["result"][:2]
return r
@classmethod
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):
@@ -210,15 +197,38 @@ def build_scheduled_prompts(graph, schedules, clip):
start_pct = 0.0
for end_pct, c in schedules:
p = c["prompt"]
p, classnames = get_function(p, "NODE", ["PCTextEncode", "text"])
classname = "PCTextEncode"
paramname = "text"
p, classnames = get_function(p, "NODE", defaults=None)
realargs = ["PCTextEncode", "text", ""]
if classnames:
classname = classnames[0][0]
paramname = classnames[0][1]
node = graph.node(classname)
# Need to explicitly expand segs for custom node
p = expand_segs(p)
args = classnames[0].args[0]
if not args.strip():
raise ValueError("NODE can't be empty!")
for i, v in enumerate(args.split(",", maxsplit=2)):
realargs[i] = v
classname, paramname, magic_spec = realargs
node = graph.node(classname.strip())
node.set_input("clip", clip)
node.set_input(paramname, p)
node.set_input(paramname.strip(), p)
magic_spec.replace(r"\;", "__ESCAPED_SEMICOLON__")
extra_inputs = magic_spec.split(";") if magic_spec.strip() else []
for e in extra_inputs:
e = e.strip()
if not e:
continue
e = e.replace("__ESCAPED_SEMICOLON__", ";")
name, jsondata = e.split(maxsplit=1)
jsondata = jsondata.strip()
if not jsondata.strip():
continue
# From helper node:
if jsondata == "__EMPTY__":
continue
try:
node.set_input(name.strip(), json.loads(jsondata.strip()))
except ValueError as e:
raise ValueError(f"Invalid JSON input: '{jsondata}'") from e
timestep = graph.node("ConditioningSetTimestepRange")
timestep.set_input("conditioning", node.out(0))
timestep.set_input("start", start_pct)
@@ -235,61 +245,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, num_steps=0, **kwargs):
schedules = parse_prompt_schedules(text, filters=tags, start=start, end=end, num_steps=num_steps)
return [(pct, s[cachekey]) for pct, s in schedules]
class PCLazyTextEncodeAdvanced:
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})},
"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}),
"num_steps": ("INT", {"min": 0, "max": 10000, "default": 0, "step": 1}),
},
"hidden": {"unique_id": "UNIQUE_ID"},
}
RETURN_TYPES = ("CONDITIONING",)
CATEGORY = "promptcontrol"
FUNCTION = "apply"
def apply(self, clip, text, unique_id, tags="", start=0.0, end=1.0, num_steps=0):
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(f"{unique_id}-")
graph = GraphBuilder()
return build_scheduled_prompts(graph, schedules, clip)
class PCLazyTextEncode(PCLazyTextEncodeAdvanced):
class PCLazyTextEncode(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP", {"rawLink": True}), "text": ("STRING", {"multiline": True})},
"hidden": {"unique_id": "UNIQUE_ID"},
}
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"),
],
)
CATEGORY = "promptcontrol"
@classmethod
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,
]
+161 -188
View File
@@ -1,166 +1,111 @@
import logging
from .parser import parse_prompt_schedules, expand_macros
from .nodes_lazy import NODE_CLASS_MAPPINGS as LAZY_NODES
import json
import folder_paths
from pathlib import Path
from comfy_execution.graph_utils import is_link
import logging
from comfy_api.latest import io
from .macros import expand_macros as macroexpand
from .macros import expand_segs as segexpand
from .macros import expand_subs as subexpand
from .macros import substitute_var
from .parser import parse_prompt_schedules
log = logging.getLogger("comfyui-prompt-control")
class PCSaveExpandedWorkflow:
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
class PCSetLogLevel(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"any": ("*", {}),
},
"hidden": {
"prompt": "DYNPROMPT",
},
}
@classmethod
def VALIDATE_INPUTS(self, input_types):
return True
OUTPUT_NODE = True
RETURN_TYPES = ()
CATEGORY = "promptcontrol/tools"
DESCRIPTION = "Saves the current expanded dynamic prompt into a JSON file"
FUNCTION = "apply"
def apply(self, any, prompt):
full_output_folder, filename, counter, subfolder, prefix = folder_paths.get_save_image_path(
"pc_workflow_debug", self.output_dir
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()],
)
p = {}
input_replace_map = {}
for node in prompt.all_node_ids():
n = prompt.get_node(node)
t = n["class_type"]
if t in LAZY_NODES:
expanded_prompt = LAZY_NODES[t]().apply(**n["inputs"], unique_id=node)
for k in expanded_prompt["expand"]:
p[k] = expanded_prompt["expand"][k]
for i, _ in enumerate(expanded_prompt["result"]):
input_replace_map[(node, i)] = [k, i]
else:
p[node] = n
for k in p:
for ik in p[k]["inputs"]:
x = p[k]["inputs"][ik]
if is_link(x) and tuple(x) in input_replace_map:
p[k]["inputs"][ik] = input_replace_map[tuple(x)]
file = f"{filename}_{counter:05}_.json"
full_path = Path(full_output_folder) / file
with open(full_path, "w") as f:
log.info(f"Saving workflow to {full_path}")
json.dump(p, f)
return ()
class PCSetLogLevel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"clip": ("CLIP",),
},
"optional": {
"level": (["INFO", "DEBUG", "WARNING", "ERROR"], {"default": "INFO"}),
},
}
def apply(self, clip, level="INFO"):
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"
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."
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"
DESCRIPTION = "Multi-input version of PCAddMaskToCLIP, for convenience"
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": {
"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"
DESCRIPTION = "Configures default values for PCTextEncode"
def apply(
self,
@classmethod
def execute(
cls,
clip,
mask_width=512,
mask_height=512,
@@ -170,7 +115,7 @@ class PCSetPCTextEncodeSettings:
sdxl_target_h=1024,
sdxl_crop_w=0,
sdxl_crop_h=0,
):
) -> io.NodeOutput:
settings = {
"mask_width": mask_width,
"mask_height": mask_height,
@@ -183,66 +128,94 @@ class PCSetPCTextEncodeSettings:
}
clip = clip.clone()
clip.patcher.model_options["x-promptcontrol.settings"] = settings
return (clip,)
return io.NodeOutput(clip)
class PCExtractScheduledPrompt:
class PCExtractScheduledPrompt(io.ComfyNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"multiline": True}),
"at": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 1.0, "step": 0.01}),
},
"optional": {"tags": ("STRING", {"default": ""})},
}
def define_schema(cls):
return io.Schema(
node_id="PCExtractScheduledPrompt",
display_name="PC: Show 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),
io.Boolean.Input("expand_segs", default=False, optional=True),
io.Boolean.Input("expand_subs", default=False, optional=True),
io.Boolean.Input("expand_macros", default=False, optional=True),
],
outputs=[io.String.Output()],
search_aliases=["extract scheduled prompt"],
)
RETURN_TYPES = ("STRING",)
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Parses the input prompt and returns the prompt scheduled at the specified point"
def apply(self, text, at, tags=""):
@classmethod
def execute(cls, text, at, tags="", expand_segs=False, expand_subs=False, expand_macros=False) -> io.NodeOutput:
if expand_macros:
text = macroexpand(text)
schedule = parse_prompt_schedules(text, filters=tags)
_, entry = schedule.at_step(at, total_steps=1)
_, entry = schedule.at_step(at)
prompt_text = entry.get("prompt", "")
return (prompt_text,)
if expand_segs:
prompt_text = segexpand(prompt_text, do_subs=expand_subs)
if expand_subs:
prompt_text = subexpand(prompt_text)
return io.NodeOutput(prompt_text)
class PCMacroExpand:
class PCMacroExpand(io.ComfyNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"multiline": True}),
},
}
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()],
)
RETURN_TYPES = ("STRING",)
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Expands DEF macros in a string and returns the result"
def apply(self, text):
return (expand_macros(text),)
@classmethod
def execute(cls, text) -> io.NodeOutput:
return io.NodeOutput(macroexpand(text))
NODE_CLASS_MAPPINGS = {
"PCSetPCTextEncodeSettings": PCSetPCTextEncodeSettings,
"PCAddMaskToCLIP": PCAddMaskToCLIP,
"PCAddMaskToCLIPMany": PCAddMaskToCLIPMany,
"PCSetLogLevel": PCSetLogLevel,
"PCExtractScheduledPrompt": PCExtractScheduledPrompt,
"PCSaveExpandedWorkflow": PCSaveExpandedWorkflow,
"PCMacroExpand": PCMacroExpand,
}
class PCLinkHelper(io.ComfyNode):
@classmethod
def define_schema(cls):
inputs: list = [io.String.Input("text", multiline=True)]
for x in "abcdefghijklmn":
inputs.append(io.AnyType.Input(x, optional=True, lazy=True, extra_dict={"rawLink": True}))
return io.Schema(
node_id="PCNODELinkHelper",
display_name="PC: Extra argument helper for NODE",
category="promptcontrol/tools",
description="Takes in arbitrary inputs and renders them as NODE-compatible values, replacing $a -> $n with JSON link values",
is_experimental=True,
inputs=inputs,
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)",
"PCExtractScheduledPrompt": "PC: Extract Scheduled Prompt",
"PCSaveExpandedWorkflow": "PC: Save Expanded Workflow (for debug)",
"PCMacroExpand": "PC: Expand Macros",
}
@classmethod
def execute(cls, text, **vars) -> io.NodeOutput:
for k in "abcdefghijklmn":
v = "__EMPTY__"
if k in vars:
v = json.dumps(vars[k])
text = substitute_var(text, k, v)
return io.NodeOutput(text)
NODES = [
PCSetPCTextEncodeSettings,
PCAddMaskToCLIP,
PCAddMaskToCLIPMany,
PCSetLogLevel,
PCExtractScheduledPrompt,
PCMacroExpand,
PCLinkHelper,
]
+6 -444
View File
@@ -1,448 +1,10 @@
# vim: sw=4 ts=4
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, find_closing_paren
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 "]"
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",
)
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, int, type(None)] 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, 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, steps 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):
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):
# 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
self.prompt = 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]
def parse_search(search):
arg_start = search.find("(")
args = ""
name = search.strip()
if arg_start > 0:
arg_end = find_closing_paren(search, arg_start)
name = search[:arg_start].strip()
args = search[arg_start + 1 : arg_end - 1]
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
if arg_start > 0:
args = [a.strip() for a in args.split(";")]
else:
args = []
return name, args
def expand_macros(text):
text, defs = get_function(text, "DEF", defaults=None)
res = text
prevres = text
replacements = []
for d in defs:
r = d.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)
res = substitute_def(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_def(text, search, replace):
search, default_args = search
for i, v in enumerate(default_args):
replace = re.sub(rf"\${i+1}\b", v, replace)
return re.sub(rf"\b{re.escape(search)}\b", replace, text)
def substitute_defcall(text, search, replace):
name, default_args = search
text, defns = get_function(text, name, defaults=None, placeholder=f"DEFNCALL{search}")
for i, defn in enumerate(defns):
ph = f"\0DEFNCALL{search}{i}\0"
paramvals = [x.strip() for x in defn.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
@lru_cache
def parse_prompt_schedules(prompt, **kwargs):
prompt = expand_macros(prompt)
return PromptSchedule(prompt, **kwargs)
if os.environ.get("PC_USE_OLD_PARSER", "0") != "1":
from .parser_parsy import parse_prompt_schedules # noqa
else:
log.warning("Using old Lark parser (UNSUPPORTED)")
from .parser_lark import parse_prompt_schedules # noqa
+359
View File
@@ -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)
+399
View File
@@ -0,0 +1,399 @@
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.many() + 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
| char_from("<>").map(Text)
)
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)
+720
View File
@@ -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__
+180 -133
View File
@@ -1,16 +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 .cutoff import process_cuts
from .parser import parse_cuts
from .attention_couple_ppm import set_cond_attnmask
from .cutoff import process_cuts
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")
@@ -20,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+")
@@ -42,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:
@@ -58,11 +73,11 @@ 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.replace("old+", "") not in AVAILABLE_STYLES:
@@ -77,8 +92,9 @@ def get_style(text, default_style="comfy", default_normalization="none"):
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":
@@ -112,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])):
@@ -128,11 +145,11 @@ def fix_word_ids(tokens):
def tokenize_chunks(clip, text, need_word_ids, can_break):
chunks = re.split(r"\bBREAK\b", text)
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"], return_func_name=True)
c, shuffles = get_function(c.strip(), "(SHIFT|SHUFFLE)", ["0", "default", "default"])
r = c
for s in shuffles:
r = shuffle_chunk(s, r)
@@ -165,12 +182,13 @@ def tokenize(clip, text, can_break, empty_tokens):
need_word_ids = True
tokens = tokenize_chunks(clip, text, need_word_ids, can_break)
per_te_prompts = {}
per_te_prompts = defaultdict(list)
if l_prompts:
log.warning("Note: CLIP_L is deprecated. Use TE(l=prompt) instead")
per_te_prompts["l"] = l_prompts
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
@@ -184,9 +202,7 @@ def tokenize(clip, text, can_break, empty_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
l = per_te_prompts.get(te, [])
l.append(prompt)
per_te_prompts[te] = l
per_te_prompts[te].append(prompt)
if per_te_prompts:
for key in per_te_prompts:
@@ -211,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)
@@ -225,27 +241,24 @@ def encode_prompt_segment(
can_break = {}
for k in empty:
tokenizer = getattr(clip.tokenizer, f"clip_{k}", getattr(clip.tokenizer, k, None))
can_break[k] = tokenizer and tokenizer.pad_to_max_length
can_break[k] = tokenizer and getattr(tokenizer, "pad_to_max_length", False)
clip = hook_te(clip, empty.keys(), style, normalization, extra)
# Chunks to ConditioningAverage:
text, averages = get_function(text, "AVG", ["0.5"], return_dict=True)
prev = 0
text, averages = split_by_function(text, "AVG", ["0.5"], require_args=False)
prompts_to_avg = []
for avg in averages:
w = safe_float(avg["args"][0], 0.5)
p = text[prev : avg["position"]], w
prompts_to_avg.append(p)
prev = avg["position"]
prompts_to_avg.append((text[prev:], 1.0))
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))
conds_to_avg = []
for prompt, weight in prompts_to_avg:
conds_to_cat = []
chunks = re.split(r"\bCAT\b", prompt)
for c in chunks:
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))
@@ -266,13 +279,22 @@ def encode_prompt_segment(
w = next_w
continue
for i in range(len(base)):
(cond,) = ConditioningAverage.addWeighted(None, [base[i]], [cond[i]], w)
(cond,) = call_node(ConditioningAverage, [base[i]], [cond[i]], w)
base[i] = cond[0]
w = next_w
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):
"""Applies weights to TE outputs"""
if not spec:
@@ -284,7 +306,7 @@ def apply_weights(output, te_name, spec):
default = spec.get("all", None)
if isinstance(output, tuple):
out, pooled = output
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)
@@ -294,16 +316,16 @@ def apply_weights(output, te_name, spec):
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 = out * w
out = calc_w(out, w)
if pooled is not None:
pooled = pooled * pooled_w
pooled = calc_w(pooled, pooled_w)
return out, pooled
return (out, pooled) + tuple(extra)
else:
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 = output * w
output = calc_w(output, w)
return output
@@ -358,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)
@@ -375,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)
@@ -385,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))
@@ -409,7 +432,8 @@ 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")
@@ -430,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
@@ -485,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
@@ -505,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]]
@@ -528,55 +551,44 @@ 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))
attnmasked_prompts = []
fill = False
for prompt in prompts:
attn_couple = False
prompt_has_fill = False
if "ATTN()" in prompt:
prompt = prompt.replace("ATTN()", "")
attn_couple = True
if "FILL()" in prompt:
prompt = prompt.replace("FILL()", "")
prompt_has_fill = True
prompt, mask, mask_weight = get_mask(prompt, mask_size, masks)
text, noise_w, generator = get_noise(text)
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 not w:
continue
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
settings["start_percent"] = start_pct
settings["end_percent"] = end_pct
x = encode_prompt_segment(clip, prompt, settings, style, normalization)
if attn_couple:
if prompt_has_fill:
if attnmasked_prompts:
log.warning("FILL() can only be used for the first prompt, ignoring")
elif mask is not None:
log.warning("MASK() and FILL() can't be used together, ignoring FILL()")
else:
fill = True
attnmasked_prompts.extend(x)
else:
conds.extend(x)
def ensure_mask(c):
if "mask" not in c[1]:
@@ -585,20 +597,55 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
c[1]["mask_strength"] = 1.0
return c
if attnmasked_prompts:
base_cond = attnmasked_prompts[0]
if not fill:
ensure_mask(base_cond)
# else, set_cond_attnmask will have the base mask fill any unspecified areas
base_cond = [base_cond]
if len(attnmasked_prompts) > 1:
base_cond = set_cond_attnmask(
base_cond,
[ensure_mask(c) for c in attnmasked_prompts[1:]],
fill=fill,
)
else:
log.warning("You must specify at least two prompt segments with ATTN() for attention couple to work")
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:
prompt, noise_w, generator = get_noise(prompt)
base_prompt, attn_couple_prompts = split_by_function(prompt, "COUPLE", defaults=None, require_args=False)
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)
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
-106
View File
@@ -1,106 +0,0 @@
import unittest
import numpy.testing as npt
clip_l = None
dual = None
def run(f, *args):
return getattr(f, f.FUNCTION)(*args)
class TestEncode(unittest.TestCase):
def tensorsEqual(self, t1, t2):
npt.assert_equal(t1.detach().numpy(), t2.detach().numpy())
def condEqual(self, c1, c2, key=None, key_assert=None):
self.assertEqual(len(c1), len(c2))
for i in range(len(c1)):
a, b = c1[i], c2[i]
if key:
(key_assert or self.assertEqual)(a[1][key], b[1][key])
else:
self.tensorsEqual(a[0], b[0])
def test_basic_encode(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
combine = nodes.ConditioningCombine()
concat = nodes.ConditioningConcat()
zeroout = nodes.ConditioningZeroOut()
for k, clip in [("l", clip_l), ("dual", dual)]:
with self.subTest(k):
with self.subTest("No exceptions"):
run(
pc,
clip,
"test AND test (test:1.2) BREAK test AND TE_WEIGHT(all=0) SDXL() AND AREA(,,) test CAT test",
)
with self.subTest("Basic"):
(c1,) = run(pc, clip, "test")
(c2,) = run(comfy, clip, "test")
c = c2 # Used in later tests
self.condEqual(c1, c2)
(c1,) = run(pc, clip, "(test:1.2)")
(c2,) = run(comfy, clip, "(test:1.2)")
with self.subTest("Concat"):
(c1,) = run(pc, clip, "test CAT test")
(c2,) = run(concat, c, c)
self.condEqual(c1, c2)
with self.subTest("Combine"):
(c1,) = run(pc, clip, "test AND test")
(c2,) = run(combine, c, c)
self.condEqual(c1, c2)
with self.subTest("Zero out"):
(c1,) = run(pc, clip, "test TE_WEIGHT(all=0)")
(c2,) = run(zeroout, c)
self.condEqual(c1, c2)
def test_styles(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
for k, clip in [("l", clip_l), ("dual", dual)]:
(no_weights,) = run(comfy, clip, "this prompt has no weights")
for style in ["comfy", "A1111", "comfy++", "compel", "down_weight", "perp"]:
with self.subTest(f"TE {k} style {style} no weights equal comfy"):
(c,) = run(pc, clip, "this prompt has no weights")
self.condEqual(no_weights, c)
with self.subTest(f"TE {k} style {style} does not fail when encoding weights"):
for normalization in ["none", "mean", "length", "mean+length", "length+mean"]:
with self.subTest(f"TE {k} style {style} normalization {normalization}"):
(c,) = run(
pc,
clip,
f"STYLE({style}, {normalization}) (this prompt) (has weights:0.9), (a:1.2) (b:1.2)",
)
def test_masks(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
solidmask = comfy_extras.nodes_mask.SolidMask()
setMask = nodes.ConditioningSetMask()
for k, clip in [("l", clip_l), ("dual", dual)]:
(c1,) = run(pc, clip, "test MASK()")
(c2,) = run(comfy, clip, "test")
(c2,) = run(setMask, c2, run(solidmask, 1.0, 512, 512)[0], "default", 1.0)
self.condEqual(c1, c2)
self.condEqual(c1, c2, "mask", self.tensorsEqual)
if __name__ == "__main__":
print("Loading ComfyUI")
import main
id(main) # get rid of flake warning
import nodes
import comfy_extras.nodes_mask
from .nodes_base import PCTextEncode
(clip_l,) = nodes.CLIPLoader().load_clip("clip_l.safetensors")
(dual,) = nodes.DualCLIPLoader().load_clip("clip_l.safetensors", "t5xxl_fp16.safetensors", "flux")
print("Starting tests")
unittest.main()
-56
View File
@@ -1,56 +0,0 @@
import unittest
import numpy.testing as npt
clip_l = None
dual = None
def run(f, *args):
return getattr(f, f.FUNCTION)(*args)
class TestEncode(unittest.TestCase):
def tensorsEqual(self, t1, t2):
npt.assert_equal(t1.detach().numpy(), t2.detach().numpy())
def condEqual(self, c1, c2, key=None, key_assert=None):
self.assertEqual(len(c1), len(c2))
for i in range(len(c1)):
a, b = c1[i], c2[i]
if key:
(key_assert or self.assertEqual)(a[1][key], b[1][key])
else:
self.tensorsEqual(a[0], b[0])
def test_styles(self):
pc = PCTextEncode()
for k, clip in [("l", clip_l), ("dual", dual)]:
for style in ["comfy++", "A1111", "comfy++", "compel", "down_weight"]:
with self.subTest(f"TE {k} style {style} does not fail when encoding weights"):
for normalization in ["none", "mean", "length", "length+mean"]:
with self.subTest(f"TE {k} style {style} normalization {normalization}"):
(c,) = run(
pc,
clip,
f"STYLE(old+{style}, {normalization}) this prompt has weights, (a:1.2) (b:1.2)",
)
(c2,) = run(
pc,
clip,
f"STYLE({style}, {normalization}) this prompt has weights, (a:1.2) (b:1.2)",
)
self.condEqual(c, c2)
if __name__ == "__main__":
print("Loading ComfyUI")
import main
id(main) # get rid of flake warning
import nodes
from .nodes_base import PCTextEncode
(clip_l,) = nodes.CLIPLoader().load_clip("clip_l.safetensors")
(dual,) = nodes.DualCLIPLoader().load_clip("clip_l.safetensors", "clip_g.safetensors", "sdxl")
print("Starting tests")
unittest.main()
-215
View File
@@ -1,215 +0,0 @@
import unittest
import unittest.mock as mock
import logging
log = logging.getLogger("comfyui-prompt-control")
def find_file(name):
names = {"test": "test.safetensors", "other": "some/other.safetensors"}
return names.get(name)
def apply(cls, text, **kwargs):
model = [0, 1]
clip = [0, 0]
return cls().apply(unique_id="UID", model=model, clip=clip, text=text, **kwargs)
@mock.patch("prompt_control.utils.lora_name_to_file", find_file)
@mock.patch("torch.cuda.current_device", lambda: "cpu")
class GraphTests(unittest.TestCase):
maxDiff = 4096
def test_textencode(self):
clip = [0, 0]
from .nodes_lazy import PCLazyTextEncode, PCLazyTextEncodeAdvanced
for p in ["test", "[test:0.2] test", "[test[test::0.5]]<lora:test:1>"]:
r1 = PCLazyTextEncode().apply(clip, p, "UID")
r2 = PCLazyTextEncodeAdvanced().apply(clip, p, "UID")
self.assertEqual(r1, r2)
r = PCLazyTextEncode().apply(clip, "test<lora:test:1>", "UID")
self.assertEqual(
r,
{
"result": (["UID-2", 0],),
"expand": {
"UID-1": {"class_type": "PCTextEncode", "inputs": {"clip": [0, 0], "text": "test"}},
"UID-2": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID-1", 0], "start": 0.0, "end": 1.0},
},
},
},
)
r = PCLazyTextEncode().apply(clip, "simple [test:0.1,0.5] prompt<lora:test:1>", "UID")
self.assertEqual(
r,
{
"result": (["UID-8", 0],),
"expand": {
"UID-1": {"class_type": "PCTextEncode", "inputs": {"clip": [0, 0], "text": "simple prompt"}},
"UID-2": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID-1", 0], "start": 0.0, "end": 0.1},
},
"UID-3": {"class_type": "PCTextEncode", "inputs": {"clip": [0, 0], "text": "simple test prompt"}},
"UID-4": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID-3", 0], "start": 0.1, "end": 0.5},
},
"UID-5": {"class_type": "PCTextEncode", "inputs": {"clip": [0, 0], "text": "simple prompt"}},
"UID-6": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID-5", 0], "start": 0.5, "end": 1.0},
},
"UID-7": {
"class_type": "ConditioningCombine",
"inputs": {"conditioning_1": ["UID-2", 0], "conditioning_2": ["UID-4", 0]},
},
"UID-8": {
"class_type": "ConditioningCombine",
"inputs": {"conditioning_1": ["UID-7", 0], "conditioning_2": ["UID-6", 0]},
},
},
},
)
@mock.patch("prompt_control.utils.lora_name_to_file", find_file)
def test_loraloader(self):
from .nodes_lazy import PCLazyLoraLoader, PCLazyLoraLoaderAdvanced
model = [0, 1]
clip = [0, 0]
with self.assertLogs(log, level="WARNING") as cm:
result = apply(PCLazyLoraLoader, "prompt here <lora:nonexistent:1.0:0.5>")["expand"]
result_adv = apply(PCLazyLoraLoaderAdvanced, "prompt here <lora:nonexistent:1.0:0.5>")["expand"]
self.assertIn("LoRA 'nonexistent' not found", cm.output[0])
self.assertEqual(result, {})
self.assertEqual(result_adv, {})
result = apply(PCLazyLoraLoader, "<lora:test:1>")["expand"]
result2 = apply(PCLazyLoraLoader, "prompt here <lora:test:1.0:0.5><lora:test:0:0.5>")["expand"]
result3 = apply(PCLazyLoraLoaderAdvanced, "prompt here <lora:test:1.0:0.5><lora:test:0:0.5>")["expand"]
self.assertEqual(result, result2)
self.assertEqual(result2, result3)
self.assertEqual(
result,
{
"UID-1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 1.0,
"lora_name": "test.safetensors",
},
}
},
)
result = apply(PCLazyLoraLoader, "<lora:test:1><lora:other:0.5>")["expand"]
self.assertEqual(
result,
{
"UID-1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 1.0,
"lora_name": "test.safetensors",
},
},
"UID-2": {
"class_type": "LoraLoader",
"inputs": {
"model": ["UID-1", 0],
"clip": ["UID-1", 1],
"strength_model": 0.5,
"strength_clip": 0.5,
"lora_name": "some/other.safetensors",
},
},
},
)
result = apply(PCLazyLoraLoader, "prompt here <lora:test:1.0:0.5>")["expand"]
self.assertEqual(
result,
{
"UID-1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 0.5,
"lora_name": "test.safetensors",
},
}
},
)
result = apply(PCLazyLoraLoader, "prompt [<lora:test:0.5>:0.5]")["expand"]
result2 = apply(PCLazyLoraLoaderAdvanced, "prompt [<lora:test:0.5>:0.5]")["expand"]
self.assertEqual(result, result2)
expected = {
"UID-1": {
"class_type": "CreateHookLora",
"inputs": {"lora_name": "test.safetensors", "strength_model": 0.5, "strength_clip": 0.5},
},
"UID-2": {
"class_type": "CreateHookKeyframe",
"inputs": {"strength_mult": 0.0, "start_percent": 0.0},
},
"UID-3": {
"class_type": "CreateHookKeyframe",
"inputs": {
"start_percent": 0.5,
"prev_hook_kf": ["UID-2", 0],
"strength_mult": 1.0,
},
},
"UID-4": {
"class_type": "SetHookKeyframes",
"inputs": {"hooks": ["UID-1", 0], "hook_kf": ["UID-3", 0]},
},
"UID-5": {
"class_type": "SetClipHooks",
"inputs": {
"clip": [0, 0],
"hooks": ["UID-4", 0],
"apply_to_conds": True,
"schedule_clip": True,
},
},
}
self.assertEqual(result, expected)
result2 = apply(PCLazyLoraLoaderAdvanced, "prompt [<lora:test:0.5>:0.5]", start=0.6)["expand"]
self.assertEqual(
result2,
{
"UID-1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 0.5,
"strength_clip": 0.5,
"lora_name": "test.safetensors",
},
}
},
)
result2 = PCLazyLoraLoaderAdvanced().apply(model, clip, "prompt [<lora:test:0.5>:0.5]", "UID", end=0.5)[
"expand"
]
self.assertEqual(result2, {})
if __name__ == "__main__":
unittest.main()
-207
View File
@@ -1,207 +0,0 @@
import unittest
from .parser import parse_prompt_schedules as parse, expand_macros
def prompt(until, text, *loras):
loras = {lora: {"weight": unet, "weight_clip": te} for lora, unet, te in loras}
return [until, {"prompt": text, "loras": loras}]
class TestParser(unittest.TestCase):
def assertPrompt(self, p, at, until, text, *loras):
self.assertEqual(p.at_step(at), prompt(until, text, *loras))
def test_no_scheduling(self):
p = parse("This is a (basic:0.6) (prompt) with [no scheduling] features")
expected = prompt(1.0, "This is a (basic:0.6) (prompt) with [no scheduling] features")
self.assertEqual(p.at_step(0), expected)
self.assertEqual(p.at_step(0.5), expected)
self.assertEqual(p.at_step(1), expected)
def test_equivalences(self):
eqs = [
[parse(p) for p in ["[a:0.1]", "[:a:0.1]", "[:a:0,0.1]", "[:a::0.1,1.0]", "[:a::0.1]"]],
[parse(p) for p in ["[before:during:after:0.1]", "[before:during:after:0.1,1.0]", "[before:during:0.1]"]],
[parse(p) for p in ["[a:0.1,0.5]", "[[a:0.1]::0.5]", "[:a::0.1,0.5]", "[a::0.1,0.5]"]],
[parse(p) for p in ["[a:b:0.5]", "[a::b:0.5,0.5]"]],
[parse(p) for p in ["[a::0.5]", "[a:::0.5,0.5]"]],
]
for group in eqs:
for p in group[1:]:
with self.subTest(p):
self.assertEqual(group[0].parsed_prompt, p.parsed_prompt)
def test_basic(self):
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]"
)
self.assertPrompt(p, 0, 0.5, "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features")
self.assertPrompt(p, 0.5, 0.5, "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features")
self.assertPrompt(p, 0.7, 0.8, "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) features")
self.assertPrompt(p, 1.0, 1.0, "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) ")
def test_lora(self):
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)
)
self.assertEqual(p.at_step(0), expected)
self.assertEqual(p.at_step(0.5), expected)
self.assertEqual(p.at_step(1), expected)
def test_scheduled_lora(self):
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>"
)
self.assertPrompt(
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),
)
self.assertPrompt(p, 0.5, 1.0, "This is a (lora:0.6) (prompt) with [scheduling] features ", ("bar", 1.0, 1.2))
def test_seq(self):
p = parse("This is a sequence of [SEQ:a:0.2::0.5:c:0.8][SEQ: and x:0.8]")
p2 = parse("This is a sequence of [[a:[c:0.5]:0.2]::0.8][ and x::0.8]")
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 ",
}
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
for k, v in prompts.items():
self.assertPrompt(p, k, k, v)
def test_shortcuts_scheduling(self):
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]")
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
self.assertEqual(p3.parsed_prompt, p4.parsed_prompt)
def test_range(self):
p = parse("test [excluded::excluded2:0.1,0.4] test")
self.assertPrompt(p, 0, 0.1, "test excluded test")
self.assertPrompt(p, 0.2, 0.4, "test test")
self.assertPrompt(p, 0.45, 1.0, "test excluded2 test")
p = parse("test [[:included::0.2,0.8]|[excluded::excluded2:0.4,0.9]:0.1] test")
self.assertPrompt(p, 0, 0.1, "test test")
self.assertPrompt(p, 0.25, 0.3, "test included test")
self.assertPrompt(p, 0.15, 0.2, "test excluded test")
self.assertPrompt(p, 0.25, 0.3, "test included test")
self.assertPrompt(p, 0.55, 0.6, "test test")
self.assertPrompt(p, 0.95, 1.0, "test excluded2 test")
def test_nested(self):
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:
self.assertEqual(p.at_step(k), [prompts[k][0], {"prompt": prompts[k][1], "loras": {}}])
self.assertPrompt(p, 0.6, 0.7, "This prompt is ", ("cool", 1.0, 1.0))
self.assertPrompt(p, 0.7, 0.7, "This prompt is ", ("cool", 1.0, 1.0))
p2 = p.with_filters(filters="hr, xyz")
self.assertEqual(p2.at_step(0), p2.at_step(1))
def test_def(self):
p = parse("DEF(X=0.5) [a:b:X] DEF(test = [c:X]) test test")
prompts = {
0.2: (0.5, "a "),
0.6: (1.0, "b c c"),
}
for k, v in prompts.items():
self.assertPrompt(p, k, v[0], v[1])
p = parse("DEF(X=[($1):($1:$2):$2])X(test;0.7)")
p2 = parse("[(test):(test:0.7):0.7]")
with self.subTest("parameters"):
self.assertEqual(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]")
with self.subTest("two functions"):
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
p = expand_macros("DEF(X(a;b)=$1 $2 $3 d)X(A) X(A;B;C)")
with self.subTest("defaults"):
self.assertEqual(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)")
with self.subTest("Empty default for $1"):
self.assertEqual(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)]")
with self.subTest("defaults, DEF=X vs DEF=X()"):
self.assertEqual(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]")
with self.subTest("defaults, nested parens"):
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
with self.assertRaises(ValueError) as c:
expand_macros("DEF(X=recurse Y) DEF(Y=recurse X) X")
self.assertTrue("Unable to resolve DEFs" in str(c.exception))
def test_misc(self):
p = parse("[[a:c:0.5]:0.7]")
p2 = parse("[:[a:c:0.5]:0.7]")
self.assertEqual(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]")
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
pf = p.with_filters(filters="hr")
self.assertEqual(pf.parsed_prompt, p2.with_filters(filters="hr").parsed_prompt)
self.assertPrompt(pf, 0, 0.5, "test a")
self.assertPrompt(pf, 0.55, 0.6, "test ")
self.assertPrompt(pf, 0.8, 1.0, "test b", ("test", 0.5, 0.5))
p = parse("[:[<lora:test:1>:c:0.5]:0.3]")
self.assertPrompt(p, 0, 0.3, "")
self.assertPrompt(p, 0.4, 0.5, "", ("test", 1.0, 1.0))
self.assertPrompt(p, 1.0, 1.0, "c")
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, v in prompts.items():
self.assertPrompt(p, k, v[0], v[1])
def test_alternating(self):
p = parse("[cat|dog|tiger]")
p2 = parse("[cat|dog|tiger:0.1]")
p3 = parse("[cat|[dog|wolf]|tiger]")
p4 = parse("[cat|[dog:wolf<lora:canine:1>:0.5]:0.2]")
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
for i, x in enumerate(["cat", "wolf", "tiger", "cat", "dog", "tiger", "cat", "wolf", "tiger", "cat"]):
step = round((i * 0.1) + 0.1, 2)
with self.subTest(step):
self.assertPrompt(p3, step, step, x)
for i, x in enumerate([["cat"], ["dog"], ["cat"], ["wolf", ("canine", 1.0, 1.0)], ["cat"]]):
step = round((i * 0.2) + 0.2, 2)
with self.subTest(step):
self.assertPrompt(p4, step, step, *x)
self.assertPrompt(p4, 0.7, 0.8, "wolf", ("canine", 1.0, 1.0))
if __name__ == "__main__":
unittest.main()
+181 -44
View File
@@ -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,68 +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, placeholder="", return_dict=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)
count = 0
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)
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:
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"
if return_dict:
instances.append(
{
"name": funcname,
"args": args,
"position": start,
"placeholder": ph,
}
)
elif return_func_name:
instances.append((funcname, args))
else:
instances.append(args)
if placeholder:
text = text[:start] + f"\0{placeholder}{count}\0" + text[end + 1 :]
else:
text = text[:start] + text[end + 1 :]
match = rex.search(text)
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
@@ -152,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:
@@ -161,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:
@@ -172,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 -4
View File
@@ -1,10 +1,10 @@
[project]
name = "comfyui-prompt-control"
description = "Nodes for convenient prompt editing, making many common operations prompt-controllable"
version = "2.0.0-rc.7"
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.8"
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"
@@ -13,3 +13,36 @@ Repository = "https://github.com/asagi4/comfyui-prompt-control"
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"]
-2
View File
@@ -1,2 +0,0 @@
# 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
View File
+5
View File
@@ -0,0 +1,5 @@
import logging
def pytest_runtest_setup(item):
logging.getLogger("comfyui-prompt-control").setLevel(logging.CRITICAL)
+26
View File
@@ -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, "-")]
+265
View File
@@ -0,0 +1,265 @@
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_neq(c1, c2, key=None, key_assert=None):
ok = False
try:
cond_equal(c1, c2, key=key, key_assert=key_assert)
except AssertionError:
ok = True
if not ok:
raise ValueError("Tensors should not be equal")
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)
def test_noise_weight0(self, text_encoder_clips, pc_text_encode, node_class_objs):
for _k, clip in text_encoder_clips:
(c1,) = run(pc_text_encode, clip, "test")
(c2,) = run(pc_text_encode, clip, "test NOISE(0, 0)")
cond_equal(c1, c2)
def test_noise(self, text_encoder_clips, pc_text_encode, node_class_objs):
for _k, clip in text_encoder_clips:
(c1,) = run(pc_text_encode, clip, "test")
(c2,) = run(pc_text_encode, clip, "test NOISE(1, 0)")
cond_neq(c1, c2)
+686
View File
@@ -0,0 +1,686 @@
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_textencode_custom():
r = te("NODE(CLIPTextEncode)simple [test:0.1,0.5] $p SEG(p) prompt")
assert r == {
"result": (["UID.0.0.8", 0],),
"expand": {
"UID.0.0.1": {
"class_type": "CLIPTextEncode",
"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": "CLIPTextEncode",
"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": "CLIPTextEncode",
"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_textencode_custom_extra():
r = te(
'NODE(CustomTextEncode, prompt, image ["1", 0]; option "test"; float [10.0:__EMPTY__:0.5])simple [test:0.1,0.5] prompt'
)
assert r == {
"result": (["UID.0.0.8", 0],),
"expand": {
"UID.0.0.1": {
"class_type": "CustomTextEncode",
"inputs": {
"clip": [0, 0],
"prompt": "simple prompt",
"image": ["1", 0],
"option": "test",
"float": 10.0,
},
},
"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": "CustomTextEncode",
"inputs": {
"clip": [0, 0],
"prompt": "simple test prompt",
"image": ["1", 0],
"option": "test",
"float": 10.0,
},
},
"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": "CustomTextEncode",
"inputs": {"clip": [0, 0], "prompt": "simple prompt", "image": ["1", 0], "option": "test"},
},
"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 == {}
def test_loraloader_segs():
result = loraloader("prompt [<lora:test:0.5>:0.5]")["expand"]
result2 = loraloader("prompt [$lora:0.5]\nSEG(lora)<lora:test:0.5>\nSEG(lora2)<lora:ignored:1>")["expand"]
assert result == result2
+72
View File
@@ -0,0 +1,72 @@
from textwrap import dedent
import pytest
from prompt_control.macros import expand_macros, expand_segs
@pytest.mark.parametrize(
"text, result",
[
("DEF(X(a;b)=$1 $2 $3 d)X(A) X(A;B;C)", "A b $3 d A B C d"),
(
"DEF(MACRO()=[empty:$1:$2])MACRO MACRO(;) MACRO(;0.5) MACRO(a;0.5)",
"[empty::$2] [empty::] [empty::0.5] [empty:a:0.5]",
),
("DEF(X=$1)DEF(Y()=$1)[X Y][X() Y()][X(1) Y(1)]", "[$1 ][ ][1 1]"),
(
"DEF(C_ANIMAL=cat)DEF(D_ANIMAL=dog)DEF(IT=It is a $1_ANIMAL $10_ANIMAL)IT(D) IT(C)",
"It is a dog $10_ANIMAL It is a cat $10_ANIMAL",
),
],
)
def test_basic_macro(text, result):
assert expand_macros(text) == result
def test_macro_recursion():
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(
"input, output",
[
(
"""\
A red $b and
a blue $a
SEG(a)
cat
SEG(b)
dog
SEG(c)""",
"A red dog and\na blue cat",
),
(
"""\
$a and $b
SEG(a)
cat, $b
SEG(b)
dog, $c
SEG(c)
tiger
""",
"cat, dog, tiger and dog, tiger",
),
(
"""\
$a
SEG(a)
a $b
SEG(b)
b $a""",
"a b a b $a",
),
],
)
def test_segments(input, output):
assert expand_segs(dedent(input)) == output
+372
View File
@@ -0,0 +1,372 @@
import os
import pytest
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
@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_ltgt_in_schedule(parse):
p = parse("This should [<parse> correctly:be <Picture 1>:0.1]<lora:test:1>")
assert_prompt(p, 0.1, 0.1, "This should <parse> correctly", ("test", 1.0, 1.0))
assert_prompt(p, 0.15, 1.0, "This should be <Picture 1>", ("test", 1.0, 1.0))
def test_floats(parse):
p = parse("[a:b:0.5] [c:d:e:0.2,0.7] <lora:test:-0.3>")
p2 = parse("[a:b:.5] [c:d:e:.2,.7] <lora:test:-.3>")
assert p.parsed_prompt == 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)
+139
View File
@@ -0,0 +1,139 @@
{
"1": {
"inputs": {
"text": "positive prompt",
"clip": [
"4",
1
]
},
"class_type": "PCLazyTextEncode",
"_meta": {
"title": "PC: Schedule prompt"
}
},
"2": {
"inputs": {
"ckpt_name": "$TEST_CHECKPOINT"
},
"class_type": "CheckpointLoaderSimple",
"_meta": {
"title": "Load Checkpoint"
}
},
"3": {
"inputs": {
"seed": 0,
"steps": 8,
"cfg": 3,
"sampler_name": "euler",
"scheduler": "simple",
"denoise": 1,
"model": [
"4",
0
],
"positive": [
"9",
0
],
"negative": [
"9",
1
],
"latent_image": [
"5",
0
]
},
"class_type": "KSampler",
"_meta": {
"title": "KSampler"
}
},
"4": {
"inputs": {
"text": "<lora:$TEST_LORA:1>",
"model": [
"2",
0
],
"clip": [
"2",
1
]
},
"class_type": "PCLazyLoraLoader",
"_meta": {
"title": "PC: Schedule LoRAs"
}
},
"5": {
"inputs": {
"width": 1024,
"height": 1024,
"batch_size": 1
},
"class_type": "EmptyLatentImage",
"_meta": {
"title": "Empty Latent Image"
}
},
"6": {
"inputs": {
"samples": [
"3",
0
],
"vae": [
"2",
2
]
},
"class_type": "VAEDecode",
"_meta": {
"title": "VAE Decode"
}
},
"7": {
"inputs": {
"images": [
"6",
0
]
},
"class_type": "PreviewImage",
"_meta": {
"title": "Preview Image"
}
},
"8": {
"inputs": {
"text": "worst quality,",
"clip": [
"4",
1
]
},
"class_type": "CLIPTextEncode",
"_meta": {
"title": "CLIP Text Encode (Prompt)"
}
},
"9": {
"inputs": {
"positive": [
"1",
0
],
"negative": [
"8",
0
]
},
"class_type": "PCAttentionCoupleBatchNegative",
"_meta": {
"title": "PC: Attention Couple (batch negative)"
}
}
}
+40
View File
@@ -0,0 +1,40 @@
import json
import os
import uuid
from time import sleep
import pytest
import requests
@pytest.fixture(scope="module", autouse=True)
def workflow(request):
with open(str(request.path).replace(".py", ".json")) as f:
data = f.read()
data = data.replace("$TEST_CHECKPOINT", os.environ["PC_TEST_CHECKPOINT"])
data = data.replace("$TEST_LORA", os.environ["PC_TEST_LORA"])
return json.loads(data)
def assert_prompt(url, p):
timeout = 60
r = requests.post(f"{url}/prompt", json={"prompt": p, "client_id": str(uuid.uuid4())}).json()
prompt_id = r["prompt_id"]
r = {"status": "pending"}
while r["status"] in ["pending", "in_progress"]:
sleep(1)
assert timeout > 0
timeout -= 1
r = requests.get(f"{url}/api/jobs/{prompt_id}").json()
assert r["status"] == "completed"
@pytest.fixture
def comfyui():
return os.environ.get("PC_TEST_COMFYUI", "http://localhost:8188")
def test_workflow(workflow, comfyui):
prompt = "DEF(blue=green)a blue dog and a cat sitting [COUPLE(0 0.5, 0 1) red (cat,:1.3) COUPLE(0.5 1, 0 1) (blue:1.2) dog,:0.1]"
workflow["1"]["inputs"]["text"] = prompt
assert_prompt(comfyui, workflow)
+3
View File
@@ -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`.
+1
View File
@@ -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)
+7
View File
@@ -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
View File
@@ -0,0 +1 @@
PCLazyLoraLoader.md
+7
View File
@@ -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
View File
@@ -0,0 +1 @@
PCLazyTextEncode.md
+5
View File
@@ -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`.
+5
View File
@@ -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.
+7
View File
@@ -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.
+5
View File
@@ -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)