Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d7be7bc29e | ||
|
|
329d4cf95f | ||
|
|
c52ace71aa | ||
|
|
4806cf5959 |
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
+43
-29
@@ -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
|
||||
@@ -591,7 +605,7 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
|
||||
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)
|
||||
|
||||
@@ -123,6 +123,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()
|
||||
|
||||
+63
-47
@@ -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
@@ -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.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"]
|
||||
|
||||
Reference in New Issue
Block a user