Compare commits

..
7 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
6 changed files with 147 additions and 89 deletions
+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)
+9 -8
View File
@@ -427,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)
@@ -452,15 +454,14 @@ 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, return_dict=True
)
text, defns = get_function(text, name, defaults=None, placeholder=f"DEFNCALL{name}", require_args=False)
for i, d in enumerate(defns):
ph = d["placeholder"]
parameters = d["args"]
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)
+46 -31
View File
@@ -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, split_quotable
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 = split_quotable(text, r"\bBREAK\b")
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 = split_quotable(prompt, r"\bCAT\b")
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))
@@ -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,7 +473,8 @@ 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:
@@ -475,7 +488,7 @@ def get_mask(text, size, input_masks):
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]
@@ -484,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
@@ -499,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
@@ -567,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 = split_quotable(text, r"\bAND\b")
prompts = list(split_quotable(text, r"\bAND\b"))
p, sdxl_opts = get_sdxl(prompts[0], defaults)
prompts[0] = p
@@ -584,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)
+27
View File
@@ -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
@@ -123,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()
@@ -161,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()
+63 -47
View File
@@ -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 == ")":
@@ -96,7 +116,9 @@ def find_closing_paren(text, start):
return -1
def find_function_spans(text, func, require_args, defaults):
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:
@@ -113,49 +135,36 @@ def find_function_spans(text, func, require_args, defaults):
if text[at_paren:after_first_paren] == "(":
end = find_closing_paren(text, after_first_paren)
if end < 0:
print("no closing paren:", text)
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, func, defaults, return_func_name=False, placeholder="", return_dict=False, require_args=True):
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:
chunks.append(text[current:start] + f"\0{placeholder}{count}\0")
else:
chunks.append(text[current:start])
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:])
@@ -163,60 +172,67 @@ def get_function(text, func, defaults, return_func_name=False, placeholder="", r
return text, instances
def spans_include(spans, s, e):
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, regexp):
res = []
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):
res.append(text[start_from:s].strip())
yield text[start_from:s].strip()
start_from = e
res.append(text[start_from:].strip())
return res
yield text[start_from:].strip()
def split_by_function(text, func, defaults=None, require_args=True):
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
@@ -224,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:
@@ -233,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
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.1.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"]