Compare commits

...
37 Commits
Author SHA1 Message Date
asagi4 34056cac19 v2.0.0-rc.5 2025-06-06 19:52:02 +03:00
asagi4 8ae436abf1 Merge pull request #114 from pamparamm/negpip_option
Use ppm_negpip option to detect NegPiP
2025-06-06 16:49:20 +03:00
Pam c2ce2ce023 Use ppm_negpip option to detect NegPiP 2025-06-06 16:24:15 +05:00
asagi4 7a9e69ec31 Fix errors found in testing
Who knew tests could be useful, too?
2025-06-05 23:49:02 +03:00
asagi4 aaff8dc7da Encoding tests 2025-06-05 23:49:02 +03:00
asagi4 3e4722278a v2.0.0-rc.4 2025-06-05 21:39:47 +03:00
asagi4 11cb430396 Support NegPiP without requiring a monkeypatch. 2025-06-05 21:38:31 +03:00
asagi4 9f1cbfd11c Add a test suite for text encoding
Can only be run manually and imports ComfyUI main to configure search paths, but
at least it works...
2025-06-05 21:05:39 +03:00
asagi4 5b3a914f1d Try not to fail in cases where masks have irregular shapes 2025-06-04 22:34:19 +03:00
asagi4 a5ffa1acd7 Documentation 2025-06-04 21:23:52 +03:00
asagi4 7f8783147b I keep forgetting expand only works on batch size 1, see #108 2025-06-04 21:12:52 +03:00
asagi4 3c9b806e5f Remove useless import 2025-06-04 21:04:15 +03:00
asagi4 a485c2655a Fix case where negative prompt size changes lcm of cond size
Also don't mutate existing hooks on input negative prompts if they exist.
2025-06-04 21:02:48 +03:00
asagi4 f888e69b00 Merge pull request #112 from pamparamm/attn_couple_batch
Add PCAttentionCoupleBatchNegative
2025-06-04 20:59:52 +03:00
asagi4 b4b0858214 Fix multi-TE failure
See #113
2025-06-04 08:48:40 +03:00
Pam 7a1cb2cf51 Fix latent masking 2025-06-04 00:30:23 +05:00
Pam fbb6b5c8fa Fix some uncond edgecases 2025-06-03 07:50:45 +05:00
Pam 6115c095cb Fix attn couple batching with multiple positive schedules 2025-06-03 07:42:43 +05:00
Pam 04d3d2e959 Missing space 2025-06-03 03:46:56 +05:00
Pam e4f64837ef Add PCAttentionCoupleBatchNegative;
Revert some optimizations in AttentionCoupleHook
2025-06-03 03:17:47 +05:00
asagi4 4a785b294b Avoid hardcoding length in adv encode
Makes these not explode on T5 at least. They seems to produce the same results still.
2025-06-03 00:47:57 +03:00
asagi4 b3195a6297 Stop using batched_clip_encode, it doesn't do anything? 2025-06-02 23:32:19 +03:00
asagi4 5c52bffc9d Docs 2025-06-02 22:16:42 +03:00
asagi4 ac2d275dfb Clarify TE lookups 2025-06-02 21:37:27 +03:00
asagi4 399f992a26 Fix BREAK while maintaining old behaviour, add CAT instead 2025-06-02 20:36:27 +03:00
asagi4 672a2a09bb BREAK is now ConditioningConcat 2025-06-02 19:36:29 +03:00
asagi4 e3629961ce Add AVG(weight) to work as ConditioningAverage 2025-06-02 19:14:20 +03:00
asagi4 ddac624ad1 Add NBREAK (Warning: unstable. Name is likely to change)
NBREAK should have the same behaviour as ComfyUI's ConditioningConcat

