Compare commits

...
20 Commits
Author SHA1 Message Date
asagi4 d7be7bc29e v2.1.2 2026-01-13 21:50:09 +02:00
asagi4 329d4cf95f Add a test for #133 2026-01-13 21:48:34 +02:00
asagi4 c52ace71aa #133 properly set function locations in get_function 2026-01-13 20:45:10 +02:00
asagi4 4806cf5959 refactor get_function to make it more consistent 2026-01-13 20:45:10 +02:00
asagi4 a0ab709f50 v2.1.1 2025-12-14 17:57:53 +02:00
asagi4 66fef1ffe8 Avoid splitting AND, CAT and others when inside quotes
See #132
2025-12-05 21:25:28 +02:00
asagi4 3341e9f81e Fix TE_WEIGHT failing with some encoders
See #131

Prompt weighting and attention couple will probably not work, but
this should prevent exceptions.
2025-12-03 16:07:31 +02:00
asagi4 aa6b4608f0 Clarify attention couple docs a bit and add a warning if IMASK is used without attaching custom masks
See #108
2025-12-01 13:19:29 +02:00
asagi4 d1cc60b00a v2.1.0 2025-11-20 20:20:22 +02:00
asagi4 efe8939250 Fix issue with T5 encoding sometimes returning NaNs in test 2025-11-20 20:14:40 +02:00
asagi4 39b353b916 Fix tests for #130 2025-11-20 20:14:37 +02:00
asagi4 485bc7f2ab Prepare for v3 conversion of imported nodes, see #12
This should prevent things from breaking, but it needs a bit of testing.
2025-11-20 15:29:33 +02:00
asagi4 1d03ded9dd Find LoRAs with partial match 2025-11-20 15:23:51 +02:00
asagi4 94a4d076e0 Remove unused attributes 2025-08-30 13:46:00 +03:00
asagi4 228dc4b22b Remove the old advanced encoding implementation 2025-08-27 20:27:30 +03:00
asagi4 db523e1f16 Remove dead code 2025-08-27 20:23:41 +03:00
asagi4 6b08c7a90e v2.0.1 2025-08-27 20:20:13 +03:00
asagi4 76142c4b7e Fix #127 2025-08-27 20:15:30 +03:00
asagi4 1d84fdaf9e Release 2.0.0 2025-08-19 19:53:34 +03:00
asagi4 1c50ae5297 Disable the cache hack for now 2025-08-19 19:53:23 +03:00
17 changed files with 290 additions and 480 deletions
+2 -2
View File
@@ -12,10 +12,10 @@ test_graph:
PYTHONPATH=../../ python -m prompt_control.test_graph
test_encode:
PYTHONPATH=../../ python -m prompt_control.test_encode
PYTHONPATH=../../ python -m prompt_control.test_encode --verbose
test_encode_both:
TEST_TE="clip_l t5" PYTHONPATH=../../ python -m prompt_control.test_encode
TEST_TE="clip_l t5" PYTHONPATH=../../ python -m prompt_control.test_encode --verbose
test_heavy: test_graph test_encode_both
-19
View File
@@ -32,21 +32,6 @@ Prompt Control uses graph generation, and tries to delegate functionality to co
If you encounter issues as a user or if you're a node developer and Prompt Control somehow breaks something, feel free to file a bug report.
## Prompt Control v2
Prompt control has been almost completely rewritten. It now uses ComfyUI's lazy execution to build graphs from the text prompt at runtime. The generated graph is often exactly equivalent to a manually built workflow using native ComfyUI nodes. There are no more weird sampling hooks that could cause problems with other nodes
### Removed features
- Prompt interpolation syntax; it was too cumbersome to maintain
- LoRA block weight integration; ditto, for now.
### Everything broke, where are the old nodes?
If you really need them, you can install the [legacy nodes](https://github.com/asagi4/comfyui-prompt-control-legacy). However, I will not fix bugs in those nodes, and I strongly recommend just migrating your workflows to the new nodes.
You can have both installed at the same time; none of the nodes conflict.
## Requirements
For LoRA scheduling to work, you'll need at least version 0.3.7 of ComfyUI (0.3.36 of ComfyUI desktop).
@@ -100,8 +85,4 @@ This node configures `PCTextEncode` default values for some functions by attachi
- ComfyUI's caching mechanism has an issue that makes it unnecessarily invalidate caches for certain inputs; you'll still get some benefit from the lazy nodes, but changing inputs that shouldn't affect downstream nodes (especially if using filtering) will still cause them to be recomputed because ComfyUI doesn't realize the inputs haven't changed.
If you want to enable a hack to fix this, set `PROMPTCONTROL_ENABLE_CACHE_HACK=1` in your environment. Unset it to disable.
It's a purely optional performance optimization that allows Prompt Control nodes to override their cache keys in a way that should not interfere with other nodes. Note that the optimization only works if the text input to the lazy nodes is a constant (so either directly on the node or from a primitive); outputs from other nodes can't be optimized.
- Cutoff does not work with models that use non-CLIP text encoders, like Flux. This might be fixable, but it's uncertain if cutoff even makes sense for those models.
-3
View File
@@ -23,9 +23,6 @@ if os.environ.get("PROMPTCONTROL_DEBUG"):
else:
log.setLevel(logging.INFO)
cache_hack = importlib.import_module(".prompt_control.cache_hack", package=__name__)
cache_hack.init()
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
+3 -1
View File
@@ -15,7 +15,7 @@ To enable batching negative prompts, run your positive and negative prompt throu
## Syntax
See also the main syntax documentation for `MASK` etc.
See also the [regional prompting documentation](/doc/regional_prompts.md) for information about `MASK` etc.
### COUPLE: Trigger Attention Couple
@@ -38,3 +38,5 @@ For example:
```
dog FILL() COUPLE(0.5 1) cat
```
Note that because the generation still sees and diffuses the full latent, attention coupling is not guaranteed to perfectly limit the effect of your prompt to the masked area.
+4
View File
@@ -23,6 +23,8 @@ cat [\:0::0.5] AND dog
```
Note that the `:` needs to be escaped with a `\` or it will be interpreted as scheduling syntax.
If `AND` is placed inside quotes (eg. `Text saying "CAT AND DOG"`) it will be treated as regular text.
## Note about processing order
Prompt operators are processed in the following order, meaning that all features "below" another can be affected by the feature above it. That is, `BREAK` can go inside a `TE()` call, but not `AND` or `CAT`.
@@ -53,6 +55,8 @@ 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.
like AND, if any function is placed inside quotes, it will *not* activate and is instead treated as regular text.
### BREAK
The keyword `BREAK` causes the prompt to be tokenized in separate chunks, padding each chunk to the text encoder's maximum size before encoding.
+8 -1
View File
@@ -89,10 +89,17 @@ You can refer to LoRAs by using the filename without extension and subdirectorie
Alternatively, the name can include the full directory path relative to ComfyUI's search paths, without extension: `<lora:XL/sdxllora:0.5>`. In this case, the *full* path must match.
You can also give the exact path (including the extension) as shown in `LoRALoader`.
If no match is found, the node will try to replace spaces with underscores and search again. That is, `<lora:cats and dogs:1>` will find `cats_and_dogs.safetensors`. This helps with some autocompletion scripts that replace underscores with spaces.
Finally, you can give the exact path (including the extension) as shown in `LoRALoader`.
Finally, if none of the above produce a match, the search term will be split by whitespace and files that contain all of the parts in any order will be considered. If this returns only a single match, it will be loaded. For example, consider LoRAs:
- `xl/red_cats.safetensors`
- `flux/blue_cats.safetensors`
- `flux/red_cats.safetensors`
Then `<lora:cats xl:1>` would match the red cats LoRA, but `cats flux` would be ambiguous and not match.
## Alternating
+25 -39
View File
@@ -3,7 +3,6 @@ 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")
@@ -26,7 +25,7 @@ def _grouper(n, iterable):
def batched_clip_encode(tokens, length, encode_func, num_chunks):
embs = []
for e in _grouper(32, tokens):
enc, pooled = encode_func(e)
enc, pooled, *_ = encode_func(e)
enc = enc.reshape((len(e), length, -1))
embs.append(enc)
@@ -101,24 +100,24 @@ def style_comfy(encoder, tokens, **kwargs):
def style_a1111(encoder, tokens, **kwargs):
base_emb, pooled = encoder.base_emb(tokens)
base_emb, pooled, *extra = encoder.base_emb(tokens)
weighted_emb = base_emb * weights_like(encoder.weights(tokens), base_emb)
weighted_emb = (base_emb.mean() / weighted_emb.mean()) * weighted_emb # renormalize
return weighted_emb, pooled
return (weighted_emb, pooled) + tuple(extra)
def style_compel(encoder, tokens, **kwargs):
pos_tokens = encoder.weighted_with(tokens, lambda w: w if w > 1.0 else 1.0)
weighted_emb, pooled = encoder.encode_fn(pos_tokens)
weighted_emb, pooled, *extra = encoder.encode_fn(pos_tokens)
weighted_emb, _, pooled = encoder.down_weight(
pos_tokens, encoder.weights(tokens), encoder.word_ids(tokens), weighted_emb, pooled
)
return weighted_emb, pooled
return (weighted_emb, pooled) + tuple(extra)
def style_comfypp(encoder, tokens, **kwargs):
unweighted_tokens = encoder.unweighted(tokens)
base_emb, pooled_base = encoder.base_emb(tokens)
base_emb, pooled_base, *extra = encoder.base_emb(tokens)
weighted_emb, tokens_down, _ = encoder.down_weight(
unweighted_tokens, encoder.weights(tokens), encoder.word_ids(tokens), base_emb, pooled_base
)
@@ -132,23 +131,23 @@ def style_comfypp(encoder, tokens, **kwargs):
)
weighted_emb += embs
return weighted_emb, pooled
return (weighted_emb, pooled) + tuple(extra)
def style_downweight(encoder, tokens, **kwargs):
weights = scale_to_norm(encoder.weights(tokens), encoder.word_ids(tokens), encoder.w_max)
base_emb, pooled_base = encoder.base_emb(tokens)
base_emb, pooled_base, *extra = encoder.base_emb(tokens)
weighted_emb, _, pooled = encoder.down_weight(
encoder.unweighted(tokens), weights, encoder.word_ids(tokens), base_emb, pooled_base
)
return weighted_emb, pooled
return (weighted_emb, pooled) + tuple(extra)
def style_perp(encoder, tokens, **kwargs):
zero_emb, zero_pooled = encoder.encode_fn(encoder.tokenizer.tokenize_with_weights(""))
base_emb, pooled = encoder.base_emb(tokens)
return perp_weight(encoder.weights(tokens), (base_emb, pooled), (zero_emb, zero_pooled))
zero_emb, zero_pooled, *_ = encoder.encode_fn(encoder.tokenizer.tokenize_with_weights(""))
base_emb, pooled, *extra = encoder.base_emb(tokens)
return perp_weight(encoder.weights(tokens), (base_emb, pooled), (zero_emb, zero_pooled)) + tuple(extra)
def apply_negpip(encoder, emb, pooled, **kwargs):
@@ -258,8 +257,8 @@ class AdvancedEncoder:
if negpip:
def _encode(t):
emb, pooled = encode_fn(t)
return emb[:, 0::2, :], pooled
emb, pooled, *extra = encode_fn(t)
return (emb[:, 0::2, :], pooled) + tuple(extra)
self.encode_fn = _encode
self.preprocessors.insert(0, lambda encoder, tokens, **kwargs: encoder.weighted_with(tokens, abs))
@@ -289,7 +288,7 @@ class AdvancedEncoder:
if w[i] >= 1:
continue
masked_current = mask_inds(masked_current, np.where(w_inv == i)[0], self.m_token)
masked, _ = self.encode_fn(masked_current)
masked, _, *extra = self.encode_fn(masked_current)
emblist.append(masked)
embs = torch.cat(emblist)
@@ -349,16 +348,16 @@ class AdvancedEncoder:
for op in self.preprocessors:
normalized_tokens = op(self, normalized_tokens)
emb, pooled = self.weight_fn(self, normalized_tokens, original_tokens=tokens)
emb, pooled, *extra = self.weight_fn(self, normalized_tokens, original_tokens=tokens)
for fn in self.postprocessors:
emb, pooled = fn(self, emb, pooled, tokens=tokens, original_tokens=tokens)
if return_pooled:
if not apply_to_pooled:
_, pooled = self.base_emb(tokens)
return emb, pooled
return emb, None
if not return_pooled:
pooled = None
elif not apply_to_pooled:
_, pooled, *_ = self.base_emb(tokens)
return (emb, pooled) + tuple(extra)
def advanced_encode_from_tokens(
@@ -373,20 +372,7 @@ def advanced_encode_from_tokens(
tokenizer=None,
**extra_args,
):
if "old+" not in weight_interpretation:
enc = AdvancedEncoder(
encode_func, weight_interpretation, token_normalization, tokenizer, m_token, w_max, **extra_args
)
return enc(tokenized, return_pooled=return_pooled, apply_to_pooled=apply_to_pooled)
else:
weight_interpretation = weight_interpretation.replace("old+", "")
log.warning("Using old implementation of %s", weight_interpretation)
return old_advanced_encode_from_tokens(
tokenized,
token_normalization,
weight_interpretation,
encode_func,
266,
return_pooled=return_pooled,
apply_to_pooled=apply_to_pooled,
)
enc = AdvancedEncoder(
encode_func, weight_interpretation, token_normalization, tokenizer, m_token, w_max, **extra_args
)
return enc(tokenized, return_pooled=return_pooled, apply_to_pooled=apply_to_pooled)
-235
View File
@@ -1,235 +0,0 @@
import torch
import numpy as np
import logging
import itertools
log = logging.getLogger("comfyui-prompt-control")
def _norm_mag(w, n):
d = w - 1
return 1 + np.sign(d) * np.sqrt(np.abs(d) ** 2 / n)
# return np.sign(w) * np.sqrt(np.abs(w)**2 / n)
def _grouper(n, iterable):
it = iter(iterable)
while True:
chunk = list(itertools.islice(it, n))
if not chunk:
return
yield chunk
def batched_clip_encode(tokens, length, encode_func, num_chunks):
embs = []
for e in _grouper(32, tokens):
enc, pooled = encode_func(e)
enc = enc.reshape((len(e), length, -1))
embs.append(enc)
embs = torch.cat(embs)
embs = embs.reshape((len(tokens) // num_chunks, length * num_chunks, -1))
return embs
def weights_like(weights, emb):
return torch.tensor(weights, dtype=emb.dtype, device=emb.device).reshape(1, -1, 1).expand(emb.shape)
def divide_length(word_ids, weights):
sums = dict(zip(*np.unique(word_ids, return_counts=True)))
sums[0] = 1
weights = [[_norm_mag(w, sums[id]) if id != 0 else 1.0 for w, id in zip(x, y)] for x, y in zip(weights, word_ids)]
return weights
def shift_mean_weight(word_ids, weights):
delta = 1 - np.mean([w for x, y in zip(weights, word_ids) for w, id in zip(x, y) if id != 0])
weights = [[w if id == 0 else w + delta for w, id in zip(x, y)] for x, y in zip(weights, word_ids)]
return weights
def scale_to_norm(weights, word_ids, w_max):
top = np.max(weights)
w_max = min(top, w_max)
weights = [[w_max if id == 0 else (w / top) * w_max for w, id in zip(x, y)] for x, y in zip(weights, word_ids)]
return weights
def mask_word_id(tokens, word_ids, target_id, mask_token):
new_tokens = [[mask_token if wid == target_id else t for t, wid in zip(x, y)] for x, y in zip(tokens, word_ids)]
mask = np.array(word_ids) == target_id
return (new_tokens, mask)
def from_masked(tokens, weights, word_ids, base_emb, length, encode_func, m_token=266):
pooled_base = base_emb[0, length - 1 : length, :]
wids, inds = np.unique(np.array(word_ids).reshape(-1), return_index=True)
weight_dict = dict((id, w) for id, w in zip(wids, np.array(weights).reshape(-1)[inds]) if w != 1.0)
if len(weight_dict) == 0:
return torch.zeros_like(base_emb), base_emb[0, length - 1 : length, :]
weight_tensor = torch.tensor(weights, dtype=base_emb.dtype, device=base_emb.device)
weight_tensor = weight_tensor.reshape(1, -1, 1).expand(base_emb.shape)
# m_token = (clip.tokenizer.end_token, 1.0) if clip.tokenizer.pad_with_end else (0,1.0)
# TODO: find most suitable masking token here
m_token = (m_token, 1.0)
ws = []
masked_tokens = []
masks = []
# create prompts
for id, w in weight_dict.items():
masked, m = mask_word_id(tokens, word_ids, id, m_token)
masked_tokens.extend(masked)
m = torch.tensor(m, dtype=base_emb.dtype, device=base_emb.device)
m = m.reshape(1, -1, 1).expand(base_emb.shape)
masks.append(m)
ws.append(w)
# batch process prompts
embs = batched_clip_encode(masked_tokens, length, encode_func, len(tokens))
masks = torch.cat(masks)
embs = base_emb.expand(embs.shape) - embs
pooled = embs[0, length - 1 : length, :]
embs *= masks
embs = embs.sum(axis=0, keepdim=True)
pooled_start = pooled_base.expand(len(ws), -1)
ws = torch.tensor(ws).reshape(-1, 1).expand(pooled_start.shape)
pooled = (pooled - pooled_start) * (ws - 1)
pooled = pooled.mean(axis=0, keepdim=True)
return ((weight_tensor - 1) * embs), pooled_base + pooled
def mask_inds(tokens, inds, mask_token):
clip_len = len(tokens[0])
inds_set = set(inds)
new_tokens = [
[mask_token if i * clip_len + j in inds_set else t for j, t in enumerate(x)] for i, x in enumerate(tokens)
]
return new_tokens
def down_weight(tokens, weights, word_ids, base_emb, length, encode_func):
w, w_inv = np.unique(weights, return_inverse=True)
if np.sum(w < 1) == 0:
return base_emb, tokens, base_emb[0, length - 1 : length, :]
# m_token = (clip.tokenizer.end_token, 1.0) if clip.tokenizer.pad_with_end else (0,1.0)
# using the comma token as a masking token seems to work better than aos tokens for SD 1.x
m_token = (266, 1.0)
masked_tokens = []
masked_current = tokens
for i in range(len(w)):
if w[i] >= 1:
continue
masked_current = mask_inds(masked_current, np.where(w_inv == i)[0], m_token)
masked_tokens.extend(masked_current)
embs = batched_clip_encode(masked_tokens, length, encode_func, len(tokens))
embs = torch.cat([base_emb, embs])
w = w[w <= 1.0]
w_mix = np.diff([0] + w.tolist())
w_mix = torch.tensor(w_mix, dtype=embs.dtype, device=embs.device).reshape((-1, 1, 1))
weighted_emb = (w_mix * embs).sum(axis=0, keepdim=True)
return weighted_emb, masked_current, weighted_emb[0, length - 1 : length, :]
def scale_emb_to_mag(base_emb, weighted_emb):
norm_base = torch.linalg.norm(base_emb)
norm_weighted = torch.linalg.norm(weighted_emb)
embeddings_final = (norm_base / norm_weighted) * weighted_emb
return embeddings_final
# For verification
def A1111_renorm(base_emb, weighted_emb):
embeddings_final = (base_emb.mean() / weighted_emb.mean()) * weighted_emb
return embeddings_final
def from_zero(weights, base_emb):
weight_tensor = torch.tensor(weights, dtype=base_emb.dtype, device=base_emb.device)
weight_tensor = weight_tensor.reshape(1, -1, 1).expand(base_emb.shape)
return base_emb * weight_tensor
def old_advanced_encode_from_tokens(
tokenized,
token_normalization,
weight_interpretation,
encode_func,
m_token=266,
w_max=1.0,
return_pooled=False,
apply_to_pooled=False,
**extra_args,
):
length = 77
tokens = [[t for t, _, _ in x] for x in tokenized]
weights = [[w for _, w, _ in x] for x in tokenized]
word_ids = [[wid for _, _, wid in x] for x in tokenized]
# weight normalization
# ====================
# distribute down/up weights over word lengths
if token_normalization.startswith("length"):
weights = divide_length(word_ids, weights)
# make mean of word tokens 1
if token_normalization.endswith("mean"):
weights = shift_mean_weight(word_ids, weights)
# weight interpretation
# =====================
pooled = None
if weight_interpretation in ["comfy", "perp"]:
weighted_tokens = [[(t, w) for t, w in zip(x, y)] for x, y in zip(tokens, weights)]
weighted_emb, pooled_base = encode_func(weighted_tokens)
pooled = pooled_base
else:
unweighted_tokens = [[(t, 1.0) for t, _, _ in x] for x in tokenized]
base_emb, pooled_base = encode_func(unweighted_tokens)
if weight_interpretation == "A1111":
weighted_emb = from_zero(weights, base_emb)
weighted_emb = A1111_renorm(base_emb, weighted_emb)
pooled = pooled_base
if weight_interpretation == "compel":
pos_tokens = [[(t, w) if w >= 1.0 else (t, 1.0) for t, w in zip(x, y)] for x, y in zip(tokens, weights)]
weighted_emb, _ = encode_func(pos_tokens)
weighted_emb, _, pooled = down_weight(pos_tokens, weights, word_ids, weighted_emb, length, encode_func)
if weight_interpretation == "comfy++":
weighted_emb, tokens_down, _ = down_weight(unweighted_tokens, weights, word_ids, base_emb, length, encode_func)
weights = [[w if w > 1.0 else 1.0 for w in x] for x in weights]
# unweighted_tokens = [[(t,1.0) for t, _,_ in x] for x in tokens_down]
embs, pooled = from_masked(unweighted_tokens, weights, word_ids, base_emb, length, encode_func)
weighted_emb += embs
if weight_interpretation == "down_weight":
weights = scale_to_norm(weights, word_ids, w_max)
weighted_emb, _, pooled = down_weight(unweighted_tokens, weights, word_ids, base_emb, length, encode_func)
if return_pooled:
if apply_to_pooled:
return weighted_emb, pooled
else:
return weighted_emb, pooled_base
return weighted_emb, None
-4
View File
@@ -76,10 +76,6 @@ class AttentionCoupleHook(TransformerOptionsHook):
self.kv = {"k": None, "v": 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]
+1 -2
View File
@@ -214,8 +214,7 @@ def build_scheduled_prompts(graph, schedules, clip):
classname = "PCTextEncode"
paramname = "text"
if classnames:
classname = classnames[0][0]
paramname = classnames[0][1]
classname, paramname = classnames[0].args
node = graph.node(classname)
node.set_input("clip", clip)
node.set_input(paramname, p)
+14 -15
View File
@@ -404,9 +404,11 @@ def parse_search(search):
args = ""
name = search.strip()
if arg_start > 0:
arg_end = find_closing_paren(search, arg_start)
arg_end = find_closing_paren(search, arg_start + 1)
if arg_end < 0:
arg_end = len(search)
name = search[:arg_start].strip()
args = search[arg_start + 1 : arg_end - 1]
args = search[arg_start + 1 : arg_end]
if not name:
return None
@@ -425,7 +427,9 @@ def expand_macros(text):
prevres = text
replacements = []
for d in defs:
r = d.split("=", 1)
if not d.args:
continue
r = d.args[0].split("=", 1)
search = parse_search(r[0].strip())
if not search or len(r) != 2:
log.warning("Ignoring invalid DEF(%s)", d)
@@ -448,21 +452,16 @@ def expand_macros(text):
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{name}")
for i, parameters in enumerate(defns):
ph = f"\0DEFNCALL{name}{i}\0"
text, defns = get_function(text, name, defaults=None, placeholder=f"DEFNCALL{name}", require_args=False)
for i, d in enumerate(defns):
ph = d.placeholder
assert ph is not None, "This is a bug"
parameters = d.args
paramvals = []
if parameters is not None:
paramvals = [x.strip() for x in parameters.split(";")]
if parameters:
paramvals = [x.strip() for x in parameters[0].split(";")]
r = replace
for i, v in enumerate(paramvals):
r = re.sub(rf"\${i+1}\b", v, r)
+70 -40
View File
@@ -1,11 +1,23 @@
from __future__ import annotations
import logging
import re
import torch
import math
from functools import partial
from comfy_extras.nodes_mask import FeatherMask, MaskComposite
from nodes import ConditioningAverage
from .utils import safe_float, get_function, split_by_function, parse_floats, smarter_split
from .utils import (
safe_float,
get_function,
split_by_function,
parse_floats,
smarter_split,
call_node,
split_quotable,
FunctionSpec,
ComfyConditioning,
)
from .adv_encode import advanced_encode_from_tokens
from .cutoff import process_cuts
from .parser import parse_cuts
@@ -25,7 +37,7 @@ def get_sdxl(text, defaults):
text, sdxl = get_function(text, "SDXL", ["none", "none", "none"])
if not sdxl:
return text, {}
args = sdxl[0]
args = sdxl[0].args
d = defaults
w, h = parse_floats(args[0], [d.get("sdxl_width", 1024), d.get("sdxl_height", 1024)], split_re="\\s+")
tw, th = parse_floats(args[1], [d.get("sdxl_twidth", 1024), d.get("sdxl_theight", 1024)], split_re="\\s+")
@@ -46,7 +58,7 @@ 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()
args = spec[0].args[0].strip()
res = {}
for arg in args.split(","):
try:
@@ -62,7 +74,7 @@ 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, normalization = styles[0].args
style = style.strip()
normalization = normalization.strip()
if style.replace("old+", "") not in AVAILABLE_STYLES:
@@ -77,8 +89,9 @@ def get_style(text, default_style="comfy", default_normalization="none"):
return style, normalization, text
def shuffle_chunk(shuffle, c):
func, shuffle = shuffle
def shuffle_chunk(func_spec: FunctionSpec, c: str) -> str:
func = func_spec.name
shuffle = func_spec.args
shuffle_count = int(safe_float(shuffle[0], 0))
_, separator, joiner = shuffle
if separator == "default":
@@ -128,11 +141,11 @@ def fix_word_ids(tokens):
def tokenize_chunks(clip, text, need_word_ids, can_break):
chunks = re.split(r"\bBREAK\b", text)
chunks = list(split_quotable(text, r"\bBREAK\b"))
token_chunks = []
shuffled_chunks = []
for c in chunks:
c, shuffles = get_function(c.strip(), "(SHIFT|SHUFFLE)", ["0", "default", "default"], return_func_name=True)
c, shuffles = get_function(c.strip(), "(SHIFT|SHUFFLE)", ["0", "default", "default"])
r = c
for s in shuffles:
r = shuffle_chunk(s, r)
@@ -168,9 +181,10 @@ def tokenize(clip, text, can_break, empty_tokens):
per_te_prompts = {}
if l_prompts:
log.warning("Note: CLIP_L is deprecated. Use TE(l=prompt) instead")
per_te_prompts["l"] = l_prompts
per_te_prompts["l"] = [x.args for x in l_prompts]
for prompt in te_prompts:
prompt = prompt.args[0]
if prompt.strip() == "help":
log.info("Encoders available for TE: %s", ", ".join(tokens.keys()))
continue
@@ -211,7 +225,7 @@ def encode_prompt_segment(
default_style="comfy",
default_normalization="none",
clip_weights=None,
) -> list[tuple[torch.Tensor, dict[str]]]:
) -> list[ComfyConditioning]:
style, normalization, text = get_style(text, default_style, default_normalization)
clip_weights, text = get_clipweights(text, clip_weights)
text, cuts = parse_cuts(text)
@@ -231,19 +245,18 @@ def encode_prompt_segment(
# Chunks to ConditioningAverage:
text, averages = split_by_function(text, "AVG", ["0.5"])
text, averages = split_by_function(text, "AVG", ["0.5"], require_args=False)
prompts_to_avg = []
for avg in averages:
w = safe_float(avg["args"][0], 0.5)
for chunk, avg in averages:
w = safe_float(avg.args[0], 0.5)
prompts_to_avg.append((text, w))
text = avg["text"]
text = chunk
prompts_to_avg.append((text, 1.0))
conds_to_avg = []
for prompt, weight in prompts_to_avg:
conds_to_cat = []
chunks = re.split(r"\bCAT\b", prompt)
for c in chunks:
for c in split_quotable(prompt, r"\bCAT\b"):
tokens = tokenize(clip, c, can_break, empty)
conds_to_cat.append(clip.encode_from_tokens_scheduled(tokens, add_dict=settings))
@@ -264,13 +277,22 @@ def encode_prompt_segment(
w = next_w
continue
for i in range(len(base)):
(cond,) = ConditioningAverage.addWeighted(None, [base[i]], [cond[i]], w)
(cond,) = call_node(ConditioningAverage, [base[i]], [cond[i]], w)
base[i] = cond[0]
w = next_w
return base
def calc_w(tensor, w):
if math.isclose(w, 0):
return torch.zeros_like(tensor)
elif math.isclose(w, 1.0):
return tensor
else:
return tensor * w
def apply_weights(output, te_name, spec):
"""Applies weights to TE outputs"""
if not spec:
@@ -282,7 +304,7 @@ def apply_weights(output, te_name, spec):
default = spec.get("all", None)
if isinstance(output, tuple):
out, pooled = output
out, pooled, *extra = output
pkey = te_name + "_pooled"
if te_name in spec or pkey in spec or default is not None:
w = spec.get(te_name, default)
@@ -292,16 +314,16 @@ def apply_weights(output, te_name, spec):
if pooled_w is None:
pooled_w = 1.0
log.info("Weighting %s output by %s, pooled by %s", te_name, w, pooled_w)
out = out * w
out = calc_w(out, w)
if pooled is not None:
pooled = pooled * pooled_w
pooled = calc_w(pooled, pooled_w)
return out, pooled
return (out, pooled) + tuple(extra)
else:
if te_name in spec or default is not None:
w = spec.get(te_name, default)
log.info("Weighting %s output by %s", te_name, w)
output = output * w
output = calc_w(output, w)
return output
@@ -356,7 +378,7 @@ def get_area(text):
if not areas:
return text, None
args = areas[0]
args = areas[0].args
x, w = parse_floats(args[0], [0.0, 1.0], split_re="\\s+")
y, h = parse_floats(args[1], [0.0, 1.0], split_re="\\s+")
weight = safe_float(args[2], 1.0)
@@ -383,7 +405,7 @@ def get_mask_size(text, defaults):
text, sizes = get_function(text, "MASK_SIZE", ["512", "512"])
if not sizes:
return text, (defaults.get("mask_width", 512), defaults.get("mask_height", 512))
w, h = sizes[0]
w, h = sizes[0].args
return text, (int(w), int(h))
@@ -429,46 +451,53 @@ def get_mask(text, size, input_masks):
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]
mask = call_node(FeatherMask, 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)
totalweight = safe_float(maskw[0].args[0], 1.0)
i = 0
for m in masks:
weight = safe_float(m[2], 1.0)
op = m[3]
nextmask = make_mask(m, size, weight)
weight = safe_float(m.args[2], 1.0)
op = m.args[3]
nextmask = make_mask(m.args, size, weight)
if i < len(feathers):
nextmask = feather(feathers[i], nextmask)
nextmask = feather(feathers[i].args, nextmask)
i += 1
if mask is not None:
log.info("MaskComposite op=%s", op)
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
mask = call_node(MaskComposite, mask, nextmask, 0, 0, op)[0]
else:
mask = nextmask
for idx, w, op in imasks:
for im in imasks:
idx, w, op = im.args
idx = int(safe_float(idx, 0.0))
w = safe_float(w, 1.0)
if input_masks is None:
log.warn(
"IMASK requires you to attach custom masks to the CLIP object using PCAddMasksToClIP before using it"
)
input_masks = []
if len(input_masks) < idx + 1:
log.warn("IMASK index %s not found, ignoring...", idx)
continue
nextmask = input_masks[idx] * w
if i < len(feathers):
nextmask = feather(feathers[i], nextmask)
nextmask = feather(feathers[i].args, nextmask)
i += 1
if mask is not None:
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
mask = call_node(MaskComposite, 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)
mask = feather(f.args, mask)
return text, mask, totalweight
@@ -483,14 +512,15 @@ def get_noise(text):
return text, None, None
w = 0
# Only take seed from first noise spec, for simplicity
seed = safe_float(noises[0][1], "none")
seed = noises[0].args[0].strip()
if seed == "none":
gen = None
else:
seed = safe_float(seed, 0)
gen = torch.Generator()
gen.manual_seed(int(seed))
for n in noises:
w += safe_float(n[0], 0.0)
w += safe_float(n.args[0], 0.0)
return text, max(min(w, 1.0), 0.0), gen
@@ -551,7 +581,7 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
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)]
prompts = list(split_quotable(text, r"\bAND\b"))
p, sdxl_opts = get_sdxl(prompts[0], defaults)
prompts[0] = p
@@ -573,9 +603,9 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
return f"MASK({args})"
for prompt in prompts:
base_prompt, attn_couple_prompts = split_by_function(prompt, "COUPLE", defaults=None)
base_prompt, attn_couple_prompts = split_by_function(prompt, "COUPLE", defaults=None, require_args=False)
prompts = [base_prompt] + [couple_mask(p["args"]) + p["text"] for p in attn_couple_prompts]
prompts = [base_prompt] + [couple_mask(f.args) + chunk for (chunk, f) in attn_couple_prompts]
encoded = []
for p in prompts:
p, settings = process_settings(p, defaults, masks, mask_size, sdxl_opts)
+38 -1
View File
@@ -14,7 +14,10 @@ logging.basicConfig()
def run(f, *args):
return getattr(f, f.FUNCTION)(*args)
if hasattr(f, "execute"):
return f.execute(*args)
else:
return getattr(f, f.FUNCTION)(*args)
@mock.patch("torch.cuda.current_device", lambda: "cpu")
@@ -80,6 +83,17 @@ class TestEncode(unittest.TestCase):
c = c2 # Used in later tests
self.condEqual(c1, c2)
with self.subTest("Quotes"):
(c1,) = run(pc, clip, 'Text saying "DOG MASK AND CAT COUPLE MASK(X)"')
(c2,) = run(comfy, clip, 'Text saying "DOG MASK AND CAT COUPLE MASK(X)"')
self.condEqual(c1, c2)
with self.subTest("Function cornercase"):
(c1,) = run(pc, clip, "test SDXL function")
(c2,) = run(comfy, clip, "test SDXL function")
(c3,) = run(pc, clip, "test SDXL() function")
self.condEqual(c1, c2)
with self.subTest("Weights"):
(c1,) = run(pc, clip, "(test:1.2) (test:0.6)")
(c2,) = run(comfy, clip, "(test:1.2) (test:0.6)")
@@ -104,8 +118,31 @@ class TestEncode(unittest.TestCase):
(c1,) = run(comfy, clip, "test1")
(c2,) = run(comfy, clip, "test2")
(c3,) = run(pc, clip, "test1 AVG() test2")
(c4,) = run(pc, clip, "test1 AVG test2")
(avg,) = run(average, c1, c2, 0.5)
self.condEqual(avg, c3)
self.condEqual(avg, c4)
with self.subTest("Average multi"):
(c1,) = run(comfy, clip, "test1")
(c2,) = run(comfy, clip, "test2")
(c3,) = run(comfy, clip, "test3")
(c4,) = run(pc, clip, "test1 AVG() test2 AVG() test3")
(c5,) = run(pc, clip, "test1 AVG test2 AVG test3")
(avg1,) = run(average, c1, c2, 0.5)
(avg,) = run(average, avg1, c3, 0.5)
self.condEqual(avg, c4)
self.condEqual(avg, c5)
@unittest.expectedFailure
def test_failure(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
for k, clip in clips:
with self.subTest(k):
(c1,) = run(comfy, clip, "test SDXL function")
(c2,) = run(pc, clip, "test SDXL() function")
self.condEqual(c1, c2)
def test_weight(self):
pc = PCTextEncode()
-71
View File
@@ -1,71 +0,0 @@
import unittest
import numpy.testing as npt
from os import environ
clips = []
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 clips:
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")
from comfy.sd import load_clip
from .nodes_base import PCTextEncode
from pathlib import Path
to_test = environ.get("TEST_TE", "clip_l").split()
model_path = environ.get("COMFYUI_MODEL_ROOT", ".")
te_root = (Path(model_path) / "text_encoders").resolve()
if "clip_l" in to_test:
clip_l = load_clip(
ckpt_paths=[str(te_root / "clip_l.safetensors")], clip_type="stable_diffusion", model_options={}
)
clips.append(("clip_l", clip_l))
if "t5" in to_test:
dual = load_clip(
[str(te_root / "clip_l.safetensors"), str(te_root / "t5xxl_fp16.safetensors")],
clip_type="flux",
model_options={},
)
clips.append(("clip_l+t5", dual))
print("Starting tests")
unittest.main()
+7
View File
@@ -18,6 +18,13 @@ class TestParser(unittest.TestCase):
self.assertEqual(p.at_step(0.5), expected)
self.assertEqual(p.at_step(1), expected)
def test_quote(self):
p = parse('This is a text with a "QUOTED DEF(X=Y)"')
expected = prompt(1.0, 'This is a text with a "QUOTED DEF(X=Y)"')
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]"]],
+117 -46
View File
@@ -1,20 +1,48 @@
from __future__ import annotations
from pathlib import Path
import re
import logging
import copy
from dataclasses import dataclass
from typing import Any, TypeAlias, Iterator, TypeVar, TYPE_CHECKING
if TYPE_CHECKING:
import torch # flakes8: noqa
FunctionArgs: TypeAlias = list[str]
ComfyConditioning: TypeAlias = tuple["torch.Tensor", dict[str, Any]]
@dataclass
class FunctionSpec:
name: str
args: FunctionArgs
position: int
placeholder: str | None
# Allow testing
try:
from folder_paths import get_filename_list
except ImportError:
def get_filename_list(x):
raise NotImplementedError("How did you get here?")
def get_filename_list(folder_name) -> list[str]:
return []
log = logging.getLogger("comfyui-prompt-control")
def call_node(cls, *args, **kwargs):
if hasattr(cls, "execute"):
# v3 node
return cls.execute(*args, **kwargs)
else:
func = getattr(cls(), cls.FUNCTION)
return func(*args, **kwargs)
def consolidate_schedule(prompt_schedule):
prev_loras = {}
not_found = []
@@ -55,10 +83,11 @@ def find_nonscheduled_loras(consolidated_schedule):
return {k: v for (k, v) in candidate_loras.items() if k not in to_remove}
def smarter_split(separator, string):
def smarter_split(separator: str, string: str) -> list[str]:
"""Does not break () when splitting"""
splits = []
prev = 0
idx = 0
stack = 0
escape = False
for idx, x in enumerate(string):
@@ -75,7 +104,7 @@ def smarter_split(separator, string):
return splits
def find_closing_paren(text, start):
def find_closing_paren(text: str, start: int) -> int:
stack = 1
for i, char in enumerate(text[start:]):
if char == ")":
@@ -84,90 +113,126 @@ def find_closing_paren(text, start):
stack += 1
if stack == 0:
return start + i
# Implicit closing paren after end
return len(text)
return -1
def get_function(text, func, defaults, return_func_name=False, placeholder="", return_dict=False):
rex = re.compile(rf"\b{func}\b", re.MULTILINE)
instances = []
def find_function_spans(
text: str, func: str, require_args: bool, defaults: FunctionArgs | None
) -> Iterator[tuple[int, int, str, FunctionArgs]]:
if require_args:
rex = re.compile(rf"\b{func}\(", re.MULTILINE)
else:
rex = re.compile(rf"\b{func}\b", re.MULTILINE)
idx = 0
match = rex.search(text)
count = 0
while match:
# Match start, content start
start, at_paren = match.span()
if require_args:
at_paren = at_paren - 1
funcname = text[start:at_paren]
after_first_paren = at_paren + 1
if text[at_paren:after_first_paren] == "(":
end = find_closing_paren(text, after_first_paren)
if end < 0:
continue
args = parse_strings(text[after_first_paren:end], defaults)
end += 1
else:
end = at_paren
args = None
args = defaults or []
yield idx + start, idx + end, funcname, args
idx = idx + end
text = text[end:]
match = rex.search(text)
def get_function(
text: str, func: str, defaults: list[str] | None, placeholder: str = "", require_args: bool = True
) -> tuple[str, list[FunctionSpec]]:
spans = [x.span() for x in re.finditer(r'".+?"', text)]
instances = []
count = 0
chunks = []
current = 0
skipped = 0
for start, end, funcname, args in find_function_spans(text, func, require_args, defaults):
ph = None
if spans_include(spans, start, end):
continue
if placeholder:
ph = f"\0{placeholder}{count}\0"
if return_dict:
instances.append(
{
"name": funcname,
"args": args,
"position": start,
"placeholder": ph,
}
)
elif return_func_name:
instances.append((funcname, args))
else:
instances.append(args)
if placeholder:
text = text[:start] + f"\0{placeholder}{count}\0" + text[end:]
else:
text = text[:start] + text[end:]
match = rex.search(text)
instances.append(FunctionSpec(funcname, args, start - skipped, ph))
skipped += end - start
chunks.append(text[current:start] + (ph or ""))
current = end
count += 1
chunks.append(text[current:])
text = "".join(chunks)
return text, instances
def split_by_function(text, func, defaults=None):
def spans_include(spans: list[tuple[int, int]], s: int, e: int) -> bool:
return any((s > a and e < b) for a, b in spans)
def split_quotable(text: str, regexp: str) -> Iterator[str]:
start_from = 0
spans = [x.span() for x in re.finditer(r'".+?"', text)]
for x in re.finditer(regexp, text):
s, e = x.span()
if not spans_include(spans, s, e):
yield text[start_from:s].strip()
start_from = e
yield text[start_from:].strip()
def split_by_function(
text: str, func: str, defaults: list[str] | None = None, require_args: bool = True
) -> tuple[str, list[tuple[str, FunctionSpec]]]:
"""
Splits a string by function calls, returning the text preceding the first call and a list of dictionaries with a "text" key with the prompt before the next split or until hthe end of the text.
Splits a string by function calls, returning the leftover text along with a list of functions with their associated text chunk.
"""
text, functions = get_function(text, func, defaults, return_dict=True)
text, functions = get_function(text, func, defaults, require_args=require_args)
chunks = []
prev = 0
for f in functions:
chunks.append(text[prev : f["position"]])
prev = f["position"]
chunks.append(text[prev : f.position])
prev = f.position
chunks.append(text[prev:])
r = []
for i, f in enumerate(functions):
f["text"] = chunks[i + 1]
return chunks[0], functions
r.append((chunks[i + 1], f))
return chunks[0], r
def parse_args(strings, arg_spec, strip=True):
T = TypeVar("T")
def parse_args(strings: list[str], arg_spec: list[tuple[Any, T]], strip: bool = True) -> list[T]:
args = [s[1] for s in arg_spec]
for i, spec in list(enumerate(arg_spec))[: len(strings)]:
try:
if strip:
strings[i] = strings[i].strip()
args[i] = spec[0](strings[i])
f = spec[0]
args[i] = f(strings[i])
except ValueError:
pass
return args
def parse_floats(string, defaults, split_re=","):
def parse_floats(string: str, defaults: list[float], split_re: str = ",") -> list[float]:
spec = [(float, d) for d in defaults]
return parse_args(re.split(split_re, string.strip()), spec)
def parse_strings(string, defaults, split_re=r"(?<!\\),", replace=(r"\,", ",")):
def parse_strings(
string: str, defaults: FunctionArgs | None, split_re: str = r"(?<!\\),", replace: tuple[str, str] = (r"\,", ",")
) -> FunctionArgs:
if defaults is None:
return string
spec = [(lambda x: x, d) for d in defaults]
return [string]
spec = [(str, d) for d in defaults]
splits = re.split(split_re, string)
if replace:
f, t = replace
@@ -175,7 +240,7 @@ def parse_strings(string, defaults, split_re=r"(?<!\\),", replace=(r"\,", ",")):
return parse_args(splits, spec, strip=False)
def safe_float(f, default):
def safe_float(f: Any, default: float) -> float:
if f is None:
return default
try:
@@ -184,7 +249,7 @@ def safe_float(f, default):
return default
def lora_name_to_file(name):
def lora_name_to_file(name: str) -> str | None:
filenames = get_filename_list("loras")
# Return exact matches as is
if name in filenames:
@@ -195,6 +260,12 @@ def lora_name_to_file(name):
p = Path(f).with_suffix("")
if p.name == n or str(p) == n:
return f
# Finally, try to find unique match from parts
parts = name.split()
search = [f for f in filenames if all(p in f for p in parts)]
if len(search) == 1:
return search[0]
return None
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-prompt-control"
description = "Provides nodes for prompt editing and LoRA scheduling, advanced regional prompting (including attention masking) and more, all controlled through your text prompt"
version = "2.0.0-rc.9"
version = "2.1.2"
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"]