Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
68766215f2 | ||
|
|
287f554a68 | ||
|
|
71465f914c | ||
|
|
d7be7bc29e | ||
|
|
329d4cf95f | ||
|
|
c52ace71aa | ||
|
|
4806cf5959 | ||
|
|
a0ab709f50 | ||
|
|
66fef1ffe8 | ||
|
|
3341e9f81e | ||
|
|
aa6b4608f0 |
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -25,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)
|
||||
|
||||
@@ -100,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
|
||||
)
|
||||
@@ -131,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):
|
||||
@@ -257,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))
|
||||
@@ -288,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)
|
||||
@@ -348,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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -451,11 +455,13 @@ def expand_macros(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)
|
||||
|
||||
+54
-33
@@ -1,3 +1,4 @@
|
||||
from __future__ import annotations
|
||||
import logging
|
||||
import re
|
||||
import torch
|
||||
@@ -6,7 +7,17 @@ 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, call_node
|
||||
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
|
||||
@@ -26,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+")
|
||||
@@ -47,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:
|
||||
@@ -63,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:
|
||||
@@ -78,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":
|
||||
@@ -129,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)
|
||||
@@ -169,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
|
||||
@@ -212,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)
|
||||
@@ -234,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))
|
||||
|
||||
@@ -292,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)
|
||||
@@ -306,7 +318,7 @@ def apply_weights(output, te_name, spec):
|
||||
if pooled is not None:
|
||||
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)
|
||||
@@ -366,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)
|
||||
@@ -393,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))
|
||||
|
||||
|
||||
@@ -446,14 +458,14 @@ def get_mask(text, size, input_masks):
|
||||
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)
|
||||
@@ -461,15 +473,22 @@ def get_mask(text, size, input_masks):
|
||||
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 = call_node(MaskComposite, mask, nextmask, 0, 0, op)[0]
|
||||
@@ -478,7 +497,7 @@ def get_mask(text, size, input_masks):
|
||||
|
||||
# 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
|
||||
|
||||
@@ -493,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
|
||||
|
||||
|
||||
@@ -561,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
|
||||
@@ -578,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)
|
||||
|
||||
@@ -20,6 +20,12 @@ def run(f, *args):
|
||||
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")
|
||||
class TestEncode(unittest.TestCase):
|
||||
@classmethod
|
||||
@@ -83,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")
|
||||
@@ -118,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()
|
||||
@@ -156,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()
|
||||
|
||||
@@ -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]"]],
|
||||
|
||||
+96
-45
@@ -1,15 +1,34 @@
|
||||
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")
|
||||
@@ -64,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):
|
||||
@@ -84,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 == ")":
|
||||
@@ -93,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
|
||||
@@ -114,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
|
||||
@@ -189,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:
|
||||
@@ -198,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:
|
||||
|
||||
+1
-1
@@ -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.1.0"
|
||||
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"]
|
||||
|
||||
Reference in New Issue
Block a user