See #111
2025-06-02 17:39:09 +03:00
asagi4 99ddfe357e DEF docs 2025-06-01 22:44:06 +03:00
asagi4 110d5248a0 DEF tests 2025-06-01 22:34:39 +03:00
asagi4 61b1ecc88e Clean up tests a bit 2025-06-01 22:34:39 +03:00
asagi4 05b0b2ad26 Set the default value of $1 to empty with DEF(MACRO()=) 2025-06-01 22:34:34 +03:00
asagi4 5831608c4e Add a macro expansion node 2025-06-01 21:41:03 +03:00
asagi4 c815bb44f1 Reduce logging verbosity 2025-05-31 01:19:37 +03:00
asagi4 e55c50e9d7 Fix the case with more than one coupled conditioning 2025-05-31 01:09:29 +03:00
asagi4 1b0ff62d10 v2.0.0-rc.3 2025-05-31 00:25:05 +03:00
asagi4 cf93093d59 Fix long prompts with Attention Couple
Broken by moving the LCM calculation outside the loop

See #108
2025-05-31 00:22:42 +03:00
16 changed files with 662 additions and 262 deletions
+3
View File
@@ -11,4 +11,7 @@ test:
test_graph:
PYTHONPATH=../../ python -m prompt_control.test_graph
test_encode:
PYTHONPATH=../../ python -m prompt_control.test_encode
.PHONY: check format all
+3 -3
View File
@@ -10,11 +10,11 @@ A `Basic Text to Image` template is included with the extension, and can be load
You can use text prompts to control the following:
- Prompt scheduling and filtering without noodle soup.
- 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, also fully schedulable.
- 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 operations like `BREAK` and `AND`
- 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`
+1 -1
View File
@@ -29,7 +29,7 @@ cache_hack.init()
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
nodes = ["base", "lazy", "tools"]
nodes = ["base", "lazy", "tools", "hooks"]
for node in nodes:
mod = importlib.import_module(f".prompt_control.nodes_{node}", package=__name__)
+39
View File
@@ -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)`
+65 -51
View File
@@ -95,16 +95,9 @@ generates a LoRA schedule based on a sinewave
# Basic prompt syntax
This syntax is also available in outside scheduled prompts, where applicable.
This syntax is also available in outside scheduled with the `PCTextEncode` node, where applicable.
## LoRA loading
The A111-style syntax `<lora:loraname:weight>` can be used to load LoRAs via the prompt. See LoRA scheduling above.
## Combining prompts, A1111-style
### BREAK
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
@@ -125,7 +118,21 @@ cat [\:0::0.5] AND dog
```
Note that the `:` needs to be escaped with a `\` or it will be interpreted as scheduling syntax.
# Functions
## 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.
@@ -137,7 +144,26 @@ Note: Whitespace is usually *not* stripped from string parameters by default. Co
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.
### STYLE: Configure prompt weighting (also known as "Advanced CLIP Encode")
### 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.
@@ -189,7 +215,7 @@ 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` inside `TE` does not do anything sensible; `TE(l=foo AND bar)` will parse as two prompts `TE(foo` and `bar)`. `BREAK`, `SHIFT` and `SHUFFLE` do work, however
- `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
@@ -218,7 +244,9 @@ Whitespace is *not* stripped and may also be used as a joiner or separator
### NOISE: Add noise to a prompt
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.
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
@@ -300,7 +328,10 @@ Experimental features are unstable and may disappear or change without warning.
## DEF: Lightweight prompt macros
You can define "prompt macros" by using `DEF`:
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]
@@ -309,7 +340,7 @@ 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])
@@ -319,7 +350,7 @@ 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.
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:
@@ -332,50 +363,33 @@ gives
[example:0,1] [test:0.2,1]
```
Note that unspecified parameters will not be substituted:
```
DEF(mything=a $1 b $2)
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 A b $2
a "$1" b "$2"
a "" b "$2"
a "A" b "$2"
```
Macros are expanded before any other parsing takes place. The expansion continues until no further changes occur. Recursion will raise an error.
## ATTN: Attention couple
## Attention Couple
Attention Couple is an attention-based implementation of regional prompting. it can often be faster and more flexible than latent-based masking.
The implementation is based on the one by [pamparamm](https://github.com/pamparamm/ComfyUI-ppm.git), but modified to use ComfyUI's hook system. This enables it to work with prompt scheduling.
The implementation produces slightly different results from Pamparamm's implementation because ComfyUI will only run the hook for conds that have it attached, unlike the ModelPatcher based implementation which has special logic to avoid messing up negative prompts with attention masks. It's also slightly slower because ComfyUI can't batch cond and uncond calculations while the hook is in use.
As a consequence of this, however, you can also use `ATTN()` in your negative prompt, and it will work correctly.
### 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)`
See [here](doc/attention_couple.md)
## TE_WEIGHT
+62 -49
View File
@@ -1,15 +1,6 @@
import torch
import numpy as np
import itertools
def _grouper(n, iterable):
it = iter(iterable)
while True:
chunk = list(itertools.islice(it, n))
if not chunk:
return
yield chunk
from math import copysign
def _norm_mag(w, n):
@@ -48,29 +39,15 @@ def mask_word_id(tokens, word_ids, target_id, mask_token):
return (new_tokens, mask)
def batched_clip_encode(tokens, length, encode_func, num_chunks):
embs = []
for e in _grouper(32, tokens):
enc, pooled = encode_func(e)
enc = enc.reshape((len(e), length, -1))
embs.append(enc)
embs = torch.cat(embs)
embs = embs.reshape((len(tokens) // num_chunks, length * num_chunks, -1))
return embs
def from_masked(tokens, weights, word_ids, base_emb, length, encode_func, m_token=266):
pooled_base = base_emb[0, length - 1 : length, :]
def from_masked(tokens, weights, word_ids, base_emb, pooled_base, max_length, encode_func, m_token):
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, :]
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)
# 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 = []
@@ -85,22 +62,22 @@ def from_masked(tokens, weights, word_ids, base_emb, length, encode_func, m_toke
ws.append(w)
# batch process prompts
embs = batched_clip_encode(masked_tokens, length, encode_func, len(tokens))
embs, pooled = encode_func(tokens)
masks = torch.cat(masks)
embs = base_emb.expand(embs.shape) - embs
pooled = embs[0, length - 1 : length, :]
if pooled is not None and max_length:
pooled = embs[0, max_length - 1 : 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
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
return ((weight_tensor - 1) * embs), pooled
def mask_inds(tokens, inds, mask_token):
@@ -112,13 +89,16 @@ def mask_inds(tokens, inds, mask_token):
return new_tokens
def down_weight(tokens, weights, word_ids, base_emb, length, encode_func, m_token=266):
def down_weight(tokens, weights, word_ids, base_emb, pooled_base, max_length, encode_func, m_token):
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
return (
base_emb,
tokens,
base_emb[0, max_length - 1 : max_length, :] if (pooled_base is not None and max_length) else None,
)
m_token = (m_token, 1.0)
masked_tokens = []
@@ -130,14 +110,16 @@ def down_weight(tokens, weights, word_ids, base_emb, length, encode_func, m_toke
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, pooled = encode_func(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, :]
if pooled and max_length:
pooled = weighted_emb[0, max_length - 1 : max_length, :]
return weighted_emb, masked_current, pooled
def scale_emb_to_mag(base_emb, weighted_emb):
@@ -179,13 +161,31 @@ def advanced_encode_from_tokens(
token_normalization,
weight_interpretation,
encode_func,
m_token=266,
length=77,
m_token="+",
w_max=1.0,
return_pooled=False,
apply_to_pooled=False,
tokenizer=None,
**extra_args
):
negpip = extra_args.get("has_negpip")
if negpip:
weights_sign = [[copysign(1, w) for _, w, _ in x] for x in tokenized]
tokenized = [[(t, abs(w), p) for t, w, p in x] for x in tokenized]
orig_encode = encode_func
def _encode(t):
emb, pooled = orig_encode(t)
return emb[:, 0::2, :], pooled
encode_func = _encode
assert tokenizer, "Must pass tokenizer"
max_length = None
if tokenizer.pad_to_max_length:
max_length = tokenizer.max_length
m_token = tokenizer.tokenize_with_weights(m_token)[0][tokenizer.tokens_start]
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]
@@ -215,25 +215,38 @@ def advanced_encode_from_tokens(
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)
weighted_emb, pooled = encode_func(pos_tokens)
weighted_emb, _, pooled = down_weight(
pos_tokens, weights, word_ids, weighted_emb, pooled, max_length, encode_func, m_token
)
if weight_interpretation == "comfy++":
weighted_emb, tokens_down, _ = down_weight(unweighted_tokens, weights, word_ids, base_emb, length, encode_func)
weighted_emb, tokens_down, _ = down_weight(
unweighted_tokens, weights, word_ids, base_emb, pooled_base, max_length, encode_func, m_token
)
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)
embs, pooled = from_masked(
unweighted_tokens, weights, word_ids, base_emb, pooled_base, max_length, encode_func, m_token
)
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)
weighted_emb, _, pooled = down_weight(
unweighted_tokens, weights, word_ids, base_emb, pooled_base, max_length, encode_func, m_token
)
if weight_interpretation == "perp":
weighted_emb, pooled = perp_weight(
weights, (base_emb, pooled_base), encode_func(extra_args["tokenizer"].tokenize_with_weights(""))
)
if negpip:
emb_negpip = torch.empty_like(weighted_emb).repeat(1, 2, 1)
emb_negpip[:, 0::2, :] = weighted_emb
emb_negpip[:, 1::2, :] = weighted_emb * weights_like(weights_sign, weighted_emb)
weighted_emb = emb_negpip
if return_pooled:
if apply_to_pooled:
return weighted_emb, pooled
+165 -73
View File
@@ -2,31 +2,46 @@
# 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 TransformerOptionsHook, HookGroup, EnumHookScope, set_hooks_for_conditioning
from comfy.hooks import EnumHookScope, HookGroup, TransformerOptionsHook, set_hooks_for_conditioning
from comfy.model_patcher import ModelPatcher
import logging
log = logging.getLogger("comfyui-prompt-control")
def set_cond_attnmask(base_cond, extra_conds, fill=False):
hook = AttentionCoupleHook(base_cond[0], extra_conds, fill=fill)
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(base_cond, hooks=group)
return set_hooks_for_conditioning(c, hooks=group)
def lcm_for_list(numbers):
current_lcm = numbers[0]
for number in numbers[1:]:
current_lcm = math.lcm(current_lcm, number)
return current_lcm
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:
@@ -42,75 +57,91 @@ class Proxy:
class AttentionCoupleHook(TransformerOptionsHook):
def __init__(self, base_cond, conds, fill):
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].pop("strength", 1.0)
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] = [cond[0] for cond in conds]
base_mask = base_cond[1].pop("mask", None)
masks = [cond[1].pop("mask") * cond[1].pop("mask_strength") 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 base_mask is None and not fill:
raise ValueError("You must specify a base mask when fill=False")
elif base_mask is None:
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)
# calculate later
self.conds_k_tensor = None
self.conds_v_tensor = None
def on_apply_hooks(self, model: ModelPatcher, transformer_options: dict[str]):
if self.conds_k_tensor is None:
attn_patches = model.model_options["transformer_options"].get("patches", {}).get("attn2_patch", [])
has_negpip = any("negpip_attn" in i.__name__ for i in attn_patches)
log.debug("AttentionCouple has_negpip=%s", has_negpip)
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)
conds_kv = (
[(cond[:, 0::2], cond[:, 1::2]) for cond in self.conds]
if has_negpip
else [(cond, cond) for cond in self.conds]
)
num_tokens_k = [cond[0].shape[1] for cond in conds_kv]
num_tokens_v = [cond[1].shape[1] for cond in conds_kv]
lcm_tokens_k = lcm_for_list(num_tokens_k)
lcm_tokens_v = lcm_for_list(num_tokens_v)
self.conds_k_tensor = torch.cat(
[
cond[0].repeat(1, lcm_tokens_k // num_tokens_k[i], 1) * self.strengths[i]
for i, cond in enumerate(conds_kv)
],
dim=0,
)
if has_negpip:
self.conds_v_tensor = torch.cat(
[
cond[1].repeat(1, lcm_tokens_v // num_tokens_v[i], 1) * self.strengths[i]
for i, cond in enumerate(conds_kv)
],
dim=0,
)
# 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_v_tensor = self.conds_k_tensor
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)
@@ -118,32 +149,93 @@ class AttentionCoupleHook(TransformerOptionsHook):
def attn2_patch(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, extra_options):
cond_or_uncond = extra_options["cond_or_uncond"]
num_chunks = len(cond_or_uncond) # should always be 1
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 = self.conds_k_tensor.expand(bs, *self.conds_k_tensor.shape[1:])
conds_v_tensor = self.conds_v_tensor.expand(bs, *self.conds_v_tensor.shape[1:])
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,
)
q = q.repeat(self.num_conds, 1, 1)
k = k.repeat(1, self.conds_k_tensor.shape[1] // k.shape[1], 1)
v = v.repeat(1, self.conds_v_tensor.shape[1] // v.shape[1], 1)
qs, ks, vs = [], [], []
cond_or_uncond_couple.clear()
k = torch.cat([k * self.base_strength, conds_k_tensor], dim=0)
v = torch.cat([v * self.base_strength, conds_v_tensor], dim=0)
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):
# out has been extended to shape [num_conds*batch_size, TOKENS, N]
# out is [b1c1 b1c2 ... b1cN, b2c1 b2c2 ... b2cn, ...]
num_conds = self.mask.shape[0]
bs = out.shape[0] // num_conds
num_tokens = out.shape[1]
mask_size = extra_options["activations_shape"][-2:]
mask_downsample = F.interpolate(self.mask, size=mask_size, mode="nearest")
mask_downsample = mask_downsample.view(num_conds, num_tokens, 1).repeat_interleave(bs, dim=0)
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
# cond_outputs is [num_conds*bs, tokens, N], output needs to be [bs, tokens, N]
cond_outputs = out * mask_downsample
cond_output = cond_outputs.view(num_conds, bs, out.shape[1], out.shape[2]).sum(0)
return cond_output
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)
+7
View File
@@ -209,6 +209,9 @@ def encode_regions(clip_regions, encode, tokenizer):
debug_tokens("region", region_prompt, tokenizer)
region_emb, _ = encode(region_prompt)
region_emb -= base_embedding_start
# NegPiP support:
if region_emb.shape[1] == 2 * region_masking.shape[1]:
region_masking = torch.repeat_interleave(region_masking, 2, dim=1)
region_emb *= region_masking
region_embeddings.append(region_emb)
@@ -217,6 +220,10 @@ def encode_regions(clip_regions, encode, tokenizer):
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
+48 -2
View File
@@ -1,9 +1,13 @@
import logging
import comfy.utils
import comfy.hooks
import comfy.utils
import folder_paths
from .utils import consolidate_schedule
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")
@@ -79,10 +83,52 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
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)]
n_hook_group = n_hook_group.clone()
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)",
}
+21 -1
View File
@@ -1,5 +1,5 @@
import logging
from .parser import parse_prompt_schedules
from .parser import parse_prompt_schedules, expand_macros
from .nodes_lazy import NODE_CLASS_MAPPINGS as LAZY_NODES
import json
import folder_paths
@@ -209,6 +209,24 @@ class PCExtractScheduledPrompt:
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,
@@ -216,6 +234,7 @@ NODE_CLASS_MAPPINGS = {
"PCSetLogLevel": PCSetLogLevel,
"PCExtractScheduledPrompt": PCExtractScheduledPrompt,
"PCSaveExpandedWorkflow": PCSaveExpandedWorkflow,
"PCMacroExpand": PCMacroExpand,
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -225,4 +244,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"PCSetLogLevel": "PC: Configure Logging (for debug)",
"PCExtractScheduledPrompt": "PC: Extract Scheduled Prompt",
"PCSaveExpandedWorkflow": "PC: Save Expanded Workflow (for debug)",
"PCMacroExpand": "PC: Expand Macros",
}
+4 -3
View File
@@ -380,14 +380,15 @@ def parse_search(search):
if not name:
return None
args = args.strip()
if args:
# 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 replace_def(text):
def expand_macros(text):
text, defs = get_function(text, "DEF", defaults=None)
res = text
prevres = text
@@ -443,5 +444,5 @@ def substitute_defcall(text, search, replace):
@lru_cache
def parse_prompt_schedules(prompt, **kwargs):
prompt = replace_def(prompt)
prompt = expand_macros(prompt)
return PromptSchedule(prompt, **kwargs)
+107 -47
View File
@@ -3,6 +3,7 @@ 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
@@ -125,9 +126,10 @@ def fix_word_ids(tokens):
return tokens
def tokenize_chunks(clip, text, need_word_ids):
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
@@ -135,40 +137,32 @@ def tokenize_chunks(clip, text, need_word_ids):
r = shuffle_chunk(s, r)
if r != c:
log.info("Shuffled prompt chunk to %s", r)
c = r
shuffled_chunks.append(r)
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])
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 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
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)
tokens = tokenize_chunks(clip, text, need_word_ids, can_break)
per_te_prompts = {}
if l_prompts:
@@ -196,29 +190,86 @@ def encode_prompt_segment(
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)[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(len(tokens[k]) for k in tokens)
empty = None
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:
if empty is None:
empty = clip.tokenize("", return_word_ids=need_word_ids)
tokens[k] += empty[k]
tokens[k] += empty_tokens[k]
tokens = fix_word_ids(tokens)
return fix_word_ids(tokens)
tes = []
for k in tokens:
if k in ["g", "l"]:
tes.append(f"clip_{k}")
else:
tes.append(k)
clip = hook_te(clip, tes, style, normalization, extra)
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
return clip.encode_from_tokens_scheduled(tokens, add_dict=settings)
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):
@@ -243,7 +294,8 @@ def apply_weights(output, te_name, spec):
pooled_w = 1.0
log.info("Weighting %s output by %s, pooled by %s", te_name, w, pooled_w)
out = out * w
pooled = pooled * pooled_w
if pooled is not None:
pooled = pooled * pooled_w
return out, pooled
else:
@@ -271,15 +323,24 @@ def hook_te(clip, te_names, style, normalization, extra):
return clip
newclip = clip.clone()
for te_name in te_names:
if hasattr(clip.patcher.model, te_name):
tokenizer = getattr(clip.tokenizer, f"clip_{te_name}", getattr(clip.tokenizer, te_name, None))
if tokenizer:
x = extra.copy()
x["tokenizer"] = getattr(clip.tokenizer, te_name)
log.debug("Hooked into %s with style=%s, normalization=%s", te_name, style, normalization)
x["tokenizer"] = tokenizer
if not hasattr(clip.patcher.model, te_name):
te_name = "clip_" + te_name
if not hasattr(clip.patcher.model, te_name):
log.warning("TE model %s not found on model patcher. Skipping...", te_name)
continue
log.debug("Hooked into te=%s with style=%s, normalization=%s", te_name, style, normalization)
encode = clip.patcher.get_model_object(f"{te_name}.encode_token_weights")
x["has_negpip"] = clip.patcher.model_options.get("ppm_negpip", False)
newclip.patcher.add_object_patch(
f"{te_name}.encode_token_weights",
make_patch(
te_name,
clip.patcher.get_model_object(f"{te_name}.encode_token_weights"),
encode,
normalization,
style,
x,
@@ -287,7 +348,7 @@ def hook_te(clip, te_names, style, normalization, extra):
)
# 'g' and 'l' exist in these are clip_g and clip_l
else:
log.debug("Tokens contain items with key %s but no TE found on object with that name.", te_name)
log.warning("Tokens contain items with key %s but no tokenizer found on object with that name.", te_name)
return newclip
@@ -353,7 +414,7 @@ def make_mask(args, size, weight):
mask = torch.full((h, w), 0, dtype=torch.float32, device="cpu")
mask[ys[0] : ys[1], xs[0] : xs[1]] = weight
mask = mask.unsqueeze(0)
log.info("Mask xs=%s, ys=%s, shape=%s, weight=%s", xs, ys, mask.shape, weight)
log.debug("Mask xs=%s, ys=%s, shape=%s, weight=%s", xs, ys, mask.shape, weight)
return mask
@@ -512,7 +573,6 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
log.warning("MASK() and FILL() can't be used together, ignoring FILL()")
else:
fill = True
log.info("Using attention masking for prompt segment")
attnmasked_prompts.extend(x)
else:
conds.extend(x)
+88
View File
@@ -0,0 +1,88 @@
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_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()
+34 -29
View File
@@ -1,5 +1,5 @@
import unittest
from .parser import parse_prompt_schedules as parse
from .parser import parse_prompt_schedules as parse, expand_macros
def prompt(until, text, *loras):
@@ -19,25 +19,17 @@ class TestParser(unittest.TestCase):
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]"]]
for p in eqs[1:]:
self.assertEqual(eqs[0].parsed_prompt, p.parsed_prompt)
eqs = [parse(p) for p in ["[before:during:after:0.1]", "[before:during:after:0.1,1.0]", "[before:during:0.1]"]]
for p in eqs[1:]:
self.assertEqual(eqs[0].parsed_prompt, p.parsed_prompt)
eqs = [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]"]]
for p in eqs[1:]:
self.assertEqual(eqs[0].parsed_prompt, p.parsed_prompt)
eqs = [parse(p) for p in ["[a:b:0.5]", "[a::b:0.5,0.5]"]]
for p in eqs[1:]:
self.assertEqual(eqs[0].parsed_prompt, p.parsed_prompt)
eqs = [parse(p) for p in ["[a::0.5]", "[a:::0.5,0.5]"]]
for p in eqs[1:]:
self.assertEqual(eqs[0].parsed_prompt, p.parsed_prompt)
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(
@@ -135,22 +127,33 @@ class TestParser(unittest.TestCase):
p = parse("DEF(X=[($1):($1:$2):$2])X(test;0.7)")
p2 = parse("[(test):(test:0.7):0.7]")
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
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]")
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
with self.subTest("two functions"):
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
p = parse("DEF(X(a;b)=$1 $2 $3 d)X(A) X(A;B;C)")
p2 = parse("A b $3 d A B C d")
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]")
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
with self.subTest("defaults, nested parens"):
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
with self.assertRaises(ValueError) as c:
parse("DEF(X=recurse Y) DEF(Y=recurse X) X")
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):
@@ -190,11 +193,13 @@ class TestParser(unittest.TestCase):
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)
self.assertPrompt(p3, step, step, x)
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)
self.assertPrompt(p4, step, step, *x)
with self.subTest(step):
self.assertPrompt(p4, step, step, *x)
self.assertPrompt(p4, 0.7, 0.8, "wolf", ("canine", 1.0, 1.0))
+14 -2
View File
@@ -87,7 +87,7 @@ def find_closing_paren(text, start):
return len(text)
def get_function(text, func, defaults, return_func_name=False, placeholder=""):
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)
@@ -98,7 +98,19 @@ def get_function(text, func, defaults, return_func_name=False, placeholder=""):
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)
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-prompt-control"
description = "Nodes for convenient prompt editing, making many common operations prompt-controllable"
version = "2.0.0-rc.2"
version = "2.0.0-rc.5"
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"]