Compare commits
29
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
88af041dce | ||
|
|
29b21f735d | ||
|
|
8348a56f32 | ||
|
|
314fd39ed4 | ||
|
|
62dde6e795 | ||
|
|
66c88ea137 | ||
|
|
61ba7fe92e | ||
|
|
ce63e1d83c | ||
|
|
af293bb7cb | ||
|
|
53d210953a | ||
|
|
2265767d80 | ||
|
|
39add37c2f | ||
|
|
c2fa8071b5 | ||
|
|
953dac842a | ||
|
|
56fbfac2a5 | ||
|
|
ddee2b9a63 | ||
|
|
a41d433719 | ||
|
|
b24d93b778 | ||
|
|
c8f925c4ea | ||
|
|
1a733e71af | ||
|
|
e36eff4356 | ||
|
|
c6137ddc49 | ||
|
|
64417230a3 | ||
|
|
0a2ceb94e9 | ||
|
|
6e3fce9dcb | ||
|
|
b581cf7f24 | ||
|
|
d7a992b96d | ||
|
|
19123570e1 | ||
|
|
e6bb57cd25 |
@@ -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.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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,
|
||||
]
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
@@ -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
@@ -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",
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user