Compare commits
24
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f9c5da7210 | ||
|
|
244ef49230 | ||
|
|
3a563e3ceb | ||
|
|
812ad90d17 | ||
|
|
7d9e8aa6ac | ||
|
|
f47428ea8c | ||
|
|
0aeeb50331 | ||
|
|
9931c6fa75 | ||
|
|
655f6ac4a1 | ||
|
|
136932de40 | ||
|
|
06a3f43a93 | ||
|
|
203a7ad45c | ||
|
|
7a566e6e9e | ||
|
|
fa288c226c | ||
|
|
e648f3bdc3 | ||
|
|
1bedc6ad53 | ||
|
|
de79f3c6af | ||
|
|
eb51fd9289 | ||
|
|
7cca9438f2 | ||
|
|
7d786dfe83 | ||
|
|
85c15ff6bd | ||
|
|
1ab1c87f74 | ||
|
|
2ea0b622b4 | ||
|
|
c4c9561a13 |
@@ -19,5 +19,5 @@ jobs:
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.11'
|
||||
- run: pip install pytest typing-extensions -r requirements.txt
|
||||
- run: pip install pytest typing-extensions
|
||||
- run: PYTHONPATH=ComfyUI pytest tests/test_parser.py
|
||||
|
||||
@@ -31,7 +31,7 @@ jobs:
|
||||
- name: install-torch
|
||||
run: pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu
|
||||
- name: install ComfyUI
|
||||
run: pip install pytest typing-extensions -r requirements.txt -r ComfyUI/requirements.txt
|
||||
run: pip install pytest typing-extensions -r ComfyUI/requirements.txt
|
||||
- name: Download clip_l.safetensors
|
||||
run: curl -LO https://huggingface.co/comfyanonymous/flux_text_encoders/resolve/main/clip_l.safetensors
|
||||
- name: Force Comfy to use the CPU
|
||||
|
||||
@@ -12,7 +12,7 @@ format:
|
||||
ruff format
|
||||
|
||||
test:
|
||||
PYTHONPATH=../../ pytest tests/test_parser.py $(ARGS)
|
||||
PYTHONPATH=../../ pytest tests/test_parser.py tests/test_cutout.py tests/test_macros.py $(ARGS)
|
||||
|
||||
test_graph:
|
||||
PYTHONPATH=../../ pytest tests/test_graph.py $(ARGS)
|
||||
@@ -20,6 +20,9 @@ test_graph:
|
||||
test_encode:
|
||||
PYTHONPATH=../../ pytest tests/test_encode.py $(ARGS)
|
||||
|
||||
test_workflow:
|
||||
PYTHONPATH=../../ pytest tests/test_workflow.py $(ARGS)
|
||||
|
||||
test_encode_both:
|
||||
TEST_TE="clip_l t5" PYTHONPATH=../../ pytest tests/test_encode.py $(ARGS)
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# ComfyUI prompt control
|
||||
|
||||
Control LoRA and prompt scheduling, advanced text encoding, regional prompting, and much more, through your text prompt. Generates dynamic graphs that are literally identical to handcrafted noodle soup.
|
||||
Control LoRA and prompt scheduling, advanced text encoding, regional prompting, and much more, through your text prompt. Prompt Control generates dynamic graphs that are literally identical to handcrafted noodle soup, condensing complicated workflows with dozens of nodes into simple text prompts.
|
||||
|
||||
Prompt Control comes with `PCTextEncode`, which provides advanced text encoding with many additional features compared to ComfyUI's base `CLIPTextEncode`.
|
||||
|
||||
@@ -14,17 +14,15 @@ A `Basic Text to Image` template is included with the extension, and can be load
|
||||
|
||||
## What can it do?
|
||||
|
||||
You can use text prompts to control the following:
|
||||
|
||||
- A1111-style prompt scheduling and filtering without noodle soup.
|
||||
- LoRA loading and [scheduling](/doc/schedules.md) via the prompt, using ComfyUI's hook system
|
||||
- LoRA loading and [scheduling](/doc/schedules.md) using ComfyUI's built-in hook system.
|
||||
- Masking, composition and area control ([regional prompting](/doc/regional_prompts.md)) with an implementation of [Attention Couple](/doc/attention_couple.md), also fully schedulable.
|
||||
- [Advanced prompt encoding](/doc/basic.md)
|
||||
- Per-encoder prompts for models with multiple text encoders, such as SDXL and Flux
|
||||
- Per-encoder prompts for models with multiple text encoders, such as SDXL and Flux.
|
||||
- Prompt combinators like `BREAK`, as well as `CAT`, `AVG()` and `AND` corresponding to ComfyUI's `ConditioningConcat`, `ConditioningAverage` and `ConditioningCombine` nodes.
|
||||
- Different weight interpretation types (ComfyUI, A1111, compel, etc.)
|
||||
- Prompt masking with an implementation of [cutoff](https://github.com/BlenderNeko/ComfyUI_Cutoff)
|
||||
- Simple [prompt macros](/doc/macros.md) with `DEF`
|
||||
- Prompt masking with an implementation of [cutoff](https://github.com/BlenderNeko/ComfyUI_Cutoff).
|
||||
- Organize complicated prompts with [segments and prompt macros](/doc/macros.md).
|
||||
|
||||
All features are fully schedulable unless otherwise stated. See the [scheduling syntax documentation](doc/schedules.md) to get started.
|
||||
|
||||
|
||||
@@ -31,6 +31,7 @@ Prompt operators are processed in the following order, meaning that all features
|
||||
|
||||
- DEF macros are expanded
|
||||
- Scheduling is expanded, and for each scheduled prompt:
|
||||
- SEGs are processed and the template is expanded
|
||||
- The prompt is split by AND, and for each:
|
||||
- Prompts are split by COUPLE. and for each:
|
||||
- Most functions (like MASK) and cutoffs are evaluated
|
||||
|
||||
@@ -58,3 +58,40 @@ a "$1" b "$2"
|
||||
a "" b "$2"
|
||||
a "A" b "$2"
|
||||
```
|
||||
|
||||
## SEG: Split your prompt into named segments
|
||||
|
||||
Syntax: `SEG(segment_name)`
|
||||
|
||||
To help with organizing prompts, you can use the `SEG` function. For example:
|
||||
```
|
||||
This is a comic
|
||||
Top panel: $CAT. $SEG3
|
||||
Bottom panel: $DOG
|
||||
|
||||
SEG(DOG)
|
||||
A dog chasing its
|
||||
tail in a living room.
|
||||
SEG(CAT)
|
||||
|
||||
a sleeping cat
|
||||
|
||||
SEG
|
||||
The cat has orange fur with white stripes
|
||||
```
|
||||
This produces:
|
||||
```
|
||||
This is a comic
|
||||
Top panel: a sleeping cat. The cat has orange fur with white stripes
|
||||
Bottom panel: A dog chasing its
|
||||
tail in a living room.
|
||||
```
|
||||
Unlike macros, SEG is processed *after* scheduling syntax.
|
||||
|
||||
In this case, the first section before any `SEG` becomes the *template* and any text after a `SEG` call becomes part of that segment. Whitespace is stripped from the start and end of segments and the template.
|
||||
|
||||
In the template, you can refer to segments by either their index (starting from 1) or the given name, prefixed with a `$SEG`, so in this example, `$SEG1` is the same as `$PANEL1`
|
||||
|
||||
Segments can also refer to each other. Recursion will terminate, but produces weird outputs.
|
||||
|
||||
Naming segments is optional, in which case you will have to refer to it by its index.
|
||||
|
||||
@@ -71,7 +71,8 @@ class AttentionCoupleHook(TransformerOptionsHook):
|
||||
}
|
||||
self.has_negpip = False
|
||||
# calculate later. All clones must refer to the same kv dict
|
||||
self.kv = {"k": [], "v": []}
|
||||
# The type is here to shut up the type checker
|
||||
self.kv: dict[str, list] = {"k": None, "v": None}
|
||||
|
||||
def initialize_regions(self, base_cond, conds, fill):
|
||||
self.num_conds = len(conds) + 1
|
||||
|
||||
@@ -1,58 +1,25 @@
|
||||
from typing import TypeAlias
|
||||
import re
|
||||
|
||||
import lark
|
||||
from .utils import parse_args
|
||||
|
||||
from .parser import flatten
|
||||
|
||||
cut_parser = lark.Lark(
|
||||
r"""
|
||||
!start: (cut | prompt | /[][:()]/+)*
|
||||
prompt: (PLAIN | WHITESPACE)+
|
||||
cut: "[CUT:" prompt ":" prompt [":" NUMBER [ ":" NUMBER [":" NUMBER [ ":" PLAIN ] ] ] ]"]"
|
||||
WHITESPACE: /\s+/
|
||||
PLAIN: /([^\[\]:])+/
|
||||
%import common.SIGNED_NUMBER -> NUMBER
|
||||
"""
|
||||
)
|
||||
CUTOFF_RE = re.compile(r"\[CUT:((.*?):(.*?))\]")
|
||||
|
||||
|
||||
class CutTransform(lark.Transformer):
|
||||
def __default__(self, data, children, meta):
|
||||
return children
|
||||
def noop(x):
|
||||
return x
|
||||
|
||||
def NUMBER(self, args):
|
||||
return float(args)
|
||||
|
||||
def cut(self, args):
|
||||
prompt, cutout, weight, strict_mask, start_from_masked, mask_token = args
|
||||
|
||||
# prompts and cutouts are always sequences of str
|
||||
return (
|
||||
"".join(prompt),
|
||||
"".join(cutout),
|
||||
weight,
|
||||
strict_mask,
|
||||
start_from_masked,
|
||||
mask_token,
|
||||
def parse_cuts(string):
|
||||
text = CUTOFF_RE.sub(r"\2", string)
|
||||
cutoffs = CUTOFF_RE.findall(string)
|
||||
cs = []
|
||||
for x, *_ in cutoffs:
|
||||
p = x.split(":")
|
||||
args = parse_args(
|
||||
p, [(str, ""), (str, ""), (float, 0), (float, None), (float, None), (noop, None)], strip=False
|
||||
)
|
||||
|
||||
def start(self, args):
|
||||
prompt = []
|
||||
cuts = []
|
||||
for a in flatten(args):
|
||||
if isinstance(a, str):
|
||||
prompt.append(a)
|
||||
else:
|
||||
prompt.append(a[0])
|
||||
cuts.append(a)
|
||||
return "".join(prompt), cuts
|
||||
|
||||
def PLAIN(self, args: str) -> str:
|
||||
return str(args)
|
||||
|
||||
|
||||
CutResult: TypeAlias = tuple[str, str, float, float, float, str]
|
||||
|
||||
|
||||
def parse_cuts(text: str) -> tuple[str, CutResult]:
|
||||
return CutTransform().transform(cut_parser.parse(text))
|
||||
args = tuple(args)
|
||||
if not args[0] or not args[1] or (args[5] is not None and not args[5].strip()):
|
||||
raise ValueError(f"Invalid CUT spec: [CUT:{x}]")
|
||||
cs.append(args)
|
||||
return text, cs
|
||||
|
||||
@@ -4,12 +4,35 @@ from __future__ import annotations
|
||||
import logging
|
||||
import re
|
||||
|
||||
from .utils import find_closing_paren, get_function
|
||||
from .utils import find_closing_paren, get_function, split_by_function
|
||||
|
||||
logging.basicConfig()
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def substitute_template(template, segments):
|
||||
def _substitute(template, segments, stack):
|
||||
for name, value in sorted(segments):
|
||||
value = substitute_var(value, name, "")
|
||||
if name not in stack:
|
||||
stack.add(name)
|
||||
value = _substitute(value, segments, stack)
|
||||
stack.remove(name)
|
||||
template = substitute_var(template, name, value)
|
||||
return template
|
||||
|
||||
return _substitute(template, segments, set())
|
||||
|
||||
|
||||
def expand_segs(text):
|
||||
template, segments = split_by_function(text, "SEG", defaults=[""], require_args=True)
|
||||
named_segs = [(f.args[0].strip() or f"SEG{i + 1}", c.strip()) for i, (c, f) in enumerate(segments)]
|
||||
|
||||
new_text = substitute_template(template, named_segs).strip()
|
||||
if new_text != text.strip():
|
||||
log.debug("Template expanded to: %s", new_text)
|
||||
return new_text
|
||||
|
||||
|
||||
def parse_search(search):
|
||||
arg_start = search.find("(")
|
||||
args = ""
|
||||
@@ -60,6 +83,11 @@ def expand_macros(text):
|
||||
return res
|
||||
|
||||
|
||||
def substitute_var(text, name, replace):
|
||||
name = re.escape(str(name))
|
||||
return re.sub(rf"\${name}\b", replace, text)
|
||||
|
||||
|
||||
def substitute_defcall(text, search, replace):
|
||||
name, default_args = search
|
||||
text, defns = get_function(text, name, defaults=None, placeholder=f"DEFNCALL{name}", require_args=False)
|
||||
@@ -72,10 +100,10 @@ def substitute_defcall(text, search, replace):
|
||||
paramvals = [x.strip() for x in parameters[0].split(";")]
|
||||
r = replace
|
||||
for i, v in enumerate(paramvals):
|
||||
r = re.sub(rf"\${i + 1}\b", v, r)
|
||||
r = substitute_var(r, i + 1, v)
|
||||
|
||||
for i, v in enumerate(default_args):
|
||||
r = re.sub(rf"\${i + 1}\b", v, r)
|
||||
r = substitute_var(r, i + 1, v)
|
||||
|
||||
text = text.replace(ph, r)
|
||||
return text
|
||||
|
||||
@@ -2,6 +2,7 @@ import logging
|
||||
|
||||
from comfy_api.latest import io
|
||||
|
||||
from .macros import expand_segs
|
||||
from .prompts import encode_prompt
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
@@ -29,6 +30,7 @@ class PCTextEncodeWithRange(io.ComfyNode):
|
||||
log.debug("PCTextEncode: Encoding '%s'", text)
|
||||
defaults = clip.patcher.model_options.get("x-promptcontrol.defaults", {})
|
||||
masks = clip.patcher.model_options.get("x-promptcontrol.masks", None)
|
||||
text = expand_segs(text)
|
||||
out = encode_prompt(clip, text, start, end, defaults, masks)
|
||||
return io.NodeOutput(out)
|
||||
|
||||
|
||||
@@ -3,20 +3,15 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
|
||||
from comfy_api.latest import io
|
||||
from comfy_execution.graph import ExecutionBlocker
|
||||
from comfy_execution.graph_utils import GraphBuilder
|
||||
|
||||
from .parser import parse_prompt_schedules
|
||||
from .utils import consolidate_schedule, find_nonscheduled_loras, get_function
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
if os.environ.get("PC_USE_OLD_PARSER", "0") != "0":
|
||||
log.info("Using new parser implementation. Set PC_USE_OLD_PARSER=1 to use old parser instead")
|
||||
from .parser_parsy import parse_prompt_schedules as parse_prompt_schedules
|
||||
else:
|
||||
from .parser import parse_prompt_schedules
|
||||
|
||||
|
||||
def create_lora_loader_nodes(graph, model, clip, loras):
|
||||
|
||||
@@ -2,7 +2,8 @@ import logging
|
||||
|
||||
from comfy_api.latest import io
|
||||
|
||||
from .parser import expand_macros, parse_prompt_schedules
|
||||
from .macros import expand_macros
|
||||
from .parser import parse_prompt_schedules
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
+6
-363
@@ -1,367 +1,10 @@
|
||||
# vim: sw=4 ts=4
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from functools import lru_cache
|
||||
from math import ceil
|
||||
import os
|
||||
|
||||
import lark
|
||||
|
||||
from .macros import expand_macros
|
||||
|
||||
logging.basicConfig()
|
||||
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 flatten(x):
|
||||
if type(x) in [str, tuple, int, type(None)] or isinstance(x, dict) and "type" in x:
|
||||
yield x
|
||||
else:
|
||||
for g in x:
|
||||
yield from flatten(g)
|
||||
|
||||
|
||||
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)
|
||||
if os.environ.get("PC_USE_OLD_PARSER", "0") != "1":
|
||||
log.info("Using new parser implementation. Set PC_USE_OLD_PARSER=1 to use old parser instead")
|
||||
from .parser_parsy import parse_prompt_schedules # noqa
|
||||
else:
|
||||
from .parser_lark import parse_prompt_schedules # noqa
|
||||
|
||||
@@ -0,0 +1,359 @@
|
||||
# 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)
|
||||
@@ -348,9 +348,14 @@ non_special = regex(r"[^:\[\]()|\\<>#]+").map(Text)
|
||||
filename = regex(r"[^:<>]+")
|
||||
|
||||
comment = string("#") >> any_char.until(eof | char_from("\n")) >> success(empty)
|
||||
escape = (string("\\") >> char_from("\\[]:#")).map(Text)
|
||||
escape = (string("\\") >> char_from("\\[]:#") | string(r"\(") | string(r"\)")).map(Text)
|
||||
emphasis = seq(lpar, (prompt | col).at_least(0), rpar)
|
||||
number = (digit.at_least(1) + string(".") * 1 + digit.many() | digit.at_least(1)).concat().map(float)
|
||||
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())
|
||||
|
||||
@@ -241,7 +241,7 @@ def encode_prompt_segment(
|
||||
can_break = {}
|
||||
for k in empty:
|
||||
tokenizer = getattr(clip.tokenizer, f"clip_{k}", getattr(clip.tokenizer, k, None))
|
||||
can_break[k] = tokenizer and tokenizer.pad_to_max_length
|
||||
can_break[k] = tokenizer and getattr(tokenizer, "pad_to_max_length", False)
|
||||
|
||||
clip = hook_te(clip, empty.keys(), style, normalization, extra)
|
||||
|
||||
|
||||
@@ -35,6 +35,14 @@ except ImportError:
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def flatten(x):
|
||||
if type(x) in [str, tuple, int, type(None)] or isinstance(x, dict) and "type" in x:
|
||||
yield x
|
||||
else:
|
||||
for g in x:
|
||||
yield from flatten(g)
|
||||
|
||||
|
||||
def call_node(cls, *args, **kwargs):
|
||||
if hasattr(cls, "execute"):
|
||||
# v3 node
|
||||
|
||||
+2
-4
@@ -1,10 +1,8 @@
|
||||
[project]
|
||||
name = "comfyui-prompt-control"
|
||||
description = "Provides nodes for prompt editing and LoRA scheduling, advanced regional prompting (including attention masking) and more, all controlled through your text prompt"
|
||||
version = "3.0.0-beta.1"
|
||||
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.4"
|
||||
license = { file = "LICENSE" }
|
||||
# some lark versions older than 1.1.9 apparently have a bug that breaks things, see https://github.com/asagi4/comfyui-prompt-control/issues/35
|
||||
dependencies = ["lark >= 1.1.9"]
|
||||
|
||||
requires-python = ">= 3.10"
|
||||
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
from textwrap import dedent
|
||||
|
||||
import pytest
|
||||
|
||||
from prompt_control.macros import expand_macros, expand_segs
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"text, result",
|
||||
[
|
||||
("DEF(X(a;b)=$1 $2 $3 d)X(A) X(A;B;C)", "A b $3 d A B C d"),
|
||||
(
|
||||
"DEF(MACRO()=[empty:$1:$2])MACRO MACRO(;) MACRO(;0.5) MACRO(a;0.5)",
|
||||
"[empty::$2] [empty::] [empty::0.5] [empty:a:0.5]",
|
||||
),
|
||||
("DEF(X=$1)DEF(Y()=$1)[X Y][X() Y()][X(1) Y(1)]", "[$1 ][ ][1 1]"),
|
||||
],
|
||||
)
|
||||
def test_basic_macro(text, result):
|
||||
assert expand_macros(text) == result
|
||||
|
||||
|
||||
def test_macro_recursion():
|
||||
with pytest.raises(ValueError) as c:
|
||||
expand_macros("DEF(X=recurse Y) DEF(Y=recurse X) X")
|
||||
assert "Unable to resolve DEFs" in str(c.value)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"input, output",
|
||||
[
|
||||
(
|
||||
"""\
|
||||
A red $b and
|
||||
a blue $a
|
||||
SEG(a)
|
||||
cat
|
||||
SEG(b)
|
||||
|
||||
dog
|
||||
SEG(c)""",
|
||||
"A red dog and\na blue cat",
|
||||
),
|
||||
(
|
||||
"""\
|
||||
$a and $b
|
||||
SEG(a)
|
||||
cat, $b
|
||||
SEG(b)
|
||||
dog, $c
|
||||
SEG(c)
|
||||
tiger
|
||||
""",
|
||||
"cat, dog, tiger and dog, tiger",
|
||||
),
|
||||
(
|
||||
"""\
|
||||
$a
|
||||
SEG(a)
|
||||
a $b
|
||||
SEG(b)
|
||||
b $a""",
|
||||
"a b a b $a",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_segments(input, output):
|
||||
assert expand_segs(dedent(input)) == output
|
||||
+13
-23
@@ -2,10 +2,6 @@ import os
|
||||
|
||||
import pytest
|
||||
|
||||
from prompt_control.parser import expand_macros
|
||||
from prompt_control.parser import parse_prompt_schedules as old_parse # noqa
|
||||
from prompt_control.parser_parsy import parse_prompt_schedules as new_parse # noqa
|
||||
|
||||
|
||||
def lora_dict(*loras):
|
||||
return {lora: {"weight": unet, "weight_clip": te} for lora, unet, te in loras}
|
||||
@@ -27,9 +23,13 @@ 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)
|
||||
|
||||
|
||||
@@ -111,9 +111,9 @@ def test_basic_ok(parse):
|
||||
|
||||
@pytest.mark.parametrize("step", [0, 0.5, 1])
|
||||
def test_lora(step, parse):
|
||||
p = parse("This is a (lora:0.6) (prompt) with [no scheduling] features <lora:foo:0.5> <lora:bar:0.5:1.0>")
|
||||
p = parse("This is a (lora:0.6) (prompt) with [no scheduling] features <lora:foo:0.5> <lora:bar:0.5:-1.0>")
|
||||
expected = prompt(
|
||||
1.0, "This is a (lora:0.6) (prompt) with [no scheduling] features ", ("foo", 0.5, 0.5), ("bar", 0.5, 1.0)
|
||||
1.0, "This is a (lora:0.6) (prompt) with [no scheduling] features ", ("foo", 0.5, 0.5), ("bar", 0.5, -1.0)
|
||||
)
|
||||
assert prompts_match(p.at_step(step), expected)
|
||||
|
||||
@@ -221,23 +221,6 @@ def test_def(parse):
|
||||
p2 = parse("[(test):(test:0.7):0.7] [(test):(test:0.5):0.5]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
p = expand_macros("DEF(X(a;b)=$1 $2 $3 d)X(A) X(A;B;C)")
|
||||
assert p == "A b $3 d A B C d"
|
||||
|
||||
p = expand_macros("DEF(MACRO()=[empty:$1:$2])MACRO MACRO(;) MACRO(;0.5) MACRO(a;0.5)")
|
||||
assert p == "[empty::$2] [empty::] [empty::0.5] [empty:a:0.5]"
|
||||
|
||||
p = expand_macros("DEF(X=$1)DEF(Y()=$1)[X Y][X() Y()][X(1) Y(1)]")
|
||||
assert p == "[$1 ][ ][1 1]"
|
||||
|
||||
p = parse("DEF(test(1)=prompt $1)DEF(test2((a); (test))=[$1:$2:0.5])test test2")
|
||||
p2 = parse("prompt 1 [(a):(prompt 1):0.5]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
with pytest.raises(ValueError) as c:
|
||||
expand_macros("DEF(X=recurse Y) DEF(Y=recurse X) X")
|
||||
assert "Unable to resolve DEFs" in str(c.value)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"text, cases",
|
||||
@@ -249,6 +232,7 @@ def test_def(parse):
|
||||
),
|
||||
(r"[a\:b\\:c:0.5]", [(0.0, 0.5, "a:b\\"), (0.55, 1, r"c")]),
|
||||
(r"[a:\#b:0.5]", [(0.0, 0.5, "a"), (0.55, 1, "#b")]),
|
||||
(r"[a:b \(test\):0.2]", [(0, 0.2, r"a"), (0.25, 1, r"b \(test\)")]),
|
||||
],
|
||||
)
|
||||
def test_escapes(text, cases, parse):
|
||||
@@ -358,6 +342,12 @@ def test_cornercase_corrected(parse):
|
||||
assert p.parsed_prompt[1:] == p2.parsed_prompt
|
||||
|
||||
|
||||
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>")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
|
||||
def test_alternating_lora(parse):
|
||||
p4 = parse("[cat|[dog:wolf<lora:canine:1>:0.5]:0.2]")
|
||||
for i, (text, *_loras) in enumerate(
|
||||
|
||||
@@ -0,0 +1,139 @@
|
||||
{
|
||||
"1": {
|
||||
"inputs": {
|
||||
"text": "positive prompt",
|
||||
"clip": [
|
||||
"4",
|
||||
1
|
||||
]
|
||||
},
|
||||
"class_type": "PCLazyTextEncode",
|
||||
"_meta": {
|
||||
"title": "PC: Schedule prompt"
|
||||
}
|
||||
},
|
||||
"2": {
|
||||
"inputs": {
|
||||
"ckpt_name": "$TEST_CHECKPOINT"
|
||||
},
|
||||
"class_type": "CheckpointLoaderSimple",
|
||||
"_meta": {
|
||||
"title": "Load Checkpoint"
|
||||
}
|
||||
},
|
||||
"3": {
|
||||
"inputs": {
|
||||
"seed": 0,
|
||||
"steps": 8,
|
||||
"cfg": 3,
|
||||
"sampler_name": "euler",
|
||||
"scheduler": "simple",
|
||||
"denoise": 1,
|
||||
"model": [
|
||||
"4",
|
||||
0
|
||||
],
|
||||
"positive": [
|
||||
"9",
|
||||
0
|
||||
],
|
||||
"negative": [
|
||||
"9",
|
||||
1
|
||||
],
|
||||
"latent_image": [
|
||||
"5",
|
||||
0
|
||||
]
|
||||
},
|
||||
"class_type": "KSampler",
|
||||
"_meta": {
|
||||
"title": "KSampler"
|
||||
}
|
||||
},
|
||||
"4": {
|
||||
"inputs": {
|
||||
"text": "<lora:$TEST_LORA:1>",
|
||||
"model": [
|
||||
"2",
|
||||
0
|
||||
],
|
||||
"clip": [
|
||||
"2",
|
||||
1
|
||||
]
|
||||
},
|
||||
"class_type": "PCLazyLoraLoader",
|
||||
"_meta": {
|
||||
"title": "PC: Schedule LoRAs"
|
||||
}
|
||||
},
|
||||
"5": {
|
||||
"inputs": {
|
||||
"width": 1024,
|
||||
"height": 1024,
|
||||
"batch_size": 1
|
||||
},
|
||||
"class_type": "EmptyLatentImage",
|
||||
"_meta": {
|
||||
"title": "Empty Latent Image"
|
||||
}
|
||||
},
|
||||
"6": {
|
||||
"inputs": {
|
||||
"samples": [
|
||||
"3",
|
||||
0
|
||||
],
|
||||
"vae": [
|
||||
"2",
|
||||
2
|
||||
]
|
||||
},
|
||||
"class_type": "VAEDecode",
|
||||
"_meta": {
|
||||
"title": "VAE Decode"
|
||||
}
|
||||
},
|
||||
"7": {
|
||||
"inputs": {
|
||||
"images": [
|
||||
"6",
|
||||
0
|
||||
]
|
||||
},
|
||||
"class_type": "PreviewImage",
|
||||
"_meta": {
|
||||
"title": "Preview Image"
|
||||
}
|
||||
},
|
||||
"8": {
|
||||
"inputs": {
|
||||
"text": "worst quality,",
|
||||
"clip": [
|
||||
"4",
|
||||
1
|
||||
]
|
||||
},
|
||||
"class_type": "CLIPTextEncode",
|
||||
"_meta": {
|
||||
"title": "CLIP Text Encode (Prompt)"
|
||||
}
|
||||
},
|
||||
"9": {
|
||||
"inputs": {
|
||||
"positive": [
|
||||
"1",
|
||||
0
|
||||
],
|
||||
"negative": [
|
||||
"8",
|
||||
0
|
||||
]
|
||||
},
|
||||
"class_type": "PCAttentionCoupleBatchNegative",
|
||||
"_meta": {
|
||||
"title": "PC: Attention Couple (batch negative)"
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
from time import sleep
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
def workflow(request):
|
||||
with open(str(request.path).replace(".py", ".json")) as f:
|
||||
data = f.read()
|
||||
data = data.replace("$TEST_CHECKPOINT", os.environ["PC_TEST_CHECKPOINT"])
|
||||
data = data.replace("$TEST_LORA", os.environ["PC_TEST_LORA"])
|
||||
return json.loads(data)
|
||||
|
||||
|
||||
def assert_prompt(url, p):
|
||||
timeout = 60
|
||||
r = requests.post(f"{url}/prompt", json={"prompt": p, "client_id": str(uuid.uuid4())}).json()
|
||||
prompt_id = r["prompt_id"]
|
||||
r = {"status": "pending"}
|
||||
while r["status"] in ["pending", "in_progress"]:
|
||||
sleep(1)
|
||||
assert timeout > 0
|
||||
timeout -= 1
|
||||
r = requests.get(f"{url}/api/jobs/{prompt_id}").json()
|
||||
assert r["status"] == "completed"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def comfyui():
|
||||
return os.environ.get("PC_TEST_COMFYUI", "http://localhost:8188")
|
||||
|
||||
|
||||
def test_workflow(workflow, comfyui):
|
||||
prompt = "DEF(blue=green)a blue dog and a cat sitting [COUPLE(0 0.5, 0 1) red (cat,:1.3) COUPLE(0.5 1, 0 1) (blue:1.2) dog,:0.1]"
|
||||
workflow["1"]["inputs"]["text"] = prompt
|
||||
assert_prompt(comfyui, workflow)
|
||||
Reference in New Issue
Block a user