Compare commits

...
Author SHA1 Message Date
asagi4 88af041dce v3.0.0-beta.10 2026-09-15 20:34:58 +03:00
asagi4 29b21f735d Remove ComfyUI caching from known issues
It seems to be fixed now. Schedules appear to be properly lazy.
2026-09-15 20:31:50 +03:00
asagi4 8348a56f32 Fix tag scheduling (perhaps the dumbest bug ever) 2026-09-15 20:29:46 +03:00
asagi4 314fd39ed4 Test for a bug with tags 2026-09-15 20:29:03 +03:00
asagi4 62dde6e795 Fix infinite loop cornercase with unterminated functions 2026-09-05 15:42:37 +03:00
asagi4 66c88ea137 Warn user if LoRA search is ambiguous 2026-09-05 15:37:01 +03:00
asagi4 61ba7fe92e Fix typo in proxy 2026-09-05 15:37:01 +03:00
asagi4 ce63e1d83c Add tests for utils 2026-09-05 15:36:58 +03:00
asagi4 af293bb7cb Clarify comment 2026-09-05 14:24:41 +03:00
asagi4 53d210953a Avoid breaking scheduling syntax in NODE helper 2026-09-02 20:51:28 +03:00
asagi4 2265767d80 Remove old lark parser 2026-08-25 23:37:39 +03:00
asagi4 39add37c2f Use a callable function instead of replacing placeholders 2026-08-25 23:32:21 +03:00
asagi4 c2fa8071b5 Expose TE hooks as an internal-only node 2026-08-25 23:01:28 +03:00
asagi4 953dac842a Add dev node for a single lazy expanded prompt 2026-08-25 23:01:25 +03:00
asagi4 56fbfac2a5 Allow predefined macros 2026-08-25 22:47:43 +03:00
asagi4 ddee2b9a63 Split out prompt building from scheduling 2026-08-25 22:47:38 +03:00
asagi4 a41d433719 Add FILTER and COMBINE (Experimental)
See #151
Two experimental functions, undocumented until I'm sure I like them.

extra args work like in NODE.

