Compare commits
282
Commits
v1.1.1
...
v2.0.0-rc.7
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8cc73a2e49 | ||
|
|
27ae5f683e | ||
|
|
68cda3663e | ||
|
|
3d46f705b6 | ||
|
|
eec4bc4da9 | ||
|
|
3d3218e831 | ||
|
|
f4e57ec514 | ||
|
|
74f65c1b31 | ||
|
|
4d94cca88f | ||
|
|
4569ecccf9 | ||
|
|
192e6d30d4 | ||
|
|
67112f11e0 | ||
|
|
f761dfac86 | ||
|
|
278e733835 | ||
|
|
2bf65720eb | ||
|
|
99f3af92b7 | ||
|
|
2437cd4daf | ||
|
|
d262a7dc7a | ||
|
|
6c319ad5b4 | ||
|
|
34056cac19 | ||
|
|
8ae436abf1 | ||
|
|
c2ce2ce023 | ||
|
|
7a9e69ec31 | ||
|
|
aaff8dc7da | ||
|
|
3e4722278a | ||
|
|
11cb430396 | ||
|
|
9f1cbfd11c | ||
|
|
5b3a914f1d | ||
|
|
a5ffa1acd7 | ||
|
|
7f8783147b | ||
|
|
3c9b806e5f | ||
|
|
a485c2655a | ||
|
|
f888e69b00 | ||
|
|
b4b0858214 | ||
|
|
7a1cb2cf51 | ||
|
|
fbb6b5c8fa | ||
|
|
6115c095cb | ||
|
|
04d3d2e959 | ||
|
|
e4f64837ef | ||
|
|
4a785b294b | ||
|
|
b3195a6297 | ||
|
|
5c52bffc9d | ||
|
|
ac2d275dfb | ||
|
|
399f992a26 | ||
|
|
672a2a09bb | ||
|
|
e3629961ce | ||
|
|
ddac624ad1 | ||
|
|
99ddfe357e | ||
|
|
110d5248a0 | ||
|
|
61b1ecc88e | ||
|
|
05b0b2ad26 | ||
|
|
5831608c4e | ||
|
|
c815bb44f1 | ||
|
|
e55c50e9d7 | ||
|
|
1b0ff62d10 | ||
|
|
cf93093d59 | ||
|
|
fc15a89a2f | ||
|
|
57c092bccf | ||
|
|
88f77a8124 | ||
|
|
a9c2487c0c | ||
|
|
75bced7d2b | ||
|
|
4b285be07e | ||
|
|
200d9f9daf | ||
|
|
d33208b1c3 | ||
|
|
ffa64816c0 | ||
|
|
6ffbf05d7d | ||
|
|
892a70d53b | ||
|
|
20711358a2 | ||
|
|
453580545c | ||
|
|
98d78df7ba | ||
|
|
bdd56410dc | ||
|
|
0289564e55 | ||
|
|
1e05d1a8cc | ||
|
|
dc6fd0fc63 | ||
|
|
a356bddcc7 | ||
|
|
b8081e5736 | ||
|
|
aa00c26365 | ||
|
|
e913bad73c | ||
|
|
5c1b739b82 | ||
|
|
5e3ab1f51a | ||
|
|
b21de76cd5 | ||
|
|
e4a27d01ee | ||
|
|
42cdfa0f5a | ||
|
|
98292e2bc8 | ||
|
|
f0c8e2e873 | ||
|
|
c4ac37333d | ||
|
|
2534e002ad | ||
|
|
fd4823fd75 | ||
|
|
a6f230ff8b | ||
|
|
633b2f05e0 | ||
|
|
a15135ddc5 | ||
|
|
0d7e2a4e60 | ||
|
|
c5495832c5 | ||
|
|
63d2cb3e0c | ||
|
|
95832e801b | ||
|
|
01bd5568d0 | ||
|
|
36b3638f4e | ||
|
|
d46000ef78 | ||
|
|
aba246a33c | ||
|
|
42ae22db83 | ||
|
|
cb6de285cb | ||
|
|
49a073bb12 | ||
|
|
e9afe779ae | ||
|
|
fa3b4f7da3 | ||
|
|
7a76cc8c72 | ||
|
|
6b1e2a5a8a | ||
|
|
9aee531c09 | ||
|
|
c08bf395a6 | ||
|
|
1964708997 | ||
|
|
1c4b5ce0c4 | ||
|
|
5eabbb419c | ||
|
|
53400a029b | ||
|
|
c39605eec4 | ||
|
|
dc62e638ed | ||
|
|
109cac16ef | ||
|
|
cf6c2b3e6a | ||
|
|
d113d4ba78 | ||
|
|
5bd1d04dcd | ||
|
|
7c10770e07 | ||
|
|
69ea298174 | ||
|
|
4ba4b28bb2 | ||
|
|
f728866b90 | ||
|
|
306c02f57b | ||
|
|
2c519310ac | ||
|
|
3f23d1b14a | ||
|
|
148776fe5d | ||
|
|
cef4a80440 | ||
|
|
cd642b5d42 | ||
|
|
2b323da9a9 | ||
|
|
79b3675c4f | ||
|
|
b8d5b7a7c4 | ||
|
|
a5da586dc5 | ||
|
|
01aa061bef | ||
|
|
2fab4be810 | ||
|
|
127acb7018 | ||
|
|
b952b2f186 | ||
|
|
2732a795fb | ||
|
|
4f78fff892 | ||
|
|
e290cb57ac | ||
|
|
d577d439e7 | ||
|
|
e59d46c8d1 | ||
|
|
bd1c69a517 | ||
|
|
cee19aea67 | ||
|
|
b93bb66aed | ||
|
|
c75d1a6651 | ||
|
|
a4c7f99cc1 | ||
|
|
b7d544c05c | ||
|
|
04c4bd0846 | ||
|
|
5365679a60 | ||
|
|
fa77c158ac | ||
|
|
525cb157ce | ||
|
|
e10950e4da | ||
|
|
91ba4c881f | ||
|
|
106ebe49aa | ||
|
|
549b4347fd | ||
|
|
4cbce5df06 | ||
|
|
21208bd733 | ||
|
|
a4065415e7 | ||
|
|
0d546f1a08 | ||
|
|
390e1ec6b8 | ||
|
|
ce0c1cd698 | ||
|
|
3745bc6879 | ||
|
|
99c815a3b6 | ||
|
|
9f6c9c11e6 | ||
|
|
15127d2466 | ||
|
|
3fbae90478 | ||
|
|
2f12069821 | ||
|
|
3baeabb8ee | ||
|
|
c85134b31e | ||
|
|
26f7e1ff24 | ||
|
|
e9e8b75d7f | ||
|
|
01526e3923 | ||
|
|
e888238625 | ||
|
|
53a6d48cb1 | ||
|
|
a0df992741 | ||
|
|
be51e0dfc4 | ||
|
|
c023956b4c | ||
|
|
b33f24e0cb | ||
|
|
a637321356 | ||
|
|
3a2d08fcf7 | ||
|
|
99d966d74a | ||
|
|
724488d20b | ||
|
|
bec28affbe | ||
|
|
08844019cc | ||
|
|
f7d78e54d5 | ||
|
|
25990cb17e | ||
|
|
de1ad8512a | ||
|
|
09925976d4 | ||
|
|
00061e18f6 | ||
|
|
06f1291727 | ||
|
|
ee914b2920 | ||
|
|
b8d002facc | ||
|
|
2f5d62b46b | ||
|
|
48c0286f09 | ||
|
|
e64a71fc6a | ||
|
|
dbd5a0e6d6 | ||
|
|
f3a4b12bc0 | ||
|
|
be921a9eea | ||
|
|
5aa6110206 | ||
|
|
48b727f4f6 | ||
|
|
5ff6d21e43 | ||
|
|
d8b850cafd | ||
|
|
24a9dd32f7 | ||
|
|
7b358e5127 | ||
|
|
85facba5d8 | ||
|
|
ca8bc3fc16 | ||
|
|
c82a041ac5 | ||
|
|
46eb8205f7 | ||
|
|
c21a31cd94 | ||
|
|
44c4ced7dd | ||
|
|
d39cc27c1d | ||
|
|
a7e48ea9fc | ||
|
|
83f859bc36 | ||
|
|
2fae4c0bc8 | ||
|
|
31e70c776b | ||
|
|
192f7e3efd | ||
|
|
f510d15f5b | ||
|
|
d71ec8d86a | ||
|
|
eab2cc09dd | ||
|
|
de1c39a74a | ||
|
|
7bfd6790df | ||
|
|
fd673a0d5b | ||
|
|
fcb63aefa6 | ||
|
|
94b066a5c2 | ||
|
|
fdc1bc4f2f | ||
|
|
579162e440 | ||
|
|
e78bf45995 | ||
|
|
94413c6d9f | ||
|
|
ffdac507ae | ||
|
|
67faf38fbe | ||
|
|
69a534eaae | ||
|
|
c3d90e874b | ||
|
|
204d990afa | ||
|
|
da502219ad | ||
|
|
df72d2c478 | ||
|
|
798c769e13 | ||
|
|
3bd170b9ca | ||
|
|
75ef59b8e0 | ||
|
|
13696a11b7 | ||
|
|
08a0a2afc5 | ||
|
|
58fe45eb87 | ||
|
|
dab719f369 | ||
|
|
5e23d3f8cc | ||
|
|
336ed5a15f | ||
|
|
d3d21f8795 | ||
|
|
bb2358e43a | ||
|
|
b222d39f5f | ||
|
|
1bafa1a6b4 | ||
|
|
7b6ff9a879 | ||
|
|
ac8de2995e | ||
|
|
a0c5c9e2fb | ||
|
|
71a6bba451 | ||
|
|
c85cb5a309 | ||
|
|
34da83b4ab | ||
|
|
4ea73cf4ec | ||
|
|
9f5a726c8a | ||
|
|
d4856a595e | ||
|
|
498a8c58f7 | ||
|
|
feb0a5c791 | ||
|
|
04819cb5c2 | ||
|
|
806b78b902 | ||
|
|
807261cb00 | ||
|
|
5523190db9 | ||
|
|
5e3764728c | ||
|
|
b81f0e653d | ||
|
|
2e60c904c8 | ||
|
|
c7427d324f | ||
|
|
96641c3e4a | ||
|
|
a86b5a9fa7 | ||
|
|
e4254828f5 | ||
|
|
acf38ad328 | ||
|
|
9c659e85c0 | ||
|
|
8a4d32ae0e | ||
|
|
81f39df673 | ||
|
|
67d41fb1b3 | ||
|
|
71e340939b | ||
|
|
751af8cabb | ||
|
|
7e9ca60dfd | ||
|
|
8b76376e56 | ||
|
|
4bbf3a895f | ||
|
|
2930f03d6c | ||
|
|
42acef7298 |
@@ -20,5 +20,7 @@ A clear and concise description of what the bug is.
|
||||
Information needed to trigger the problem.
|
||||
If possible, attach a workflow to reproduce the problem
|
||||
|
||||
If a workflow works, but isn't producing the correct output, please enable debug logging with the `PCSetLogLevel` node (from `promptcontrol/tools`) and run your workflow with debug logging enabled, and copy the outputs here.
|
||||
|
||||
**Expected behavior**
|
||||
A description of what you expected to happen.
|
||||
|
||||
@@ -7,15 +7,23 @@ on:
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
tests:
|
||||
uses: ./.github/workflows/tests.yml
|
||||
tests_with_comfy:
|
||||
uses: ./.github/workflows/tests_with_comfy.yml
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'asagi4' }}
|
||||
needs: [tests, tests_with_comfy]
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
name: Run parser tests
|
||||
on:
|
||||
- workflow_call
|
||||
- workflow_dispatch
|
||||
- push
|
||||
|
||||
jobs:
|
||||
run-parser-tests:
|
||||
name: Run parser tests
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.11'
|
||||
- run: pip install -r requirements.txt
|
||||
- run: python -m prompt_control.test_parser
|
||||
@@ -0,0 +1,28 @@
|
||||
name: Run tests requiring ComfyUI
|
||||
on:
|
||||
workflow_call:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
paths:
|
||||
- prompt_control/nodes_lazy.py
|
||||
- prompt_control/utils.py
|
||||
|
||||
|
||||
jobs:
|
||||
run-graph-tests:
|
||||
name: Run graph tests
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Check out ComfyUI
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
repository: comfyanonymous/ComfyUI
|
||||
path: ComfyUI
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.11'
|
||||
- run: pip install 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
|
||||
@@ -1,8 +1,20 @@
|
||||
all: format check
|
||||
all: format check test
|
||||
@echo "Done"
|
||||
check:
|
||||
pyflakes *.py */*.py
|
||||
find . -name "*.py" | xargs pyflakes
|
||||
format:
|
||||
black -l 120 *.py */*.py
|
||||
find . -name "*.py" | xargs black -l 120
|
||||
|
||||
test:
|
||||
python -m prompt_control.test_parser
|
||||
|
||||
test_graph:
|
||||
PYTHONPATH=../../ python -m prompt_control.test_graph
|
||||
|
||||
test_encode:
|
||||
PYTHONPATH=../../ python -m prompt_control.test_encode
|
||||
|
||||
manual_test:
|
||||
PYTHONPATH=../../ python -im prompt_control.manual_test
|
||||
|
||||
.PHONY: check format all
|
||||
|
||||
@@ -1,25 +1,57 @@
|
||||
# ComfyUI prompt control
|
||||
|
||||
Nodes for LoRA and prompt scheduling that make basic operations in ComfyUI completely prompt-controllable.
|
||||
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.
|
||||
|
||||
LoRA and prompt scheduling should produce identical output to the equivalent ComfyUI workflow using multiple samplers or the various conditioning manipulation nodes. If you find situations where this is not the case, please report a bug.
|
||||
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.
|
||||
|
||||
## What can it do?
|
||||
|
||||
Things you can control via the prompt:
|
||||
- Prompt editing and filtering without multiple samplers
|
||||
- LoRA loading and scheduling (including LoRA block weights)
|
||||
- Prompt masking and area control, combining prompts and interpolation
|
||||
- SDXL parameters
|
||||
- Other miscellaneous things
|
||||
You can use text prompts to control the following:
|
||||
|
||||
[This example workflow](workflows/example.json?raw=1) implements a two-pass workflow illustrating most scheduling features.
|
||||
- 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
|
||||
|
||||
The tools in this repository combine well with the macro and wildcard functionality in [comfyui-utility-nodes](https://github.com/asagi4/comfyui-utility-nodes)
|
||||
All features are fully schedulable unless otherwise stated. See the [syntax documentation](doc/syntax.md) for details on how to use each feature.
|
||||
|
||||
If you find prompt scheduling inconvenient for some reason, `PCTextEncode` can be used as a drop-in replacement for `CLIPTextEncode` to get everything else.
|
||||
|
||||
[This workflow](example_workflows/Workflow%20Comparison.json?raw=1) shows LoRA scheduling and prompt editing and compares it with the same prompt implemented with built-in ComfyUI nodes. You can also find it in the template library.
|
||||
|
||||
## Compatibility
|
||||
|
||||
Prompt Control uses graph generation, and tries to delegate functionality to core ComfyUI wherever possible, implementing any hooks and patches in a way that is maximally compatible. This means that it should just work in most cases, even with models and nodes not explicitly supported.
|
||||
|
||||
If you encounter issues as a user or if you're a node developer and Prompt Control somehow breaks something, feel free to file a bug report.
|
||||
|
||||
## 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
|
||||
|
||||
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)
|
||||
For LoRA scheduling to work, you'll need at least version 0.3.7 of ComfyUI (0.3.36 of ComfyUI desktop).
|
||||
|
||||
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:
|
||||
```
|
||||
@@ -28,305 +60,48 @@ If you use the portable version of ComfyUI on Windows with its embedded Python,
|
||||
|
||||
Then restart ComfyUI afterwards.
|
||||
|
||||
## Notable changes
|
||||
# Core nodes
|
||||
|
||||
I try to avoid behavioural changes that break old prompts, but they may happen occasionally.
|
||||
**Note**: The documentation refers to the nodes with their internal names for consistency. The display name may change, but ComfyUI's search will always find the nodes with the internal name. `PCLazyTextEncode` and `PCLazyLoraLoader` are the main ones you'll want to use, also known as `PC: Schedule Prompt` and `PC: Schedule LoRas`.
|
||||
|
||||
- 2024-02-02 The node will now automatically enable offloading LoRA backup weights to the CPU if you run out of memory during LoRA operations, even when `--highvram` is specified. This change persists until ComfyUI is restarted.
|
||||
- 2024-01-14 Multiple `CLIP_L` instances are now joined with a space separator instead of concatenated.
|
||||
- 2024-01-09 AITemplate support dropped. I don't recommend or test AITemplate anymore. Use Stable-Fast instead (see below for info)
|
||||
- 2024-01-08 Prompt control now enables in-place weight updates on the model. This shouldn't affect anything, but increases performance slightly. You can disable this by setting the environment variable `PC_NO_INPLACE_UPDATE` to any non-empty value.
|
||||
- 2023-12-28 MASK now uses ComfyUI's `mask_strength` attribute instead of calculating it on its own. This changes its behaviour slightly.
|
||||
- 2023-12-06: Removed `JinjaRender`, `SimpleWildcard`, `ConditioningCutoff`, `CondLinearInterpolate` and `StringConcat`. For the first two, see [this repository](https://github.com/asagi4/comfyui-utility-nodes) for mostly-compatible implementations.
|
||||
- 2023-10-04: `STYLE:...` syntax changed to `STYLE(...)`
|
||||
## PCLazyTextEncode and PCLazyTextEncodeAdvanced
|
||||
|
||||
## Note on how schedules work
|
||||
`PCLazyTextEncode` uses ComfyUI's lazy graph execution mechanism to generate a graph of `PCTextEncode` and `SetConditioningTimestepRange` nodes from a prompt with schedules. This has the advantage that if a part of the schedule doesn't change, ComfyUI's caching mechanism allows you to avoid re-encoding the non-changed part.
|
||||
|
||||
ComfyUI does not use the step number to determine whether to apply conds; instead, it uses the sampler's timestep value which is affected by the scheduler you're using. This means that when the sampler scheduler isn't linear, the schedules generated by prompt control will not be either.
|
||||
for example, if you first encode `[cat:dog:0.1]` and later change that to `[cat:dog:0.5]`, no re-encoding takes place.
|
||||
|
||||
Currently there doesn't seem to be a good way to change this.
|
||||
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`.
|
||||
|
||||
You can try using the `PCSplitSampling` node to enable an alternative method of sampling.
|
||||
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.
|
||||
|
||||
# Scheduling syntax
|
||||
The advanced node enables filtering the prompt for multi-pass workflows.
|
||||
|
||||
Syntax is like A1111 for now, but only fractions are supported for steps.
|
||||
## PCLazyLoraLoader and PCLazyLoraLoaderAdvanced
|
||||
|
||||
```
|
||||
a [large::0.1] [cat|dog:0.05] [<lora:somelora:0.5:0.6>::0.5]
|
||||
[in a park:in space:0.4]
|
||||
```
|
||||
This node reads LoRA expressions from the scheduled prompt and constructs a graph of `LoraLoader`s and `CreateHookLora`s as necessary to provide the necessary LoRA scheduling. Just use it in place of a `LoRALoader` and use the output normally.
|
||||
|
||||
You can also use `a [b:c:0.3,0.7]` as a shortcut. The prompt be `a` until 0.3, `a b` until 0.7, and then `a c`. `[a:0.1,0.4]` is equivalent to `[a::0.1,0.4]`
|
||||
The Advanced node gives you access to the generated hooks. If you have `apply_hooks` set to true, you **do not** need to apply the `HOOKS` output to a CLIP model separately; it's provided in case you want to use it elsewhere. The advanced node also enables filtering the prompt for multi-pass workflows.
|
||||
|
||||
## LoRA loading
|
||||
## PCTextEncode
|
||||
|
||||
LoRAs can be loaded by referring to 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.
|
||||
Encodes a single prompt with advanced (non-scheduling) syntax enabled. This is what actually does most of the work under the hood.
|
||||
|
||||
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.
|
||||
Note: `PCTextEncode` **does not** ignore `<lora:...:1>` and will treat it as part of the prompt. To use a combined prompt for LoRAs and your input, use `PCLazyTextEncode` and `PCLazyLoraLoader`
|
||||
|
||||
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.
|
||||
## PCAddMaskToCLIP
|
||||
|
||||
Finally, you can give the exact path (including the extension) as shown in `LoRALoader`.
|
||||
This node attaches masks to a `CLIP` model so that they can be referred to when using the `IMASK` custom mask function of `PCTextEncode`.
|
||||
|
||||
## PCSetTextEncodeSettings
|
||||
|
||||
## 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
|
||||
|
||||
## 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`
|
||||
|
||||
## Prompt interpolation
|
||||
|
||||
`a red [INT:dog:cat:0.2,0.8:0.05]` will attempt to interpolate the tensors for `a red dog` and `a red cat` between the specified range in as many steps of 0.05 as will fit.
|
||||
|
||||
|
||||
## SDXL
|
||||
|
||||
The nodes do not treat SDXL models specially, but there are some utilities that enable SDXL specific functionality.
|
||||
|
||||
You can use the function `SDXL(width height, target_width target_height, crop_w crop_h)` to set SDXL prompt parameters. `SDXL()` is equivalent to `SDXL(1024 1024, 1024 1024, 0 0)` unless the default values have been overridden by `PCScheduleSettings`.
|
||||
|
||||
To set the `clip_l` prompt, as with `CLIPTextEncodeSDXL`, use the function `CLIP_L(prompt text goes here)`.
|
||||
|
||||
Things to note:
|
||||
- Multiple instances of `CLIP_L` are joined with a space. That is, `CLIP_L(foo)CLIP_L(bar)` is the same as `CLIP_L(foo bar)`
|
||||
- Using `BREAK` isn't supported in it; it'll just parse as the plain word BREAK.
|
||||
- similarly, `AND` inside `CLIP_L` does not do anything sensible; `CLIP_L(foo AND bar)` will parse as two prompts `CLIP_L(foo` and `bar)`
|
||||
- `CLIP_L` and `SDXL` have no effect on SD 1.5.
|
||||
- The rest of the prompt becomes the `clip_g` prompt.
|
||||
- If there is no `CLIP_L` or `SDXL`, the prompts will work as with `CLIPTextEncode`.
|
||||
|
||||
# Other syntax:
|
||||
|
||||
- `<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`.
|
||||
|
||||
- The keyword `BREAK` causes the prompt to be tokenized in separate chunks, which results in each chunk being individually padded to the text encoder's maximum token length. This is mostly equivalent to the `ConditioningConcat` node.
|
||||
|
||||
## Combining prompts
|
||||
`AND` can be used to combine prompts. You can also use a weight at the end. It does a weighted sum of each prompt,
|
||||
```
|
||||
cat :1 AND dog :2
|
||||
```
|
||||
The weight defaults to 1 and are normalized so that `a:2 AND b:2` is equal to `a AND b`. `AND` is processed after schedule parsing, so you can change the weight mid-prompt: `cat:[1:2:0.5] AND dog`
|
||||
|
||||
if there is `COMFYAND()` in the prompt, the behaviour of `AND` will change to work like `ConditioningCombine`, but in practice this seems to be just slower while producing the same output.
|
||||
|
||||
## Functions
|
||||
|
||||
There are some "functions" that can be included in a prompt to do various things.
|
||||
|
||||
Functions have the form `FUNCNAME(param1, param2, ...)`. How parameters are interpreted is up to the function.
|
||||
Note: Whitespace is *not* stripped from string parameters by default. Commas can be escaped with `\,`
|
||||
|
||||
Like `AND`, these functions are parsed after regular scheduling syntax has been expanded, allowing things like `[AREA:MASK:0.3](...)`, in case that's somehow useful.
|
||||
|
||||
### SHUFFLE and SHIFT
|
||||
|
||||
Default parameters: `SHUFFLE(seed=0, separator=,, joiner=,)`, `SHIFT(steps=0, separator=,, joiner=,)`
|
||||
|
||||
`SHIFT` moves elements to the left by `steps`. The default is 0 so `SHIFT()` does nothing
|
||||
`SHUFFLE` generates a random permutation with `seed` as its seed.
|
||||
|
||||
These functions are applied to each prompt chunk **after** `BREAK`, `AND` etc. have been parsed. The prompt is split by `separator`, the operation is applied, and it's then joined back by `joiner`.
|
||||
|
||||
Multiple instances of these functions are applied in the order they appear in the prompt.
|
||||
|
||||
**NOTE:** These functions are *not* smart about syntax and will break emphasis if the separator occurs inside parentheses. I might fix this at some point, but for now, keep this in mind.
|
||||
|
||||
For example:
|
||||
- `SHIFT(1) cat, dog, tiger, mouse` does a shift and results in `dog, tiger, mouse, cat`. (whitespace may vary)
|
||||
- `SHIFT(1,;) cat, dog ; tiger, mouse` results in `tiger, mouse, cat, dog`
|
||||
- `SHUFFLE() cat, dog, tiger, mouse` results in `cat, dog, mouse, tiger`
|
||||
- `SHUFFLE() SHIFT(1) cat, dog, tiger, mouse` results in `dog, mouse, tiger, cat`
|
||||
|
||||
- `SHIFT(1) cat,dog BREAK tiger,mouse` results in `dog,cat BREAK tiger,mouse`
|
||||
- `SHIFT(1) cat, dog AND SHIFT(1) tiger, mouse` results in `dog, cat BREAK mouse, tiger`
|
||||
|
||||
Whitespace is *not* stripped and may also be used as a joiner or separator
|
||||
- `SHIFT(1,, ) cat,dog` results in `dog cat`
|
||||
|
||||
|
||||
### NOISE
|
||||
|
||||
The function `NOISE(weight, seed)` adds some random noise into the prompt. The seed is optional, and if not specified, the global RNG is used. `weight` should be between 0 and 1.
|
||||
|
||||
### MASK, IMASK and AREA
|
||||
|
||||
You can use `MASK(x1 x2, y1 y2, weight, op)` to specify a region mask for a prompt. The values are specified as a percentage with a float between `0` and `1`, or as absolute pixel values (these can't be mixed). `1` will be interpreted as a percentage instead of a pixel value.
|
||||
|
||||
Similarly, you can use `AREA(x1 x2, y1 y2, weight)` to specify an area for the prompt (see ComfyUI's area composition examples). The area is calculated by ComfyUI relative to your latent size.
|
||||
|
||||
#### Custom masks: IMASK and `PCScheduleAddMasks`
|
||||
|
||||
You can attach custom masks to a `PROMPT_SCHEDULE` with the `PCScheduleAddMasks` node 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 `PCScheduleAddMasks` multiple times *appends* masks to a schedule 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 `PCScheduleSettings` and pixel values will be relative to that. ComfyUI will scale the mask to match the image resolution. You can change it manually by using `MASK_SIZE(width, height)` anywhere in the prompt,
|
||||
|
||||
These are handled per `AND`-ed prompt, so in `prompt1 AND MASK(...) prompt2`, the mask will only affect prompt2.
|
||||
|
||||
The default values are `MASK(0 1, 0 1, 1)` and you can omit unnecessary ones, that is, `MASK(0 0.5, 0.3)` is `MASK(0 0.5, 0.3 1, 1)`
|
||||
|
||||
Note that because the default values are percentages, `MASK(0 256, 64 512)` is valid, but `MASK(0 200)` will raise an error.
|
||||
|
||||
Masking does not affect LoRA scheduling unless you set unet weights to 0 for a LoRA.
|
||||
|
||||
### FEATHER
|
||||
|
||||
When you use `MASK` or `IMASK`, you can also call `FEATHER(left top right bottom)` to apply feathering using ComfyUI's `FeatherMask` node. The values are in pixels and default to `0`.
|
||||
|
||||
If multiple masks are used, `FEATHER` is applied *before compositing* in the order they appear in the prompt, and any leftovers are applied to the combined mask. If you want to skip feathering a mask while compositing, just use `FEATHER()` with no arguments.
|
||||
|
||||
For example:
|
||||
```
|
||||
MASK(1) MASK(2) MASK(3) FEATHER(1) FEATHER() FEATHER(3) weirdmask FEATHER(4)
|
||||
```
|
||||
|
||||
gives you a mask that is a combination of 1, 2 and 3, where 1 and 3 are feathered before compositing and then `FEATHER(4)` is applied to the composite.
|
||||
|
||||
The order of the `FEATHER` and `MASK` calls doesn't matter; you can have `FEATHER` before `MASK` or even interleave them.
|
||||
|
||||
# Schedulable LoRAs
|
||||
The `ScheduleToModel` node patches a model so that when sampling, it'll switch LoRAs between steps. You can apply the LoRA's effect separately to CLIP conditioning and the unet (model).
|
||||
|
||||
Swapping LoRAs often can be quite slow without the `--highvram` switch because ComfyUI will shuffle things between the CPU and GPU. When things stay on the GPU, it's quite fast.
|
||||
|
||||
If you run out of VRAM during a LoRA swap, the node will attempt to save VRAM by enabling CPU offloading for future generations even in highvram mode. This persists until ComfyUI is restarted.
|
||||
|
||||
You can also set the `PC_RETRY_ON_OOM` environment variable to any non-empty value to automatically retry sampling once if VRAM runs out.
|
||||
|
||||
## LoRA Block Weight
|
||||
|
||||
If you have [ComfyUI Inspire Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) installed, you can use its Lora Block Weight syntax, for example:
|
||||
|
||||
```
|
||||
a prompt <lora:cars:1:LBW=SD-OUTALL;A=1.0;B=0.0;>
|
||||
```
|
||||
The `;` is optional if there is only 1 parameter.
|
||||
The syntax is the same as in the `ImpactWildcard` node, documented [here](https://github.com/ltdrdata/ComfyUI-extension-tutorials/blob/Main/ComfyUI-Impact-Pack/tutorial/ImpactWildcard.md)
|
||||
|
||||
# Other integrations
|
||||
## Advanced CLIP encoding
|
||||
You can use the syntax `STYLE(weight_interpretation, normalization)` in a prompt to affect how prompts are interpreted.
|
||||
|
||||
Without any extra nodes, only `perp` is available, which does the same as [ComfyUI_PerpWeight](https://github.com/bvhari/ComfyUI_PerpWeight) extension.
|
||||
|
||||
If you have [Advanced CLIP Encoding nodes](https://github.com/BlenderNeko/ComfyUI_ADV_CLIP_emb/tree/master) cloned into your `custom_nodes`, more options will be available.
|
||||
|
||||
The style can be specified separately for each AND:ed prompt, but the first prompt is special; later prompts will "inherit" it as default. For example:
|
||||
|
||||
```
|
||||
STYLE(A1111) a (red:1.1) cat with (brown:0.9) spots and a long tail AND an (old:0.5) dog AND a (green:1.4) (balloon:1.1)
|
||||
```
|
||||
will interpret everything as A1111, but
|
||||
```
|
||||
a (red:1.1) cat with (brown:0.9) spots and a long tail AND STYLE(A1111) an (old:0.5) dog AND a (green:1.4) (balloon:1.1)
|
||||
```
|
||||
Will interpret the first one using the default ComfyUI behaviour, the second prompt with A1111 and the last prompt with the default again
|
||||
|
||||
For things (ie. the code imports) to work, the nodes must be cloned in a directory named exactly `ComfyUI_ADV_CLIP_emb`.
|
||||
|
||||
## Cutoff node integration
|
||||
|
||||
If you have [ComfyUI Cutoff](https://github.com/BlenderNeko/ComfyUI_Cutoff) cloned into your `custom_nodes`, you can use the `CUT` keyword to use cutoff functionality
|
||||
|
||||
The syntax is
|
||||
```
|
||||
a group of animals, [CUT:white cat:white], [CUT:brown dog:brown:0.5:1.0:1.0:_]
|
||||
```
|
||||
the parameters in the `CUT` section are `region_text:target_text:weight;strict_mask:start_from_masked:padding_token` of which only the first two are required.
|
||||
If `strict_mask`, `start_from_masked` or `padding_token` are specified in more than one section, the last one takes effect for the whole prompt
|
||||
|
||||
## Stable-Fast
|
||||
|
||||
The prompt control node works well with [ComfyUI_stable_fast](https://github.com/gameltb/ComfyUI_stable_fast). However, you should apply `ScheduleToModel` **after** applying `Apply StableFast Unet` to prevent constant recompilations.
|
||||
|
||||
# Nodes
|
||||
|
||||
## PromptToSchedule
|
||||
Parses a schedule from a text prompt. A schedule is essentially an array of `(valid_until, prompt)` pairs that the other nodes can use.
|
||||
|
||||
## FilterSchedule
|
||||
Filters a schedule according to its parameters, removing any *changes* that do not occur within `[start, end)`.
|
||||
|
||||
The node also does tag filtering if any tags are specified.
|
||||
|
||||
Always returns at least the last prompt in the schedule if everything would otherwise be filtered.
|
||||
|
||||
`start=0, end=0` returns the prompt at the start and `start=1.0, end=1.0` returns the prompt at the end.
|
||||
|
||||
## ScheduleToCond
|
||||
Produces a combined conditioning for the appropriate timesteps. From a schedule. Also applies LoRAs to the CLIP model according to the schedule.
|
||||
|
||||
## ScheduleToModel
|
||||
Produces a model that'll cause the sampler to reapply LoRAs at specific steps according to the schedule.
|
||||
|
||||
This depends on a callback handled by a monkeypatch of the ComfyUI sampler function, so it might not work with custom samplers, but it shouldn't interfere with them either.
|
||||
|
||||
## PCSplitSampling
|
||||
Causes sampling to be split into multiple sampler calls instead of relying on timesteps for scheduling. This makes the schedules more accurate, but seems to cause weird behaviour with SDE samplers. (Upstream bug?)
|
||||
|
||||
## PCScheduleSettings
|
||||
Returns an object representing **default values** for the `SDXL` function and allows configuring `MASK_SIZE` outside the prompt. You need to apply them to a schedule with `PCApplySettings`. Note that for the SDXL settings to apply, you still need to have `SDXL()` in the prompt.
|
||||
|
||||
The "steps" parameter currently does nothing; it's for future features.
|
||||
|
||||
## PCApplySettings
|
||||
Applies the give default values from `PCScheduleSettings` to a schedule
|
||||
|
||||
## PCPromptFromSchedule
|
||||
|
||||
Extracts a text prompt from a schedule; also logs it to the console.
|
||||
LoRAs are *not* included in the text prompt, though they are logged.
|
||||
|
||||
## PCScheduleAddMasks
|
||||
|
||||
Attaches custom masks to a `PROMPT_SCHEDULE` that can then be used in a prompt.
|
||||
|
||||
## PromptControlSimple
|
||||
This node exists purely for convenience. It's a combination of `PromptToSchedule`, `ScheduleToCond`, `ScheduleToModel` and `FilterSchedule` such that it provides as output a model, positive conds and negative conds, both with and without any specified filters applied.
|
||||
|
||||
This makes it handy for quick one- or two-pass workflows.
|
||||
|
||||
## Older nodes
|
||||
|
||||
- `EditableCLIPEncode`: A combination of `PromptToSchedule` and `ScheduleToCond`
|
||||
- `LoRAScheduler`: A combination of `PromptToSchedule`, `FilterSchedule` and `ScheduleToModel`
|
||||
This node configures `PCTextEncode` default values for some functions by attaching the information to a `CLIP` model.
|
||||
|
||||
# Known issues
|
||||
|
||||
- If you use LoRA scheduling in a workflow with `LoRALoader` nodes, you might get inconsistent results. For now, just avoid mixing `ScheduleToModel` or `LoRAScheduler` with `LoRALoader`. See https://github.com/asagi4/comfyui-prompt-control/issues/36
|
||||
- Workflows using `SamplerCustom` will calculate LoRA schedules based on the number of sigmas given to the sampler instead of the number of steps, since that information isn't available.
|
||||
- `CUT` does not work with `STYLE:perp`
|
||||
- `PCSplitSampling` overrides ComfyUI's `BrownianTreeNoiseSampler` noise sampling behaviour so that each split segment doesn't add crazy amounts of noise to the result with some samplers.
|
||||
- Split sampling may have weird behaviour if your step percentages go below 1 step.
|
||||
- Interpolation is probably buggy and will likely change behaviour whenever code gets refactored.
|
||||
- If execution is interrupted and LoRA scheduling is used, your models might be left in an undefined state until you restart ComfyUI
|
||||
- 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.
|
||||
|
||||
+20
-29
@@ -1,46 +1,37 @@
|
||||
"""
|
||||
@author: asagi4
|
||||
@title: ComfyUI Prompt Control
|
||||
@nickname: ComfyUI Prompt Control
|
||||
@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 os
|
||||
import sys
|
||||
import logging
|
||||
import importlib
|
||||
|
||||
from .prompt_control.node_clip import EditableCLIPEncode, ScheduleToCond
|
||||
from .prompt_control.node_lora import LoRAScheduler, ScheduleToModel, PCSplitSampling, PCWrapGuider
|
||||
from .prompt_control.node_other import (
|
||||
PromptToSchedule,
|
||||
FilterSchedule,
|
||||
PCScheduleSettings,
|
||||
PCScheduleAddMasks,
|
||||
PCApplySettings,
|
||||
PCPromptFromSchedule,
|
||||
)
|
||||
from .prompt_control.node_aio import PromptControlSimple
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
log.propagate = False
|
||||
if not log.handlers:
|
||||
h = logging.StreamHandler(sys.stdout)
|
||||
h.setFormatter(logging.Formatter("[%(levelname)s] PromptControl: %(message)s"))
|
||||
h.setFormatter(logging.Formatter("[PromptControl] %(levelname)s: %(message)s"))
|
||||
log.addHandler(h)
|
||||
|
||||
if os.environ.get("COMFYUI_PC_DEBUG"):
|
||||
if os.environ.get("PROMPTCONTROL_DEBUG"):
|
||||
log.setLevel(logging.DEBUG)
|
||||
else:
|
||||
log.setLevel(logging.INFO)
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy"))
|
||||
cache_hack = importlib.import_module(".prompt_control.cache_hack", package=__name__)
|
||||
cache_hack.init()
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"PromptControlSimple": PromptControlSimple,
|
||||
"PromptToSchedule": PromptToSchedule,
|
||||
"PCSplitSampling": PCSplitSampling,
|
||||
"PCScheduleSettings": PCScheduleSettings,
|
||||
"PCScheduleAddMasks": PCScheduleAddMasks,
|
||||
"PCApplySettings": PCApplySettings,
|
||||
"PCPromptFromSchedule": PCPromptFromSchedule,
|
||||
"PCWrapGuider": PCWrapGuider,
|
||||
"FilterSchedule": FilterSchedule,
|
||||
"ScheduleToCond": ScheduleToCond,
|
||||
"ScheduleToModel": ScheduleToModel,
|
||||
"EditableCLIPEncode": EditableCLIPEncode,
|
||||
"LoRAScheduler": LoRAScheduler,
|
||||
}
|
||||
nodes = ["base", "lazy", "tools", "hooks"]
|
||||
|
||||
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)
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
# 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.
|
||||
|
||||
To enable batching negative prompts, run your positive and negative prompt through the `PPCAttentionCoupleBatchNegative` node. This will make the outputs identical to pamparamm's implementation and will also improve performance. It will fall back to the default behaviour in cases where batching can't be done, so it should always be safe to use.
|
||||
|
||||
|
||||
## Syntax
|
||||
|
||||
See also the main syntax documentation for `MASK` etc.
|
||||
|
||||
### ATTN: 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.
|
||||
|
||||
For example:
|
||||
```
|
||||
dog FILL() ATTN() AND cat MASK(0.5 1) ATTN()
|
||||
```
|
||||
|
||||
If typing `ATTN() MASK()` feels bothersome, try the following macro:
|
||||
```
|
||||
DEF(AM=ATTN() MASK($1))
|
||||
```
|
||||
and then use it like `MASK`: `AM(0 1, 0.5 1)`
|
||||
+400
@@ -0,0 +1,400 @@
|
||||
# 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)`
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 102 KiB |
@@ -0,0 +1,672 @@
|
||||
{
|
||||
"id": "e820c2fb-9502-45b7-a864-684757dddcdf",
|
||||
"revision": 0,
|
||||
"last_node_id": 18,
|
||||
"last_link_id": 20,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 1,
|
||||
"type": "CheckpointLoaderSimple",
|
||||
"pos": [
|
||||
-135,
|
||||
-930
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
98
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "MODEL",
|
||||
"type": "MODEL",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
2
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "CLIP",
|
||||
"type": "CLIP",
|
||||
"slot_index": 1,
|
||||
"links": [
|
||||
3
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "VAE",
|
||||
"type": "VAE",
|
||||
"slot_index": 2,
|
||||
"links": [
|
||||
18
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.18",
|
||||
"Node name for S&R": "CheckpointLoaderSimple"
|
||||
},
|
||||
"widgets_values": [
|
||||
"NoobAI-XL-Vpred-v1.0.safetensors"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"type": "PCLazyTextEncode",
|
||||
"pos": [
|
||||
555,
|
||||
-720
|
||||
],
|
||||
"size": [
|
||||
252,
|
||||
78
|
||||
],
|
||||
"flags": {
|
||||
"collapsed": true
|
||||
},
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "clip",
|
||||
"type": "CLIP",
|
||||
"link": 5
|
||||
},
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "text"
|
||||
},
|
||||
"link": 7
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "CONDITIONING",
|
||||
"type": "CONDITIONING",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
12
|
||||
]
|
||||
}
|
||||
],
|
||||
"title": "PC: Schedule Prompt (positive)",
|
||||
"properties": {
|
||||
"cnr_id": "comfyui-prompt-control",
|
||||
"ver": "2.0.0-beta.7",
|
||||
"Node name for S&R": "PCLazyTextEncode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"STYLE(A1111) 1girl, [painting \\(medium\\), realistic,::0.2] fennec fox girl, animal ear fluff, [[purple:white pupils, purple:0.2] eyes:sparkling eyes:0.85], cargo pants, long sleeves, cardigan, winter, snow, steaming cup, coffee mug, [thermos,:0.1] [long hair,:0.25] [BREAK:0.3]\n[(masterpiece, best quality, newest, very awa,):0.1], night sky, full moon, star \\(sky\\),"
|
||||
],
|
||||
"color": "#232",
|
||||
"bgcolor": "#353"
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"type": "PCLazyLoraLoader",
|
||||
"pos": [
|
||||
257.5,
|
||||
-745
|
||||
],
|
||||
"size": [
|
||||
210,
|
||||
78
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"shape": 7,
|
||||
"type": "MODEL",
|
||||
"link": 2
|
||||
},
|
||||
{
|
||||
"name": "clip",
|
||||
"shape": 7,
|
||||
"type": "CLIP",
|
||||
"link": 3
|
||||
},
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "text"
|
||||
},
|
||||
"link": 6
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "MODEL",
|
||||
"type": "MODEL",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
17
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "CLIP",
|
||||
"type": "CLIP",
|
||||
"slot_index": 1,
|
||||
"links": [
|
||||
5,
|
||||
9
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfyui-prompt-control",
|
||||
"ver": "2.0.0-beta.7",
|
||||
"Node name for S&R": "PCLazyLoraLoader"
|
||||
},
|
||||
"widgets_values": [
|
||||
"STYLE(A1111) 1girl, [painting \\(medium\\), realistic,::0.2] fennec fox girl, animal ear fluff, [[purple:white pupils, purple:0.2] eyes:sparkling eyes:0.85], cargo pants, long sleeves, cardigan, winter, snow, steaming cup, coffee mug, [thermos,:0.1] [long hair,:0.25] [BREAK:0.3]\n[(masterpiece, best quality, newest, very awa,):0.1], night sky, full moon, star \\(sky\\),"
|
||||
],
|
||||
"color": "#223",
|
||||
"bgcolor": "#335"
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"type": "KSampler",
|
||||
"pos": [
|
||||
930,
|
||||
-780
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
474
|
||||
],
|
||||
"flags": {},
|
||||
"order": 9,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "MODEL",
|
||||
"link": 17
|
||||
},
|
||||
{
|
||||
"name": "positive",
|
||||
"type": "CONDITIONING",
|
||||
"link": 12
|
||||
},
|
||||
{
|
||||
"name": "negative",
|
||||
"type": "CONDITIONING",
|
||||
"link": 13
|
||||
},
|
||||
{
|
||||
"name": "latent_image",
|
||||
"type": "LATENT",
|
||||
"link": 14
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "LATENT",
|
||||
"type": "LATENT",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
15
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.18",
|
||||
"Node name for S&R": "KSampler"
|
||||
},
|
||||
"widgets_values": [
|
||||
2,
|
||||
"fixed",
|
||||
25,
|
||||
1.4000000000000001,
|
||||
"euler_cfg_pp",
|
||||
"simple",
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "PrimitiveNode",
|
||||
"pos": [
|
||||
-270,
|
||||
-780
|
||||
],
|
||||
"size": [
|
||||
495,
|
||||
225
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "text"
|
||||
},
|
||||
"links": [
|
||||
6,
|
||||
7
|
||||
]
|
||||
}
|
||||
],
|
||||
"title": "Positive prompt (with LoRAs)",
|
||||
"properties": {
|
||||
"Run widget replace on values": false
|
||||
},
|
||||
"widgets_values": [
|
||||
"STYLE(A1111) 1girl, [painting \\(medium\\), realistic,::0.2] fennec fox girl, animal ear fluff, [[purple:white pupils, purple:0.2] eyes:sparkling eyes:0.85], cargo pants, long sleeves, cardigan, winter, snow, steaming cup, coffee mug, [thermos,:0.1] [long hair,:0.25] [BREAK:0.3]\n[(masterpiece, best quality, newest, very awa,):0.1], night sky, full moon, star \\(sky\\),"
|
||||
],
|
||||
"color": "#232",
|
||||
"bgcolor": "#353"
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "PrimitiveNode",
|
||||
"pos": [
|
||||
-270,
|
||||
-510
|
||||
],
|
||||
"size": [
|
||||
480,
|
||||
225
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "STRING",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "text"
|
||||
},
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
8
|
||||
]
|
||||
}
|
||||
],
|
||||
"title": "Negative prompt",
|
||||
"properties": {
|
||||
"Run widget replace on values": false
|
||||
},
|
||||
"widgets_values": [
|
||||
"chibi, [bad hands,low quality, worst quality,:0.05], simple background, blurry, sketch, unfinished, [holding two cups,no pupils,:0.1]"
|
||||
],
|
||||
"color": "#322",
|
||||
"bgcolor": "#533"
|
||||
},
|
||||
{
|
||||
"id": 7,
|
||||
"type": "PCLazyTextEncode",
|
||||
"pos": [
|
||||
555,
|
||||
-675
|
||||
],
|
||||
"size": [
|
||||
252,
|
||||
78
|
||||
],
|
||||
"flags": {
|
||||
"collapsed": true
|
||||
},
|
||||
"order": 8,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "clip",
|
||||
"type": "CLIP",
|
||||
"link": 9
|
||||
},
|
||||
{
|
||||
"name": "text",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "text"
|
||||
},
|
||||
"link": 8
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "CONDITIONING",
|
||||
"type": "CONDITIONING",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
13
|
||||
]
|
||||
}
|
||||
],
|
||||
"title": "PC: Schedule Prompt (negative)",
|
||||
"properties": {
|
||||
"cnr_id": "comfyui-prompt-control",
|
||||
"ver": "2.0.0-beta.7",
|
||||
"Node name for S&R": "PCLazyTextEncode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"chibi, [bad hands,low quality, worst quality,:0.05], simple background, blurry, sketch, unfinished, [holding two cups,no pupils,:0.1]"
|
||||
],
|
||||
"color": "#322",
|
||||
"bgcolor": "#533"
|
||||
},
|
||||
{
|
||||
"id": 9,
|
||||
"type": "EmptyLatentImage",
|
||||
"pos": [
|
||||
525,
|
||||
-615
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
106
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "LATENT",
|
||||
"type": "LATENT",
|
||||
"links": [
|
||||
14
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.18",
|
||||
"Node name for S&R": "EmptyLatentImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
896,
|
||||
1152,
|
||||
1
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 10,
|
||||
"type": "VAEDecode",
|
||||
"pos": [
|
||||
1290,
|
||||
-780
|
||||
],
|
||||
"size": [
|
||||
210,
|
||||
46
|
||||
],
|
||||
"flags": {},
|
||||
"order": 10,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "samples",
|
||||
"type": "LATENT",
|
||||
"link": 15
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"type": "VAE",
|
||||
"link": 18
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
20
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.18",
|
||||
"Node name for S&R": "VAEDecode"
|
||||
},
|
||||
"widgets_values": []
|
||||
},
|
||||
{
|
||||
"id": 13,
|
||||
"type": "MarkdownNote",
|
||||
"pos": [
|
||||
240,
|
||||
-615
|
||||
],
|
||||
"size": [
|
||||
240,
|
||||
105
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"If you do not need LoRA scheduling, you can simply skip this node."
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 15,
|
||||
"type": "MarkdownNote",
|
||||
"pos": [
|
||||
240,
|
||||
-450
|
||||
],
|
||||
"size": [
|
||||
600,
|
||||
210
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"`PC: Schedule prompt` will expand into instances of `PCTextEncode`. `PC: Schedule LoRAs` will expand into the required `LoRALoader`s and `CLIP` hooks required to schedule LoRAs in the prompt.\n\nYou can pass the same prompt to both nodes; `PC: Schedule Prompt` will simply ignore any `<lora:xyz:1>` elements, so they will not affect the prompt.\nSee the [full syntax available in the prompts](https://github.com/asagi4/comfyui-prompt-control/blob/master/doc/syntax.md) on GitHub"
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 18,
|
||||
"type": "SaveImage",
|
||||
"pos": [
|
||||
1290,
|
||||
-690
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
405
|
||||
],
|
||||
"flags": {},
|
||||
"order": 11,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 20
|
||||
}
|
||||
],
|
||||
"outputs": [],
|
||||
"properties": {
|
||||
"cnr_id": "comfy-core",
|
||||
"ver": "0.3.18"
|
||||
},
|
||||
"widgets_values": [
|
||||
"PromptControl"
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
2,
|
||||
1,
|
||||
0,
|
||||
3,
|
||||
0,
|
||||
"MODEL"
|
||||
],
|
||||
[
|
||||
3,
|
||||
1,
|
||||
1,
|
||||
3,
|
||||
1,
|
||||
"CLIP"
|
||||
],
|
||||
[
|
||||
5,
|
||||
3,
|
||||
1,
|
||||
2,
|
||||
0,
|
||||
"CLIP"
|
||||
],
|
||||
[
|
||||
6,
|
||||
5,
|
||||
0,
|
||||
3,
|
||||
2,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
7,
|
||||
5,
|
||||
0,
|
||||
2,
|
||||
1,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
8,
|
||||
6,
|
||||
0,
|
||||
7,
|
||||
1,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
9,
|
||||
3,
|
||||
1,
|
||||
7,
|
||||
0,
|
||||
"CLIP"
|
||||
],
|
||||
[
|
||||
12,
|
||||
2,
|
||||
0,
|
||||
4,
|
||||
1,
|
||||
"CONDITIONING"
|
||||
],
|
||||
[
|
||||
13,
|
||||
7,
|
||||
0,
|
||||
4,
|
||||
2,
|
||||
"CONDITIONING"
|
||||
],
|
||||
[
|
||||
14,
|
||||
9,
|
||||
0,
|
||||
4,
|
||||
3,
|
||||
"LATENT"
|
||||
],
|
||||
[
|
||||
15,
|
||||
4,
|
||||
0,
|
||||
10,
|
||||
0,
|
||||
"LATENT"
|
||||
],
|
||||
[
|
||||
17,
|
||||
3,
|
||||
0,
|
||||
4,
|
||||
0,
|
||||
"MODEL"
|
||||
],
|
||||
[
|
||||
20,
|
||||
10,
|
||||
0,
|
||||
18,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
18,
|
||||
1,
|
||||
2,
|
||||
10,
|
||||
1,
|
||||
"VAE"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.8,
|
||||
"offset": [
|
||||
591.75,
|
||||
1235
|
||||
]
|
||||
},
|
||||
"linkExtensions": [
|
||||
{
|
||||
"id": 18,
|
||||
"parentId": 1
|
||||
}
|
||||
],
|
||||
"reroutes": [
|
||||
{
|
||||
"id": 1,
|
||||
"pos": [
|
||||
1273.75,
|
||||
-879.5
|
||||
],
|
||||
"linkIds": [
|
||||
18
|
||||
]
|
||||
}
|
||||
]
|
||||
},
|
||||
"version": 0.4,
|
||||
"models": [{
|
||||
"name": "NoobAI-XL-Vpred-v1.0.safetensors",
|
||||
"url": "https://huggingface.co/Laxhar/noobai-XL-Vpred-1.0/resolve/main/NoobAI-XL-Vpred-v1.0.safetensors",
|
||||
"directory": "checkpoints"
|
||||
}]
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,392 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
from math import copysign
|
||||
import logging
|
||||
import itertools
|
||||
from .adv_encode_old import old_advanced_encode_from_tokens
|
||||
|
||||
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 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 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 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
|
||||
|
||||
|
||||
def perp_weight(weights, unweighted_embs, empty_embs):
|
||||
unweighted, unweighted_pooled = unweighted_embs
|
||||
zero, zero_pooled = empty_embs
|
||||
|
||||
weights = weights_like(weights, unweighted)
|
||||
|
||||
if zero.shape != unweighted.shape:
|
||||
zero = zero.repeat(1, unweighted.shape[1] // zero.shape[1], 1)
|
||||
|
||||
perp = (
|
||||
torch.mul(zero, unweighted).sum(dim=-1, keepdim=True) / (unweighted.norm(dim=-1, keepdim=True) ** 2)
|
||||
) * unweighted
|
||||
|
||||
over1 = weights.abs() > 1.0
|
||||
result = unweighted + weights * perp
|
||||
result[~over1] = (unweighted - (1 - weights) * perp)[~over1]
|
||||
result[weights == 0.0] = zero[weights == 0.0]
|
||||
|
||||
# Not sure if this is an implementation bug or if this just doesn't make sense with T5
|
||||
nans = result.isnan()
|
||||
if nans.any():
|
||||
log.warning("perp weight returned NaNs (known to happen with T5), replacing with 0")
|
||||
result[nans] = 0.0
|
||||
|
||||
return result, unweighted_pooled
|
||||
|
||||
|
||||
def style_comfy(encoder, tokens, **kwargs):
|
||||
tokens = encoder.without_word_ids(tokens)
|
||||
return encoder.encode_fn(tokens)
|
||||
|
||||
|
||||
def style_a1111(encoder, tokens, **kwargs):
|
||||
base_emb, pooled = 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
|
||||
|
||||
|
||||
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 = encoder.down_weight(
|
||||
pos_tokens, encoder.weights(tokens), encoder.word_ids(tokens), weighted_emb, pooled
|
||||
)
|
||||
return weighted_emb, pooled
|
||||
|
||||
|
||||
def style_comfypp(encoder, tokens, **kwargs):
|
||||
unweighted_tokens = encoder.unweighted(tokens)
|
||||
base_emb, pooled_base = encoder.base_emb(tokens)
|
||||
weighted_emb, tokens_down, _ = encoder.down_weight(
|
||||
unweighted_tokens, encoder.weights(tokens), encoder.word_ids(tokens), base_emb, pooled_base
|
||||
)
|
||||
weights = encoder.weights(encoder.weighted_with(tokens, lambda w: w if w > 1.0 else 1.0))
|
||||
embs, pooled = encoder.from_masked(
|
||||
unweighted_tokens,
|
||||
weights,
|
||||
encoder.word_ids(tokens),
|
||||
base_emb,
|
||||
pooled_base,
|
||||
)
|
||||
weighted_emb += embs
|
||||
|
||||
return weighted_emb, pooled
|
||||
|
||||
|
||||
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)
|
||||
weighted_emb, _, pooled = encoder.down_weight(
|
||||
encoder.unweighted(tokens), weights, encoder.word_ids(tokens), base_emb, pooled_base
|
||||
)
|
||||
|
||||
return weighted_emb, pooled
|
||||
|
||||
|
||||
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))
|
||||
|
||||
|
||||
def apply_negpip(encoder, emb, pooled, **kwargs):
|
||||
original_tokens = kwargs["original_tokens"]
|
||||
emb_negpip = torch.empty_like(emb).repeat(1, 2, 1)
|
||||
emb_negpip[:, 0::2, :] = emb
|
||||
emb_negpip[:, 1::2, :] = emb * weights_like(encoder.signs(original_tokens), emb)
|
||||
return emb_negpip, pooled
|
||||
|
||||
|
||||
def norm_length(encoder, tokens, **kwargs):
|
||||
word_ids = encoder.word_ids(tokens)
|
||||
sums = dict(zip(*np.unique(word_ids, return_counts=True)))
|
||||
sums[0] = 1
|
||||
tokens = [[(t, _norm_mag(w, sums[id]) if id != 0 else 1.0, id) for (t, w, id) in x] for x in tokens]
|
||||
return tokens
|
||||
|
||||
|
||||
def norm_mean(encoder, tokens, **kwargs):
|
||||
weights = encoder.weights(tokens)
|
||||
word_ids = encoder.word_ids(tokens)
|
||||
delta = 1 - np.mean([w for x, y in zip(weights, word_ids) for w, id in zip(x, y) if id != 0])
|
||||
tokens = [[(t, w if id == 0 else w + delta, id) for (t, w, id) in x] for x in tokens]
|
||||
return tokens
|
||||
|
||||
|
||||
def norm_none(encoder, tokens, **kwargs):
|
||||
return tokens
|
||||
|
||||
|
||||
class AdvancedEncoder:
|
||||
STYLES = {
|
||||
"A1111": style_a1111,
|
||||
"comfy": style_comfy,
|
||||
"comfy++": style_comfypp,
|
||||
"compel": style_compel,
|
||||
"down_weight": style_downweight,
|
||||
"perp": style_perp,
|
||||
}
|
||||
NORMALIZATION_OPS = {
|
||||
"none": norm_none,
|
||||
"length": norm_length,
|
||||
"mean": norm_mean,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def add_encoder(cls, name, fn):
|
||||
cls.STYLES[name] = fn
|
||||
|
||||
def add_normalization_op(cls, name, fn):
|
||||
cls.NORMALIZATION_OPS[name] = fn
|
||||
|
||||
@classmethod
|
||||
def weighted_with(cls, tokens, fn=id, word_ids=True):
|
||||
w = ([(t, fn(w), id) for t, w, id in x] for x in tokens)
|
||||
if not word_ids:
|
||||
w = cls.without_word_ids(w)
|
||||
return list(w)
|
||||
|
||||
@classmethod
|
||||
def unweighted(cls, tokens, word_ids=False):
|
||||
return cls.weighted_with(tokens, fn=lambda w: 1.0, word_ids=word_ids)
|
||||
|
||||
@classmethod
|
||||
def tokens_only(cls, tokens):
|
||||
return list([t[0] for t in x] for x in tokens)
|
||||
|
||||
@classmethod
|
||||
def weights(cls, tokens):
|
||||
return list([t[1] for t in x] for x in tokens)
|
||||
|
||||
@classmethod
|
||||
def word_ids(cls, tokens):
|
||||
return list([t[2] for t in x] for x in tokens)
|
||||
|
||||
@classmethod
|
||||
def signs(cls, tokens):
|
||||
return list([copysign(1, t[1]) for t in x] for x in tokens)
|
||||
|
||||
@classmethod
|
||||
def without_word_ids(cls, tokens):
|
||||
return list([(t, w) for t, w, _ in x] for x in tokens)
|
||||
|
||||
def __init__(self, encode_fn, style, normalization, tokenizer, m_token="+", w_max=1.0, **extra_args):
|
||||
self.encode_fn = encode_fn
|
||||
self.preprocessors = []
|
||||
self.postprocessors = []
|
||||
self.tokenizer = tokenizer
|
||||
self.extra_args = extra_args
|
||||
self.m_token = tokenizer.tokenize_with_weights(m_token)[0][tokenizer.tokens_start]
|
||||
self.max_length = tokenizer.max_length if tokenizer.pad_to_max_length else None
|
||||
self.w_max = w_max
|
||||
|
||||
if style == "comfy++" and not self.max_length:
|
||||
log.warning("comfy++ does not work with tokenizer %s, using default weighting", tokenizer)
|
||||
style = "comfy"
|
||||
|
||||
norms = normalization.split("+")
|
||||
assert style in self.STYLES, f"Invalid weight interpretation: {style}"
|
||||
self.weight_fn = self.STYLES[style]
|
||||
for n in norms:
|
||||
n = n.strip()
|
||||
assert n in self.NORMALIZATION_OPS, f"Invalid normalization: {normalization}"
|
||||
self.preprocessors.append(self.NORMALIZATION_OPS[n])
|
||||
|
||||
negpip = extra_args.get("has_negpip")
|
||||
if negpip:
|
||||
|
||||
def _encode(t):
|
||||
emb, pooled = encode_fn(t)
|
||||
return emb[:, 0::2, :], pooled
|
||||
|
||||
self.encode_fn = _encode
|
||||
self.preprocessors.insert(0, lambda encoder, tokens, **kwargs: encoder.weighted_with(tokens, abs))
|
||||
self.postprocessors.insert(0, apply_negpip)
|
||||
|
||||
def base_emb(self, tokens):
|
||||
unweighted = self.unweighted(tokens)
|
||||
return self.encode_fn(unweighted)
|
||||
|
||||
def down_weight(self, tokens, weights, word_ids, base_emb, pooled_base):
|
||||
w, w_inv = np.unique(weights, return_inverse=True)
|
||||
|
||||
if np.sum(w < 1) == 0:
|
||||
return (
|
||||
base_emb,
|
||||
tokens,
|
||||
(
|
||||
base_emb[0, self.max_length - 1 : self.max_length, :]
|
||||
if (pooled_base is not None and self.max_length)
|
||||
else None
|
||||
),
|
||||
)
|
||||
|
||||
masked_current = tokens
|
||||
emblist = [base_emb]
|
||||
for i in range(len(w)):
|
||||
if w[i] >= 1:
|
||||
continue
|
||||
masked_current = mask_inds(masked_current, np.where(w_inv == i)[0], self.m_token)
|
||||
masked, _ = self.encode_fn(masked_current)
|
||||
emblist.append(masked)
|
||||
|
||||
embs = torch.cat(emblist)
|
||||
w = w[w <= 1.0]
|
||||
w_mix = np.diff([0] + w.tolist())
|
||||
w_mix = torch.tensor(w_mix, dtype=embs.dtype, device=embs.device).reshape((-1, 1, 1))
|
||||
|
||||
weighted_emb = (w_mix * embs).sum(axis=0, keepdim=True)
|
||||
pooled = pooled_base
|
||||
if pooled is not None and self.max_length:
|
||||
pooled = weighted_emb[0, self.max_length - 1 : self.max_length, :]
|
||||
return weighted_emb, masked_current, pooled
|
||||
|
||||
def from_masked(self, tokens, weights, word_ids, base_emb, pooled_base):
|
||||
wids, inds = np.unique(np.array(word_ids).reshape(-1), return_index=True)
|
||||
weight_dict = dict((id, w) for id, w in zip(wids, np.array(weights).reshape(-1)[inds]) if w != 1.0)
|
||||
|
||||
if len(weight_dict) == 0:
|
||||
return torch.zeros_like(base_emb), torch.zeros_like(pooled_base) if pooled_base is not None else None
|
||||
|
||||
weight_tensor = weights_like(weights, base_emb)
|
||||
|
||||
ws = []
|
||||
masked_tokens = []
|
||||
masks = []
|
||||
|
||||
# create prompts
|
||||
for id, w in weight_dict.items():
|
||||
masked, m = mask_word_id(tokens, word_ids, id, self.m_token)
|
||||
masks.append(weights_like(m, base_emb))
|
||||
masked_tokens.extend(masked)
|
||||
|
||||
ws.append(w)
|
||||
|
||||
# TODO: figure out how to get rid of this
|
||||
embs = batched_clip_encode(masked_tokens, self.max_length, self.encode_fn, len(tokens))
|
||||
masks = torch.cat(masks)
|
||||
|
||||
embs = base_emb.expand(embs.shape) - embs
|
||||
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_base + pooled
|
||||
|
||||
if embs.shape[0] != masks.shape[0]:
|
||||
embs = embs.repeat(masks.shape[0], 1, 1)
|
||||
embs *= masks
|
||||
embs = embs.sum(axis=0, keepdim=True)
|
||||
|
||||
return ((weight_tensor - 1) * embs), pooled
|
||||
|
||||
def __call__(self, tokens, apply_to_pooled=False, return_pooled=False):
|
||||
normalized_tokens = tokens
|
||||
for op in self.preprocessors:
|
||||
normalized_tokens = op(self, normalized_tokens)
|
||||
|
||||
emb, pooled = 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
|
||||
|
||||
|
||||
def advanced_encode_from_tokens(
|
||||
tokenized,
|
||||
token_normalization,
|
||||
weight_interpretation,
|
||||
encode_func,
|
||||
m_token="+",
|
||||
w_max=1.0,
|
||||
return_pooled=False,
|
||||
apply_to_pooled=False,
|
||||
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,
|
||||
)
|
||||
@@ -0,0 +1,235 @@
|
||||
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
|
||||
@@ -0,0 +1,241 @@
|
||||
# Lifted from https://github.com/pamparamm/ComfyUI-ppm/blob/c3e6b673ee2d424405dcb99aeed89f21943c89ac/nodes_ppm/attention_couple_ppm.py
|
||||
# Original implementation by laksjdjf, hako-mikan, Haoming02 licensed under GPL-3.0
|
||||
# https://github.com/laksjdjf/cgem156-ComfyUI/blob/1f5533f7f31345bafe4b833cbee15a3c4ad74167/scripts/attention_couple/node.py
|
||||
# https://github.com/Haoming02/sd-forge-couple/blob/e8e258e982a8d149ba59a4bc43b945467604311c/scripts/attention_couple.py
|
||||
import itertools
|
||||
import logging
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from comfy.hooks import EnumHookScope, HookGroup, TransformerOptionsHook, set_hooks_for_conditioning
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def set_cond_attnmask(base_cond, extra_conds, fill=False):
|
||||
hook = AttentionCoupleHook()
|
||||
c = [base_cond[0][0], base_cond[0][1].copy()]
|
||||
# hook uses these, remove them to avoid doing latent masking
|
||||
c[1].pop("mask", None)
|
||||
c[1].pop("strength", None)
|
||||
c[1].pop("mask_strength", None)
|
||||
c = [c]
|
||||
c.extend(base_cond[1:])
|
||||
|
||||
hook.initialize_regions(base_cond[0], extra_conds, fill=fill)
|
||||
group = HookGroup()
|
||||
group.add(hook)
|
||||
|
||||
return set_hooks_for_conditioning(c, hooks=group)
|
||||
|
||||
|
||||
def get_mask(mask, batch_size, num_tokens, extra_options):
|
||||
activations_shape = extra_options["activations_shape"]
|
||||
size = activations_shape[-2:]
|
||||
|
||||
num_conds = mask.shape[0]
|
||||
mask_downsample = F.interpolate(mask, size=size, mode="nearest")
|
||||
mask_downsample_reshaped = mask_downsample.view(num_conds, num_tokens, 1).repeat_interleave(batch_size, dim=0)
|
||||
|
||||
return mask_downsample_reshaped
|
||||
|
||||
|
||||
class Proxy:
|
||||
def __init__(self, function):
|
||||
self.function = function
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
self.function.__self__.to(*args, **kwargs)
|
||||
return self
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
return self.function(*args, *kwargs)
|
||||
|
||||
|
||||
class AttentionCoupleHook(TransformerOptionsHook):
|
||||
COND_UNCOND_COUPLE_OPTION = "cond_or_uncond_hook_couple"
|
||||
COND = 0
|
||||
UNCOND = 1
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(hook_scope=EnumHookScope.HookedOnly)
|
||||
|
||||
self.transformers_dict = {
|
||||
"patches": {
|
||||
"attn2_output_patch": [Proxy(self.attn2_output_patch)],
|
||||
"attn2_patch": [Proxy(self.attn2_patch)],
|
||||
}
|
||||
}
|
||||
self.has_negpip = False
|
||||
|
||||
# calculate later
|
||||
self.conds_k: list[torch.Tensor] = None
|
||||
self.conds_v: list[torch.Tensor] = None
|
||||
|
||||
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.conds: list[torch.Tensor] = [base_cond[0]] + [cond[0] for cond in conds]
|
||||
base_mask = base_cond[1].get("mask", None)
|
||||
masks = [cond[1].get("mask") * cond[1].get("mask_strength") for cond in conds]
|
||||
if len(masks) < 1:
|
||||
raise ValueError("Attention Couple hook makes no sense without masked conds")
|
||||
|
||||
if any(m is None for m in masks):
|
||||
raise ValueError("All conds given to Attention Couple must have masks")
|
||||
|
||||
if any(m.shape != masks[0].shape for m in masks) or (
|
||||
base_mask is not None and base_mask.shape != masks[0].shape
|
||||
):
|
||||
largest_shape = max(m.shape for m in masks)
|
||||
if base_mask is not None:
|
||||
largest_shape = max(largest_shape, base_mask.shape)
|
||||
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)
|
||||
|
||||
if base_mask is not None:
|
||||
base_mask = F.interpolate(base_mask.unsqueeze(1), size=largest_shape[1:], mode="nearest-exact").squeeze(
|
||||
1
|
||||
)
|
||||
|
||||
if base_mask is None:
|
||||
if not fill:
|
||||
raise ValueError("You must specify a base mask when fill=False")
|
||||
sum = torch.stack(masks, dim=0).sum(dim=0)
|
||||
base_mask = torch.zeros_like(sum)
|
||||
base_mask[sum <= 0] = 1.0
|
||||
|
||||
mask = [base_mask] + masks
|
||||
mask = torch.stack(mask, dim=0)
|
||||
if mask.sum(dim=0).min() <= 0 and not fill:
|
||||
raise ValueError("Masks contain non-filled areas")
|
||||
|
||||
self.mask = mask / mask.sum(dim=0, keepdim=True)
|
||||
|
||||
def on_apply_hooks(self, model: ModelPatcher, transformer_options: dict[str, Any]):
|
||||
if self.conds_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:]]
|
||||
else:
|
||||
self.conds_k = self.conds_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)
|
||||
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)
|
||||
return self
|
||||
|
||||
def attn2_patch(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, extra_options):
|
||||
cond_or_uncond = extra_options["cond_or_uncond"]
|
||||
cond_or_uncond_couple = extra_options[self.COND_UNCOND_COUPLE_OPTION] = list(cond_or_uncond)
|
||||
num_chunks = len(cond_or_uncond)
|
||||
|
||||
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))
|
||||
q_chunks = q.chunk(num_chunks, dim=0)
|
||||
k_chunks = k.chunk(num_chunks, dim=0)
|
||||
v_chunks = v.chunk(num_chunks, dim=0)
|
||||
|
||||
bs = q.shape[0] // num_chunks
|
||||
|
||||
conds_k_tensor = conds_v_tensor = torch.cat(
|
||||
[
|
||||
cond.repeat(bs, lcm_tokens_k // cond.shape[1], 1) * self.strengths[i]
|
||||
for i, cond in enumerate(self.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)
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
|
||||
qs, ks, vs = [], [], []
|
||||
cond_or_uncond_couple.clear()
|
||||
|
||||
for i, cond_type in enumerate(cond_or_uncond):
|
||||
q_target = q_chunks[i]
|
||||
k_target = k_chunks[i].repeat(1, lcm_tokens_k // k.shape[1], 1)
|
||||
v_target = v_chunks[i].repeat(1, lcm_tokens_v // v.shape[1], 1)
|
||||
if cond_type == self.UNCOND:
|
||||
qs.append(q_target)
|
||||
ks.append(k_target)
|
||||
vs.append(v_target)
|
||||
cond_or_uncond_couple.append(self.UNCOND)
|
||||
else:
|
||||
qs.append(q_target.repeat(self.num_conds, 1, 1))
|
||||
ks.append(
|
||||
torch.cat(
|
||||
[
|
||||
k_target * self.base_strength,
|
||||
conds_k_tensor,
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
vs.append(
|
||||
torch.cat(
|
||||
[
|
||||
v_target * self.base_strength,
|
||||
conds_v_tensor,
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
cond_or_uncond_couple.extend(itertools.repeat(self.COND, self.num_conds))
|
||||
|
||||
q = torch.cat(qs, dim=0)
|
||||
k = torch.cat(ks, dim=0)
|
||||
v = torch.cat(vs, dim=0)
|
||||
|
||||
return q, k, v
|
||||
|
||||
def attn2_output_patch(self, out, extra_options):
|
||||
cond_or_uncond = extra_options[self.COND_UNCOND_COUPLE_OPTION]
|
||||
bs = out.shape[0] // len(cond_or_uncond)
|
||||
mask_downsample = get_mask(self.mask, bs, out.shape[1], extra_options)
|
||||
outputs = []
|
||||
cond_outputs = []
|
||||
i_cond = 0
|
||||
for i, cond_type in enumerate(cond_or_uncond):
|
||||
pos, next_pos = i * bs, (i + 1) * bs
|
||||
|
||||
if cond_type == self.UNCOND:
|
||||
outputs.append(out[pos:next_pos])
|
||||
else:
|
||||
pos_cond, next_pos_cond = i_cond * bs, (i_cond + 1) * bs
|
||||
masked_output = out[pos:next_pos] * mask_downsample[pos_cond:next_pos_cond]
|
||||
cond_outputs.append(masked_output)
|
||||
i_cond += 1
|
||||
|
||||
if len(cond_outputs) > 0:
|
||||
cond_output = torch.stack(cond_outputs).sum(0)
|
||||
outputs.append(cond_output)
|
||||
|
||||
return torch.cat(outputs, dim=0)
|
||||
@@ -0,0 +1,47 @@
|
||||
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
|
||||
)
|
||||
@@ -0,0 +1,229 @@
|
||||
import torch
|
||||
import copy
|
||||
import re
|
||||
|
||||
import numpy as np
|
||||
import logging
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def replace_embeddings(max_token, prompt, replacements=None):
|
||||
"""Replaces embedding tensors in a token array and replaces them with increasing IDs past max_token"""
|
||||
|
||||
if replacements is None:
|
||||
emb_lookup = []
|
||||
else:
|
||||
emb_lookup = replacements.copy()
|
||||
max_token += len(emb_lookup)
|
||||
|
||||
def get_replacement(embedding):
|
||||
for e, n in emb_lookup:
|
||||
if torch.equal(embedding, e):
|
||||
return n
|
||||
return None
|
||||
|
||||
tokens = []
|
||||
for x in prompt:
|
||||
row = []
|
||||
for i in range(len(x)):
|
||||
emb = x[i][0]
|
||||
if not torch.is_tensor(emb):
|
||||
row.append(emb)
|
||||
else:
|
||||
n = get_replacement(emb)
|
||||
if n is not None:
|
||||
row.append(n)
|
||||
else:
|
||||
max_token += 1
|
||||
row.append(max_token)
|
||||
emb_lookup.append((emb, max_token))
|
||||
tokens.append(row)
|
||||
tokens = np.array(tokens)[:, 1:-1].reshape(-1)
|
||||
return (tokens, emb_lookup)
|
||||
|
||||
|
||||
def unpad_prompt(pad_token, prompt):
|
||||
res = np.trim_zeros(prompt, "b")
|
||||
return np.trim_zeros(res - pad_token, "b") + pad_token
|
||||
|
||||
|
||||
def get_sublists(super_list, sub_list):
|
||||
positions = []
|
||||
for candidate_ind in (i for i, e in enumerate(super_list) if e == sub_list[0]):
|
||||
if super_list[candidate_ind : candidate_ind + len(sub_list)] == sub_list:
|
||||
positions.append(candidate_ind)
|
||||
return positions
|
||||
|
||||
|
||||
def cutoff_add_region(
|
||||
clip_regions, tokenizer, region_text, target_text, weight, strict_mask, start_from_masked, mask_token
|
||||
):
|
||||
"""Adds a cut region to the clip_regions dictionary. It is modified in place"""
|
||||
base_tokens = clip_regions["base_tokens"]
|
||||
region_outputs = []
|
||||
target_outputs = []
|
||||
if strict_mask is not None:
|
||||
clip_regions["strict_mask"] = float(strict_mask)
|
||||
if start_from_masked is not None:
|
||||
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)
|
||||
|
||||
region_text = region_text.strip()
|
||||
target_text = target_text.strip()
|
||||
|
||||
strict_mask = clip_regions["strict_mask"]
|
||||
start_from_masked = clip_regions["start_from_masked"]
|
||||
mask_token = clip_regions["mask_token"]
|
||||
log.info(f"CUT region {region_text=} {target_text=} {weight=} {strict_mask=} {start_from_masked=} {mask_token=}")
|
||||
|
||||
pad_token = tokenizer.end_token
|
||||
|
||||
prompt_tokens, emb_lookup = replace_embeddings(pad_token, base_tokens)
|
||||
|
||||
for rt in region_text.split("\n"):
|
||||
region_tokens = tokenizer.tokenize_with_weights(rt)
|
||||
region_tokens, _ = replace_embeddings(pad_token, region_tokens, emb_lookup)
|
||||
region_tokens = unpad_prompt(pad_token, region_tokens).tolist()
|
||||
|
||||
# calc region mask
|
||||
region_length = len(region_tokens)
|
||||
regions = get_sublists(list(prompt_tokens), region_tokens)
|
||||
|
||||
region_mask = np.zeros(len(prompt_tokens))
|
||||
for r in regions:
|
||||
region_mask[r : r + region_length] = 1
|
||||
region_mask = region_mask.reshape(-1, tokenizer.max_length - 2)
|
||||
region_mask = np.pad(region_mask, pad_width=((0, 0), (1, 1)), mode="constant", constant_values=0)
|
||||
region_mask = region_mask.reshape(1, -1)
|
||||
region_outputs.append(region_mask)
|
||||
|
||||
# calc target mask
|
||||
targets = []
|
||||
for target in target_text.split(" "):
|
||||
# deal with underscores
|
||||
target = re.sub(r"(?<!\\)_", " ", target)
|
||||
target = re.sub(r"\\_", "_", target)
|
||||
|
||||
target_tokens = tokenizer.tokenize_with_weights(target)
|
||||
target_tokens, _ = replace_embeddings(pad_token, target_tokens, emb_lookup)
|
||||
target_tokens = unpad_prompt(pad_token, target_tokens).tolist()
|
||||
|
||||
targets.extend([(x, len(target_tokens)) for x in get_sublists(region_tokens, target_tokens)])
|
||||
targets = [(t_start + r, t_start + t_end + r) for r in regions for t_start, t_end in targets]
|
||||
|
||||
targets_mask = np.zeros(len(prompt_tokens))
|
||||
for t_start, t_end in targets:
|
||||
targets_mask[t_start:t_end] = 1
|
||||
targets_mask = targets_mask.reshape(-1, tokenizer.max_length - 2)
|
||||
targets_mask = np.pad(targets_mask, pad_width=((0, 0), (1, 1)), mode="constant", constant_values=0)
|
||||
targets_mask = targets_mask.reshape(1, -1)
|
||||
target_outputs.append(targets_mask)
|
||||
|
||||
# prepare output
|
||||
region_mask_list = clip_regions["regions"].copy()
|
||||
region_mask_list.extend(region_outputs)
|
||||
target_mask_list = clip_regions["targets"].copy()
|
||||
target_mask_list.extend(target_outputs)
|
||||
weight_list = clip_regions["weights"].copy()
|
||||
weight_list.extend([weight] * len(region_outputs))
|
||||
|
||||
clip_regions["regions"] = region_mask_list
|
||||
clip_regions["targets"] = target_mask_list
|
||||
clip_regions["weights"] = weight_list
|
||||
|
||||
|
||||
def create_masked_prompt(weighted_tokens, mask, mask_token):
|
||||
mask_ids = list(zip(*np.nonzero(mask.reshape((len(weighted_tokens), -1)))))
|
||||
new_prompt = copy.deepcopy(weighted_tokens)
|
||||
for x, y in mask_ids:
|
||||
new_prompt[x][y] = (mask_token,) + new_prompt[x][y][1:]
|
||||
return new_prompt
|
||||
|
||||
|
||||
def process_cuts(encode, extra, tokens):
|
||||
if not extra.get("cuts"):
|
||||
return encode(tokens)
|
||||
|
||||
base = {
|
||||
"base_tokens": tokens,
|
||||
"regions": [],
|
||||
"targets": [],
|
||||
"weights": [],
|
||||
"strict_mask": 1.0,
|
||||
"start_from_masked": 1.0,
|
||||
"mask_token": extra["tokenizer"].tokenizer("+")["input_ids"][1],
|
||||
}
|
||||
|
||||
for cut in extra["cuts"]:
|
||||
cutoff_add_region(base, extra["tokenizer"], *cut)
|
||||
|
||||
return encode_regions(base, encode, extra["tokenizer"])
|
||||
|
||||
|
||||
def debug_tokens(label, prompt, tokenizer):
|
||||
log.debug("Tokens for %s", label)
|
||||
for tokens in prompt:
|
||||
tokens = (t for t in tokens if not torch.is_tensor(t[0]))
|
||||
log.debug(" ".join(f"{x[0][0]} {x[1]}" for x in tokenizer.untokenize(tokens) if x[0][0] != tokenizer.end_token))
|
||||
|
||||
|
||||
def encode_regions(clip_regions, encode, tokenizer):
|
||||
base_weighted_tokens = clip_regions["base_tokens"]
|
||||
start_from_masked = clip_regions["start_from_masked"]
|
||||
mask_token = clip_regions["mask_token"]
|
||||
strict_mask = clip_regions["strict_mask"]
|
||||
|
||||
# calc base embedding
|
||||
base_embedding_full, pool = encode(base_weighted_tokens)
|
||||
|
||||
# Avoid numpy value error and passthrough base embeddings if no regions are set.
|
||||
|
||||
# calc global target mask
|
||||
global_target_mask = np.any(np.stack(clip_regions["targets"]), axis=0).astype(int)
|
||||
|
||||
# calc global region mask
|
||||
global_region_mask = np.any(np.stack(clip_regions["regions"]), axis=0).astype(float)
|
||||
regions_sum = np.sum(np.stack(clip_regions["regions"]), axis=0)
|
||||
regions_normalized = np.divide(1, regions_sum, out=np.zeros_like(regions_sum), where=regions_sum != 0)
|
||||
|
||||
# mask base embeddings
|
||||
base_masked_prompt = create_masked_prompt(base_weighted_tokens, global_target_mask, mask_token)
|
||||
debug_tokens("base_masked", base_masked_prompt, tokenizer)
|
||||
base_embedding_masked, _ = encode(base_masked_prompt)
|
||||
base_embedding_start = base_embedding_full * (1 - start_from_masked) + base_embedding_masked * start_from_masked
|
||||
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"]):
|
||||
region_masking = torch.tensor(
|
||||
regions_normalized * region * weight, dtype=base_embedding_full.dtype, device=base_embedding_full.device
|
||||
).unsqueeze(-1)
|
||||
|
||||
region_prompt = create_masked_prompt(base_weighted_tokens, global_target_mask - target, mask_token)
|
||||
debug_tokens("region", region_prompt, tokenizer)
|
||||
region_emb, _ = encode(region_prompt)
|
||||
region_emb -= base_embedding_start
|
||||
# NegPiP support:
|
||||
if region_emb.shape[1] == 2 * region_masking.shape[1]:
|
||||
region_masking = torch.repeat_interleave(region_masking, 2, dim=1)
|
||||
region_emb *= region_masking
|
||||
|
||||
region_embeddings.append(region_emb)
|
||||
region_embeddings = torch.stack(region_embeddings).sum(axis=0)
|
||||
|
||||
embeddings_final_mask = torch.tensor(
|
||||
global_region_mask, dtype=base_embedding_full.dtype, device=base_embedding_full.device
|
||||
).unsqueeze(-1)
|
||||
# NegPiP support:
|
||||
if region_embeddings.shape[1] == 2 * embeddings_final_mask.shape[1]:
|
||||
embeddings_final_mask = torch.repeat_interleave(embeddings_final_mask, 2, dim=1)
|
||||
|
||||
embeddings_final = base_embedding_start * embeddings_final_mask + base_embedding_outer * (1 - embeddings_final_mask)
|
||||
embeddings_final += region_embeddings
|
||||
return embeddings_final, pool
|
||||
@@ -1,160 +0,0 @@
|
||||
from .utils import get_callback, unpatch_model
|
||||
import sys
|
||||
|
||||
import logging
|
||||
import gc
|
||||
import comfy.model_management
|
||||
import os
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def has_hijack(obj):
|
||||
return hasattr(obj, "pc_hijack_done")
|
||||
|
||||
|
||||
def hijack(obj, attr, replacement):
|
||||
setattr(obj, attr, replacement)
|
||||
setattr(replacement, "pc_hijack_done", True)
|
||||
|
||||
|
||||
def hijack_sampler(module, function, is_custom):
|
||||
mod = sys.modules[module]
|
||||
orig_sampler = getattr(mod, function)
|
||||
if has_hijack(orig_sampler):
|
||||
return
|
||||
|
||||
from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler
|
||||
|
||||
def pc_sample(*args, **kwargs):
|
||||
model = args[0]
|
||||
cb = get_callback(model)
|
||||
BrownianTreeNoiseSampler.pc_reset(
|
||||
model.model_options.get("pc_split_sampling"),
|
||||
kwargs.get("force_full_denoise") or kwargs.get("denoise", 1.0) >= 1.0,
|
||||
)
|
||||
if cb:
|
||||
try:
|
||||
try:
|
||||
r = cb(orig_sampler, is_custom, *args, **kwargs)
|
||||
except comfy.model_management.OOM_EXCEPTION:
|
||||
if not os.environ.get("PC_RETRY_ON_OOM"):
|
||||
raise
|
||||
log.error("Got OOM while sampling, freeing memory and retrying once...")
|
||||
unpatch_model(model)
|
||||
BrownianTreeNoiseSampler.pc_reset(False)
|
||||
gc.collect()
|
||||
comfy.model_management.soft_empty_cache()
|
||||
r = cb(orig_sampler, is_custom, *args, **kwargs)
|
||||
except Exception:
|
||||
log.error("Exception occurred during callback, unpatching model.")
|
||||
unpatch_model(model)
|
||||
BrownianTreeNoiseSampler.pc_reset(False)
|
||||
raise
|
||||
else:
|
||||
r = orig_sampler(*args, **kwargs)
|
||||
BrownianTreeNoiseSampler.pc_reset()
|
||||
return r
|
||||
|
||||
hijack(mod, function, pc_sample)
|
||||
|
||||
|
||||
def hijack_ksampler(module, cls):
|
||||
mod = sys.modules[module]
|
||||
orig_sampler = getattr(mod, cls)
|
||||
if has_hijack(orig_sampler):
|
||||
return
|
||||
|
||||
class HijackedKSampler(orig_sampler):
|
||||
def sample(
|
||||
self,
|
||||
noise,
|
||||
positive,
|
||||
negative,
|
||||
cfg,
|
||||
latent_image=None,
|
||||
start_step=None,
|
||||
last_step=None,
|
||||
force_full_denoise=False,
|
||||
denoise_mask=None,
|
||||
sigmas=None,
|
||||
callback=None,
|
||||
disable_pbar=False,
|
||||
seed=None,
|
||||
):
|
||||
if sigmas is None:
|
||||
sigmas = self.sigmas
|
||||
|
||||
from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler
|
||||
|
||||
BrownianTreeNoiseSampler.set_global_sigmas(self.sigmas)
|
||||
|
||||
return super().sample(
|
||||
noise,
|
||||
positive,
|
||||
negative,
|
||||
cfg,
|
||||
latent_image,
|
||||
start_step,
|
||||
last_step,
|
||||
force_full_denoise,
|
||||
denoise_mask,
|
||||
sigmas,
|
||||
callback,
|
||||
disable_pbar,
|
||||
seed,
|
||||
)
|
||||
|
||||
hijack(mod, cls, HijackedKSampler)
|
||||
|
||||
|
||||
def hijack_browniannoisesampler(module, cls):
|
||||
mod = sys.modules[module]
|
||||
orig_sampler = getattr(mod, cls)
|
||||
if has_hijack(orig_sampler):
|
||||
return
|
||||
|
||||
class PCBrownianTreeNoiseSampler(orig_sampler):
|
||||
global_instance = None
|
||||
use_global_sigmas = False
|
||||
global_sigmas = None
|
||||
force_full_denoise = False
|
||||
|
||||
@classmethod
|
||||
def pc_reset(cls, use_global_sigmas=False, force_full_denoise=False):
|
||||
cls.global_instance = None
|
||||
cls.global_sigmas = None
|
||||
cls.use_global_sigmas = use_global_sigmas
|
||||
cls.force_full_denoise = force_full_denoise
|
||||
|
||||
@classmethod
|
||||
def set_global_sigmas(cls, sigmas):
|
||||
if cls.global_sigmas is None and cls.use_global_sigmas:
|
||||
cls.global_sigmas = (0 if cls.force_full_denoise else sigmas[sigmas > 0].min(), sigmas.max())
|
||||
log.info(
|
||||
"Initializing BrownianTreeNoiseSampler instance with global sigmas %s, %s",
|
||||
cls.global_sigmas,
|
||||
cls.force_full_denoise,
|
||||
)
|
||||
|
||||
def __init__(self, x, sigma_min, sigma_max, **kwargs):
|
||||
if self.global_sigmas is not None:
|
||||
sigma_min, sigma_max = self.global_sigmas
|
||||
if not self.global_instance:
|
||||
super().__init__(x, sigma_min, sigma_max, **kwargs)
|
||||
PCBrownianTreeNoiseSampler.global_instance = self
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
if self.global_instance and self != self.global_instance:
|
||||
return self.global_instance(*args, **kwargs)
|
||||
else:
|
||||
return super().__call__(*args, **kwargs)
|
||||
|
||||
hijack(mod, cls, PCBrownianTreeNoiseSampler)
|
||||
|
||||
|
||||
def do_hijack():
|
||||
hijack_browniannoisesampler("comfy.k_diffusion.sampling", "BrownianTreeNoiseSampler")
|
||||
hijack_sampler("comfy.sample", "sample", False)
|
||||
hijack_sampler("comfy.sample", "sample_custom", True)
|
||||
hijack_ksampler("comfy.samplers", "KSampler")
|
||||
@@ -0,0 +1,46 @@
|
||||
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
|
||||
@@ -1,48 +0,0 @@
|
||||
from .node_clip import control_to_clip_common
|
||||
from .node_lora import schedule_lora_common
|
||||
from .parser import parse_prompt_schedules
|
||||
|
||||
|
||||
class PromptControlSimple:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"clip": ("CLIP",),
|
||||
"positive": ("STRING", {"multiline": True}),
|
||||
"negative": ("STRING", {"multiline": True}),
|
||||
},
|
||||
"optional": {
|
||||
"tags": ("STRING", {"default": ""}),
|
||||
"start": ("FLOAT", {"min": 0.0, "max": 1.0, "step": 0.1, "default": 0.0}),
|
||||
"end": ("FLOAT", {"min": 0.0, "max": 1.0, "step": 0.1, "default": 1.0}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CONDITIONING", "CONDITIONING", "MODEL", "CONDITIONING", "CONDITIONING")
|
||||
RETURN_NAMES = ("model", "positive", "negative", "model_filtered", "pos_filtered", "neg_filtered")
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, model, clip, positive, negative, tags="", start=0.0, end=1.0):
|
||||
lora_cache = {}
|
||||
cond_cache = {}
|
||||
pos_sched = parse_prompt_schedules(positive)
|
||||
pos_cond = pos_filtered = control_to_clip_common(clip, pos_sched, lora_cache, cond_cache)
|
||||
|
||||
neg_sched = parse_prompt_schedules(negative)
|
||||
neg_cond = neg_filtered = control_to_clip_common(clip, neg_sched, lora_cache, cond_cache)
|
||||
|
||||
new_model = model_filtered = schedule_lora_common(model, pos_sched, lora_cache)
|
||||
|
||||
if [tags.strip(), start, end] != ["", 0.0, 1.0]:
|
||||
pos_filtered = control_to_clip_common(
|
||||
clip, pos_sched.with_filters(tags, start, end), lora_cache, cond_cache
|
||||
)
|
||||
neg_filtered = control_to_clip_common(
|
||||
clip, neg_sched.with_filters(tags, start, end), lora_cache, cond_cache
|
||||
)
|
||||
model_filtered = schedule_lora_common(model, pos_sched.with_filters(tags, start, end), lora_cache)
|
||||
|
||||
return (new_model, pos_cond, neg_cond, model_filtered, pos_filtered, neg_filtered)
|
||||
@@ -1,703 +0,0 @@
|
||||
import logging
|
||||
import re
|
||||
import torch
|
||||
from . import utils as utils
|
||||
from .parser import parse_prompt_schedules, parse_cuts
|
||||
from .utils import Timer, equalize, safe_float, get_function, parse_floats
|
||||
from .perp_weight import perp_encode
|
||||
from comfy_extras.nodes_mask import FeatherMask, MaskComposite
|
||||
from node_helpers import conditioning_set_values
|
||||
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
try:
|
||||
from custom_nodes.ComfyUI_ADV_CLIP_emb.adv_encode import (
|
||||
advanced_encode_from_tokens,
|
||||
encode_token_weights_l,
|
||||
encode_token_weights_g,
|
||||
prepareXL,
|
||||
encode_token_weights,
|
||||
)
|
||||
|
||||
have_advanced_encode = True
|
||||
AVAILABLE_STYLES = ["comfy", "A1111", "compel", "comfy++", "down_weight"]
|
||||
AVAILABLE_NORMALIZATIONS = ["none", "mean", "length", "length+mean"]
|
||||
except ImportError:
|
||||
have_advanced_encode = False
|
||||
AVAILABLE_STYLES = ["comfy"]
|
||||
AVAILABLE_NORMALIZATIONS = ["none"]
|
||||
|
||||
try:
|
||||
from custom_nodes.Vector_Sculptor_ComfyUI.nodes import vector_sculptor_tokens
|
||||
|
||||
can_sculpt = True
|
||||
log.info("Vector sculptor extension detected, can use SCULPT()")
|
||||
except ImportError:
|
||||
can_sculpt = False
|
||||
|
||||
|
||||
AVAILABLE_STYLES.append("perp")
|
||||
log.info("Use STYLE(weight_interpretation, normalization) at the start of a prompt to use advanced encodings")
|
||||
log.info("Weight interpretations available: %s", ",".join(AVAILABLE_STYLES))
|
||||
log.info("Normalization types available: %s", ",".join(AVAILABLE_NORMALIZATIONS))
|
||||
|
||||
|
||||
def linear_interpolate_cond(
|
||||
start, end, from_step=0.0, to_step=1.0, step=0.1, start_at=None, end_at=None, prompt_start="N/A", prompt_end="N/A"
|
||||
):
|
||||
count = min(len(start), len(end))
|
||||
if len(start) != len(end):
|
||||
log.info(
|
||||
"Length of conds to interpolate does not match (start=%s != end=%s), interpolating up to %s.",
|
||||
len(start),
|
||||
len(end),
|
||||
count,
|
||||
)
|
||||
|
||||
all_res = []
|
||||
for idx in range(count):
|
||||
res = []
|
||||
from_cond, to_cond = equalize(start[idx][0], end[idx][0])
|
||||
from_pooled = start[idx][1].get("pooled_output")
|
||||
to_pooled = end[idx][1].get("pooled_output")
|
||||
start_at = start_at if start_at is not None else from_step
|
||||
end_at = end_at if end_at is not None else to_step
|
||||
total_steps = int(round((to_step - from_step) / step, 0))
|
||||
num_steps = int(round((end_at - from_step) / step, 0))
|
||||
start_on = int(round((start_at - from_step) / step, 0))
|
||||
start_pct = start_at
|
||||
log.debug(
|
||||
f"interpolate_cond {idx=} {from_step=} {to_step=} {start_at=} {end_at=} {total_steps=} {num_steps=} {start_on=} {step=}"
|
||||
)
|
||||
x = 1 / (total_steps + 1)
|
||||
for s in range(start_on, num_steps):
|
||||
factor = round((s + 1) * x, 2)
|
||||
new_cond = from_cond + (to_cond - from_cond) * factor
|
||||
if from_pooled is not None and to_pooled is not None:
|
||||
from_pooled, to_pooled = equalize(from_pooled, to_pooled)
|
||||
new_pooled = from_pooled + (to_pooled - from_pooled) * factor
|
||||
elif from_pooled is not None:
|
||||
new_pooled = from_pooled
|
||||
|
||||
n = [new_cond, start[idx][1].copy()]
|
||||
if new_pooled is not None:
|
||||
n[1]["pooled_output"] = new_pooled
|
||||
n[1]["start_percent"] = round(start_pct, 2)
|
||||
n[1]["end_percent"] = min(round((start_pct + step), 2), 1.0)
|
||||
start_pct += step
|
||||
start_pct = round(start_pct, 2)
|
||||
if prompt_start:
|
||||
n[1]["prompt"] = f"linear:{round(1.0 - factor, 2)} / {factor}"
|
||||
log.debug(
|
||||
"Interpolating at step %s with factor %s (%s, %s)...",
|
||||
s,
|
||||
factor,
|
||||
n[1]["start_percent"],
|
||||
n[1]["end_percent"],
|
||||
)
|
||||
res.append(n)
|
||||
if res:
|
||||
res[-1][1]["end_percent"] = round(end_at, 2)
|
||||
all_res.extend(res)
|
||||
return all_res
|
||||
|
||||
|
||||
def get_control_points(schedule, steps, encoder):
|
||||
assert len(steps) > 1
|
||||
new_steps = set(steps)
|
||||
|
||||
for step in (s[0] for s in schedule if s[0] >= steps[0] and s[0] <= steps[-1]):
|
||||
new_steps.add(step)
|
||||
control_points = [(s, encoder(schedule.at_step(s)[1])) for s in new_steps]
|
||||
log.debug("Actual control points for interpolation: %s (from %s)", new_steps, steps)
|
||||
return sorted(control_points, key=lambda x: x[0])
|
||||
|
||||
|
||||
def linear_interpolator(control_points, step, start_pct, end_pct):
|
||||
o_start, start = control_points[0]
|
||||
o_end, _ = control_points[-1]
|
||||
t_start = o_start
|
||||
conds = []
|
||||
for t_end, end in control_points[1:]:
|
||||
if t_start < start_pct:
|
||||
t_start, start = t_end, end
|
||||
continue
|
||||
if t_start >= end_pct:
|
||||
break
|
||||
cs = linear_interpolate_cond(start, end, o_start, o_end, step, start_at=t_start, end_at=end_pct)
|
||||
if cs:
|
||||
conds.extend(cs)
|
||||
else:
|
||||
break
|
||||
t_start = t_end
|
||||
start = end
|
||||
return conds
|
||||
|
||||
|
||||
class ScheduleToCond:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP",), "prompt_schedule": ("PROMPT_SCHEDULE",)},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, clip, prompt_schedule):
|
||||
with Timer("ScheduleToCond"):
|
||||
r = (control_to_clip_common(clip, prompt_schedule),)
|
||||
return r
|
||||
|
||||
|
||||
class EditableCLIPEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"clip": ("CLIP",),
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
},
|
||||
"optional": {"filter_tags": ("STRING", {"default": ""})},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "promptcontrol/old"
|
||||
FUNCTION = "parse"
|
||||
|
||||
def parse(self, clip, text, filter_tags=""):
|
||||
parsed = parse_prompt_schedules(text).with_filters(filter_tags)
|
||||
return (control_to_clip_common(clip, parsed),)
|
||||
|
||||
|
||||
def get_sdxl(text, defaults):
|
||||
# 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]
|
||||
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+")
|
||||
cropw, croph = parse_floats(args[2], [d.get("sdxl_cwidth", 0), d.get("sdxl_cheight", 0)], split_re="\\s+")
|
||||
|
||||
opts = {
|
||||
"width": int(w),
|
||||
"height": int(h),
|
||||
"target_width": int(tw),
|
||||
"target_height": int(th),
|
||||
"crop_w": int(cropw),
|
||||
"crop_h": int(croph),
|
||||
}
|
||||
return text, opts
|
||||
|
||||
|
||||
def get_style(text, default_style="comfy", default_normalization="none"):
|
||||
text, styles = get_function(text, "STYLE", [default_style, default_normalization])
|
||||
if not styles:
|
||||
return default_style, default_normalization, text
|
||||
style, normalization = styles[0]
|
||||
style = style.strip()
|
||||
normalization = normalization.strip()
|
||||
if style not in AVAILABLE_STYLES:
|
||||
log.warning("Unrecognized prompt style: %s. Using %s", style, default_style)
|
||||
style = default_style
|
||||
|
||||
if normalization not in AVAILABLE_NORMALIZATIONS:
|
||||
log.warning("Unrecognized prompt normalization: %s. Using %s", normalization, default_normalization)
|
||||
normalization = default_normalization
|
||||
|
||||
return style, normalization, text
|
||||
|
||||
|
||||
def encode_regions(clip, tokens, regions, weight_interpretation="comfy", token_normalization="none"):
|
||||
from custom_nodes.ComfyUI_Cutoff.cutoff import CLIPSetRegion, finalize_clip_regions
|
||||
|
||||
clip_regions = {
|
||||
"clip": clip,
|
||||
"base_tokens": tokens,
|
||||
"regions": [],
|
||||
"targets": [],
|
||||
"weights": [],
|
||||
}
|
||||
|
||||
strict_mask = 1.0
|
||||
start_from_masked = 1.0
|
||||
mask_token = ""
|
||||
|
||||
for region in regions:
|
||||
region_text, target_text, w, sm, sfm, mt = region
|
||||
if w is not None:
|
||||
w = safe_float(w, 0)
|
||||
else:
|
||||
w = 1.0
|
||||
if sm is not None:
|
||||
strict_mask = safe_float(sm, 1.0)
|
||||
if sfm is not None:
|
||||
start_from_masked = safe_float(sfm, 1.0)
|
||||
if mt is not None:
|
||||
mask_token = mt
|
||||
log.info("Region: text %s, target %s, weight %s", region_text.strip(), target_text.strip(), w)
|
||||
(clip_regions,) = CLIPSetRegion.add_clip_region(None, clip_regions, region_text, target_text, w)
|
||||
log.info("Regions: mask_token=%s strict_mask=%s start_from_masked=%s", mask_token, strict_mask, start_from_masked)
|
||||
|
||||
(r,) = finalize_clip_regions(
|
||||
clip_regions, mask_token, strict_mask, start_from_masked, token_normalization, weight_interpretation
|
||||
)
|
||||
cond, pooled = r[0][0], r[0][1].get("pooled_output")
|
||||
return cond, pooled
|
||||
|
||||
|
||||
SHUFFLE_GEN = torch.Generator(device="cpu")
|
||||
|
||||
|
||||
def shuffle_chunk(shuffle, c):
|
||||
func, shuffle = shuffle
|
||||
shuffle_count = int(safe_float(shuffle[0], 0))
|
||||
_, separator, joiner = shuffle
|
||||
if separator == "default":
|
||||
separator = ","
|
||||
|
||||
if not separator:
|
||||
separator = ","
|
||||
|
||||
joiner = {
|
||||
"default": ",",
|
||||
"separator": separator,
|
||||
}.get(joiner, joiner)
|
||||
|
||||
log.info("%s arg=%s sep=%s join=%s", func, shuffle_count, separator, joiner)
|
||||
separated = c.split(separator)
|
||||
if func == "SHIFT":
|
||||
shuffle_count = shuffle_count % len(separated)
|
||||
permutation = separated[shuffle_count:] + separated[:shuffle_count]
|
||||
elif func == "SHUFFLE":
|
||||
SHUFFLE_GEN.manual_seed(shuffle_count)
|
||||
permutation = [separated[i] for i in torch.randperm(len(separated), generator=SHUFFLE_GEN)]
|
||||
else:
|
||||
# ??? should never get here
|
||||
permutation = separated
|
||||
|
||||
permutation = [p for p in permutation if p.strip()]
|
||||
if permutation != separated:
|
||||
c = joiner.join(permutation)
|
||||
return 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"""
|
||||
for key in tokens:
|
||||
max_idx = 0
|
||||
for group in range(len(tokens[key])):
|
||||
for i, token in enumerate(tokens[key][group]):
|
||||
if len(token) < 3:
|
||||
# No need to fix ids when they don't exist
|
||||
return tokens
|
||||
# Ignore zeros, they represent the padding token
|
||||
if token[2] != 0 and token[2] < max_idx:
|
||||
tokens[key][group][i] = (token[0], token[1], token[2] + max_idx)
|
||||
max_idx = max(max_idx, max(x for _, _, x in tokens[key][group]))
|
||||
return tokens
|
||||
|
||||
|
||||
def encode_prompt(clip, text, default_style="comfy", default_normalization="none"):
|
||||
style, normalization, text = get_style(text, default_style, default_normalization)
|
||||
sculpts = []
|
||||
if can_sculpt:
|
||||
text, sculpts = get_function(text, "SCULPT", ["1.0", "forward", "none"])
|
||||
text, regions = parse_cuts(text)
|
||||
# defaults=None means there is no argument parsing at all
|
||||
text, l_prompts = get_function(text, "CLIP_L", defaults=None)
|
||||
chunks = re.split(r"\bBREAK\b", text)
|
||||
token_chunks = []
|
||||
need_word_ids = len(regions) > 0 or (have_advanced_encode and style != "perp")
|
||||
for c in chunks:
|
||||
c, shuffles = get_function(c.strip(), "(SHIFT|SHUFFLE)", ["0", "default", "default"], return_func_name=True)
|
||||
r = c
|
||||
for s in shuffles:
|
||||
r = shuffle_chunk(s, r)
|
||||
if r != c:
|
||||
log.info("Shuffled prompt chunk to %s", r)
|
||||
c = r
|
||||
if sculpts:
|
||||
w, method, norm = sculpts[0]
|
||||
log.info("Using vector sculptor with method=%s norm=%s w=%s", method, norm, w)
|
||||
w = safe_float(w, 1.0)
|
||||
t = vector_sculptor_tokens(clip, c, method, norm, w)
|
||||
else:
|
||||
# Tokenizer returns padded results
|
||||
t = clip.tokenize(c, return_word_ids=need_word_ids)
|
||||
token_chunks.append(t)
|
||||
tokens = token_chunks[0]
|
||||
|
||||
for key in tokens:
|
||||
for c in token_chunks[1:]:
|
||||
tokens[key].extend(c[key])
|
||||
|
||||
# Non-SDXL has only "l"
|
||||
if "g" in tokens and l_prompts:
|
||||
text_l = " ".join(l_prompts)
|
||||
log.info("Encoded SDXL CLIP_L prompt: %s", text_l)
|
||||
tokens["l"] = clip.tokenize(text_l, return_word_ids=need_word_ids)["l"]
|
||||
|
||||
if "g" in tokens and "l" in tokens and len(tokens["l"]) != len(tokens["g"]):
|
||||
empty = clip.tokenize(text_l, return_word_ids=need_word_ids)
|
||||
while len(tokens["l"]) < len(tokens["g"]):
|
||||
tokens["l"] += empty["l"]
|
||||
while len(tokens["l"]) > len(tokens["g"]):
|
||||
tokens["g"] += empty["g"]
|
||||
|
||||
tokens = fix_word_ids(tokens)
|
||||
|
||||
if len(regions) > 0:
|
||||
return encode_regions(clip, tokens, regions, style, normalization)
|
||||
|
||||
if style == "perp":
|
||||
if normalization != "none":
|
||||
log.warning("Normalization is not supported with perp style weighting. Ignored '%s'", normalization)
|
||||
return perp_encode(clip, tokens)
|
||||
|
||||
if have_advanced_encode and not sculpts:
|
||||
if "g" in tokens:
|
||||
embs_l = None
|
||||
embs_g = None
|
||||
pooled = None
|
||||
if "l" in tokens:
|
||||
embs_l, _ = advanced_encode_from_tokens(
|
||||
tokens["l"],
|
||||
normalization,
|
||||
style,
|
||||
lambda x: encode_token_weights(clip, x, encode_token_weights_l),
|
||||
return_pooled=False,
|
||||
)
|
||||
if "g" in tokens:
|
||||
embs_g, pooled = advanced_encode_from_tokens(
|
||||
tokens["g"],
|
||||
normalization,
|
||||
style,
|
||||
lambda x: encode_token_weights(clip, x, encode_token_weights_g),
|
||||
return_pooled=True,
|
||||
apply_to_pooled=False,
|
||||
)
|
||||
# Hardcoded clip_balance
|
||||
return prepareXL(embs_l, embs_g, pooled, 0.5)
|
||||
return advanced_encode_from_tokens(
|
||||
tokens["l"],
|
||||
normalization,
|
||||
style,
|
||||
lambda x: clip.encode_from_tokens({"l": x}, return_pooled=True),
|
||||
return_pooled=True,
|
||||
apply_to_pooled=True,
|
||||
)
|
||||
else:
|
||||
return clip.encode_from_tokens(tokens, return_pooled=True)
|
||||
|
||||
|
||||
def get_area(text):
|
||||
text, areas = get_function(text, "AREA", ["0 1", "0 1", "1"])
|
||||
if not areas:
|
||||
return text, None
|
||||
|
||||
args = areas[0]
|
||||
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)
|
||||
|
||||
def is_pct(f):
|
||||
return f >= 0.0 and f <= 1.0
|
||||
|
||||
def is_pixel(f):
|
||||
return f == 0 or f > 1
|
||||
|
||||
if all(is_pct(v) for v in [h, w, y, x]):
|
||||
area = ("percentage", h, w, y, x)
|
||||
elif all(is_pixel(v) for v in [h, w, y, x]):
|
||||
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"
|
||||
)
|
||||
|
||||
return text, (area, weight)
|
||||
|
||||
|
||||
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]
|
||||
return text, (int(w), int(h))
|
||||
|
||||
|
||||
def make_mask(args, size, weight):
|
||||
x1, x2 = parse_floats(args[0], [0.0, 1.0], split_re="\\s+")
|
||||
y1, y2 = parse_floats(args[1], [0.0, 1.0], split_re="\\s+")
|
||||
|
||||
def is_pct(f):
|
||||
return f >= 0.0 and f <= 1.0
|
||||
|
||||
def is_pixel(f):
|
||||
return f == 0 or f > 1
|
||||
|
||||
if all(is_pct(v) for v in [x1, x2, y1, y2]):
|
||||
w, h = size
|
||||
xs = int(w * x1), int(w * x2)
|
||||
ys = int(h * y1), int(h * y2)
|
||||
elif all(is_pixel(v) for v in [x1, x2, y1, y2]):
|
||||
w, h = size
|
||||
xs = int(x1), int(x2)
|
||||
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"
|
||||
)
|
||||
|
||||
mask = torch.full((h, w), 0, dtype=torch.float32, device="cpu")
|
||||
mask[ys[0] : ys[1], xs[0] : xs[1]] = weight
|
||||
mask = mask.unsqueeze(0)
|
||||
log.info("Mask xs=%s, ys=%s, shape=%s, weight=%s", xs, ys, mask.shape, weight)
|
||||
return mask
|
||||
|
||||
|
||||
def get_mask(text, size, input_masks):
|
||||
"""Parse MASK(x1 x2, y1 y2, weight), IMASK(i, weight) and FEATHER(left top right bottom)"""
|
||||
# TODO: combine multiple masks
|
||||
text, masks = get_function(text, "MASK", ["0 1", "0 1", "1", "multiply"])
|
||||
text, imasks = get_function(text, "IMASK", ["0", "1", "multiply"])
|
||||
text, feathers = get_function(text, "FEATHER", ["0 0 0 0"])
|
||||
text, maskw = get_function(text, "MASKW", ["1.0"])
|
||||
if not masks and not imasks:
|
||||
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)
|
||||
return mask
|
||||
|
||||
mask = None
|
||||
totalweight = 1.0
|
||||
if maskw:
|
||||
totalweight = safe_float(maskw[0][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)
|
||||
if i < len(feathers):
|
||||
nextmask = feather(feathers[i], nextmask)
|
||||
i += 1
|
||||
if mask is not None:
|
||||
log.info("MaskComposite op=%s", op)
|
||||
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
|
||||
else:
|
||||
mask = nextmask
|
||||
|
||||
for idx, w, op in imasks:
|
||||
idx = int(safe_float(idx, 0.0))
|
||||
w = safe_float(w, 1.0)
|
||||
if len(input_masks) < idx + 1:
|
||||
log.warn("IMASK index %s not found, ignoring...", idx)
|
||||
continue
|
||||
nextmask = input_masks[idx] * w
|
||||
if i < len(feathers):
|
||||
nextmask = feather(feathers[i], nextmask)
|
||||
i += 1
|
||||
if mask is not None:
|
||||
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
|
||||
else:
|
||||
mask = nextmask
|
||||
|
||||
# apply leftover FEATHER() specs to the whole
|
||||
for f in feathers[i:]:
|
||||
mask = feather(f, mask)
|
||||
|
||||
return text, mask, totalweight
|
||||
|
||||
|
||||
def get_noise(text):
|
||||
text, noises = get_function(
|
||||
text,
|
||||
"NOISE",
|
||||
["0.0", "none"],
|
||||
)
|
||||
if not noises:
|
||||
return text, None, None
|
||||
w = 0
|
||||
# Only take seed from first noise spec, for simplicity
|
||||
seed = safe_float(noises[0][1], "none")
|
||||
if seed == "none":
|
||||
gen = None
|
||||
else:
|
||||
gen = torch.Generator()
|
||||
gen.manual_seed(int(seed))
|
||||
for n in noises:
|
||||
w += safe_float(n[0], 0.0)
|
||||
return text, max(min(w, 1.0), 0.0), gen
|
||||
|
||||
|
||||
def apply_noise(cond, weight, gen):
|
||||
if cond is None or not weight:
|
||||
return cond
|
||||
|
||||
n = torch.randn(cond.size(), generator=gen).to(cond)
|
||||
|
||||
return cond * (1 - weight) + n * weight
|
||||
|
||||
|
||||
def do_encode(clip, text, 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)
|
||||
|
||||
# Don't sum ANDs if this is in prompt
|
||||
alt_method = "COMFYAND()" in text
|
||||
text = text.replace("COMFYAND()", "")
|
||||
|
||||
prompts = [p.strip() for p in re.split(r"\bAND\b", text)]
|
||||
|
||||
p, sdxl_opts = get_sdxl(prompts[0], defaults)
|
||||
prompts[0] = p
|
||||
|
||||
def weight(t):
|
||||
opts = {}
|
||||
m = re.search(r":(-?\d\.?\d*)(![A-Za-z]+)?$", t)
|
||||
if not m:
|
||||
return (1.0, opts, t)
|
||||
w = float(m[1])
|
||||
tag = m[2]
|
||||
t = t[: m.span()[0]]
|
||||
if tag == "!noscale":
|
||||
opts["scale"] = 1
|
||||
|
||||
return w, opts, t
|
||||
|
||||
conds = []
|
||||
res = []
|
||||
scale = sum(abs(weight(p)[0]) for p in prompts if not ("AREA(" in p or "MASK(" in p))
|
||||
for prompt in prompts:
|
||||
prompt, mask, mask_weight = get_mask(prompt, mask_size, masks)
|
||||
w, opts, prompt = weight(prompt)
|
||||
text, noise_w, generator = get_noise(text)
|
||||
if not w:
|
||||
continue
|
||||
prompt, area = get_area(prompt)
|
||||
prompt, local_sdxl_opts = get_sdxl(prompt, defaults)
|
||||
cond, pooled = encode_prompt(clip, prompt, style, normalization)
|
||||
cond = apply_noise(cond, noise_w, generator)
|
||||
pooled = apply_noise(pooled, noise_w, generator)
|
||||
|
||||
settings = {"prompt": prompt}
|
||||
if alt_method:
|
||||
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
|
||||
|
||||
if mask is not None or area or alt_method or local_sdxl_opts:
|
||||
if pooled is not None:
|
||||
settings["pooled_output"] = pooled
|
||||
conds.append([cond, settings])
|
||||
else:
|
||||
s = opts.get("scale", scale)
|
||||
res.append((cond, pooled, w / s))
|
||||
|
||||
sumconds = [r[0] * r[2] for r in res]
|
||||
pooleds = [r[1] for r in res if r[1] is not None]
|
||||
|
||||
if len(res) > 0:
|
||||
opts = sdxl_opts
|
||||
if pooleds:
|
||||
opts["pooled_output"] = sum(equalize(*pooleds))
|
||||
sumcond = sum(equalize(*sumconds))
|
||||
conds.append([sumcond, opts])
|
||||
return conds
|
||||
|
||||
|
||||
def debug_conds(conds):
|
||||
r = []
|
||||
for i, c in enumerate(conds):
|
||||
x = c[1].copy()
|
||||
if "pooled_output" in x:
|
||||
del x["pooled_output"]
|
||||
r.append((i, x))
|
||||
return r
|
||||
|
||||
|
||||
def control_to_clip_common(clip, schedules, lora_cache=None, cond_cache=None):
|
||||
orig_clip = clip.clone()
|
||||
current_loras = {}
|
||||
if lora_cache is None:
|
||||
lora_cache = {}
|
||||
start_pct = 0.0
|
||||
conds = []
|
||||
cond_cache = cond_cache if cond_cache is not None else {}
|
||||
|
||||
def c_str(c):
|
||||
r = [c["prompt"]]
|
||||
loras = c["loras"]
|
||||
for k in sorted(loras.keys()):
|
||||
r.append(k)
|
||||
r.append(loras[k]["weight_clip"])
|
||||
for lbw, val in loras[k].get("lbw", {}).items():
|
||||
r.append(lbw)
|
||||
r.append(val)
|
||||
return "".join(str(i) for i in r)
|
||||
|
||||
def encode(c):
|
||||
nonlocal clip
|
||||
nonlocal current_loras
|
||||
prompt = c["prompt"]
|
||||
loras = c["loras"]
|
||||
cachekey = c_str(c)
|
||||
cond = cond_cache.get(cachekey)
|
||||
if cond is None:
|
||||
if loras != current_loras:
|
||||
_, clip = utils.apply_loras_from_spec(
|
||||
loras, clip=orig_clip, cache=lora_cache, applied_loras=current_loras
|
||||
)
|
||||
current_loras = loras
|
||||
cond_cache[cachekey] = do_encode(clip, prompt, schedules.defaults, schedules.masks)
|
||||
return cond_cache[cachekey]
|
||||
|
||||
for end_pct, c in schedules:
|
||||
interpolations = [
|
||||
i
|
||||
for i in schedules.interpolations
|
||||
if (start_pct >= i[0][0] and start_pct < i[0][-1]) or (end_pct > i[0][0] and start_pct < i[0][-1])
|
||||
]
|
||||
new_start_pct = start_pct
|
||||
if interpolations:
|
||||
min_step = min(i[1] for i in interpolations)
|
||||
for i in interpolations:
|
||||
control_points, _ = i
|
||||
interpolation_end_pct = min(control_points[-1], end_pct)
|
||||
interpolation_start_pct = max(control_points[0], start_pct)
|
||||
|
||||
control_points = get_control_points(schedules, control_points, encode)
|
||||
cs = linear_interpolator(control_points, min_step, interpolation_start_pct, interpolation_end_pct)
|
||||
conds.extend(cs)
|
||||
new_start_pct = max(new_start_pct, interpolation_end_pct)
|
||||
start_pct = new_start_pct
|
||||
|
||||
if start_pct < end_pct:
|
||||
cond = encode(c)
|
||||
# Node functions return lists of cond
|
||||
cond = conditioning_set_values(
|
||||
cond, {"start_percent": round(start_pct, 2), "end_percent": round(end_pct, 2), "prompt": c["prompt"]}
|
||||
)
|
||||
conds.extend(cond)
|
||||
|
||||
start_pct = end_pct
|
||||
log.debug("Conds at the end: %s", debug_conds(conds))
|
||||
|
||||
log.debug("Final cond info: %s", debug_conds(conds))
|
||||
return conds
|
||||
@@ -1,247 +0,0 @@
|
||||
import logging
|
||||
import torch
|
||||
|
||||
from .utils import unpatch_model, clone_model, set_callback, apply_loras_from_spec
|
||||
from .parser import parse_prompt_schedules
|
||||
from .hijack import do_hijack
|
||||
from comfy.samplers import CFGGuider
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def apply_lora_for_step(schedules, step, total_steps, state, original_model, lora_cache, patch=True):
|
||||
# zero-indexed steps, 0 = first step, but schedules are 1-indexed
|
||||
sched = schedules.at_step(step + 1, total_steps)
|
||||
lora_spec = sched[1]["loras"]
|
||||
|
||||
if state["applied_loras"] != lora_spec:
|
||||
log.debug("At step %s, applying lora_spec %s", step, lora_spec)
|
||||
m, _ = apply_loras_from_spec(
|
||||
lora_spec,
|
||||
model=state["model"],
|
||||
orig_model=original_model,
|
||||
cache=lora_cache,
|
||||
patch=patch,
|
||||
applied_loras=state["applied_loras"],
|
||||
)
|
||||
state["model"] = m
|
||||
state["applied_loras"] = lora_spec
|
||||
|
||||
|
||||
def schedule_lora_common(model, schedules, lora_cache=None):
|
||||
do_hijack()
|
||||
orig_model = clone_model(model)
|
||||
orig_model.model_options["pc_schedules"] = schedules
|
||||
|
||||
if lora_cache is None:
|
||||
lora_cache = {}
|
||||
|
||||
def sampler_cb(orig_sampler, is_custom, *args, **kwargs):
|
||||
split_sampling = args[0].model_options.get("pc_split_sampling")
|
||||
state = {}
|
||||
if is_custom:
|
||||
steps = len(args[4])
|
||||
log.info(
|
||||
"SamplerCustom detected, number of steps not available. LoRA schedules will be calculated based on the number of sigmas (%s)",
|
||||
steps,
|
||||
)
|
||||
else:
|
||||
log.debug("Normal sampler detected, using steps from parameter")
|
||||
steps = args[2]
|
||||
start_step = kwargs.get("start_step") or 0
|
||||
# The model patcher may change if LoRAs are applied
|
||||
state["model"] = args[0]
|
||||
state["applied_loras"] = {}
|
||||
|
||||
orig_cb = kwargs["callback"]
|
||||
|
||||
def step_callback(*args, **kwargs):
|
||||
current_step = args[0] + start_step
|
||||
apply_lora_for_step(schedules, current_step, steps, state, orig_model, lora_cache, patch=True)
|
||||
if orig_cb:
|
||||
return orig_cb(*args, **kwargs)
|
||||
|
||||
kwargs["callback"] = step_callback
|
||||
|
||||
apply_lora_for_step(schedules, start_step, steps, state, orig_model, lora_cache, patch=True)
|
||||
|
||||
def filter_conds(conds, t, start_t, end_t):
|
||||
r = []
|
||||
for c in conds:
|
||||
x = c[1].copy()
|
||||
start_at = round(x["start_percent"], 2)
|
||||
end_at = round(x["end_percent"], 2)
|
||||
# Take any cond that has any effect before end_t, since the percentages may not perfectly match
|
||||
if end_t > start_at and end_t <= end_at:
|
||||
del x["start_percent"]
|
||||
del x["end_percent"]
|
||||
r.append([c[0].clone(), x])
|
||||
else:
|
||||
log.debug("Rejecting cond (%s, %s) between (%s, %s)", start_at, end_at, start_t, end_t)
|
||||
if len(r) == 0:
|
||||
log.error("No %s conds between (%s, %s); Try adjusting your steps", t, start_t, end_t)
|
||||
return r
|
||||
|
||||
def get_steps(conds):
|
||||
for c in conds:
|
||||
yield round(c[1].get("end_percent", 0), 2)
|
||||
|
||||
if split_sampling:
|
||||
actual_end_step = kwargs["last_step"] or steps
|
||||
first_step = True
|
||||
s = args[8]
|
||||
all_steps = sorted(set(int(steps * i) for i in [1.0] + list(get_steps(args[6])) + list(get_steps(args[7]))))
|
||||
for end_step in all_steps:
|
||||
if end_step <= start_step:
|
||||
continue
|
||||
start_t = round(start_step / steps, 2)
|
||||
end_t = round(end_step / steps, 2)
|
||||
new_kwargs = kwargs.copy()
|
||||
new_args = list(args)
|
||||
new_args[0] = state["model"]
|
||||
new_args[6] = filter_conds(new_args[6], "positive", start_t, end_t)
|
||||
new_args[7] = filter_conds(new_args[7], "negative", start_t, end_t)
|
||||
new_args[8] = s
|
||||
log.info("Sampling from %s to %s (total: %s)", start_step, end_step, actual_end_step)
|
||||
new_kwargs["start_step"] = start_step
|
||||
new_kwargs["last_step"] = end_step
|
||||
if end_step >= min(steps, actual_end_step):
|
||||
new_kwargs["force_full_denoise"] = kwargs["force_full_denoise"]
|
||||
else:
|
||||
new_kwargs["force_full_denoise"] = False
|
||||
|
||||
if not first_step:
|
||||
# disable_noise apparently does nothing currently, we need to override noise in args
|
||||
new_kwargs["disable_noise"] = True
|
||||
new_args[1] = torch.zeros_like(s)
|
||||
|
||||
s = orig_sampler(*new_args, **new_kwargs)
|
||||
start_step = end_step
|
||||
first_step = False
|
||||
else:
|
||||
args = list(args)
|
||||
args[0] = state["model"]
|
||||
s = orig_sampler(*args, **kwargs)
|
||||
|
||||
unpatch_model(state["model"])
|
||||
|
||||
return s
|
||||
|
||||
set_callback(orig_model, sampler_cb)
|
||||
|
||||
return orig_model
|
||||
|
||||
|
||||
class PCWrapGuider:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"guider": ("GUIDER",),
|
||||
},
|
||||
}
|
||||
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
RETURN_TYPES = ("GUIDER",)
|
||||
|
||||
def apply(self, guider):
|
||||
return (PCGuider(guider),)
|
||||
|
||||
|
||||
class PCGuider(CFGGuider):
|
||||
def __init__(self, original_guider):
|
||||
if "pc_schedules" not in original_guider.model_patcher.model_options:
|
||||
raise ValueError(
|
||||
"The guider passed to PCWrapGuider must contain a Model that has schedules applied. Use ScheduleToModel"
|
||||
)
|
||||
self.schedules = original_guider.model_patcher.model_options["pc_schedules"]
|
||||
self.guider = original_guider
|
||||
self.lora_cache = {}
|
||||
# sets self.model_patcher
|
||||
super().__init__(original_guider.model_patcher)
|
||||
|
||||
def sample(self, *args, **kwargs):
|
||||
orig_cb = kwargs["callback"]
|
||||
sigmas = args[3]
|
||||
state = {"model": self.guider.model_patcher, "applied_loras": {}}
|
||||
|
||||
def step_callback(*args, **kwargs):
|
||||
apply_lora_for_step(
|
||||
self.schedules,
|
||||
args[0],
|
||||
len(sigmas),
|
||||
state,
|
||||
self.guider.model_patcher,
|
||||
self.lora_cache,
|
||||
patch=True,
|
||||
)
|
||||
if orig_cb:
|
||||
return orig_cb(*args, **kwargs)
|
||||
|
||||
kwargs["callback"] = step_callback
|
||||
apply_lora_for_step(
|
||||
self.schedules, 0, len(sigmas), state, self.guider.model_patcher, self.lora_cache, patch=True
|
||||
)
|
||||
try:
|
||||
r = self.guider.sample(*args, **kwargs)
|
||||
finally:
|
||||
unpatch_model(state["model"])
|
||||
return r
|
||||
|
||||
|
||||
class ScheduleToModel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"prompt_schedule": ("PROMPT_SCHEDULE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, model, prompt_schedule):
|
||||
return (schedule_lora_common(model, prompt_schedule),)
|
||||
|
||||
|
||||
class PCSplitSampling:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"split_sampling": (["enable", "disable"],),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, model, split_sampling):
|
||||
model = clone_model(model)
|
||||
model.model_options["pc_split_sampling"] = split_sampling == "enable"
|
||||
return (model,)
|
||||
|
||||
|
||||
class LoRAScheduler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
CATEGORY = "promptcontrol/old"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, model, text):
|
||||
schedules = parse_prompt_schedules(text)
|
||||
return (schedule_lora_common(model, schedules),)
|
||||
@@ -1,153 +0,0 @@
|
||||
import logging
|
||||
from .parser import parse_prompt_schedules
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
class FilterSchedule:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"prompt_schedule": ("PROMPT_SCHEDULE",)},
|
||||
"optional": {
|
||||
"tags": ("STRING", {"default": ""}),
|
||||
"start": ("FLOAT", {"min": 0.00, "max": 1.00, "default": 0.0, "step": 0.01}),
|
||||
"end": ("FLOAT", {"min": 0.00, "max": 1.00, "default": 1.0, "step": 0.01}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PROMPT_SCHEDULE",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, prompt_schedule, tags="", start=0.0, end=1.0):
|
||||
p = prompt_schedule.with_filters(tags, start=start, end=end)
|
||||
log.debug(
|
||||
f"Filtered {prompt_schedule.parsed_prompt} with: ({tags}, {start}, {end}); the result is %s",
|
||||
p.parsed_prompt,
|
||||
)
|
||||
return (p,)
|
||||
|
||||
|
||||
class PCApplySettings:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"prompt_schedule": ("PROMPT_SCHEDULE",), "settings": ("SCHEDULE_SETTINGS",)}}
|
||||
|
||||
RETURN_TYPES = ("PROMPT_SCHEDULE",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, prompt_schedule, settings):
|
||||
return (prompt_schedule.with_filters(defaults=settings),)
|
||||
|
||||
|
||||
class PCScheduleAddMasks:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"prompt_schedule": ("PROMPT_SCHEDULE",)},
|
||||
"optional": {
|
||||
"mask1": ("MASK",),
|
||||
"mask2": ("MASK",),
|
||||
"mask3": ("MASK",),
|
||||
"mask4": ("MASK",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PROMPT_SCHEDULE",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, prompt_schedule, mask1=None, mask2=None, mask3=None, mask4=None):
|
||||
p = prompt_schedule.clone()
|
||||
p.add_masks(mask1, mask2, mask3, mask4)
|
||||
return (p,)
|
||||
|
||||
|
||||
class PCScheduleSettings:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {},
|
||||
"optional": {
|
||||
"steps": ("INT", {"default": 0, "min": 0, "max": 10000}),
|
||||
"mask_width": ("INT", {"default": 512, "min": 64, "max": 4096 * 4}),
|
||||
"mask_height": ("INT", {"default": 512, "min": 64, "max": 4096 * 4}),
|
||||
"sdxl_width": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
|
||||
"sdxl_height": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
|
||||
"sdxl_target_w": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
|
||||
"sdxl_target_h": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
|
||||
"sdxl_crop_w": ("INT", {"default": 0, "min": 0, "max": 4096 * 4}),
|
||||
"sdxl_crop_h": ("INT", {"default": 0, "min": 0, "max": 4096 * 4}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SCHEDULE_SETTINGS",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(
|
||||
self,
|
||||
steps=0,
|
||||
mask_width=512,
|
||||
mask_height=512,
|
||||
sdxl_width=1024,
|
||||
sdxl_height=1024,
|
||||
sdxl_target_w=1024,
|
||||
sdxl_target_h=1024,
|
||||
sdxl_crop_w=0,
|
||||
sdxl_crop_h=0,
|
||||
):
|
||||
settings = {
|
||||
"steps": steps,
|
||||
"mask_width": mask_width,
|
||||
"mask_height": mask_height,
|
||||
"sdxl_width": sdxl_width,
|
||||
"sdxl_height": sdxl_height,
|
||||
"sdxl_twidth": sdxl_target_w,
|
||||
"sdxl_theight": sdxl_target_h,
|
||||
"sdxl_cwidth": sdxl_crop_w,
|
||||
"sdxl_cheight": sdxl_crop_h,
|
||||
}
|
||||
return (settings,)
|
||||
|
||||
|
||||
class PCPromptFromSchedule:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"prompt_schedule": ("PROMPT_SCHEDULE",),
|
||||
"at": ("FLOAT", {"min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
},
|
||||
"optional": {"tags": ("STRING", {"default": ""})},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, prompt_schedule, at, tags=""):
|
||||
p = prompt_schedule.with_filters(tags, start=at, end=at).parsed_prompt[-1][1]
|
||||
log.info("Prompt at %s:\n%s", at, p["prompt"])
|
||||
log.info("LoRAs: %s", p["loras"])
|
||||
return (p["prompt"],)
|
||||
|
||||
|
||||
class PromptToSchedule:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PROMPT_SCHEDULE",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "parse"
|
||||
|
||||
def parse(self, text, settings=None):
|
||||
schedules = parse_prompt_schedules(text)
|
||||
return (schedules,)
|
||||
@@ -0,0 +1,51 @@
|
||||
import logging
|
||||
from .prompts import encode_prompt
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
class PCTextEncodeWithRange:
|
||||
@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}),
|
||||
},
|
||||
}
|
||||
|
||||
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):
|
||||
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),)
|
||||
|
||||
|
||||
class PCTextEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP",), "text": ("STRING", {"multiline": True})},
|
||||
}
|
||||
|
||||
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)
|
||||
|
||||
|
||||
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)",
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
import logging
|
||||
|
||||
import comfy.hooks
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
from comfy.comfy_types.node_typing import IO, ComfyNodeABC, InputTypeDict
|
||||
|
||||
from .attention_couple_ppm import AttentionCoupleHook
|
||||
from .parser import parse_prompt_schedules
|
||||
from .utils import consolidate_schedule
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
class PCLoraHooksFromText:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"text": ("STRING",)},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("HOOKS",)
|
||||
OUTPUT_TOOLTIPS = ("set of hooks created from the prompt schedule",)
|
||||
CATEGORY = "promptcontrol/v2"
|
||||
FUNCTION = "apply"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def apply(self, text):
|
||||
prompt_schedule = parse_prompt_schedules(text)
|
||||
consolidated = consolidate_schedule(prompt_schedule)
|
||||
hooks = lora_hooks_from_schedule(consolidated, {})
|
||||
return (hooks,)
|
||||
|
||||
|
||||
def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
start_pct = 0.0
|
||||
lora_cache = {}
|
||||
all_hooks = []
|
||||
|
||||
def create_hook(loraspec, start_pct, end_pct, non_scheduled):
|
||||
nonlocal lora_cache
|
||||
hooks = []
|
||||
hook_kf = comfy.hooks.HookKeyframeGroup()
|
||||
for path, info in loras.items():
|
||||
if non_scheduled.get(path) == info:
|
||||
log.info("Skipping %s from hook, it's loaded directly on model", path)
|
||||
continue
|
||||
if path not in lora_cache:
|
||||
lora_cache[path] = comfy.utils.load_torch_file(
|
||||
folder_paths.get_full_path("loras", path), safe_load=True
|
||||
)
|
||||
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']}"
|
||||
hooks.append(new_hook)
|
||||
if start_pct > 0.0:
|
||||
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=0.0)
|
||||
hook_kf.add(kf)
|
||||
kf = comfy.hooks.HookKeyframe(strength=1.0, start_percent=start_pct)
|
||||
hook_kf.add(kf)
|
||||
if end_pct < 1.0:
|
||||
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=end_pct)
|
||||
hook_kf.add(kf)
|
||||
hooks = comfy.hooks.HookGroup.combine_all_hooks(hooks)
|
||||
if hooks:
|
||||
hooks.set_keyframes_on_hooks(hook_kf=hook_kf)
|
||||
return hooks
|
||||
|
||||
for end_pct, loras in schedules:
|
||||
log.info("Creating LoRA hook from %s to %s: %s", start_pct, end_pct, loras)
|
||||
hook = create_hook(loras, start_pct, end_pct, 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):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> InputTypeDict:
|
||||
return {
|
||||
"required": {
|
||||
"positive": (IO.CONDITIONING, {}),
|
||||
"negative": (IO.CONDITIONING, {}),
|
||||
},
|
||||
}
|
||||
|
||||
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):
|
||||
if len(negative) != 1:
|
||||
log.warning("Batching scheduled negatives is not supported yet")
|
||||
return (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())
|
||||
attn_couple = [hook for hook in p_hook_group.hooks if isinstance(hook, AttentionCoupleHook)]
|
||||
for hook in attn_couple:
|
||||
n_hook_group.add(hook)
|
||||
n[1]["hooks"] = p_hook_group if n_hook_group.hooks == p_hook_group.hooks else n_hook_group
|
||||
n[1]["start_percent"] = p[1].get("start_percent", 0.0)
|
||||
n[1]["end_percent"] = p[1].get("end_percent", 1.0)
|
||||
negative_batch.append(n)
|
||||
|
||||
return (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)",
|
||||
}
|
||||
@@ -0,0 +1,295 @@
|
||||
import logging
|
||||
from .parser import parse_prompt_schedules
|
||||
from comfy_execution.graph_utils import GraphBuilder, is_link
|
||||
|
||||
from comfy_execution.graph import ExecutionBlocker
|
||||
|
||||
from .utils import 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():
|
||||
log.info("Creating LoraLoader for %s", path)
|
||||
loader = graph.node("LoraLoader")
|
||||
loader.set_input("model", model)
|
||||
loader.set_input("clip", clip)
|
||||
loader.set_input("strength_model", info["weight"])
|
||||
loader.set_input("strength_clip", info["weight_clip"])
|
||||
loader.set_input("lora_name", path)
|
||||
model = loader.out(0)
|
||||
clip = loader.out(1)
|
||||
return model, clip
|
||||
|
||||
|
||||
def create_hook_nodes_for_lora(graph, path, info, existing_node, start_pct, end_pct):
|
||||
prev_keyframe = None
|
||||
next_keyframe = None
|
||||
if not existing_node:
|
||||
log.debug("Creating hook for %s, weight=%s, weight_clip=%s", path, info["weight"], info["weight_clip"])
|
||||
hook_node = graph.node("CreateHookLora")
|
||||
hook_node.set_input("lora_name", path)
|
||||
hook_node.set_input("strength_model", info["weight"])
|
||||
hook_node.set_input("strength_clip", info["weight_clip"])
|
||||
prev_hook_kf = None
|
||||
if start_pct > 0:
|
||||
log.debug("Creating KF (0, %s) for %s", start_pct, path)
|
||||
prev_keyframe = graph.node("CreateHookKeyframe")
|
||||
prev_keyframe.set_input("strength_mult", 0.0)
|
||||
prev_keyframe.set_input("start_percent", 0.0)
|
||||
prev_hook_kf = prev_keyframe.out(0)
|
||||
else:
|
||||
log.debug("Hook already created for %s", path)
|
||||
hook_node, prev_keyframe = existing_node
|
||||
prev_hook_kf = prev_keyframe.out(0)
|
||||
|
||||
if (
|
||||
prev_keyframe
|
||||
and prev_keyframe.get_input("start_pct") == start_pct
|
||||
and prev_keyframe.get_input("strength_mult") == 0.0
|
||||
):
|
||||
next_keyframe = prev_keyframe
|
||||
log.debug("Previous keyframe for %s starts at %s and has 0 strength, overriding", path, start_pct)
|
||||
else:
|
||||
log.debug("Creating keyframe for %s, start=%s ", path, start_pct)
|
||||
next_keyframe = graph.node("CreateHookKeyframe")
|
||||
next_keyframe.set_input("start_percent", start_pct)
|
||||
next_keyframe.set_input("prev_hook_kf", prev_hook_kf)
|
||||
|
||||
next_keyframe.set_input("strength_mult", 1.0)
|
||||
prev_hook_kf = next_keyframe.out(0)
|
||||
if end_pct < 1.0:
|
||||
log.debug("Creating end keyframe for %s, start=%s", path, end_pct)
|
||||
next_keyframe = graph.node("CreateHookKeyframe")
|
||||
next_keyframe.set_input("strength_mult", 0.0)
|
||||
next_keyframe.set_input("start_percent", end_pct)
|
||||
next_keyframe.set_input("prev_hook_kf", prev_hook_kf)
|
||||
return hook_node, next_keyframe
|
||||
|
||||
|
||||
def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True):
|
||||
# This gets rid of non-existent LoRAs
|
||||
consolidated = consolidate_schedule(schedule)
|
||||
if model is not None:
|
||||
non_scheduled = find_nonscheduled_loras(consolidated)
|
||||
model, clip = create_lora_loader_nodes(graph, model, clip, non_scheduled)
|
||||
else:
|
||||
non_scheduled = {}
|
||||
model = ExecutionBlocker("No model provided to PCLazyLoRALoader or PCLazyLoRALoaderAdvanced")
|
||||
|
||||
hook_nodes = {}
|
||||
start_pct = 0.0
|
||||
|
||||
def key(lora, info):
|
||||
return f"{lora}-{info['weight']}-{info['weight_clip']}"
|
||||
|
||||
for end_pct, loras in consolidated:
|
||||
for lora, info in loras.items():
|
||||
if non_scheduled.get(lora) == info:
|
||||
continue
|
||||
k = key(lora, info)
|
||||
existing_node = hook_nodes.get(k)
|
||||
hook_nodes[k] = create_hook_nodes_for_lora(graph, lora, info, existing_node, start_pct, end_pct)
|
||||
start_pct = end_pct
|
||||
|
||||
hooks = []
|
||||
# Attach the keyframe chain to the hook node
|
||||
for hook, kfs in hook_nodes.values():
|
||||
n = graph.node("SetHookKeyframes")
|
||||
n.set_input("hooks", hook.out(0))
|
||||
n.set_input("hook_kf", kfs.out(0))
|
||||
hooks.append(n)
|
||||
|
||||
res = None
|
||||
# Finally, combine all hooks and optionally apply
|
||||
if len(hooks) > 0:
|
||||
res = hooks[0]
|
||||
for h in hooks[1:]:
|
||||
n = graph.node("CombineHooks2")
|
||||
n.set_input("hooks_A", res.out(0))
|
||||
n.set_input("hooks_B", h.out(0))
|
||||
res = n
|
||||
res = res.out(0)
|
||||
if clip is not None and apply_hooks:
|
||||
n = graph.node("SetClipHooks")
|
||||
n.set_input("clip", clip)
|
||||
n.set_input("hooks", res)
|
||||
n.set_input("apply_to_conds", True)
|
||||
n.set_input("schedule_clip", True)
|
||||
clip = n.out(0)
|
||||
|
||||
if clip is None:
|
||||
clip = ExecutionBlocker("No clip model provided to PCLazyLoRALoader or PCLazyLoRALoaderAdvanced")
|
||||
r = graph.finalize()
|
||||
log.debug("LazyLoraLoader built graph: %s", json.dumps(r))
|
||||
|
||||
ret = (model, clip, res)
|
||||
|
||||
return {"result": ret, "expand": r}
|
||||
|
||||
|
||||
class PCLazyLoraLoaderAdvanced:
|
||||
CACHE_KEY = cache_key_lora
|
||||
|
||||
@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
|
||||
):
|
||||
schedule = parse_prompt_schedules(text, filters=tags, start=start, end=end, num_steps=num_steps)
|
||||
graph = GraphBuilder(f"{unique_id}-")
|
||||
r = build_lora_schedule(graph, schedule, model, clip, apply_hooks=apply_hooks)
|
||||
return r
|
||||
|
||||
|
||||
class PCLazyLoraLoader(PCLazyLoraLoaderAdvanced):
|
||||
@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"},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (
|
||||
"MODEL",
|
||||
"CLIP",
|
||||
)
|
||||
CATEGORY = "promptcontrol"
|
||||
|
||||
def apply(self, *args, **kwargs):
|
||||
r = super().apply(*args, **kwargs)
|
||||
r["result"] = r["result"][:2]
|
||||
return r
|
||||
|
||||
|
||||
def build_scheduled_prompts(graph, schedules, clip):
|
||||
nodes = []
|
||||
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"
|
||||
if classnames:
|
||||
classname = classnames[0][0]
|
||||
paramname = classnames[0][1]
|
||||
node = graph.node(classname)
|
||||
node.set_input("clip", clip)
|
||||
node.set_input(paramname, p)
|
||||
timestep = graph.node("ConditioningSetTimestepRange")
|
||||
timestep.set_input("conditioning", node.out(0))
|
||||
timestep.set_input("start", start_pct)
|
||||
timestep.set_input("end", end_pct)
|
||||
nodes.append(timestep)
|
||||
start_pct = end_pct
|
||||
node = nodes[0]
|
||||
for othernode in nodes[1:]:
|
||||
combiner = graph.node("ConditioningCombine")
|
||||
combiner.set_input("conditioning_1", node.out(0))
|
||||
combiner.set_input("conditioning_2", othernode.out(0))
|
||||
node = combiner
|
||||
|
||||
g = graph.finalize()
|
||||
log.debug("Built graph: %s", json.dumps(g))
|
||||
|
||||
return {"result": (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
|
||||
|
||||
@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):
|
||||
schedules = parse_prompt_schedules(text, filters=tags, start=start, end=end, num_steps=num_steps)
|
||||
graph = GraphBuilder(f"{unique_id}-")
|
||||
return build_scheduled_prompts(graph, schedules, clip)
|
||||
|
||||
|
||||
class PCLazyTextEncode(PCLazyTextEncodeAdvanced):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP", {"rawLink": True}), "text": ("STRING", {"multiline": True})},
|
||||
"hidden": {"unique_id": "UNIQUE_ID"},
|
||||
}
|
||||
|
||||
CATEGORY = "promptcontrol"
|
||||
|
||||
|
||||
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)",
|
||||
}
|
||||
@@ -0,0 +1,248 @@
|
||||
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
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
class PCSaveExpandedWorkflow:
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_output_directory()
|
||||
|
||||
@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
|
||||
)
|
||||
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"):
|
||||
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"
|
||||
|
||||
|
||||
class PCAddMaskToCLIP:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP",)},
|
||||
"optional": {
|
||||
"mask": ("MASK",),
|
||||
},
|
||||
}
|
||||
|
||||
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",),
|
||||
},
|
||||
}
|
||||
|
||||
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):
|
||||
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,)
|
||||
|
||||
|
||||
class PCSetPCTextEncodeSettings:
|
||||
@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}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CLIP",)
|
||||
CATEGORY = "promptcontrol/tools"
|
||||
FUNCTION = "apply"
|
||||
DESCRIPTION = "Configures default values for PCTextEncode"
|
||||
|
||||
def apply(
|
||||
self,
|
||||
clip,
|
||||
mask_width=512,
|
||||
mask_height=512,
|
||||
sdxl_width=1024,
|
||||
sdxl_height=1024,
|
||||
sdxl_target_w=1024,
|
||||
sdxl_target_h=1024,
|
||||
sdxl_crop_w=0,
|
||||
sdxl_crop_h=0,
|
||||
):
|
||||
settings = {
|
||||
"mask_width": mask_width,
|
||||
"mask_height": mask_height,
|
||||
"sdxl_width": sdxl_width,
|
||||
"sdxl_height": sdxl_height,
|
||||
"sdxl_twidth": sdxl_target_w,
|
||||
"sdxl_theight": sdxl_target_h,
|
||||
"sdxl_cwidth": sdxl_crop_w,
|
||||
"sdxl_cheight": sdxl_crop_h,
|
||||
}
|
||||
clip = clip.clone()
|
||||
clip.patcher.model_options["x-promptcontrol.settings"] = settings
|
||||
return (clip,)
|
||||
|
||||
|
||||
class PCExtractScheduledPrompt:
|
||||
@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": ""})},
|
||||
}
|
||||
|
||||
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=""):
|
||||
schedule = parse_prompt_schedules(text, filters=tags)
|
||||
_, entry = schedule.at_step(at, total_steps=1)
|
||||
prompt_text = entry.get("prompt", "")
|
||||
return (prompt_text,)
|
||||
|
||||
|
||||
class PCMacroExpand:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
},
|
||||
}
|
||||
|
||||
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),)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"PCSetPCTextEncodeSettings": PCSetPCTextEncodeSettings,
|
||||
"PCAddMaskToCLIP": PCAddMaskToCLIP,
|
||||
"PCAddMaskToCLIPMany": PCAddMaskToCLIPMany,
|
||||
"PCSetLogLevel": PCSetLogLevel,
|
||||
"PCExtractScheduledPrompt": PCExtractScheduledPrompt,
|
||||
"PCSaveExpandedWorkflow": PCSaveExpandedWorkflow,
|
||||
"PCMacroExpand": PCMacroExpand,
|
||||
}
|
||||
|
||||
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",
|
||||
}
|
||||
+163
-109
@@ -1,23 +1,40 @@
|
||||
# vim: sw=4 ts=4
|
||||
import lark
|
||||
import logging
|
||||
from math import ceil
|
||||
|
||||
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 | interpolate | loraspec | PLAIN | /</ | />/ | WHITESPACE)+
|
||||
prompt: (emphasized | embedding | scheduled | alternate | sequence | loraspec | PLAIN | /</ | />/ | WHITESPACE)+
|
||||
!emphasized: "(" prompt? ")"
|
||||
| "(" prompt ":" prompt ")"
|
||||
| "[" prompt "]"
|
||||
scheduled: "[" [prompt ":"] [prompt] ":" _WS? NUMBER ["," NUMBER] "]"
|
||||
| "[" [prompt ":"] [prompt] ":" _WS? TAG "]"
|
||||
sequence: "[SEQ" ":" [prompt] ":" NUMBER (":" [prompt] ":" NUMBER)+ "]"
|
||||
interpolate.100: "[INT" ":" interp_prompts ":" interp_steps "]"
|
||||
interp_prompts: prompt (":" [prompt])+
|
||||
interp_steps: NUMBER ("," NUMBER)+ [":" NUMBER]
|
||||
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
|
||||
@@ -33,6 +50,7 @@ TAG: /[A-Z_]+/
|
||||
lexer="dynamic",
|
||||
)
|
||||
|
||||
|
||||
cut_parser = lark.Lark(
|
||||
r"""
|
||||
!start: (prompt | /[][:()]/+)*
|
||||
@@ -74,7 +92,7 @@ def parse_cuts(text):
|
||||
|
||||
|
||||
def flatten(x):
|
||||
if type(x) in [str, tuple] or isinstance(x, dict) and "type" in 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:
|
||||
@@ -86,14 +104,25 @@ def clamp(a, b, c):
|
||||
return min(max(a, b), c)
|
||||
|
||||
|
||||
def get_steps(tree):
|
||||
res = [100]
|
||||
interpolation_steps = []
|
||||
def get_steps(tree, num_steps):
|
||||
res = [num_steps or 100]
|
||||
|
||||
def tostep(s):
|
||||
w = float(s) * 100
|
||||
w = int(clamp(0, w, 100))
|
||||
return w
|
||||
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):
|
||||
@@ -110,55 +139,60 @@ def get_steps(tree):
|
||||
for i, _ in enumerate(tree.children[:-1]):
|
||||
tree.children[i] = tostep(tree.children[i])
|
||||
|
||||
interpolation_steps.append((tuple(tree.children[:-1]), tree.children[-1]))
|
||||
res.extend(tree.children[:-1])
|
||||
|
||||
def sequence(self, tree):
|
||||
steps = tree.children[1::2]
|
||||
for i, steps in enumerate(steps):
|
||||
w = float(tree.children[i * 2 + 1]) * 100
|
||||
tree.children[i * 2 + 1] = clamp(0, w, 100)
|
||||
w = tostep(tree.children[i * 2 + 1])
|
||||
tree.children[i * 2 + 1] = w
|
||||
res.append(w)
|
||||
|
||||
def alternate(self, tree):
|
||||
step_size = int(round(float(tree.children[-1] or 0.1), 2) * 100)
|
||||
step_size = clamp(1, step_size, 100)
|
||||
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, 100, step_size)])
|
||||
res.extend([x for x in range(step_size, num_steps or 100, step_size)])
|
||||
|
||||
CollectSteps().visit(tree)
|
||||
|
||||
return sorted(set(interpolation_steps)), sorted(set(res))
|
||||
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
|
||||
before, after, when, *rest = args
|
||||
if isinstance(when, str):
|
||||
return before or "" if when not in filters else after or ""
|
||||
|
||||
pl, when, *rest = args
|
||||
if rest:
|
||||
when_end = rest[0]
|
||||
|
||||
if when_end is not None and step <= when and before is not None:
|
||||
return ""
|
||||
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 when_end is not None and (step > when and step <= when_end):
|
||||
# handle [a:0,1]
|
||||
if before is None:
|
||||
return after or ""
|
||||
return before or ""
|
||||
if isinstance(when, str):
|
||||
return before or "" if when not in filters else after or ""
|
||||
|
||||
if when_end is not None and step >= when_end:
|
||||
# handle [a:0,1]
|
||||
if before is None:
|
||||
return ""
|
||||
return 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 ""
|
||||
|
||||
@@ -174,24 +208,6 @@ def at_step(step, filters, tree):
|
||||
previous_step = s
|
||||
return ""
|
||||
|
||||
def interpolate(self, args):
|
||||
prompts, starts = args
|
||||
starts = starts[:-1]
|
||||
prev_prompt = None
|
||||
if step < starts[0]:
|
||||
return prompts[0]
|
||||
for i, x in enumerate(starts):
|
||||
prev_prompt = prompts[i]
|
||||
if x >= step:
|
||||
break
|
||||
return prev_prompt
|
||||
|
||||
def interp_steps(self, args):
|
||||
return list(args)
|
||||
|
||||
def interp_prompts(self, args):
|
||||
return ["".join(flatten(a or [])) for a in args]
|
||||
|
||||
def alternate(self, args):
|
||||
step_size = args[-1]
|
||||
idx = ceil(step / step_size)
|
||||
@@ -262,68 +278,46 @@ def at_step(step, filters, tree):
|
||||
|
||||
|
||||
class PromptSchedule(object):
|
||||
def __init__(self, prompt, filters="", start=0.0, end=1.0, defaults=None, masks=None):
|
||||
# 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 = {}
|
||||
if defaults:
|
||||
self.defaults = defaults
|
||||
self.loaded_loras = {}
|
||||
|
||||
self.interpolations = None
|
||||
self.parsed_prompt = None
|
||||
self.interpolations, self.parsed_prompt = self._parse()
|
||||
self.masks = masks
|
||||
if masks is None:
|
||||
self.masks = []
|
||||
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):
|
||||
def _parse(self, num_steps):
|
||||
filters = [x.strip() for x in self.filters.upper().split(",")]
|
||||
try:
|
||||
parsed = []
|
||||
interpolations = set()
|
||||
tree = prompt_parser.parse(self.prompt)
|
||||
interpolation_steps, steps = get_steps(tree)
|
||||
log.debug("Interpolation steps: %s", interpolation_steps)
|
||||
steps = get_steps(tree, num_steps=num_steps)
|
||||
|
||||
def f(x):
|
||||
return round(x / 100, 2)
|
||||
return round(x / (num_steps or 100), 2)
|
||||
|
||||
for t in steps:
|
||||
p = at_step(t, filters, tree)
|
||||
for control_points, step in interpolation_steps:
|
||||
interp_start = None
|
||||
interp_end = None
|
||||
if t == control_points[-1]:
|
||||
interp_start = max(control_points[0], int(self.start * 100))
|
||||
interp_end = min(control_points[-1], int(self.end * 100))
|
||||
control_points = tuple(
|
||||
sorted(set(f(c) for c in control_points if c >= interp_start or c <= interp_end))
|
||||
)
|
||||
if interp_start is not None and interp_end is not None and interp_end > interp_start:
|
||||
interpolations.add((control_points, f(step)))
|
||||
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_p = None
|
||||
prev_end = -1
|
||||
|
||||
for end_at, p in parsed:
|
||||
# Preserve prompt if it ends at the start of an interpolation, otherwise bump its end time
|
||||
if p == prev_p and res[-1][0] not in [x[0][0] for x in interpolations]:
|
||||
res[-1][0] = end_at
|
||||
continue
|
||||
if end_at < self.start:
|
||||
continue
|
||||
elif end_at <= self.end:
|
||||
@@ -332,18 +326,20 @@ class PromptSchedule(object):
|
||||
elif end_at > self.end and prev_end < self.end:
|
||||
res.append([end_at, p])
|
||||
break
|
||||
prev_p = p
|
||||
|
||||
# Always use the last prompt if everything was filtered
|
||||
if len(res) == 0:
|
||||
res = [[1.0, parsed[-1][1]]]
|
||||
|
||||
return interpolations, res
|
||||
final = [res[0]]
|
||||
|
||||
def add_masks(self, *masks):
|
||||
for mask in masks:
|
||||
if mask is not None:
|
||||
self.masks.append(mask)
|
||||
# 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()
|
||||
@@ -357,8 +353,7 @@ class PromptSchedule(object):
|
||||
filters=ifspecified(filters, self.filters),
|
||||
start=ifspecified(start, self.start),
|
||||
end=ifspecified(end, self.end),
|
||||
defaults=ifspecified(defaults, self.defaults),
|
||||
masks=self.masks[:],
|
||||
num_steps=self.num_steps,
|
||||
)
|
||||
return p
|
||||
|
||||
@@ -372,23 +367,82 @@ class PromptSchedule(object):
|
||||
return i, x
|
||||
return len(self.parsed_prompt) - 1, self.parsed_prompt[-1]
|
||||
|
||||
def interpolation_at(self, step, total_steps=1):
|
||||
i, x = self.at_step_idx(step, total_steps)
|
||||
for y in self.parsed_prompt[i:]:
|
||||
step = min(y[0], 1.0)
|
||||
if x[1]["prompt"] != y[1]["prompt"]:
|
||||
return step, y
|
||||
return 1.0, self.parsed_prompt[-1]
|
||||
|
||||
def load_loras(self, lora_cache=None):
|
||||
from .utils import Timer, load_loras_from_schedule
|
||||
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 lora_cache is not None:
|
||||
self.loaded_loras = lora_cache
|
||||
with Timer("PromptSchedule.load_loras()"):
|
||||
self.loaded_loras = load_loras_from_schedule(self.parsed_prompt, self.loaded_loras)
|
||||
return self.loaded_loras
|
||||
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 parse_prompt_schedules(prompt):
|
||||
return PromptSchedule(prompt)
|
||||
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)
|
||||
|
||||
@@ -1,70 +0,0 @@
|
||||
import torch
|
||||
|
||||
|
||||
# Copied and adapted from https://github.com/bvhari/ComfyUI_PerpWeight/blob/main/clipperpweight.py
|
||||
def perp_encode(clip, tokens):
|
||||
empty_tokens = clip.tokenize("")
|
||||
sdxl_flag = "g" in tokens
|
||||
empty_cond, empty_cond_pooled = clip.encode_from_tokens(empty_tokens, return_pooled=True)
|
||||
unweighted_tokens = {}
|
||||
for k in ["l", "g"]:
|
||||
if k not in tokens:
|
||||
continue
|
||||
unweighted_tokens[k] = [[(t, 1.0) for t, _ in x] for x in tokens[k]]
|
||||
unweighted_cond, unweighted_pooled = clip.encode_from_tokens(unweighted_tokens, return_pooled=True)
|
||||
cond = torch.clone(unweighted_cond)
|
||||
|
||||
if sdxl_flag:
|
||||
for i in range(unweighted_cond.shape[0]):
|
||||
for j in range(unweighted_cond.shape[1]):
|
||||
weight_l = tokens["l"][(j // 77)][(j % 77)][1]
|
||||
if weight_l != 1.0:
|
||||
token_vector_l = unweighted_cond[i][j][:768]
|
||||
zero_vector_l = empty_cond[0][(j % 77)][:768]
|
||||
perp_l = (
|
||||
(torch.mul(zero_vector_l, token_vector_l).sum()) / (torch.norm(token_vector_l) ** 2)
|
||||
) * token_vector_l
|
||||
if weight_l > 1.0:
|
||||
cond[i][j][:768] = token_vector_l + (weight_l * perp_l)
|
||||
elif (weight_l > 0.0) and (weight_l < 1.0):
|
||||
cond[i][j][:768] = token_vector_l - ((1 - weight_l) * perp_l)
|
||||
elif weight_l < 0.0:
|
||||
cond[i][j][:768] = token_vector_l + (weight_l * perp_l)
|
||||
elif weight_l == 0.0:
|
||||
cond[i][j][:768] = empty_cond[0][(j % 77)][:768]
|
||||
|
||||
weight_g = tokens["g"][(j // 77)][(j % 77)][1]
|
||||
if weight_g != 1.0:
|
||||
token_vector_g = unweighted_cond[i][j][768:]
|
||||
zero_vector_g = empty_cond[0][(j % 77)][768:]
|
||||
perp_g = (
|
||||
(torch.mul(zero_vector_g, token_vector_g).sum()) / (torch.norm(token_vector_g) ** 2)
|
||||
) * token_vector_g
|
||||
if weight_g > 1.0:
|
||||
cond[i][j][768:] = token_vector_g + (weight_g * perp_g)
|
||||
elif (weight_g > 0.0) and (weight_g < 1.0):
|
||||
cond[i][j][768:] = token_vector_g - ((1 - weight_g) * perp_g)
|
||||
elif weight_g < 0.0:
|
||||
cond[i][j][768:] = token_vector_g + (weight_g * perp_g)
|
||||
elif weight_g == 0.0:
|
||||
cond[i][j][768:] = empty_cond[0][(j % 77)][768:]
|
||||
else:
|
||||
tokens = tokens["l"]
|
||||
for i in range(unweighted_cond.shape[0]):
|
||||
for j in range(unweighted_cond.shape[1]):
|
||||
weight = tokens[(j // 77)][(j % 77)][1]
|
||||
if weight != 1.0:
|
||||
token_vector = unweighted_cond[i][j]
|
||||
zero_vector = empty_cond[0][(j % 77)]
|
||||
perp = (
|
||||
(torch.mul(zero_vector, token_vector).sum()) / (torch.norm(token_vector) ** 2)
|
||||
) * token_vector
|
||||
if weight > 1.0:
|
||||
cond[i][j] = token_vector + (weight * perp)
|
||||
elif (weight > 0.0) and (weight < 1.0):
|
||||
cond[i][j] = token_vector - ((1 - weight) * perp)
|
||||
elif weight < 0.0:
|
||||
cond[i][j] = token_vector + (weight * perp)
|
||||
elif weight == 0.0:
|
||||
cond[i][j] = empty_cond[0][(j % 77)]
|
||||
return cond, unweighted_pooled
|
||||
@@ -0,0 +1,604 @@
|
||||
import logging
|
||||
import re
|
||||
import torch
|
||||
from functools import partial
|
||||
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
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
AVAILABLE_STYLES = ["comfy", "perp", "A1111", "compel", "comfy++", "down_weight"]
|
||||
AVAILABLE_NORMALIZATIONS = ["none", "mean", "length", "length+mean"]
|
||||
|
||||
SHUFFLE_GEN = torch.Generator(device="cpu")
|
||||
|
||||
|
||||
def get_sdxl(text, defaults):
|
||||
# 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]
|
||||
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+")
|
||||
cropw, croph = parse_floats(args[2], [d.get("sdxl_cwidth", 0), d.get("sdxl_cheight", 0)], split_re="\\s+")
|
||||
|
||||
opts = {
|
||||
"width": int(w),
|
||||
"height": int(h),
|
||||
"target_width": int(tw),
|
||||
"target_height": int(th),
|
||||
"crop_w": int(cropw),
|
||||
"crop_h": int(croph),
|
||||
}
|
||||
return text, opts
|
||||
|
||||
|
||||
def get_clipweights(text, existing_spec=None):
|
||||
text, spec = get_function(text, "TE_WEIGHT", defaults=None)
|
||||
if not spec:
|
||||
return existing_spec or {}, text
|
||||
args = spec[0].strip()
|
||||
res = {}
|
||||
for arg in args.split(","):
|
||||
try:
|
||||
te, val = arg.strip().split("=")
|
||||
te, val = te.strip(), float(val.strip())
|
||||
res[te] = val
|
||||
except ValueError:
|
||||
log.warning("Invalid TE weight spec '%s', ignoring...", arg.strip())
|
||||
return res, text
|
||||
|
||||
|
||||
def get_style(text, default_style="comfy", default_normalization="none"):
|
||||
text, styles = get_function(text, "STYLE", [default_style, default_normalization])
|
||||
if not styles:
|
||||
return default_style, default_normalization, text
|
||||
style, normalization = styles[0]
|
||||
style = style.strip()
|
||||
normalization = normalization.strip()
|
||||
if style.replace("old+", "") not in AVAILABLE_STYLES:
|
||||
log.warning("Unrecognized prompt style: %s. Using %s", style, default_style)
|
||||
|
||||
for part in normalization.split("+"):
|
||||
if part not in AVAILABLE_NORMALIZATIONS:
|
||||
log.warning("Unrecognized prompt normalization: %s. Using %s", normalization, default_normalization)
|
||||
normalization = default_normalization
|
||||
break
|
||||
|
||||
return style, normalization, text
|
||||
|
||||
|
||||
def shuffle_chunk(shuffle, c):
|
||||
func, shuffle = shuffle
|
||||
shuffle_count = int(safe_float(shuffle[0], 0))
|
||||
_, separator, joiner = shuffle
|
||||
if separator == "default":
|
||||
separator = ","
|
||||
|
||||
if not separator:
|
||||
separator = ","
|
||||
|
||||
joiner = {
|
||||
"default": ",",
|
||||
"separator": separator,
|
||||
}.get(joiner, joiner)
|
||||
|
||||
log.debug("%s arg=%s sep=%s join=%s", func, shuffle_count, separator, joiner)
|
||||
separated = smarter_split(separator, c)
|
||||
log.debug("Prompt split into %s", separated)
|
||||
if func == "SHIFT":
|
||||
shuffle_count = shuffle_count % len(separated)
|
||||
permutation = separated[shuffle_count:] + separated[:shuffle_count]
|
||||
elif func == "SHUFFLE":
|
||||
SHUFFLE_GEN.manual_seed(shuffle_count)
|
||||
permutation = [separated[i] for i in torch.randperm(len(separated), generator=SHUFFLE_GEN)]
|
||||
else:
|
||||
# ??? should never get here
|
||||
permutation = separated
|
||||
|
||||
permutation = [p for p in permutation if p.strip()]
|
||||
if permutation != separated:
|
||||
c = joiner.join(permutation)
|
||||
return 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"""
|
||||
for key in tokens:
|
||||
max_idx = 0
|
||||
for group in range(len(tokens[key])):
|
||||
for i, token in enumerate(tokens[key][group]):
|
||||
if len(token) < 3:
|
||||
# No need to fix ids when they don't exist
|
||||
return tokens
|
||||
# Ignore zeros, they represent the padding token
|
||||
if token[2] != 0 and token[2] < max_idx:
|
||||
tokens[key][group][i] = (token[0], token[1], token[2] + max_idx)
|
||||
max_idx = max(max_idx, max(x for _, _, x in tokens[key][group]))
|
||||
return tokens
|
||||
|
||||
|
||||
def tokenize_chunks(clip, text, need_word_ids, can_break):
|
||||
chunks = re.split(r"\bBREAK\b", text)
|
||||
token_chunks = []
|
||||
shuffled_chunks = []
|
||||
for c in chunks:
|
||||
c, shuffles = get_function(c.strip(), "(SHIFT|SHUFFLE)", ["0", "default", "default"], return_func_name=True)
|
||||
r = c
|
||||
for s in shuffles:
|
||||
r = shuffle_chunk(s, r)
|
||||
if r != c:
|
||||
log.info("Shuffled prompt chunk to %s", r)
|
||||
shuffled_chunks.append(r)
|
||||
t = clip.tokenize(c, return_word_ids=need_word_ids)
|
||||
token_chunks.append(t)
|
||||
|
||||
tokens = token_chunks[0]
|
||||
full_prompt = "".join(shuffled_chunks)
|
||||
full_tokenized = tokens
|
||||
if len(chunks) > 1:
|
||||
full_tokenized = clip.tokenize(full_prompt, return_word_ids=need_word_ids)
|
||||
for key in tokens:
|
||||
if not can_break.get(key):
|
||||
log.warning("BREAK does not make sense for %s, tokenizing as one chunk. Use CAT instead.", key)
|
||||
tokens[key] = full_tokenized[key]
|
||||
continue
|
||||
for c in token_chunks[1:]:
|
||||
tokens[key].extend(c[key])
|
||||
|
||||
return tokens
|
||||
|
||||
|
||||
def tokenize(clip, text, can_break, empty_tokens):
|
||||
# defaults=None means there is no argument parsing at all
|
||||
text, l_prompts = get_function(text, "CLIP_L", defaults=None)
|
||||
text, te_prompts = get_function(text, "TE", defaults=None)
|
||||
need_word_ids = True
|
||||
tokens = tokenize_chunks(clip, text, need_word_ids, can_break)
|
||||
|
||||
per_te_prompts = {}
|
||||
if l_prompts:
|
||||
log.warning("Note: CLIP_L is deprecated. Use TE(l=prompt) instead")
|
||||
per_te_prompts["l"] = l_prompts
|
||||
|
||||
for prompt in te_prompts:
|
||||
if prompt.strip() == "help":
|
||||
log.info("Encoders available for TE: %s", ", ".join(tokens.keys()))
|
||||
continue
|
||||
params = prompt.split("=", 1)
|
||||
if len(params) != 2:
|
||||
log.warning("Invalid TE call, ignoring: %s", prompt)
|
||||
continue
|
||||
te = params[0].strip()
|
||||
prompt = params[1].strip()
|
||||
if te not in tokens:
|
||||
log.warning("Invalid TE call, no TE with key '%s', ignoring: %s", te)
|
||||
log.info("Encoders available for TE: %s", ", ".join(tokens.keys()))
|
||||
continue
|
||||
l = per_te_prompts.get(te, [])
|
||||
l.append(prompt)
|
||||
per_te_prompts[te] = l
|
||||
|
||||
if per_te_prompts:
|
||||
for key in per_te_prompts:
|
||||
prompt = " ".join(per_te_prompts[key])
|
||||
tokens[key] = tokenize_chunks(clip, prompt, need_word_ids, can_break)[key]
|
||||
log.info("Encoded prompt with TE '%s': %s", key, prompt)
|
||||
|
||||
maxlen = max([0] + [len(tokens[k]) for k in tokens if can_break[k]])
|
||||
for k in tokens:
|
||||
if not can_break[k]:
|
||||
continue
|
||||
while len(tokens[k]) < maxlen:
|
||||
tokens[k] += empty_tokens[k]
|
||||
|
||||
return fix_word_ids(tokens)
|
||||
|
||||
|
||||
def encode_prompt_segment(
|
||||
clip,
|
||||
text,
|
||||
settings,
|
||||
default_style="comfy",
|
||||
default_normalization="none",
|
||||
clip_weights=None,
|
||||
) -> list[tuple[torch.Tensor, dict[str]]]:
|
||||
style, normalization, text = get_style(text, default_style, default_normalization)
|
||||
clip_weights, text = get_clipweights(text, clip_weights)
|
||||
text, cuts = parse_cuts(text)
|
||||
extra = {}
|
||||
if clip_weights:
|
||||
extra["clip_weights"] = clip_weights
|
||||
if cuts:
|
||||
extra["cuts"] = cuts
|
||||
|
||||
empty = clip.tokenize("", return_word_ids=True)
|
||||
can_break = {}
|
||||
for k in empty:
|
||||
tokenizer = getattr(clip.tokenizer, f"clip_{k}", getattr(clip.tokenizer, k, None))
|
||||
can_break[k] = tokenizer and tokenizer.pad_to_max_length
|
||||
|
||||
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
|
||||
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))
|
||||
|
||||
conds_to_avg = []
|
||||
for prompt, weight in prompts_to_avg:
|
||||
conds_to_cat = []
|
||||
chunks = re.split(r"\bCAT\b", prompt)
|
||||
for c in chunks:
|
||||
tokens = tokenize(clip, c, can_break, empty)
|
||||
conds_to_cat.append(clip.encode_from_tokens_scheduled(tokens, add_dict=settings))
|
||||
|
||||
base = conds_to_cat[0]
|
||||
for cond in conds_to_cat[1:]:
|
||||
assert len(cond) == len(base), "Conditioning length mismatch"
|
||||
# Pooled gets ignored
|
||||
for i in range(len(base)):
|
||||
c1 = base[i][0]
|
||||
c2 = cond[i][0]
|
||||
base[i][0] = torch.cat((c1, c2), 1)
|
||||
conds_to_avg.append((base, weight))
|
||||
|
||||
base, w = conds_to_avg[0]
|
||||
for cond, next_w in conds_to_avg[1:]:
|
||||
assert len(base) == len(cond), "Conditioning length mismatch"
|
||||
if w == 1.0:
|
||||
w = next_w
|
||||
continue
|
||||
for i in range(len(base)):
|
||||
(cond,) = ConditioningAverage.addWeighted(None, [base[i]], [cond[i]], w)
|
||||
base[i] = cond[0]
|
||||
w = next_w
|
||||
|
||||
return base
|
||||
|
||||
|
||||
def apply_weights(output, te_name, spec):
|
||||
"""Applies weights to TE outputs"""
|
||||
if not spec:
|
||||
return output
|
||||
|
||||
if te_name.startswith("clip_"):
|
||||
te_name = te_name[5:]
|
||||
|
||||
default = spec.get("all", None)
|
||||
|
||||
if isinstance(output, tuple):
|
||||
out, pooled = output
|
||||
pkey = te_name + "_pooled"
|
||||
if te_name in spec or pkey in spec or default is not None:
|
||||
w = spec.get(te_name, default)
|
||||
pooled_w = spec.get(pkey, w)
|
||||
if w is None:
|
||||
w = 1.0
|
||||
if pooled_w is None:
|
||||
pooled_w = 1.0
|
||||
log.info("Weighting %s output by %s, pooled by %s", te_name, w, pooled_w)
|
||||
out = out * w
|
||||
if pooled is not None:
|
||||
pooled = pooled * pooled_w
|
||||
|
||||
return out, pooled
|
||||
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
|
||||
return output
|
||||
|
||||
|
||||
def make_patch(te_name, orig_fn, normalization, style, extra):
|
||||
def encode(t):
|
||||
r = advanced_encode_from_tokens(
|
||||
t, normalization, style, orig_fn, return_pooled=True, apply_to_pooled=False, **extra
|
||||
)
|
||||
return apply_weights(r, te_name, extra.get("clip_weights"))
|
||||
|
||||
if "cuts" in extra:
|
||||
return partial(process_cuts, encode, extra)
|
||||
return encode
|
||||
|
||||
|
||||
def hook_te(clip, te_names, style, normalization, extra):
|
||||
if style == "comfy" and normalization == "none" and not extra:
|
||||
return clip
|
||||
newclip = clip.clone()
|
||||
for te_name in te_names:
|
||||
tokenizer = getattr(clip.tokenizer, f"clip_{te_name}", getattr(clip.tokenizer, te_name, None))
|
||||
if tokenizer:
|
||||
x = extra.copy()
|
||||
x["tokenizer"] = tokenizer
|
||||
if not hasattr(clip.patcher.model, te_name):
|
||||
te_name = "clip_" + te_name
|
||||
if not hasattr(clip.patcher.model, te_name):
|
||||
log.warning("TE model %s not found on model patcher. Skipping...", te_name)
|
||||
continue
|
||||
|
||||
log.debug("Hooked into te=%s with style=%s, normalization=%s", te_name, style, normalization)
|
||||
encode = clip.patcher.get_model_object(f"{te_name}.encode_token_weights")
|
||||
x["has_negpip"] = clip.patcher.model_options.get("ppm_negpip", False)
|
||||
newclip.patcher.add_object_patch(
|
||||
f"{te_name}.encode_token_weights",
|
||||
make_patch(
|
||||
te_name,
|
||||
encode,
|
||||
normalization,
|
||||
style,
|
||||
x,
|
||||
),
|
||||
)
|
||||
# 'g' and 'l' exist in these are clip_g and clip_l
|
||||
else:
|
||||
log.warning("Tokens contain items with key %s but no tokenizer found on object with that name.", te_name)
|
||||
return newclip
|
||||
|
||||
|
||||
def get_area(text):
|
||||
text, areas = get_function(text, "AREA", ["0 1", "0 1", "1"])
|
||||
if not areas:
|
||||
return text, None
|
||||
|
||||
args = areas[0]
|
||||
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)
|
||||
|
||||
def is_pct(f):
|
||||
return f >= 0.0 and f <= 1.0
|
||||
|
||||
def is_pixel(f):
|
||||
return f == 0 or f > 1
|
||||
|
||||
if all(is_pct(v) for v in [h, w, y, x]):
|
||||
area = ("percentage", h, w, y, x)
|
||||
elif all(is_pixel(v) for v in [h, w, y, x]):
|
||||
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"
|
||||
)
|
||||
|
||||
return text, (area, weight)
|
||||
|
||||
|
||||
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]
|
||||
return text, (int(w), int(h))
|
||||
|
||||
|
||||
def make_mask(args, size, weight):
|
||||
x1, x2 = parse_floats(args[0], [0.0, 1.0], split_re="\\s+")
|
||||
y1, y2 = parse_floats(args[1], [0.0, 1.0], split_re="\\s+")
|
||||
|
||||
def is_pct(f):
|
||||
return f >= 0.0 and f <= 1.0
|
||||
|
||||
def is_pixel(f):
|
||||
return f == 0 or f > 1
|
||||
|
||||
if all(is_pct(v) for v in [x1, x2, y1, y2]):
|
||||
w, h = size
|
||||
xs = int(w * x1), int(w * x2)
|
||||
ys = int(h * y1), int(h * y2)
|
||||
elif all(is_pixel(v) for v in [x1, x2, y1, y2]):
|
||||
w, h = size
|
||||
xs = int(x1), int(x2)
|
||||
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"
|
||||
)
|
||||
|
||||
mask = torch.full((h, w), 0, dtype=torch.float32, device="cpu")
|
||||
mask[ys[0] : ys[1], xs[0] : xs[1]] = weight
|
||||
mask = mask.unsqueeze(0)
|
||||
log.debug("Mask xs=%s, ys=%s, shape=%s, weight=%s", xs, ys, mask.shape, weight)
|
||||
return mask
|
||||
|
||||
|
||||
def get_mask(text, size, input_masks):
|
||||
"""Parse MASK(x1 x2, y1 y2, weight), IMASK(i, weight) and FEATHER(left top right bottom)"""
|
||||
# TODO: combine multiple masks
|
||||
text, masks = get_function(text, "MASK", ["0 1", "0 1", "1", "multiply"])
|
||||
text, imasks = get_function(text, "IMASK", ["0", "1", "multiply"])
|
||||
text, feathers = get_function(text, "FEATHER", ["0 0 0 0"])
|
||||
text, maskw = get_function(text, "MASKW", ["1.0"])
|
||||
if not masks and not imasks:
|
||||
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)
|
||||
return mask
|
||||
|
||||
mask = None
|
||||
totalweight = 1.0
|
||||
if maskw:
|
||||
totalweight = safe_float(maskw[0][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)
|
||||
if i < len(feathers):
|
||||
nextmask = feather(feathers[i], nextmask)
|
||||
i += 1
|
||||
if mask is not None:
|
||||
log.info("MaskComposite op=%s", op)
|
||||
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
|
||||
else:
|
||||
mask = nextmask
|
||||
|
||||
for idx, w, op in imasks:
|
||||
idx = int(safe_float(idx, 0.0))
|
||||
w = safe_float(w, 1.0)
|
||||
if len(input_masks) < idx + 1:
|
||||
log.warn("IMASK index %s not found, ignoring...", idx)
|
||||
continue
|
||||
nextmask = input_masks[idx] * w
|
||||
if i < len(feathers):
|
||||
nextmask = feather(feathers[i], nextmask)
|
||||
i += 1
|
||||
if mask is not None:
|
||||
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
|
||||
else:
|
||||
mask = nextmask
|
||||
|
||||
# apply leftover FEATHER() specs to the whole
|
||||
for f in feathers[i:]:
|
||||
mask = feather(f, mask)
|
||||
|
||||
return text, mask, totalweight
|
||||
|
||||
|
||||
def get_noise(text):
|
||||
text, noises = get_function(
|
||||
text,
|
||||
"NOISE",
|
||||
["0.0", "none"],
|
||||
)
|
||||
if not noises:
|
||||
return text, None, None
|
||||
w = 0
|
||||
# Only take seed from first noise spec, for simplicity
|
||||
seed = safe_float(noises[0][1], "none")
|
||||
if seed == "none":
|
||||
gen = None
|
||||
else:
|
||||
gen = torch.Generator()
|
||||
gen.manual_seed(int(seed))
|
||||
for n in noises:
|
||||
w += safe_float(n[0], 0.0)
|
||||
return text, max(min(w, 1.0), 0.0), gen
|
||||
|
||||
|
||||
def apply_noise(cond, weight, gen):
|
||||
if cond is None or not weight:
|
||||
return cond
|
||||
|
||||
n = torch.randn(cond.size(), generator=gen).to(cond)
|
||||
|
||||
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 weight(t):
|
||||
opts = {}
|
||||
m = re.search(r":(-?\d\.?\d*)(![A-Za-z]+)?$", t)
|
||||
if not m:
|
||||
return (1.0, opts, t)
|
||||
w = float(m[1])
|
||||
tag = m[2]
|
||||
t = t[: m.span()[0]]
|
||||
if tag == "!noscale":
|
||||
opts["scale"] = 1
|
||||
|
||||
return w, opts, t
|
||||
|
||||
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]:
|
||||
_, mask, _ = get_mask("MASK()", mask_size, masks)
|
||||
c[1]["mask"] = mask
|
||||
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")
|
||||
conds.extend(base_cond)
|
||||
|
||||
return conds
|
||||
@@ -0,0 +1,106 @@
|
||||
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()
|
||||
@@ -0,0 +1,56 @@
|
||||
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()
|
||||
@@ -0,0 +1,215 @@
|
||||
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()
|
||||
@@ -0,0 +1,207 @@
|
||||
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()
|
||||
+82
-253
@@ -1,60 +1,77 @@
|
||||
from collections import namedtuple
|
||||
from os import environ
|
||||
from pathlib import Path
|
||||
import re
|
||||
from math import lcm
|
||||
import time
|
||||
import logging
|
||||
import torch
|
||||
|
||||
# Allow testing
|
||||
try:
|
||||
from folder_paths import get_filename_list
|
||||
except ImportError:
|
||||
|
||||
import nodes
|
||||
import folder_paths
|
||||
def get_filename_list(x):
|
||||
raise NotImplementedError("How did you get here?")
|
||||
|
||||
import comfy.model_management
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
FORCE_CPU_OFFLOAD = bool(environ.get("COMFYUI_PC_CPU_OFFLOAD"))
|
||||
|
||||
def consolidate_schedule(prompt_schedule):
|
||||
prev_loras = {}
|
||||
not_found = []
|
||||
consolidated = []
|
||||
for end_pct, c in reversed(list(prompt_schedule)):
|
||||
loras = {}
|
||||
for k, v in c["loras"].items():
|
||||
if k in not_found:
|
||||
continue
|
||||
path = lora_name_to_file(k)
|
||||
if path is None:
|
||||
not_found.append(k)
|
||||
continue
|
||||
loras[path] = v
|
||||
|
||||
if loras != prev_loras:
|
||||
consolidated.append((end_pct, loras))
|
||||
prev_loras = loras
|
||||
for k in not_found:
|
||||
log.warning("LoRA '%s' not found, ignoring...", k)
|
||||
return list(reversed(consolidated))
|
||||
|
||||
|
||||
# Minimal Modelpatcher that doesn't do anything, for LoRA loading when not
|
||||
# interested in either CLIP or unet
|
||||
class DummyModelPatcher:
|
||||
class DummyTorchModel:
|
||||
def __init__(self):
|
||||
dummyconf = {
|
||||
"num_res_blocks": [],
|
||||
"channel_mult": [],
|
||||
"transformer_depth": [],
|
||||
"transformer_depth_output": [],
|
||||
"transformer_depth_middle": 0,
|
||||
}
|
||||
self.model_config = namedtuple("DummyConfig", ["unet_config"])(dummyconf)
|
||||
|
||||
def state_dict(self):
|
||||
return {}
|
||||
|
||||
def __init__(self):
|
||||
self.model = self.DummyTorchModel()
|
||||
self.cond_stage_model = self.DummyTorchModel()
|
||||
self.weight_inplace_update = True
|
||||
self.model_options = {}
|
||||
|
||||
def add_patches(self, patches, *args, **kwargs):
|
||||
return []
|
||||
|
||||
def patch_model(self):
|
||||
pass
|
||||
|
||||
def unpatch_model(self):
|
||||
pass
|
||||
|
||||
def clone(self):
|
||||
return self
|
||||
def find_nonscheduled_loras(consolidated_schedule):
|
||||
consolidated_schedule = list(consolidated_schedule)
|
||||
if not consolidated_schedule:
|
||||
return {}
|
||||
last_end, candidate_loras = consolidated_schedule[0]
|
||||
to_remove = set()
|
||||
for candidate, weights in candidate_loras.items():
|
||||
for end, loras in consolidated_schedule[1:]:
|
||||
last_end = end
|
||||
if loras.get(candidate) != weights:
|
||||
to_remove.add(candidate)
|
||||
# No candidates if the schedule does not span full time
|
||||
if last_end < 1.0:
|
||||
return {}
|
||||
return {k: v for (k, v) in candidate_loras.items() if k not in to_remove}
|
||||
|
||||
|
||||
DUMMY_MODEL = DummyModelPatcher()
|
||||
def smarter_split(separator, string):
|
||||
"""Does not break () when splitting"""
|
||||
splits = []
|
||||
prev = 0
|
||||
stack = 0
|
||||
escape = False
|
||||
for idx, x in enumerate(string):
|
||||
if x == "(" and not escape:
|
||||
stack += 1
|
||||
elif x == ")" and not escape:
|
||||
stack = max(0, stack - 1)
|
||||
elif x == separator and stack == 0:
|
||||
splits.append(string[prev:idx])
|
||||
prev = idx + 1
|
||||
escape = x == "\\"
|
||||
|
||||
splits.append(string[prev : idx + 1])
|
||||
return splits
|
||||
|
||||
|
||||
def find_closing_paren(text, start):
|
||||
@@ -70,23 +87,40 @@ def find_closing_paren(text, start):
|
||||
return len(text)
|
||||
|
||||
|
||||
def get_function(text, func, defaults, return_func_name=False):
|
||||
def get_function(text, func, defaults, return_func_name=False, placeholder="", return_dict=False):
|
||||
rex = re.compile(rf"\b{func}\(", re.MULTILINE)
|
||||
instances = []
|
||||
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)
|
||||
if return_func_name:
|
||||
ph = None
|
||||
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)
|
||||
|
||||
text = text[:start] + text[end + 1 :]
|
||||
if placeholder:
|
||||
text = text[:start] + f"\0{placeholder}{count}\0" + text[end + 1 :]
|
||||
else:
|
||||
text = text[:start] + text[end + 1 :]
|
||||
match = rex.search(text)
|
||||
count += 1
|
||||
return text, instances
|
||||
|
||||
|
||||
@@ -118,15 +152,6 @@ def parse_strings(string, defaults, split_re=r"(?<!\\),", replace=(r"\,", ",")):
|
||||
return parse_args(splits, spec, strip=False)
|
||||
|
||||
|
||||
def equalize(*tensors):
|
||||
if all(t.shape[1] == tensors[0].shape[1] for t in tensors):
|
||||
return tensors
|
||||
|
||||
x = lcm(*(t.shape[1] for t in tensors))
|
||||
|
||||
return (t.repeat(1, x // t.shape[1], 1) for t in tensors)
|
||||
|
||||
|
||||
def safe_float(f, default):
|
||||
if f is None:
|
||||
return default
|
||||
@@ -136,86 +161,8 @@ def safe_float(f, default):
|
||||
return default
|
||||
|
||||
|
||||
def unpatch_model(model):
|
||||
if model:
|
||||
log.info("Unpatching model")
|
||||
model.unpatch_model()
|
||||
|
||||
|
||||
def clone_model(model):
|
||||
if not model:
|
||||
return None
|
||||
model = model.clone()
|
||||
if not environ.get("PC_NO_INPLACE_UPDATE"):
|
||||
model.weight_inplace_update = True
|
||||
return model
|
||||
|
||||
|
||||
def add_patches(model, patches, weight):
|
||||
model.add_patches(patches, weight)
|
||||
|
||||
|
||||
def patch_model(model, forget=False, orig=None):
|
||||
global FORCE_CPU_OFFLOAD
|
||||
try:
|
||||
return _patch_model(model, forget, orig, FORCE_CPU_OFFLOAD)
|
||||
except comfy.model_management.OOM_EXCEPTION:
|
||||
FORCE_CPU_OFFLOAD = True
|
||||
log.error("Ran out of memory while applying LoRAs, Forcing CPU offload from now on")
|
||||
# Unpatch to restore partially applied weights
|
||||
unpatch_model(model)
|
||||
raise
|
||||
|
||||
|
||||
def _patch_model(model, forget=False, orig=None, offload_to_cpu=False):
|
||||
if not model:
|
||||
return None
|
||||
if offload_to_cpu:
|
||||
saved_offload = model.offload_device
|
||||
model.offload_device = torch.device("cpu")
|
||||
log.info("Patching model, cpu_offload=%s", model.offload_device == torch.device("cpu"))
|
||||
if orig:
|
||||
model.backup = orig.backup
|
||||
model.patch_model()
|
||||
if offload_to_cpu:
|
||||
model.offload_device = saved_offload
|
||||
if forget:
|
||||
model.patches = {}
|
||||
model.object_patches = {}
|
||||
return model
|
||||
|
||||
|
||||
def get_callback(model):
|
||||
return model.model_options.get("prompt_control_callback")
|
||||
|
||||
|
||||
def set_callback(model, cb):
|
||||
model.model_options["prompt_control_callback"] = cb
|
||||
|
||||
|
||||
# Hack to temporarily override printing to stdout to stop log spam
|
||||
def suppress_print(f):
|
||||
def noop(*args):
|
||||
pass
|
||||
|
||||
p = print
|
||||
__builtins__["print"] = noop
|
||||
rootlogger = logging.getLogger()
|
||||
oldlevel = rootlogger.level
|
||||
try:
|
||||
rootlogger.setLevel(logging.ERROR)
|
||||
x = f()
|
||||
except BaseException:
|
||||
__builtins__["print"] = p
|
||||
rootlogger.setLevel(oldlevel)
|
||||
raise
|
||||
__builtins__["print"] = p
|
||||
rootlogger.setLevel(oldlevel)
|
||||
return x
|
||||
|
||||
|
||||
def lora_name_to_file(name):
|
||||
filenames = folder_paths.get_filename_list("loras")
|
||||
filenames = get_filename_list("loras")
|
||||
# Return exact matches as is
|
||||
if name in filenames:
|
||||
return name
|
||||
@@ -226,121 +173,3 @@ def lora_name_to_file(name):
|
||||
if p.name == n or str(p) == n:
|
||||
return f
|
||||
return None
|
||||
|
||||
|
||||
def load_lbw():
|
||||
return nodes.NODE_CLASS_MAPPINGS.get("LoraLoaderBlockWeight //Inspire")
|
||||
|
||||
|
||||
def make_loader(filename, lbw):
|
||||
if not lbw:
|
||||
l = nodes.LoraLoader()
|
||||
|
||||
def loader(model, clip, model_weight, clip_weight, lbw):
|
||||
return suppress_print(lambda: l.load_lora(model, clip, filename, model_weight, clip_weight))
|
||||
|
||||
else:
|
||||
# This is already checked before calling make_loader
|
||||
l = load_lbw()()
|
||||
|
||||
def loader(model, clip, model_weight, clip_weight, lbw):
|
||||
spec = lbw["LBW"]
|
||||
lbw_a = safe_float(lbw.get("A"), 4.0)
|
||||
lbw_b = safe_float(lbw.get("B"), 1.0)
|
||||
m = model or DUMMY_MODEL
|
||||
c = clip or DUMMY_MODEL
|
||||
m, c, _ = suppress_print(
|
||||
lambda: l.doit(m, c, filename, model_weight, clip_weight, False, 0, lbw_a, lbw_b, "", spec)
|
||||
)
|
||||
if m is DUMMY_MODEL:
|
||||
m = None
|
||||
if c is DUMMY_MODEL:
|
||||
c = None
|
||||
return m, c
|
||||
|
||||
return loader
|
||||
|
||||
|
||||
def apply_loras_from_spec(
|
||||
loraspec, model=None, clip=None, orig_model=None, orig_clip=None, patch=False, cache=None, applied_loras=None
|
||||
):
|
||||
if applied_loras is None:
|
||||
applied_loras = {}
|
||||
actual_loraspec = {}
|
||||
additive = True
|
||||
for key in loraspec:
|
||||
if key in applied_loras and applied_loras[key] == loraspec[key]:
|
||||
continue
|
||||
if key in applied_loras and applied_loras[key] != loraspec[key]:
|
||||
additive = False
|
||||
actual_loraspec[key] = loraspec[key]
|
||||
|
||||
for key in applied_loras:
|
||||
if key not in loraspec:
|
||||
actual_loraspec = loraspec
|
||||
additive = False
|
||||
|
||||
backup_model = model
|
||||
if not additive:
|
||||
unpatch_model(model)
|
||||
# Reset clip to unpatched
|
||||
if clip:
|
||||
clip = orig_clip or clip
|
||||
|
||||
if cache is None:
|
||||
cache = {}
|
||||
if not loraspec:
|
||||
return model, clip
|
||||
|
||||
for name, params in actual_loraspec.items():
|
||||
m, c = model, clip
|
||||
w, w_clip = params["weight"], params["weight_clip"]
|
||||
if w == 0:
|
||||
m = None
|
||||
if w_clip == 0:
|
||||
c = None
|
||||
if not w and not c:
|
||||
continue
|
||||
|
||||
lbw = params.get("lbw")
|
||||
if lbw and not load_lbw():
|
||||
log.warning("LoraBlockWeight not available, ignoring LBW parameters")
|
||||
lbw = None
|
||||
|
||||
# Cache the loader instance so that it doesn't reload the LoRA from disk all the time
|
||||
cache_key = name, bool(lbw)
|
||||
loader = cache.get(cache_key)
|
||||
if not loader:
|
||||
f = lora_name_to_file(name)
|
||||
if not f:
|
||||
log.warning("Lora %s not found", name)
|
||||
continue
|
||||
log.info("Loading LoRA: %s", f)
|
||||
loader = make_loader(f, bool(lbw))
|
||||
cache[cache_key] = loader
|
||||
|
||||
m, c = loader(m, c, w, w_clip, lbw)
|
||||
model = m or model
|
||||
clip = c or clip
|
||||
if model:
|
||||
log.info("Applying LoRA: %s:%s, LBW=%s, additive=%s", name, params["weight"], bool(lbw), additive)
|
||||
if clip:
|
||||
log.info("Applying CLIP LoRA: %s:%s, LBW=%s, additive=%s", name, params["weight_clip"], bool(lbw), additive)
|
||||
|
||||
# forget patches so we don't double-patch
|
||||
model = patch_model(model, forget=True, orig=backup_model)
|
||||
return model, clip
|
||||
|
||||
|
||||
class Timer:
|
||||
def __init__(self, name):
|
||||
self.name = name
|
||||
self.start = None
|
||||
|
||||
def __enter__(self):
|
||||
self.start = time.time()
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
elapsed = time.time() - self.start
|
||||
if environ.get("PC_SHOW_TIMINGS"):
|
||||
log.info("Executed %s in %s seconds", self.name, elapsed)
|
||||
|
||||
+1
-2
@@ -1,14 +1,13 @@
|
||||
[project]
|
||||
name = "comfyui-prompt-control"
|
||||
description = "Nodes for convenient prompt editing, making many common operations prompt-controllable"
|
||||
version = "1.1.1"
|
||||
version = "2.0.0-rc.7"
|
||||
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"]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/asagi4/comfyui-prompt-control"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "asagi4"
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user