Compare commits

..
19 Commits
Author SHA1 Message Date
asagi4 68766215f2 v2.1.3 2026-02-04 16:32:26 +02:00
asagi4 287f554a68 Test COUPLE mask shortcut 2026-01-16 18:20:46 +02:00
asagi4 71465f914c Properly parse COUPLE(), see #134 2026-01-16 18:05:30 +02:00
asagi4 d7be7bc29e v2.1.2 2026-01-13 21:50:09 +02:00
asagi4 329d4cf95f Add a test for #133 2026-01-13 21:48:34 +02:00
asagi4 c52ace71aa #133 properly set function locations in get_function 2026-01-13 20:45:10 +02:00
asagi4 4806cf5959 refactor get_function to make it more consistent 2026-01-13 20:45:10 +02:00
asagi4 a0ab709f50 v2.1.1 2025-12-14 17:57:53 +02:00
asagi4 66fef1ffe8 Avoid splitting AND, CAT and others when inside quotes
See #132
2025-12-05 21:25:28 +02:00
asagi4 3341e9f81e Fix TE_WEIGHT failing with some encoders
See #131

Prompt weighting and attention couple will probably not work, but
this should prevent exceptions.
2025-12-03 16:07:31 +02:00
asagi4 aa6b4608f0 Clarify attention couple docs a bit and add a warning if IMASK is used without attaching custom masks
See #108
2025-12-01 13:19:29 +02:00
asagi4 d1cc60b00a v2.1.0 2025-11-20 20:20:22 +02:00
asagi4 efe8939250 Fix issue with T5 encoding sometimes returning NaNs in test 2025-11-20 20:14:40 +02:00
asagi4 39b353b916 Fix tests for #130 2025-11-20 20:14:37 +02:00
asagi4 485bc7f2ab Prepare for v3 conversion of imported nodes, see #12
This should prevent things from breaking, but it needs a bit of testing.
2025-11-20 15:29:33 +02:00
asagi4 1d03ded9dd Find LoRAs with partial match 2025-11-20 15:23:51 +02:00
asagi4 94a4d076e0 Remove unused attributes 2025-08-30 13:46:00 +03:00
asagi4 228dc4b22b Remove the old advanced encoding implementation 2025-08-27 20:27:30 +03:00
asagi4 db523e1f16 Remove dead code 2025-08-27 20:23:41 +03:00
14 changed files with 280 additions and 454 deletions
+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)
+13 -14
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}", require_args=False)
for i, parameters in enumerate(defns):
ph = f"\0DEFNCALL{name}{i}\0"
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)
+71 -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)
@@ -233,17 +247,16 @@ def encode_prompt_segment(
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
@@ -568,14 +598,15 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
return c
def couple_mask(args):
if args is None:
assert len(args) <= 1, "Argument parsing failure. This is a bug in Prompt Control"
if not args:
return ""
return f"MASK({args})"
return f"MASK({args[0]})"
for prompt in prompts:
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)
+36 -1
View File
@@ -14,7 +14,16 @@ 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)
def compare_hookgroup_mask(h1, h2):
assert len(h1.hooks) == len(h2.hooks)
for a, b in zip(h1.hooks, h2.hooks):
assert (a.mask == b.mask).all()
@mock.patch("torch.cuda.current_device", lambda: "cpu")
@@ -80,6 +89,11 @@ 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")
@@ -115,6 +129,17 @@ class TestEncode(unittest.TestCase):
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()
@@ -153,6 +178,16 @@ class TestEncode(unittest.TestCase):
(c2,) = run(pc, clip, "test COUPLE prompt1 COUPLE test2 COUPLE prompt2")
self.assertTrue(len(c) == 2)
self.assertTrue(len(c2) == 1)
with self.subTest(f"Testing {k} mask shortcut"):
(c,) = run(pc, clip, "test COUPLE() prompt1")
(c2,) = run(pc, clip, "test COUPLE MASK() prompt1")
self.condEqual(c, c2)
self.condEqual(c, c2, "hooks", compare_hookgroup_mask)
with self.subTest(f"Testing {k} mask shortcut 2"):
(c,) = run(pc, clip, "test COUPLE(0 0.2, 0.5) prompt1")
(c2,) = run(pc, clip, "test COUPLE MASK(0 0.2, 0.5) prompt1")
self.condEqual(c, c2)
self.condEqual(c, c2, "hooks", compare_hookgroup_mask)
def test_styles(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]"]],
+111 -45
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,20 +113,20 @@ 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, require_args=True):
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)
instances = []
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
@@ -105,74 +134,105 @@ def get_function(text, func, defaults, return_func_name=False, placeholder="", r
after_first_paren = at_paren + 1
if text[at_paren:after_first_paren] == "(":
end = find_closing_paren(text, after_first_paren)
if end < 0:
continue
args = parse_strings(text[after_first_paren:end], defaults)
end += 1
else:
end = at_paren
args = defaults
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, require_args=True):
def spans_include(spans: list[tuple[int, int]], s: int, e: int) -> bool:
return any((s > a and e < b) for a, b in spans)
def split_quotable(text: str, regexp: str) -> Iterator[str]:
start_from = 0
spans = [x.span() for x in re.finditer(r'".+?"', text)]
for x in re.finditer(regexp, text):
s, e = x.span()
if not spans_include(spans, s, e):
yield text[start_from:s].strip()
start_from = e
yield text[start_from:].strip()
def split_by_function(
text: str, func: str, defaults: list[str] | None = None, require_args: bool = True
) -> tuple[str, list[tuple[str, FunctionSpec]]]:
"""
Splits a string by function calls, returning the 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, require_args=require_args)
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
@@ -180,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:
@@ -189,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:
@@ -200,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.1"
version = "2.1.3"
license = { file = "LICENSE" }
# some lark versions older than 1.1.9 apparently have a bug that breaks things, see https://github.com/asagi4/comfyui-prompt-control/issues/35
dependencies = ["lark >= 1.1.9"]