Examples:
DEF(AND=COMBINE(ConditioningCombine, conditioning_1, conditioning_2))
DEF(MUL=FILTER(ConditioningMultiply, conditioning, multiplier $1))
FILTER(SetReferenceLatent, conditioning, latent ["123", 0])
2026-08-24 19:57:52 +03:00
asagi4 b24d93b778 v3.0.0-beta.9 2026-08-22 18:56:51 +03:00
asagi4 c8f925c4ea Lazify NODE helper 2026-08-22 18:53:02 +03:00
asagi4 1a733e71af Make NODE helper autogrow 2026-08-22 18:52:58 +03:00
asagi4 e36eff4356 Note whitespace change in README 2026-08-20 16:08:09 +03:00
asagi4 c6137ddc49 Strip whitespace from scheduled prompts by default 2026-08-20 16:04:01 +03:00
asagi4 64417230a3 SEGs need to be expanded *before* NODE is evaluated 2026-08-19 23:07:12 +03:00
asagi4 0a2ceb94e9 doc: Note caveat about missing functionality when using NODE 2026-08-19 14:02:21 +03:00
asagi4 6e3fce9dcb v3.0.0-beta.8 2026-08-19 13:42:56 +03:00
asagi4 b581cf7f24 NODE: Fix SEG expansion
See: #150
2026-08-19 13:33:46 +03:00
asagi4 d7a992b96d NODE: Test for SEG expansion 2026-08-19 13:32:04 +03:00
asagi4 19123570e1 Fix parsing of < and > in schedules, see #149 2026-08-16 17:17:38 +03:00
asagi4 e6bb57cd25 Formatting 2026-08-14 18:13:49 +03:00
18 changed files with 935 additions and 1246 deletions
+3 -2
View File
@@ -11,6 +11,9 @@ A `Basic Text to Image` template is included with the extension, and can be load
> The parser was rewritten using parsy. It is intended to have the same behaviour as the old parser, but is **significantly** faster.
> Please report any bugs or incompatibilities you find.
## Notable changes
- `PC: Schedule Prompt` now strips surrounding whitespace by default, which may change some prompts. Add `NOSTRIP()` to your prompt to restore previous behaviour.
## What can it do?
@@ -77,6 +80,4 @@ This node configures `PCTextEncode` default values for some functions by attachi
# Known issues
- ComfyUI's caching mechanism has an issue that makes it unnecessarily invalidate caches for certain inputs; you'll still get some benefit from the lazy nodes, but changing inputs that shouldn't affect downstream nodes (especially if using filtering) will still cause them to be recomputed because ComfyUI doesn't realize the inputs haven't changed.
- Cutoff does not work with models that use non-CLIP text encoders, like Flux. This might be fixable, but it's uncertain if cutoff even makes sense for those models.
+3
View File
@@ -2,6 +2,9 @@
The `NODE` function allows you to use any other text encoding node within `PC: Schedule Prompt`, replacing the default `PCTextEncode` and allowing for example video model scheduling.
> [!NOTE]
> When using NODE, you lose access to *all* special syntax provided by `PCTextEncode`. Only SEGs, macros and scheduling will continue to work since those are processed at graph expansion time before the text prompt is passed into the node.
## Basic usage
Use `NODE(NodeClassName, textinputname)` in a prompt to generate a graph using any node that's compatible. The requirements are as follows:
+3
View File
@@ -14,6 +14,9 @@ Besides the syntax documented below, the [basic syntax](/doc/basic.md) and [prom
a [large::0.1] [cat|dog:0.05] [<lora:somelora:0.5:0.6>::0.5]
[in a park:in space:0.4]
```
## Note on whitespace
`PC: Schedule Prompt` will strip leading and following whitespace from the prompt automatically. If you really want whitespace in your prompt, include `NOSTRIP()` in your prompt.
## Comments and escaping
In schedules, any text on a line following a `#` is considered a comment and removed, including the `#` character.
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -52,7 +52,7 @@ class Proxy:
return self
def __call__(self, *args, **kwargs):
return self.function(*args, *kwargs)
return self.function(*args, **kwargs)
class AttentionCoupleHook(TransformerOptionsHook):
+12 -9
View File
@@ -69,8 +69,12 @@ def parse_search(search):
return name, args
def expand_macros(text):
text, defs = get_function(text, "DEF", defaults=None)
def expand_macros(text, defs=None):
silent = False
if defs is None:
text, defs = get_function(text, "DEF", defaults=None)
else:
silent = True
res = text
prevres = text
replacements = []
@@ -95,7 +99,8 @@ def expand_macros(text):
prevres = res
if res.strip() != text.strip():
res = res.strip()
log.debug("DEFs expanded to: %s", res)
if not silent:
log.debug("DEFs expanded to: %s", res)
return res
@@ -108,11 +113,8 @@ def substitute_var(text, name, replace, boundary=r"\b"):
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, d in enumerate(defns):
ph = d.placeholder
assert ph is not None, "This is a bug"
parameters = d.args
def run_macro(*parameters):
paramvals = []
if parameters:
paramvals = [x.strip() for x in parameters[0].split(";")]
@@ -123,6 +125,7 @@ def substitute_defcall(text, search, replace):
for i, v in enumerate(default_args):
r = substitute_var(r, i + 1, v, boundary=end_re)
return r
text = text.replace(ph, r)
text, _ = get_function(text, name, defaults=None, processor=run_macro, require_args=False)
return text
+29 -5
View File
@@ -3,7 +3,7 @@ import logging
from comfy_api.latest import io
from .macros import expand_segs
from .prompts import encode_prompt
from .prompts import encode_prompt, hook_te
log = logging.getLogger("comfyui-prompt-control")
@@ -56,7 +56,31 @@ class PCTextEncode(io.ComfyNode):
return PCTextEncodeWithRange.execute(clip, text, 0.0, 1.0)
NODES = [
PCTextEncodeWithRange,
PCTextEncode,
]
class PCHookEncoderModsInternal(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="PCHookTextEncoderModsInternal",
display_name="PC: Apply Text Encoder Mods",
category="promptcontrol",
description="Apply TE modifications (internal)",
is_experimental=True,
is_dev_only=True,
inputs=[
io.Clip.Input("clip"),
io.String.Input("te_names"),
io.String.Input("style"),
io.String.Input("normalization"),
io.Custom("PC_EXTRA_DATA").Input("extra", optional=True),
],
outputs=[io.Clip.Output()],
)
@classmethod
def execute(cls, clip, te_names, style, normalization, extra) -> io.NodeOutput:
te_names = [x.strip() for x in te_names.split(",")]
clip = hook_te(clip, te_names, style, normalization, extra)
return io.NodeOutput(clip)
NODES = [PCTextEncodeWithRange, PCTextEncode, PCHookEncoderModsInternal]
+135 -36
View File
@@ -10,7 +10,7 @@ from comfy_execution.graph_utils import GraphBuilder
from .macros import expand_macros, expand_segs
from .parser import parse_prompt_schedules
from .utils import consolidate_schedule, find_nonscheduled_loras, get_function
from .utils import consolidate_schedule, find_nonscheduled_loras, get_function, split_by_function
log = logging.getLogger("comfyui-prompt-control")
@@ -192,47 +192,106 @@ class PCLazyLoraLoader(io.ComfyNode):
return io.NodeOutput(*no.args[:2], expand=no.expand)
def parse_extra_inputs(args, defaults):
params = {}
if not args.strip():
return defaults + [{}]
defaults = defaults[:]
defaults.append("")
for i, v in enumerate(args.split(",", maxsplit=len(defaults) - 1)):
defaults[i] = v
# We should strip extra whitespace so that people don't have to worry about functions.
magic_spec = defaults[-1]
magic_spec.replace(r"\;", "__ESCAPED_SEMICOLON__")
extra_inputs = magic_spec.split(";") if magic_spec.strip() else []
for e in extra_inputs:
e = e.strip()
if not e:
continue
e = e.replace("__ESCAPED_SEMICOLON__", ";")
name, jsondata = e.split(maxsplit=1)
jsondata = jsondata.strip()
if not jsondata.strip():
continue
# From helper node:
if jsondata == "__EMPTY__":
continue
try:
params[name.strip()] = json.loads(jsondata.strip())
except ValueError as e:
raise ValueError(f"Invalid JSON input: '{jsondata}'") from e
return [x.strip() for x in defaults[:-1]] + [params]
def make_node(graph, p, clip, strip):
p, classnames = get_function(p, "NODE", defaults=None)
p, filters = get_function(p, "FILTER", defaults=None)
args = ""
if len(classnames) > 1:
log.warning("You have more than one NODE call in your prompt. Only the first one will be used")
if classnames:
args = classnames[0].args[0]
if not args.strip():
raise ValueError("NODE can't be empty!")
classname, paramname, extras = parse_extra_inputs(args, ["PCTextEncode", "text"])
# We should strip extra whitespace so that people don't have to worry about functions.
node = graph.node(classname.strip())
node.set_input("clip", clip)
node.set_input(paramname.strip(), p.strip() if strip else p)
for e, v in extras.items():
node.set_input(e, v)
for f in filters:
classname, paramname, extras = parse_extra_inputs(f.args[0], ["", "conditioning"])
if not classname:
raise ValueError("FILTER requires a Node class name")
extras[paramname] = node.out(0)
node = graph.node(classname)
for e, v in extras.items():
node.set_input(e, v)
return node
def build_prompt(graph, prompt, clip, start=None, end=None):
p = prompt
strip = "NOSTRIP()" not in p
p = p.replace("NOSTRIP()", "")
# Need to explicitly expand SEGs here *before* NODE is processed
p = expand_segs(p)
p, combines = split_by_function(p, "COMBINE")
current_cond = make_node(graph, p, clip, strip)
for text, f in combines:
classname, param1, param2, extra = parse_extra_inputs(f.args[0], ["", "conditioning_1", "conditioning_2"])
if classname.strip() == "":
raise ValueError("Can't use COMBINE without a class name")
combiner = graph.node(classname.strip())
c2 = make_node(graph, text, clip, strip)
extra[param1] = current_cond.out(0)
extra[param2] = c2.out(0)
for e, v in extra.items():
combiner.set_input(e, v)
current_cond = combiner
node = current_cond
if start is not None and end is not None:
node = graph.node("ConditioningSetTimestepRange")
node.set_input("conditioning", current_cond.out(0))
node.set_input("start", start)
node.set_input("end", end)
return node
def build_scheduled_prompts(graph, schedules, clip):
nodes = []
start_pct = 0.0
for end_pct, c in schedules:
p = c["prompt"]
p, classnames = get_function(p, "NODE", defaults=None)
realargs = ["PCTextEncode", "text", ""]
if classnames:
args = classnames[0].args[0]
if not args.strip():
raise ValueError("NODE can't be empty!")
for i, v in enumerate(args.split(",", maxsplit=2)):
realargs[i] = v
classname, paramname, magic_spec = realargs
node = graph.node(classname.strip())
node.set_input("clip", clip)
node.set_input(paramname.strip(), p)
magic_spec.replace(r"\;", "__ESCAPED_SEMICOLON__")
extra_inputs = magic_spec.split(";") if magic_spec.strip() else []
for e in extra_inputs:
e = e.strip()
if not e:
continue
e = e.replace("__ESCAPED_SEMICOLON__", ";")
name, jsondata = e.split(maxsplit=1)
jsondata = jsondata.strip()
if not jsondata.strip():
continue
# From helper node:
if jsondata == "__EMPTY__":
continue
try:
node.set_input(name.strip(), json.loads(jsondata.strip()))
except ValueError as e:
raise ValueError(f"Invalid JSON input: '{jsondata}'") from e
timestep = graph.node("ConditioningSetTimestepRange")
timestep.set_input("conditioning", node.out(0))
timestep.set_input("start", start_pct)
timestep.set_input("end", end_pct)
nodes.append(timestep)
node = build_prompt(graph, p, clip, start_pct, end_pct)
nodes.append(node)
start_pct = end_pct
node = nodes[0]
for othernode in nodes[1:]:
combiner = graph.node("ConditioningCombine")
@@ -296,9 +355,49 @@ class PCLazyTextEncode(io.ComfyNode):
return PCLazyTextEncodeAdvanced.execute(clip, text)
predefined_macros = get_function(
"""
DEF(AND=COMBINE(ConditioningCombine, conditioning_1, conditioning_2))
DEF(CAT=COMBINE(ConditioningConcat, conditioning_to, conditioning_from))
DEF(AVG(0.5)=COMBINE(ConditioningAverage, conditioning_from, conditioning_to, conditioning_to_strength $1))
""",
"DEF",
defaults=None,
)
class PCLazyTextEncodeSingle(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="PCLazyTextEncodeSingle",
display_name="PC: Prompt (without scheduling)",
is_experimental=True,
is_dev_only=True,
enable_expand=True,
category="promptcontrol",
inputs=[
io.Clip.Input("clip", raw_link=True),
io.String.Input("text", multiline=True, default=""),
],
outputs=[
io.Conditioning.Output("conditioning"),
],
)
@classmethod
def execute(cls, clip, text):
graph = GraphBuilder()
text = expand_macros(text, predefined_macros)
node = build_prompt(graph, text, clip)
g = graph.finalize()
return io.NodeOutput(node.out(0), expand=g)
NODES = [
PCLazyTextEncode,
PCLazyTextEncodeAdvanced,
PCLazyTextEncodeSingle,
PCLazyLoraLoader,
PCLazyLoraLoaderAdvanced,
]
+47 -9
View File
@@ -185,28 +185,66 @@ class PCMacroExpand(io.ComfyNode):
class PCLinkHelper(io.ComfyNode):
# a-z
NAMES = [chr(97 + i) for i in range(26)]
@classmethod
def define_schema(cls):
inputs: list = [io.String.Input("text", multiline=True)]
for x in "abcdefghijklmn":
inputs.append(io.AnyType.Input(x, optional=True, lazy=True, extra_dict={"rawLink": True}))
t1 = io.Autogrow.TemplateNames(io.AnyType.Input("link", raw_link=True), min=0, names=cls.NAMES)
t2 = io.Autogrow.TemplateNames(
io.AnyType.Input("value", lazy=True), min=0, names=[f"var{i + 1}" for i in range(50)]
)
return io.Schema(
node_id="PCNODELinkHelper",
display_name="PC: Extra argument helper for NODE",
category="promptcontrol/tools",
description="Takes in arbitrary inputs and renders them as NODE-compatible values, replacing $a -> $n with JSON link values",
description="Takes in arbitrary inputs and renders them as NODE-compatible values, replacing $a -> $z with JSON link values.",
is_experimental=True,
inputs=inputs,
inputs=[
io.Autogrow.Input("links", template=t1),
io.Autogrow.Input(
"vars",
template=t2,
),
io.String.Input(
"template",
tooltip="The variables $a to $z will be replaced in this text with their corresponding input's JSON link value",
placeholder="In this text you can refer to the input links as $a, $b etc. and the var inputs as either $var1 or $json1 etc. (the latter will be rendered through Python's json.dumps function which will cause strings to be quoted)",
multiline=True,
),
],
outputs=[io.String.Output()],
)
# This requires https://github.com/Comfy-Org/ComfyUI/pull/15103 to work properly
# Without that PR, all inputs will be evaluated non-lazily
@classmethod
def execute(cls, text, **vars) -> io.NodeOutput:
for k in "abcdefghijklmn":
def check_lazy_status(cls, template, links, vars):
r = []
for name, (v, input_name) in vars.items():
if v is None and f"${name}" in template or v is None and f"$json{name[3:]}" in template:
r.append(input_name)
return r
@classmethod
def execute(cls, template, links, vars) -> io.NodeOutput:
text = template
for k in cls.NAMES:
v = "__EMPTY__"
if k in vars:
v = json.dumps(vars[k])
if k in links:
# Replace : with \: to avoid breaking scheduling syntax when linking subgraphs. Any function that consumes this should replace \: with :
v = json.dumps(links[k]).replace(":", r"\:")
text = substitute_var(text, k, v)
for i in range(50):
v = "__EMPTY__"
k = f"var{i + 1}"
if k in vars:
v = vars[k]
text = substitute_var(text, k, str(v))
if f"$json{i + 1}" in text:
v = v if v == "__EMPTY__" else json.dumps(v)
text = substitute_var(text, f"json{i + 1}", v)
return io.NodeOutput(text)
+395 -8
View File
@@ -1,10 +1,397 @@
import logging
import os
from __future__ import annotations
log = logging.getLogger("comfyui-prompt-control")
import itertools as it
from dataclasses import dataclass
from math import ceil
from typing import Any, TypeAlias
if os.environ.get("PC_USE_OLD_PARSER", "0") != "1":
from .parser_parsy import parse_prompt_schedules # noqa
else:
log.warning("Using old Lark parser (UNSUPPORTED)")
from .parser_lark import parse_prompt_schedules # noqa
from typing_extensions import override
from .macros import expand_macros
from .parsy import any_char, char_from, digit, eof, forward_declaration, generate, regex, seq, string, success
FOREVER = float("inf")
EvalResult: TypeAlias = tuple[float, str, list["LoRA"]]
def merge_until(i: EvalResult, minimum: float):
until, p, loras = i
until = min(until, minimum)
return until, p, loras
def batched(iterable, n, *, strict=False):
# batched('ABCDEFG', 2) → AB CD EF G
if n < 1:
raise ValueError("n must be at least one")
iterator = iter(iterable)
while batch := tuple(it.islice(iterator, n)):
if strict and len(batch) != n:
raise ValueError("batched(): incomplete batch")
yield batch
EvalResult: TypeAlias = tuple[float, str, list["LoRA"]]
class Expression:
def eval(self, step: float, tags: list[str]) -> EvalResult:
return (FOREVER, "", [])
def required_steps(self, max_steps: float) -> set[float]:
return set()
@dataclass
class Text(Expression):
string: str
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
assert isinstance(self.string, str)
return FOREVER, self.string, []
@dataclass
class Alternate(Expression):
prompts: list[Expression]
step: float = 0.1
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
SCALE = 10_000
step = max(step, self.step)
position = (step * SCALE) / (self.step * SCALE)
idx = (ceil(position) - 1) % len(self.prompts)
r = self.prompts[max(0, idx)].eval(step, tags)
r = merge_until(r, max(self.step, ceil(position) * self.step))
return r
@override
def required_steps(self, max_steps: float):
r = set()
for x in self.prompts:
r.update(x.required_steps(max_steps))
r.update(set(x / 100 for x in range(0, int(max_steps * 100), int(self.step * 100))))
return r
@dataclass
class Sequence(Expression):
prompts: list[tuple[Expression, float]]
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
item = Text("")
found_step = FOREVER
for prompt, switch_step in self.prompts:
if step <= switch_step:
found_step = switch_step
item = prompt
break
return merge_until(item.eval(step, tags), found_step)
@override
def required_steps(self, max_steps: float):
return set(step for _, step in self.prompts if step <= max_steps)
@dataclass
class Schedule(Expression):
before: Prompt
during: Prompt
after: Prompt
start: float
end: float
tag: str | None
def tag_matches(self, tags: list[str]):
return self.tag in tags
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
if self.tag is not None and not self.tag_matches(tags):
return self.before.eval(step, tags)
if self.tag_matches(tags):
return self.during.eval(step, tags)
if step <= self.start:
return merge_until(self.before.eval(step, tags), self.start)
if self.start < step <= self.end:
return merge_until(self.during.eval(step, tags), self.end)
if step > self.end:
return self.after.eval(step, tags)
raise AssertionError("How are you here?")
@override
def required_steps(self, max_steps: float):
r = set()
if self.start < max_steps:
r.add(self.start)
if self.end < max_steps:
r.add(self.end)
r.update(self.before.required_steps(max_steps))
r.update(self.during.required_steps(max_steps))
r.update(self.after.required_steps(max_steps))
return r
@dataclass
class Prompt(Expression):
data: list[Expression]
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
evals = [x.eval(step, tags) for x in self.data]
text = "".join(x[1] for x in evals)
untils = [x[0] for x in evals]
loras = []
for x in evals:
loras.extend(x[2])
until = FOREVER if not untils else min(untils)
return until, text, loras
@override
def required_steps(self, max_steps):
r = set()
for x in self.data:
r.update(x.required_steps(max_steps))
return r
@dataclass
class LoRA(Expression):
filename: str
w_model: float = 1.0
w_te: float = 1.0
def eval(self, step: float, tags: list[str]) -> EvalResult:
return FOREVER, "", [self]
def find_weight_at(weights: list[tuple[float, float]], step: float, until: float):
res_w = 0
for this, next in zip(weights, it.chain(weights[1:], [(0, FOREVER)]), strict=False):
w, start = this
_, next_start = next
if start > step or next_start < step:
until = min(until, start)
continue
res_w = w
return until, res_w
@dataclass
class LoRACTL(Expression):
filename: str
w_model: list[tuple[float, float]]
w_te: list[tuple[float, float]]
def eval(self, step: float, tags: list[str]) -> EvalResult:
until, w1 = find_weight_at(self.w_model, step, FOREVER)
until, w2 = find_weight_at(self.w_te, step, until)
lora = []
if w1 != 0 or w1 != 0:
lora = [LoRA(self.filename, w1, w2)]
return until, "", lora
def required_steps(self, max_steps):
r = set(x[1] for x in self.w_model)
r.update(set(x[1] for x in self.w_te))
return r
def combine_arglist(prompts, start_end) -> Schedule:
a, b, c = prompts
start_or_tag, end = start_end
empty = Prompt([])
start = start_or_tag
# Handle [a:b:TAG]
if isinstance(start_or_tag, str):
if b is None:
before = empty
during = a # [a:TAG] produces a when tag is active
else:
before, during = a, b # [a:b:TAG] changes from a to b when tag is active
return Schedule(before, during, empty, start=0.0, end=FOREVER, tag=start_or_tag)
during = before = after = empty
if end is not None:
if b is None: # [a:0,0.5] == [:a:0,0.5]
during = a
before = after = empty
elif c is None: # [a:b:0,0.5]
before = empty
during = a
after = b
else:
before, during, after = a, b, c
else:
end = FOREVER
if b is None: # [a:0.5] == [::a:0.5,0.5]
before = empty
during = a
after = a
else:
before = a
during = b
after = b
# c always gets ignored
start = float(start) # for typechecking
return Schedule(before, during, after, start, end, tag=None)
def token(s: str):
return string(s).map(Text)
def combine_prompt(*prompts):
p = prompts
if len(p) == 1:
p = p[0]
if isinstance(p, Prompt):
p = p.data[0] if len(p.data) == 1 else combine_prompt(*p.data)
if isinstance(p, Expression):
return p
p = [combine_prompt(x) for x in p]
return Prompt(p)
@dataclass
class PromptSchedule:
parse_tree: Expression
filters: list[str]
start: float
end: float
num_steps: int
def at_step(self, step: float) -> tuple[float, dict[str, Any]]:
max_step = self.num_steps or 1.0
if max_step > 1 and step < 1:
step = step * max_step
until, p, lora_list = self.parse_tree.eval(step, self.filters)
loras = {}
for lora in lora_list:
d = loras.get(lora.filename, {})
d["weight"] = d.get("weight", 0) + lora.w_model
d["weight_clip"] = d.get("weight_clip", 0) + lora.w_te
loras[lora.filename] = d
if max_step > 0 and until > 1:
# TODO: better logic for this?
until = min(until / max_step, 1.0)
return (min(max_step, round(until, 2)), {"prompt": p, "loras": loras})
def with_filters(self, filters: str | None = None, start: float | None = None, end: float | None = None):
return PromptSchedule(
self.parse_tree,
self.filters if filters is None else parse_filters(filters),
self.start if start is None else start,
self.end if end is None else end,
self.num_steps,
)
def clone(self):
return self.with_filters()
def __iter__(self):
return (x for x in self.parsed_prompt if x[0] != 0)
@property
def parsed_prompt(self):
max_step = self.num_steps or 1.0
required_steps = self.parse_tree.required_steps(max_step).union({max_step})
prompts = list(sorted((self.at_step(step) for step in required_steps), key=lambda x: x[0]))
res = []
prev_end = -1
for end_at, p in prompts:
if end_at < self.start:
continue
elif end_at < self.end and prev_end < end_at:
res.append([end_at, p])
prev_end = end_at
elif end_at >= self.end and prev_end < self.end:
res.append([end_at, p])
break
if len(res) == 0:
res = [[1.0], prompts[-1][1]]
return res
def lora_weights(p):
@generate
def parser():
w_model = yield col >> p
w_te = yield (col >> p).optional(w_model)
return [w_model, w_te]
return parser.desc("lora_weights")
prompt = forward_declaration()
empty = Text("")
comma = token(",")
col = token(":")
lsq = token("[")
rsq = token("]")
lpar = token("(")
rpar = token(")")
tag = regex(r"[A-Z_]+")
non_special = regex(r"[^:\[\]()|\\<>#]+").map(Text)
filename = regex(r"[^:<>]+")
comment = string("#") >> any_char.until(eof | char_from("\n")) >> success(empty)
escape = (string("\\") >> char_from("\\[]:#") | string(r"\(") | string(r"\)")).map(Text)
emphasis = seq(lpar, (prompt | col).at_least(0), rpar)
sign = string("+") | string("-")
number = (
(sign.optional("") + (digit.many() + string(".") * 1 + digit.many() | digit.at_least(1)).concat())
.concat()
.map(float)
)
opt_prompt = prompt.optional(empty)
step_range = seq(number | tag, (comma >> number).optional())
arglist = seq((opt_prompt << col).optional() * 3, step_range)
schedule = lsq >> arglist.combine(combine_arglist) << rsq
alternate = (lsq >> seq(prompt.sep_by(string("|"), min=1), (col >> number).optional(0.1)) << rsq).combine(Alternate)
sequence = (lsq >> string("SEQ") >> seq(col >> opt_prompt << col, number).at_least(1) << rsq).map(Sequence)
bracketed = seq(lsq, prompt.at_least(0), rsq) | sequence | schedule | alternate
lora = (string("<lora:") >> filename * 1 + lora_weights(number) << string(">")).combine(LoRA)
ctlweight = seq(number, (string("@") >> number).optional(0)).sep_by(comma, min=1)
loractl = (string("<loractl:") >> filename * 1 + lora_weights(ctlweight) << string(">")).combine(LoRACTL)
emb = (string("<emb:") >> filename << string(">")).map(lambda f: Text(f"embedding:{f}"))
expr = (
escape
| comment
| non_special
| bracketed
| emphasis.combine(combine_prompt)
| lora
| loractl
| emb
| char_from("<>").map(Text)
)
prompt_ = expr.at_least(1).combine(combine_prompt)
prompt.become(prompt_)
# Treat any character that isn't valid prompt syntax as just text
all = (prompt | any_char.map(Text)).at_least(0).combine(combine_prompt)
def parse_filters(filters: str):
return [x.strip().upper() for x in filters.split(",") if x.strip()]
def parse(text):
return combine_prompt(all.parse(text))
def parse_prompt_schedules(text, filters="", start=0, end=1.0, num_steps=0):
return PromptSchedule(parse(expand_macros(text.strip())), parse_filters(filters), start, end, num_steps)
-359
View File
@@ -1,359 +0,0 @@
# vim: sw=4 ts=4
from __future__ import annotations
import logging
from functools import lru_cache
from math import ceil
import lark
from .macros import expand_macros
from .utils import flatten
log = logging.getLogger("comfyui-prompt-control")
if lark.__version__ == "0.12.0":
from sys import executable
x = "\n".join(
[
"Your lark package reports an ancient version (0.12.0) and will not work.",
"If you have the 'lark-parser' package in your Python environment, remove that and *reinstall* lark!",
f"{executable} -m pip uninstall lark-parser lark",
f"{executable} -m pip install lark",
]
)
log.error(x)
raise ImportError(x)
ESCAPES = [
("XxPCBackslashESCAPExX", "\\"),
("XxPCColonESCAPExX", ":"),
("XxPCCommentESCAPExX", "#"),
]
def escape_specials(string: str) -> str:
for ph, c in ESCAPES:
string = string.replace(rf"\{c}", ph)
return string
def restore_escaped(string: str) -> str:
for ph, c in ESCAPES:
string = string.replace(ph, c)
return string
def remove_comments(string: str) -> str:
r = []
for line in string.split("\n"):
comment = line.find("#")
if comment >= 0:
r.append(line[:comment])
else:
r.append(line)
return "\n".join(r)
prompt_parser = lark.Lark(
r"""
!start: (prompt | /[][():|]/+)*
prompt: (emphasized | embedding | scheduled | alternate | sequence | loraspec | PLAIN | | /\\:/ | /</ | />/ | WHITESPACE)+
!emphasized: "(" prompt? ")"
| "(" prompt ":" prompt ")"
| "[" prompt "]"
promptlist: ([prompt] ":")~1..3
scheduled: "[" promptlist _WS? NUMBER ["," NUMBER] "]"
| "[" promptlist _WS? TAG "]"
sequence.5: "[SEQ" ":" [prompt] ":" NUMBER (":" [prompt] ":" NUMBER)* "]"
alternate: "[" [prompt] ("|" [prompt])+ [":" NUMBER] "]"
loraspec.99: "<lora:" FILENAME lora_weights [lora_block_weights] ">"
lora_weights.1: (":" _WS? NUMBER)~1..2
lora_block_weights.-1: ":" PLAIN
embedding.100: "<emb:" FILENAME ">"
WHITESPACE: /\s+/
_WS: WHITESPACE
PLAIN: /([^<>\\\[\]():|]|\\.)+/
FILENAME: /[^<>:]+/
TAG: /[A-Z_]+/
%import common.SIGNED_NUMBER -> NUMBER
""",
lexer="dynamic",
)
def clamp(a, b, c):
"""clamp b between a and c"""
return min(max(a, b), c)
def get_steps(tree, num_steps):
res = [num_steps or 100]
def tostep(s):
steps = num_steps or 100
if "." in str(s) or not num_steps:
w = float(s)
value = w * steps
else:
w = int(s)
value = w
if w > 1 and not num_steps:
log.warning(
"You haven't configured the number of steps for Prompt Control to use, %s will be clipped to 1.0", w
)
value = steps
return int(clamp(0, value, steps))
class CollectSteps(lark.Visitor):
def scheduled(self, tree):
i = tree.children[-1]
if i and i.type == "TAG":
return
for i in [-1, -2]:
if tree.children[i] is not None:
tree.children[i] = tostep(tree.children[i])
res.append(tree.children[i])
def interp_steps(self, tree):
tree.children[-1] = tostep(tree.children[-1] or 0.1)
for i, _ in enumerate(tree.children[:-1]):
tree.children[i] = tostep(tree.children[i])
res.extend(tree.children[:-1])
def sequence(self, tree):
steps = tree.children[1::2]
for i, _ in enumerate(steps):
w = tostep(tree.children[i * 2 + 1])
tree.children[i * 2 + 1] = w
res.append(w)
def alternate(self, tree):
step_size = tostep(round(float(tree.children[-1] or 0.1), 2))
tree.children[-1] = step_size
res.extend([x for x in range(step_size, num_steps or 100, step_size)])
CollectSteps().visit(tree)
return sorted(set(res))
def at_step(step, filters, tree):
class AtStep(lark.Transformer):
def scheduled(self, args):
before = None
during = None
after = None
when_end = None
pl, when, *rest = args
if rest:
when_end = rest[0]
pl = list(pl)
if len(pl) == 1:
(during,) = pl # [after:0.5] == [::after:0.5,0.5]
if when_end is None:
when_end = when
after = during
elif len(pl) == 2:
during, after = pl # [during:after:0.5] = [before::after:0.5,0.5]
if when_end is None:
when_end = when
before = during
else:
before, during, after = pl # [before:during:after:0.5,0.8]
if isinstance(when, str):
return before or "" if when not in filters else after or ""
if when_end is None:
when_end = 1000_000
if step <= when:
return before or ""
if when < step <= when_end:
return during or ""
else:
return after or ""
def sequence(self, args):
previous_step = 0.0
prompts = args[::2]
steps = args[1::2]
for s, p in zip(steps, prompts, strict=False):
if s >= step and step >= previous_step:
previous_step = step
return p or ""
else:
previous_step = s
return ""
def alternate(self, args):
step_size = args[-1]
idx = ceil(step / step_size)
return args[(idx - 1) % (len(args) - 1)] or ""
def start(self, args):
prompt = []
loraspecs = {}
args = flatten(args)
for a in args:
if isinstance(a, str):
prompt.append(a)
elif isinstance(a, tuple):
# sum identical specs together
n = a[0]
# if clip weight is not provided, use unet weight
w, w_clip = a[1][0], a[1][1 % len(a[1])]
e = loraspecs.get(n, {})
loraspecs[n] = {
"weight": round(e.get("weight", 0.0) + w, 2),
"weight_clip": round(e.get("weight_clip", 0.0) + w_clip, 2),
}
lbw = a[2]
if lbw:
loraspecs[n]["lbw"] = lbw
if loraspecs[n]["weight"] == 0 and loraspecs[n]["weight_clip"] == 0 and not lbw:
del loraspecs[n]
else:
pass
p = "".join(prompt)
return {"prompt": p, "loras": loraspecs}
def PLAIN(self, args):
return restore_escaped(args)
def FILENAME(self, value):
return str(value)
def embedding(self, args):
return "embedding:" + str(args[0])
def lora_weights(self, args):
return [float(str(a)) for a in args]
def lora_block_weights(self, args):
vals = args[0].split(";")
r = {}
for v in vals:
x = v.split("=", 2)
if len(x) != 2:
continue
k, v = x[0].strip().upper(), x[1].strip()
r[k] = v
return r
def loraspec(self, args):
name = args[0]
params = args[1]
lbw = args[2]
return name, params, lbw
def __default__(self, data, children, meta):
return children
return AtStep().transform(tree)
class PromptSchedule:
# 0 num_steps means unconfigured
def __init__(self, prompt, filters="", start=0.0, end=1.0, num_steps=0):
self.filters = filters
self.start = start
self.end = end
self.num_steps = num_steps
# placeholder is restored on parse
self.prompt = remove_comments(escape_specials(prompt.strip()))
self.defaults = {}
self.loaded_loras = {}
self.parsed_prompt = self._parse(num_steps)
def __iter__(self):
# Filter out zero, it's only useful for interpolation
return (x for x in self.parsed_prompt if x[0] != 0)
def _parse(self, num_steps):
filters = [x.strip() for x in self.filters.upper().split(",")]
try:
parsed = []
tree = prompt_parser.parse(self.prompt)
steps = get_steps(tree, num_steps=num_steps)
def f(x):
return round(x / (num_steps or 100), 2)
for t in steps:
p = at_step(t, filters, tree)
parsed.append([f(t), p])
except lark.exceptions.LarkError as e:
log.error("Prompt editing parse error: %s", e)
parsed = [[1.0, {"prompt": self.prompt, "loras": {}}]]
raise
# Tag filtering may return redundant prompts, so filter them out here
res = []
prev_end = -1
for end_at, p in parsed:
if end_at < self.start:
continue
elif end_at <= self.end:
res.append([end_at, p])
prev_end = end_at
elif end_at > self.end and prev_end < self.end:
res.append([end_at, p])
break
# Always use the last prompt if everything was filtered
if len(res) == 0:
res = [[1.0, parsed[-1][1]]]
final = [res[0]]
# Clean up duplicates
for p in res[1:]:
if p[1] != final[-1][1]:
final.append(p)
else:
final[-1][0] = p[0]
return final
def clone(self):
return self.with_filters()
def with_filters(self, filters=None, start=None, end=None, defaults=None):
def ifspecified(x, defval):
return x if x is not None else defval
p = PromptSchedule(
self.prompt,
filters=ifspecified(filters, self.filters),
start=ifspecified(start, self.start),
end=ifspecified(end, self.end),
num_steps=self.num_steps,
)
return p
def at_step(self, step, total_steps=1):
_, x = self.at_step_idx(step, total_steps)
return x
def at_step_idx(self, step, total_steps=1):
for i, x in enumerate(self.parsed_prompt):
if x[0] * total_steps >= step:
return i, x
return len(self.parsed_prompt) - 1, self.parsed_prompt[-1]
@lru_cache
def parse_prompt_schedules(prompt, **kwargs):
prompt = expand_macros(prompt)
return PromptSchedule(prompt, **kwargs)
-389
View File
@@ -1,389 +0,0 @@
from __future__ import annotations
import itertools as it
from dataclasses import dataclass
from math import ceil
from typing import Any, TypeAlias
from typing_extensions import override
from .macros import expand_macros
from .parsy import any_char, char_from, digit, eof, forward_declaration, generate, regex, seq, string, success
FOREVER = float("inf")
EvalResult: TypeAlias = tuple[float, str, list["LoRA"]]
def merge_until(i: EvalResult, minimum: float):
until, p, loras = i
until = min(until, minimum)
return until, p, loras
def batched(iterable, n, *, strict=False):
# batched('ABCDEFG', 2) → AB CD EF G
if n < 1:
raise ValueError("n must be at least one")
iterator = iter(iterable)
while batch := tuple(it.islice(iterator, n)):
if strict and len(batch) != n:
raise ValueError("batched(): incomplete batch")
yield batch
EvalResult: TypeAlias = tuple[float, str, list["LoRA"]]
class Expression:
def eval(self, step: float, tags: list[str]) -> EvalResult:
return (FOREVER, "", [])
def required_steps(self, max_steps: float) -> set[float]:
return set()
@dataclass
class Text(Expression):
string: str
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
assert isinstance(self.string, str)
return FOREVER, self.string, []
@dataclass
class Alternate(Expression):
prompts: list[Expression]
step: float = 0.1
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
SCALE = 10_000
step = max(step, self.step)
position = (step * SCALE) / (self.step * SCALE)
idx = (ceil(position) - 1) % len(self.prompts)
r = self.prompts[max(0, idx)].eval(step, tags)
r = merge_until(r, max(self.step, ceil(position) * self.step))
return r
@override
def required_steps(self, max_steps: float):
r = set()
for x in self.prompts:
r.update(x.required_steps(max_steps))
r.update(set(x / 100 for x in range(0, int(max_steps * 100), int(self.step * 100))))
return r
@dataclass
class Sequence(Expression):
prompts: list[tuple[Expression, float]]
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
item = Text("")
found_step = FOREVER
for prompt, switch_step in self.prompts:
if step <= switch_step:
found_step = switch_step
item = prompt
break
return merge_until(item.eval(step, tags), found_step)
@override
def required_steps(self, max_steps: float):
return set(step for _, step in self.prompts if step <= max_steps)
@dataclass
class Schedule(Expression):
before: Prompt
during: Prompt
after: Prompt
start: float
end: float
tag: str | None
def tag_matches(self, tags: list[str]):
return self.tag in tags
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
if self.tag is not None and not self.tag_matches(tags):
return self.before.eval(step, tags)
if self.tag_matches(tags):
return self.during.eval(step, tags)
if step <= self.start:
return merge_until(self.before.eval(step, tags), self.start)
if self.start < step <= self.end:
return merge_until(self.during.eval(step, tags), self.end)
if step > self.end:
return self.after.eval(step, tags)
raise AssertionError("How are you here?")
@override
def required_steps(self, max_steps: float):
r = set()
if self.tag is not None:
return r
if self.start < max_steps:
r.add(self.start)
if self.end < max_steps:
r.add(self.end)
r.update(self.before.required_steps(max_steps))
r.update(self.during.required_steps(max_steps))
r.update(self.after.required_steps(max_steps))
return r
@dataclass
class Prompt(Expression):
data: list[Expression]
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
evals = [x.eval(step, tags) for x in self.data]
text = "".join(x[1] for x in evals)
untils = [x[0] for x in evals]
loras = []
for x in evals:
loras.extend(x[2])
until = FOREVER if not untils else min(untils)
return until, text, loras
@override
def required_steps(self, max_steps):
r = set()
for x in self.data:
r.update(x.required_steps(max_steps))
return r
@dataclass
class LoRA(Expression):
filename: str
w_model: float = 1.0
w_te: float = 1.0
def eval(self, step: float, tags: list[str]) -> EvalResult:
return FOREVER, "", [self]
def find_weight_at(weights: list[tuple[float, float]], step: float, until: float):
res_w = 0
for this, next in zip(weights, it.chain(weights[1:], [(0, FOREVER)]), strict=False):
w, start = this
_, next_start = next
if start > step or next_start < step:
until = min(until, start)
continue
res_w = w
return until, res_w
@dataclass
class LoRACTL(Expression):
filename: str
w_model: list[tuple[float, float]]
w_te: list[tuple[float, float]]
def eval(self, step: float, tags: list[str]) -> EvalResult:
until, w1 = find_weight_at(self.w_model, step, FOREVER)
until, w2 = find_weight_at(self.w_te, step, until)
lora = []
if w1 != 0 or w1 != 0:
lora = [LoRA(self.filename, w1, w2)]
return until, "", lora
def required_steps(self, max_steps):
r = set(x[1] for x in self.w_model)
r.update(set(x[1] for x in self.w_te))
return r
def combine_arglist(prompts, start_end) -> Schedule:
a, b, c = prompts
start_or_tag, end = start_end
empty = Prompt([])
start = start_or_tag
# Handle [a:b:TAG]
if isinstance(start_or_tag, str):
if b is None:
before = empty
during = a # [a:TAG] produces a when tag is active
else:
before, during = a, b # [a:b:TAG] changes from a to b when tag is active
return Schedule(before, during, empty, start=0.0, end=FOREVER, tag=start_or_tag)
during = before = after = empty
if end is not None:
if b is None: # [a:0,0.5] == [:a:0,0.5]
during = a
before = after = empty
elif c is None: # [a:b:0,0.5]
before = empty
during = a
after = b
else:
before, during, after = a, b, c
else:
end = FOREVER
if b is None: # [a:0.5] == [::a:0.5,0.5]
before = empty
during = a
after = a
else:
before = a
during = b
after = b
# c always gets ignored
start = float(start) # for typechecking
return Schedule(before, during, after, start, end, tag=None)
def token(s: str):
return string(s).map(Text)
def combine_prompt(*prompts):
p = prompts
if len(p) == 1:
p = p[0]
if isinstance(p, Prompt):
p = p.data[0] if len(p.data) == 1 else combine_prompt(*p.data)
if isinstance(p, Expression):
return p
p = [combine_prompt(x) for x in p]
return Prompt(p)
@dataclass
class PromptSchedule:
parse_tree: Expression
filters: list[str]
start: float
end: float
num_steps: int
def at_step(self, step: float) -> tuple[float, dict[str, Any]]:
max_step = self.num_steps or 1.0
if max_step > 1 and step < 1:
step = step * max_step
until, p, lora_list = self.parse_tree.eval(step, self.filters)
loras = {}
for lora in lora_list:
d = loras.get(lora.filename, {})
d["weight"] = d.get("weight", 0) + lora.w_model
d["weight_clip"] = d.get("weight_clip", 0) + lora.w_te
loras[lora.filename] = d
if max_step > 0 and until > 1:
# TODO: better logic for this?
until = min(until / max_step, 1.0)
return (min(max_step, round(until, 2)), {"prompt": p, "loras": loras})
def with_filters(self, filters: str | None = None, start: float | None = None, end: float | None = None):
return PromptSchedule(
self.parse_tree,
self.filters if filters is None else parse_filters(filters),
self.start if start is None else start,
self.end if end is None else end,
self.num_steps,
)
def clone(self):
return self.with_filters()
def __iter__(self):
return (x for x in self.parsed_prompt if x[0] != 0)
@property
def parsed_prompt(self):
max_step = self.num_steps or 1.0
required_steps = self.parse_tree.required_steps(max_step).union({max_step})
prompts = list(sorted((self.at_step(step) for step in required_steps), key=lambda x: x[0]))
res = []
prev_end = -1
for end_at, p in prompts:
if end_at < self.start:
continue
elif end_at < self.end and prev_end < end_at:
res.append([end_at, p])
prev_end = end_at
elif end_at >= self.end and prev_end < self.end:
res.append([end_at, p])
break
if len(res) == 0:
res = [[1.0], prompts[-1][1]]
return res
def lora_weights(p):
@generate
def parser():
w_model = yield col >> p
w_te = yield (col >> p).optional(w_model)
return [w_model, w_te]
return parser.desc("lora_weights")
prompt = forward_declaration()
empty = Text("")
comma = token(",")
col = token(":")
lsq = token("[")
rsq = token("]")
lpar = token("(")
rpar = token(")")
tag = regex(r"[A-Z_]+")
non_special = regex(r"[^:\[\]()|\\<>#]+").map(Text)
filename = regex(r"[^:<>]+")
comment = string("#") >> any_char.until(eof | char_from("\n")) >> success(empty)
escape = (string("\\") >> char_from("\\[]:#") | string(r"\(") | string(r"\)")).map(Text)
emphasis = seq(lpar, (prompt | col).at_least(0), rpar)
sign = string("+") | string("-")
number = (
(sign.optional("") + (digit.many() + string(".") * 1 + digit.many() | digit.at_least(1)).concat())
.concat()
.map(float)
)
opt_prompt = prompt.optional(empty)
step_range = seq(number | tag, (comma >> number).optional())
arglist = seq((opt_prompt << col).optional() * 3, step_range)
schedule = lsq >> arglist.combine(combine_arglist) << rsq
alternate = (lsq >> seq(prompt.sep_by(string("|"), min=1), (col >> number).optional(0.1)) << rsq).combine(Alternate)
sequence = (lsq >> string("SEQ") >> seq(col >> opt_prompt << col, number).at_least(1) << rsq).map(Sequence)
bracketed = seq(lsq, prompt.at_least(0), rsq) | sequence | schedule | alternate
lora = (string("<lora:") >> filename * 1 + lora_weights(number) << string(">")).combine(LoRA)
ctlweight = seq(number, (string("@") >> number).optional(0)).sep_by(comma, min=1)
loractl = (string("<loractl:") >> filename * 1 + lora_weights(ctlweight) << string(">")).combine(LoRACTL)
emb = (string("<emb:") >> filename << string(">")).map(lambda f: Text(f"embedding:{f}"))
expr = escape | comment | non_special | bracketed | emphasis.combine(combine_prompt) | lora | loractl | emb
prompt_ = expr.at_least(1).combine(combine_prompt)
prompt.become(prompt_)
# Treat any character that isn't valid prompt syntax as just text
all = (prompt | any_char.map(Text)).at_least(0).combine(combine_prompt)
def parse_filters(filters: str):
return [x.strip().upper() for x in filters.split(",") if x.strip()]
def parse(text):
return combine_prompt(all.parse(text))
def parse_prompt_schedules(text, filters="", start=0, end=1.0, num_steps=0):
return PromptSchedule(parse(expand_macros(text.strip())), parse_filters(filters), start, end, num_steps)
+18 -8
View File
@@ -3,7 +3,7 @@ from __future__ import annotations
import copy
import logging
import re
from collections.abc import Iterator
from collections.abc import Callable, Iterator
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING, Any, TypeAlias, TypeVar
@@ -20,7 +20,6 @@ class FunctionSpec:
name: str
args: FunctionArgs
position: int
placeholder: str | None
# Allow testing
@@ -142,6 +141,10 @@ def find_function_spans(
if text[at_paren:after_first_paren] == "(":
end = find_closing_paren(text, after_first_paren)
if end < 0:
# Unclosed paren: skip past this match so the loop terminates
idx += match.end()
text = text[match.end() :]
match = rex.search(text)
continue
args = parse_strings(text[after_first_paren:end], defaults)
end += 1
@@ -155,7 +158,11 @@ def find_function_spans(
def get_function(
text: str, func: str, defaults: list[str] | None, placeholder: str = "", require_args: bool = True
text: str,
func: str,
defaults: list[str] | None,
processor: Callable[..., str] | None = None,
require_args: bool = True,
) -> tuple[str, list[FunctionSpec]]:
spans = [x.span() for x in re.finditer(r'".+?"', text)]
instances = []
@@ -164,14 +171,13 @@ def get_function(
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"
instances.append(FunctionSpec(funcname, args, start - skipped, ph))
instances.append(FunctionSpec(funcname, args, start - skipped))
skipped += end - start
chunks.append(text[current:start] + (ph or ""))
chunks.append(text[current:start])
if processor:
chunks.append(processor(*args))
current = end
count += 1
chunks.append(text[current:])
@@ -273,6 +279,10 @@ def lora_name_to_file(name: str) -> str | None:
search = [f for f in filenames if all(p in f for p in parts)]
if len(search) == 1:
return search[0]
elif len(search) > 1:
if len(search) > 4:
search[4] = "..."
log.warning("Ignored LoRA search 's%'; matched more than one file: %s", name, ", ".join(search[:5]))
return None
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-prompt-control"
description = "Nodes for prompt editing and LoRA scheduling, advanced regional prompting (including attention masking) and advanced prompt encoding, all controlled through your text prompt. Feature keywords: comfyui-prompt-control, schedule, macros, attention couple, loractl, A1111"
version = "3.0.0-beta.7"
version = "3.0.0-beta.10"
license = { file = "LICENSE" }
requires-python = ">= 3.10"
+5 -10
View File
@@ -461,7 +461,7 @@ def test_textencode_lora_with_schedule():
def test_textencode_custom():
r = te("NODE(CLIPTextEncode)simple [test:0.1,0.5] prompt")
r = te("NODE(CLIPTextEncode)simple [test:0.1,0.5] $p SEG(p) prompt")
assert r == {
"result": (["UID.0.0.8", 0],),
"expand": {
@@ -503,7 +503,7 @@ def test_textencode_custom():
def test_textencode_custom_extra():
r = te(
'NODE(CustomTextEncode, prompt, image ["1", 0]; option "test"; float [10.0:__EMPTY__:0.5])simple [test:0.1,0.5] prompt'
'NODE(CustomTextEncode, prompt, image ["1\:1", 0]; option "test"; float [10.0:__EMPTY__:0.5])simple [test:0.1,0.5] prompt'
)
assert r == {
"result": (["UID.0.0.8", 0],),
@@ -513,7 +513,7 @@ def test_textencode_custom_extra():
"inputs": {
"clip": [0, 0],
"prompt": "simple prompt",
"image": ["1", 0],
"image": ["1:1", 0],
"option": "test",
"float": 10.0,
},
@@ -527,7 +527,7 @@ def test_textencode_custom_extra():
"inputs": {
"clip": [0, 0],
"prompt": "simple test prompt",
"image": ["1", 0],
"image": ["1:1", 0],
"option": "test",
"float": 10.0,
},
@@ -538,12 +538,7 @@ def test_textencode_custom_extra():
},
"UID.0.0.5": {
"class_type": "CustomTextEncode",
"inputs": {
"clip": [0, 0],
"prompt": "simple prompt",
"image": ["1", 0],
"option": "test"
},
"inputs": {"clip": [0, 0], "prompt": "simple prompt", "image": ["1:1", 0], "option": "test"},
},
"UID.0.0.6": {
"class_type": "ConditioningSetTimestepRange",
+5
View File
@@ -30,6 +30,11 @@ def test_macro_recursion():
assert "Unable to resolve DEFs" in str(c.value)
def test_parsing_cornercase():
r = expand_macros("This should not get stuck DEF(")
assert r == "This should not get stuck DEF("
@pytest.mark.parametrize(
"input, output",
[
+15 -12
View File
@@ -1,7 +1,7 @@
import os
import pytest
from prompt_control.parser import parse_prompt_schedules as parse
def lora_dict(*loras):
return {lora: {"weight": unet, "weight_clip": te} for lora, unet, te in loras}
@@ -19,18 +19,9 @@ def assert_prompt(p, at, until, text, *loras):
assert prompts_match(p.at_step(at), prompt(until, text, *loras))
parsers_to_test = os.environ.get("PC_PARSERS_TO_TEST", "new").split()
params = []
if "old" in parsers_to_test:
from prompt_control.parser_lark import parse_prompt_schedules as old_parse # noqa
params.append(old_parse)
if "new" in parsers_to_test:
from prompt_control.parser_parsy import parse_prompt_schedules as new_parse # noqa
params.append(new_parse)
params.append(parse)
@pytest.fixture(scope="module", autouse=True, params=params)
@@ -342,6 +333,12 @@ def test_cornercase_corrected(parse):
assert p.parsed_prompt[1:] == p2.parsed_prompt
def test_ltgt_in_schedule(parse):
p = parse("This should [<parse> correctly:be <Picture 1>:0.1]<lora:test:1>")
assert_prompt(p, 0.1, 0.1, "This should <parse> correctly", ("test", 1.0, 1.0))
assert_prompt(p, 0.15, 1.0, "This should be <Picture 1>", ("test", 1.0, 1.0))
def test_floats(parse):
p = parse("[a:b:0.5] [c:d:e:0.2,0.7] <lora:test:-0.3>")
p2 = parse("[a:b:.5] [c:d:e:.2,.7] <lora:test:-.3>")
@@ -364,3 +361,9 @@ def test_alternating_nested(parse):
for i, x in enumerate(catdogtigers):
step = round((i * 0.1) + 0.1, 2)
assert_prompt(p3, step, step, x)
def test_alternating_with_tags(parse):
p1 = parse("[[a|b]:HR]", filters="HR")
p2 = parse("[a|b]")
assert p1.parsed_prompt == p2.parsed_prompt
+7
View File
@@ -0,0 +1,7 @@
from prompt_control import utils
def test_smart_split():
assert utils.smarter_split(",", "foo,bar") == ["foo", "bar"]
assert utils.smarter_split(",", "(foo,bar),zonk") == ["(foo,bar)", "zonk"]
assert utils.smarter_split(",", r"\(foo,bar),zonk") == [r"\(foo", "bar)", "zonk"]