Compare commits

...
36 Commits
Author SHA1 Message Date
asagi4 9857c662af tests: COUPLE mask shortcut 2026-02-02 19:33:34 +02:00
asagi4 32f18ef338 Tests for step count 2026-02-02 19:26:13 +02:00
asagi4 d86d9f902c test: Verify that basic paren escapes don't get stripped 2026-02-02 19:26:13 +02:00
asagi4 f600b1ef90 Fix type complaint 2026-02-02 19:26:13 +02:00
asagi4 72484d1034 test only new parser 2026-02-02 19:26:13 +02:00
asagi4 e7b444316f Test multiple averages 2026-02-02 19:26:13 +02:00
asagi4 b9db1c7c4c Make github tests work again 2026-02-02 19:26:13 +02:00
asagi4 3ee001a3fa Disable method override complaint 2026-02-02 19:26:13 +02:00
asagi4 3fcd708dc6 Working importing for v3 nodes 2026-02-02 19:26:13 +02:00
asagi4 cb5ef8d97a v3: nodes_lazy.py 2026-02-02 19:26:13 +02:00
asagi4 24fe09ba4e v3: nodes_tools 2026-02-02 19:26:13 +02:00
asagi4 1c7149a811 v3: nodes_hooks.py 2026-02-02 19:26:13 +02:00
asagi4 b01bdc3d50 v3: nodes_base.py 2026-02-02 19:26:13 +02:00
asagi4 5ffd94d7d2 Initial v3 migration 2026-02-02 19:26:13 +02:00
asagi4 a433b246a7 handle setting steps 2026-02-02 19:26:13 +02:00
asagi4 bdd0c62665 parsy ruff fixes 2026-02-02 19:26:13 +02:00
asagi4 1c763f743c parsy loractl 2026-02-02 19:23:55 +02:00
asagi4 c223603100 simplify LoRA weight parsing 2026-02-02 19:23:55 +02:00
asagi4 c9efa87291 Rewrite parser to use parsy 2026-02-02 19:23:50 +02:00
asagi4 e15d4d693c Remove broken expand_graph.py 2026-02-02 19:19:21 +02:00
asagi4 26d5bfc961 Split cutoff parser to its own file 2026-02-02 19:19:21 +02:00
asagi4 cdb6c418eb split macros out of parser 2026-02-02 19:19:21 +02:00
asagi4 dc08fa93b7 Fix NOISE doing nothing 2026-02-02 19:19:21 +02:00
asagi4 e29f3ca2d0 Fix github workflows 2026-02-02 19:19:21 +02:00
asagi4 ee186a73e5 Remove old tests 2026-02-02 19:19:17 +02:00
asagi4 7440b309fc Fix LazyLoraLoader tests 2026-02-02 19:18:36 +02:00
asagi4 2ccc697311 Test for discovered corner case behaviour 2026-02-02 19:18:36 +02:00
asagi4 4559060535 Add a graph test for alternating 2026-02-02 19:18:36 +02:00
asagi4 60ceddbd73 Refactor tests for new parser 2026-02-02 19:18:36 +02:00
asagi4 6fd28cfca7 Use pytest tests 2026-02-02 19:18:36 +02:00
asagi4 11b1313ed4 Convert tests to pytest
Not 100% sure these fully work yet
2026-02-02 19:18:36 +02:00
asagi4 9be7133e48 Ruff fixes etc. 2026-02-02 19:18:36 +02:00
asagi4 22e78f08eb Typing fixes etc 2026-02-02 19:18:36 +02:00
asagi4 39befe2a34 Remove cache hack, it's broken anyway 2026-02-02 19:18:36 +02:00
asagi4 ea9017862f Add ruff and ty 2026-02-02 19:18:36 +02:00
asagi4 6ae6cf0d65 Add pyright 2026-02-02 19:18:36 +02:00
32 changed files with 2981 additions and 1422 deletions
+7 -2
View File
@@ -11,8 +11,13 @@ jobs:
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Check out ComfyUI
uses: actions/checkout@v4
with:
repository: comfyanonymous/ComfyUI
path: ComfyUI
- uses: actions/setup-python@v5
with:
python-version: '3.11'
- run: pip install -r requirements.txt
- run: python -m prompt_control.test_parser
- run: pip install pytest typing-extensions -r requirements.txt
- run: PYTHONPATH=ComfyUI pytest tests/test_parser.py
+2 -4
View File
@@ -31,12 +31,10 @@ jobs:
- name: install-torch
run: pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu
- name: install ComfyUI
run: pip install -r requirements.txt -r ComfyUI/requirements.txt
run: pip install pytest typing-extensions -r requirements.txt -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
run: sed -i "s/^cpu_state = CPUState.GPU/cpu_state = CPUState.CPU/g" ComfyUI/comfy/model_management.py
- name: Run graph tests
run: PYTHONPATH=ComfyUI python -m prompt_control.test_graph
- name: Run encoder tests (clip_l only)
run: PYTHONPATH=ComfyUI python -m prompt_control.test_encode
run: PYTHONPATH=ComfyUI pytest tests/test_graph.py tests/test_encode.py
+1
View File
@@ -1 +1,2 @@
__pycache__
.pyre
+12 -6
View File
@@ -1,21 +1,27 @@
ARGS=
all: format check test
@echo "Done"
check:
find . -name "*.py" | xargs pyflakes
ty check && ruff check
fix:
ruff check --fix
format:
find . -name "*.py" | xargs black -l 120
ruff format
test:
python -m prompt_control.test_parser
PYTHONPATH=../../ pytest tests/test_parser.py $(ARGS)
test_graph:
PYTHONPATH=../../ python -m prompt_control.test_graph
PYTHONPATH=../../ pytest tests/test_graph.py $(ARGS)
test_encode:
PYTHONPATH=../../ python -m prompt_control.test_encode --verbose
PYTHONPATH=../../ pytest tests/test_encode.py $(ARGS)
test_encode_both:
TEST_TE="clip_l t5" PYTHONPATH=../../ python -m prompt_control.test_encode --verbose
TEST_TE="clip_l t5" PYTHONPATH=../../ pytest tests/test_encode.py $(ARGS)
test_heavy: test_graph test_encode_both
+25 -16
View File
@@ -5,32 +5,41 @@
@description: 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.
"""
import logging
import os
import sys
import logging
import importlib
log = logging.getLogger("comfyui-prompt-control")
log.propagate = False
if not log.handlers:
h = logging.StreamHandler(sys.stdout)
h.setFormatter(logging.Formatter("[PromptControl] %(levelname)s: %(message)s"))
log.addHandler(h)
if os.environ.get("PROMPTCONTROL_DEBUG"):
log.setLevel(logging.DEBUG)
else:
log.setLevel(logging.INFO)
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
WEB_DIRECTORY = "web"
nodes = ["base", "lazy", "tools", "hooks"]
v1_modules = []
v3_modules = []
# Importing things here breaks pytest for whatever reason...
if "PYTEST_CURRENT_TEST" not in os.environ:
import importlib
for node in nodes:
mod = importlib.import_module(f".prompt_control.nodes_{node}", package=__name__)
NODE_CLASS_MAPPINGS.update(mod.NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(mod.NODE_DISPLAY_NAME_MAPPINGS)
from comfy_api.latest import ComfyExtension
if not log.handlers:
h = logging.StreamHandler(sys.stdout)
h.setFormatter(logging.Formatter("[PromptControl] %(levelname)s: %(message)s"))
log.addHandler(h)
for node in ["base", "hooks", "tools", "lazy"]:
mod = importlib.import_module(f".prompt_control.nodes_{node}", package=__name__)
v3_modules.append(mod)
class PromptControlExtension(ComfyExtension):
async def get_node_list(self):
r = []
for m in v3_modules:
r.extend(m.NODES)
return r
async def comfy_entrypoint():
return PromptControlExtension()
+24 -11
View File
@@ -1,8 +1,9 @@
import torch
import numpy as np
from math import copysign
import logging
import itertools
import logging
from math import copysign
import numpy as np
import torch
log = logging.getLogger("comfyui-prompt-control")
@@ -41,12 +42,18 @@ def weights_like(weights, emb):
def scale_to_norm(weights, word_ids, w_max):
top = np.max(weights)
w_max = min(top, w_max)
weights = [[w_max if id == 0 else (w / top) * w_max for w, id in zip(x, y)] for x, y in zip(weights, word_ids)]
weights = [
[w_max if id == 0 else (w / top) * w_max for w, id in zip(x, y, strict=False)]
for x, y in zip(weights, word_ids, strict=False)
]
return weights
def mask_word_id(tokens, word_ids, target_id, mask_token):
new_tokens = [[mask_token if wid == target_id else t for t, wid in zip(x, y)] for x, y in zip(tokens, word_ids)]
new_tokens = [
[mask_token if wid == target_id else t for t, wid in zip(x, y, strict=False)]
for x, y in zip(tokens, word_ids, strict=False)
]
mask = np.array(word_ids) == target_id
return (new_tokens, mask)
@@ -160,7 +167,7 @@ def apply_negpip(encoder, emb, pooled, **kwargs):
def norm_length(encoder, tokens, **kwargs):
word_ids = encoder.word_ids(tokens)
sums = dict(zip(*np.unique(word_ids, return_counts=True)))
sums = dict(zip(*np.unique(word_ids, return_counts=True), strict=False))
sums[0] = 1
tokens = [[(t, _norm_mag(w, sums[id]) if id != 0 else 1.0, id) for (t, w, id) in x] for x in tokens]
return tokens
@@ -169,7 +176,9 @@ def norm_length(encoder, tokens, **kwargs):
def norm_mean(encoder, tokens, **kwargs):
weights = encoder.weights(tokens)
word_ids = encoder.word_ids(tokens)
delta = 1 - np.mean([w for x, y in zip(weights, word_ids) for w, id in zip(x, y) if id != 0])
delta = 1 - np.mean(
[w for x, y in zip(weights, word_ids, strict=False) for w, id in zip(x, y, strict=False) if id != 0]
)
tokens = [[(t, w if id == 0 else w + delta, id) for (t, w, id) in x] for x in tokens]
return tokens
@@ -197,6 +206,7 @@ class AdvancedEncoder:
def add_encoder(cls, name, fn):
cls.STYLES[name] = fn
@classmethod
def add_normalization_op(cls, name, fn):
cls.NORMALIZATION_OPS[name] = fn
@@ -296,7 +306,7 @@ class AdvancedEncoder:
w_mix = np.diff([0] + w.tolist())
w_mix = torch.tensor(w_mix, dtype=embs.dtype, device=embs.device).reshape((-1, 1, 1))
weighted_emb = (w_mix * embs).sum(axis=0, keepdim=True)
weighted_emb = (w_mix * embs).sum(dim=0, keepdim=True)
pooled = pooled_base
if pooled is not None and self.max_length:
pooled = weighted_emb[0, self.max_length - 1 : self.max_length, :]
@@ -304,7 +314,9 @@ class AdvancedEncoder:
def from_masked(self, tokens, weights, word_ids, base_emb, pooled_base):
wids, inds = np.unique(np.array(word_ids).reshape(-1), return_index=True)
weight_dict = dict((id, w) for id, w in zip(wids, np.array(weights).reshape(-1)[inds]) if w != 1.0)
weight_dict = dict(
(id, w) for id, w in zip(wids, np.array(weights).reshape(-1)[inds], strict=False) if w != 1.0
)
if len(weight_dict) == 0:
return torch.zeros_like(base_emb), torch.zeros_like(pooled_base) if pooled_base is not None else None
@@ -328,12 +340,13 @@ class AdvancedEncoder:
masks = torch.cat(masks)
embs = base_emb.expand(embs.shape) - embs
pooled = None
if pooled_base is not None and self.max_length:
pooled = embs[0, self.max_length - 1 : self.max_length, :]
pooled_start = pooled_base.expand(len(ws), -1)
ws = torch.tensor(ws).reshape(-1, 1).expand(pooled_start.shape)
pooled = (pooled - pooled_start) * (ws - 1)
pooled = pooled.mean(axis=0, keepdim=True)
pooled = pooled.mean(dim=0, keepdim=True)
pooled = pooled_base + pooled
if embs.shape[0] != masks.shape[0]:
+3 -4
View File
@@ -9,7 +9,6 @@ from typing import Any
import torch
import torch.nn.functional as F
from comfy.hooks import EnumHookScope, HookGroup, TransformerOptionsHook, set_hooks_for_conditioning
from comfy.model_patcher import ModelPatcher
@@ -71,14 +70,13 @@ class AttentionCoupleHook(TransformerOptionsHook):
}
}
self.has_negpip = False
# calculate later. All clones must refer to the same kv dict
self.kv = {"k": None, "v": None}
self.kv = {"k": [], "v": []}
def initialize_regions(self, base_cond, conds, fill):
self.num_conds = len(conds) + 1
self.base_strength = base_cond[1].get("strength", 1.0)
self.strengths = [cond[1].get("strength", 1.0) for cond in conds]
self.strengths: list[float] = [cond[1].get("strength", 1.0) for cond in conds]
self.conds: list[torch.Tensor] = [base_cond[0]] + [cond[0] for cond in conds]
base_mask = base_cond[1].get("mask", None)
masks = [cond[1].get("mask") * cond[1].get("mask_strength") for cond in conds]
@@ -215,6 +213,7 @@ class AttentionCoupleHook(TransformerOptionsHook):
dim=0,
)
)
assert self.num_conds is not None, "this is a bug"
cond_or_uncond_couple.extend(itertools.repeat(self.COND, self.num_conds))
q = torch.cat(qs, dim=0)
-47
View File
@@ -1,47 +0,0 @@
import comfy_execution.caching
from comfy_execution.graph_utils import is_link
import nodes
from os import environ
import logging
log = logging.getLogger("comfyui-prompt-control")
include_unique_id_in_input = comfy_execution.caching.include_unique_id_in_input
def promptcontrol_get_immediate_node_signature(self, dynprompt, node_id, ancestor_order_mapping):
if not dynprompt.has_node(node_id):
# This node doesn't exist -- we can't cache it.
return [float("NaN")]
node = dynprompt.get_node(node_id)
class_type = node["class_type"]
class_def = nodes.NODE_CLASS_MAPPINGS[class_type]
inputs = node["inputs"]
if hasattr(class_def, "CACHE_KEY"):
inputs = getattr(class_def, "CACHE_KEY")(inputs)
signature = [class_type, self.is_changed_cache.get(node_id)]
if (
self.include_node_id_in_input()
or (hasattr(class_def, "NOT_IDEMPOTENT") and class_def.NOT_IDEMPOTENT)
or include_unique_id_in_input(class_type)
):
signature.append(node_id)
for key in sorted(inputs.keys()):
if is_link(inputs[key]):
(ancestor_id, ancestor_socket) = inputs[key]
ancestor_index = ancestor_order_mapping[ancestor_id]
signature.append((key, ("ANCESTOR", ancestor_index, ancestor_socket)))
else:
signature.append((key, inputs[key]))
return signature
def init():
if environ.get("PROMPTCONTROL_ENABLE_CACHE_HACK") != "1":
return
log.warning("Enabling Prompt Control cache hack")
comfy_execution.caching.CacheKeySetInputSignature.get_immediate_node_signature = (
promptcontrol_get_immediate_node_signature
)
+8 -9
View File
@@ -1,9 +1,9 @@
import torch
import copy
import logging
import re
import numpy as np
import logging
import torch
log = logging.getLogger("comfyui-prompt-control")
@@ -69,10 +69,7 @@ def cutoff_add_region(
clip_regions["start_from_masked"] = float(start_from_masked)
if mask_token is not None:
clip_regions["mask_token"] = tokenizer.tokenizer(mask_token)["input_ids"][1]
if weight is None:
weight = 1.0
else:
weight = float(weight)
weight = 1.0 if weight is None else float(weight)
region_text = region_text.strip()
target_text = target_text.strip()
@@ -139,7 +136,7 @@ def cutoff_add_region(
def create_masked_prompt(weighted_tokens, mask, mask_token):
mask_ids = list(zip(*np.nonzero(mask.reshape((len(weighted_tokens), -1)))))
mask_ids = list(zip(*np.nonzero(mask.reshape((len(weighted_tokens), -1))), strict=False))
new_prompt = copy.deepcopy(weighted_tokens)
for x, y in mask_ids:
new_prompt[x][y] = (mask_token,) + new_prompt[x][y][1:]
@@ -200,7 +197,9 @@ def encode_regions(clip_regions, encode, tokenizer):
base_embedding_outer = base_embedding_full * (1 - strict_mask) + base_embedding_masked * strict_mask
region_embeddings = []
for region, target, weight in zip(clip_regions["regions"], clip_regions["targets"], clip_regions["weights"]):
for region, target, weight in zip(
clip_regions["regions"], clip_regions["targets"], clip_regions["weights"], strict=False
):
region_masking = torch.tensor(
regions_normalized * region * weight, dtype=base_embedding_full.dtype, device=base_embedding_full.device
).unsqueeze(-1)
@@ -215,7 +214,7 @@ def encode_regions(clip_regions, encode, tokenizer):
region_emb *= region_masking
region_embeddings.append(region_emb)
region_embeddings = torch.stack(region_embeddings).sum(axis=0)
region_embeddings = torch.stack(region_embeddings).sum(dim=0)
embeddings_final_mask = torch.tensor(
global_region_mask, dtype=base_embedding_full.dtype, device=base_embedding_full.device
+58
View File
@@ -0,0 +1,58 @@
from typing import TypeAlias
import lark
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
"""
)
class CutTransform(lark.Transformer):
def __default__(self, data, children, meta):
return children
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 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))
+81
View File
@@ -0,0 +1,81 @@
# vim: sw=4 ts=4
from __future__ import annotations
import logging
import re
from .utils import find_closing_paren, get_function
logging.basicConfig()
log = logging.getLogger("comfyui-prompt-control")
def parse_search(search):
arg_start = search.find("(")
args = ""
name = search.strip()
if arg_start > 0:
arg_end = find_closing_paren(search, arg_start + 1)
if arg_end < 0:
arg_end = len(search)
name = search[:arg_start].strip()
args = search[arg_start + 1 : arg_end]
if not name:
return None
args = args.strip()
# If using the form DEF(F()=$1) then the default value of $1 is the empty string
args = [a.strip() for a in args.split(";")] if arg_start > 0 else []
return name, args
def expand_macros(text):
text, defs = get_function(text, "DEF", defaults=None)
res = text
prevres = text
replacements = []
for d in defs:
if not d.args:
continue
r = d.args[0].split("=", 1)
search = parse_search(r[0].strip())
if not search or len(r) != 2:
log.warning("Ignoring invalid DEF(%s)", d)
continue
replacements.append((search, r[1].strip()))
iterations = 0
while True:
iterations += 1
if iterations > 10:
raise ValueError("Unable to resolve DEFs, make sure there are no cycles!")
return text
for search, replace in replacements:
res = substitute_defcall(res, search, replace)
if res == prevres:
break
prevres = res
if res.strip() != text.strip():
res = res.strip()
log.info("DEFs expanded to: %s", res)
return res
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
paramvals = []
if parameters:
paramvals = [x.strip() for x in parameters[0].split(";")]
r = replace
for i, v in enumerate(paramvals):
r = re.sub(rf"\${i + 1}\b", v, r)
for i, v in enumerate(default_args):
r = re.sub(rf"\${i + 1}\b", v, r)
text = text.replace(ph, r)
return text
-46
View File
@@ -1,46 +0,0 @@
import main
import nodes
import prompt_control.adv_encode
(l,) = nodes.CLIPLoader.load_clip(None, "clip_l.safetensors")
(t5,) = nodes.CLIPLoader.load_clip(None, "t5base.safetensors")
id(main) # get rid of warning
def adv(t, text, style="A1111", norm="none", new=True, **kwargs):
c = t.tokenize(text, return_word_ids=True)
if new:
style = "new+" + style
if t is t5:
te = t.patcher.model.t5base.encode_token_weights
token = t.tokenizer.clip_t5base
tok = c["t5base"]
else:
te = t.patcher.model.clip_l.encode_token_weights
token = t.tokenizer.clip_l
tok = c["l"]
return prompt_control.adv_encode.advanced_encode_from_tokens(tok, norm, style, te, tokenizer=token)
def adv_all(t, text, styles=[], **kwargs):
r = []
for s in styles or prompt_control.adv_encode.AdvancedEncoder.STYLES:
print("Testing", s, kwargs)
r.append([s, adv(t, text, style=s, **kwargs)])
return r
def replacenan(t):
t[t.isnan()] = 42.123321
return t
def adv_equal(t, text, **kwargs):
old = adv_all(t, text, new=False, **kwargs)
new = adv_all(t, text, new=True, **kwargs)
r = {}
for i, o in enumerate(old):
n = new[i]
r[n[0]] = (replacenan(n[1][0]) == replacenan(o[1][0])).all()
return r
+43 -34
View File
@@ -1,51 +1,60 @@
import logging
from comfy_api.latest import io
from .prompts import encode_prompt
log = logging.getLogger("comfyui-prompt-control")
class PCTextEncodeWithRange:
class PCTextEncodeWithRange(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP",), "text": ("STRING", {"multiline": True})},
"optional": {
"start": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 0.0, "step": 0.01}),
"end": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 1.0, "step": 0.01}),
},
}
def define_schema(cls):
return io.Schema(
node_id="PCTextEncodeWithRange",
display_name="PC: Text Encode with Range (no scheduling)",
category="promptcontrol/tools",
description="Like PCTextEncode, but if you know the range you need for a prompt, can be slightly more efficient when you have LoRAs scheduled on a CLIP model.",
inputs=[
io.Clip.Input("clip"),
io.String.Input("text", multiline=True),
io.Float.Input("start", default=0.0, min=0.0, max=1.0, step=0.01, optional=True),
io.Float.Input("end", default=1.0, min=0.0, max=1.0, step=0.01, optional=True),
],
outputs=[io.Conditioning.Output()],
)
RETURN_TYPES = ("CONDITIONING",)
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Like PCTextEncode, but if you know the range you need for a prompt, can be slightly more efficient when you have LoRAs scheduled on a CLIP model"
def apply(self, clip, text, start=0.0, end=1.0):
@classmethod
def execute(cls, clip, text, start=0.0, end=1.0) -> io.NodeOutput: # ty: ignore[invalid-method-override]
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)
return (encode_prompt(clip, text, start, end, defaults, masks),)
out = encode_prompt(clip, text, start, end, defaults, masks)
return io.NodeOutput(out)
class PCTextEncode:
class PCTextEncode(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP",), "text": ("STRING", {"multiline": True})},
}
def define_schema(cls):
return io.Schema(
node_id="PCTextEncode",
display_name="PC: Text Encode (no scheduling)",
category="promptcontrol",
description="Encodes a prompt with extra goodies from Prompt Control. This node does *not* support scheduling.",
inputs=[
io.Clip.Input("clip"),
io.String.Input("text", multiline=True),
],
outputs=[io.Conditioning.Output()],
)
RETURN_TYPES = ("CONDITIONING",)
CATEGORY = "promptcontrol"
FUNCTION = "apply"
DESCRIPTION = "Encodes a prompt with extra goodies from Prompt Control. This node does *not* support scheduling"
def apply(self, clip, text):
return PCTextEncodeWithRange.apply(self, clip, text, 0.0, 1.0)
@classmethod
def execute(cls, clip, text) -> io.NodeOutput: # ty: ignore[invalid-method-override]
# Use the WithRange node for the range 0.0, 1.0
return PCTextEncodeWithRange.execute(clip, text, 0.0, 1.0)
NODE_CLASS_MAPPINGS = {"PCTextEncode": PCTextEncode, "PCTextEncodeWithRange": PCTextEncodeWithRange}
NODE_DISPLAY_NAME_MAPPINGS = {
"PCTextEncode": "PC: Text Encode (no scheduling)",
"PCTextEncodeWithRange": "PC: Text Encode with Range (no scheduling)",
}
NODES = [
PCTextEncodeWithRange,
PCTextEncode,
]
+49 -51
View File
@@ -3,7 +3,8 @@ import logging
import comfy.hooks
import comfy.utils
import folder_paths
from comfy.comfy_types.node_typing import IO, ComfyNodeABC, InputTypeDict
from comfy_api.latest import io
from typing_extensions import override
from .attention_couple_ppm import AttentionCoupleHook
from .parser import parse_prompt_schedules
@@ -12,24 +13,27 @@ from .utils import consolidate_schedule
log = logging.getLogger("comfyui-prompt-control")
class PCLoraHooksFromText:
class PCLoraHooksFromText(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {
"required": {"text": ("STRING",)},
}
def define_schema(cls):
return io.Schema(
node_id="PCLoraHooksFromText",
display_name="PC: LoRA Hooks From Text (non-lazy)",
category="promptcontrol/v2",
description="set of hooks created from the prompt schedule",
is_experimental=True,
inputs=[
io.String.Input("text", multiline=True),
],
outputs=[io.Hooks.Output()],
)
RETURN_TYPES = ("HOOKS",)
OUTPUT_TOOLTIPS = ("set of hooks created from the prompt schedule",)
CATEGORY = "promptcontrol/v2"
FUNCTION = "apply"
EXPERIMENTAL = True
def apply(self, text):
@classmethod
def execute(cls, text) -> io.NodeOutput: # ty: ignore[invalid-method-override]
prompt_schedule = parse_prompt_schedules(text)
consolidated = consolidate_schedule(prompt_schedule)
hooks = lora_hooks_from_schedule(consolidated, {})
return (hooks,)
return io.NodeOutput(hooks)
def lora_hooks_from_schedule(schedules, non_scheduled):
@@ -37,8 +41,7 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
lora_cache = {}
all_hooks = []
def create_hook(loraspec, start_pct, end_pct, non_scheduled):
nonlocal lora_cache
def create_hook(loras, start_pct, end_pct, non_scheduled):
hooks = []
hook_kf = comfy.hooks.HookKeyframeGroup()
for path, info in loras.items():
@@ -52,8 +55,8 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
new_hook = comfy.hooks.create_hook_lora(
lora_cache[path], strength_model=info["weight"], strength_clip=info["weight_clip"]
)
# Set hook_ref so that identical hooks compare equal
new_hook.hooks[0].hook_ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
new_hook.hooks[0].hook_ref = ref
hooks.append(new_hook)
if start_pct > 0.0:
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=0.0)
@@ -74,43 +77,43 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
all_hooks.append(hook)
start_pct = end_pct
del lora_cache
all_hooks = [x for x in all_hooks if x]
if all_hooks:
hooks = comfy.hooks.HookGroup.combine_all_hooks(all_hooks)
return hooks
class PCAttentionCoupleBatchNegative(ComfyNodeABC):
class PCAttentionCoupleBatchNegative(io.ComfyNode):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
return {
"required": {
"positive": (IO.CONDITIONING, {}),
"negative": (IO.CONDITIONING, {}),
},
}
def define_schema(cls):
return io.Schema(
node_id="PCAttentionCoupleBatchNegative",
display_name="PC: Attention Couple (batch negative)",
category="promptcontrol/v2",
description="Batch negatives, carrying over Attention Couple hooks",
is_experimental=True,
inputs=[
io.Conditioning.Input("positive"),
io.Conditioning.Input("negative"),
],
outputs=[
io.Conditioning.Output("positive"),
io.Conditioning.Output("negative"),
],
)
RETURN_TYPES = (IO.CONDITIONING, IO.CONDITIONING)
RETURN_NAMES = ("positive", "negative")
CATEGORY = "promptcontrol/v2"
FUNCTION = "batch"
EXPERIMENTAL = True
# May cause side-effects?
# TODO: Support scheduling in negative prompt
def batch(self, positive, negative):
@classmethod
@override
def execute(cls, positive, negative) -> io.NodeOutput: # ty: ignore[invalid-method-override]
if len(negative) != 1:
log.warning("Batching scheduled negatives is not supported yet")
return (positive, negative)
return io.NodeOutput(positive, negative)
negative_batch = []
for p in positive:
n = [negative[0][0], negative[0][1].copy()]
n_hook_group: comfy.hooks.HookGroup = n[1].get("hooks", comfy.hooks.HookGroup()).clone()
p_hook_group: comfy.hooks.HookGroup = p[1].get("hooks", comfy.hooks.HookGroup())
n_hook_group = n[1].get("hooks", comfy.hooks.HookGroup()).clone()
p_hook_group = p[1].get("hooks", comfy.hooks.HookGroup())
attn_couple = [hook for hook in p_hook_group.hooks if isinstance(hook, AttentionCoupleHook)]
for hook in attn_couple:
n_hook_group.add(hook)
@@ -119,15 +122,10 @@ class PCAttentionCoupleBatchNegative(ComfyNodeABC):
n[1]["end_percent"] = p[1].get("end_percent", 1.0)
negative_batch.append(n)
return (positive, negative_batch)
return io.NodeOutput(positive, negative_batch)
NODE_CLASS_MAPPINGS = {
"PCLoraHooksFromText": PCLoraHooksFromText,
"PCAttentionCoupleBatchNegative": PCAttentionCoupleBatchNegative,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"PCLoraHooksFromText": "PC: LoRA Hooks From Text (non-lazy)",
"PCAttentionCoupleBatchNegative": "PC: Attention Couple (batch negative)",
}
NODES = [
PCLoraHooksFromText,
PCAttentionCoupleBatchNegative,
]
+104 -113
View File
@@ -1,31 +1,22 @@
# pyright: reportSelfClsParameterName=false
from __future__ import annotations
import json
import logging
from .parser import parse_prompt_schedules
from comfy_execution.graph_utils import GraphBuilder, is_link
import os
from comfy_api.latest import io
from comfy_execution.graph import ExecutionBlocker
from comfy_execution.graph_utils import GraphBuilder
from .utils import get_function
from .utils import consolidate_schedule, find_nonscheduled_loras, get_function
log = logging.getLogger("comfyui-prompt-control")
from .utils import consolidate_schedule, find_nonscheduled_loras
import json
def _cache_key(cachekey, inputs):
out = inputs.copy()
text = inputs.get("text")
if text is not None and not is_link(text):
out["text"] = cache_key_from_inputs(cachekey, **inputs)
return out
def cache_key_prompt(inputs):
return _cache_key("prompt", inputs)
def cache_key_lora(inputs):
return _cache_key("loras", inputs)
if os.environ.get("PC_USE_NEW_PARSER", "0") == "1":
log.info("Using new parsy parser")
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):
@@ -145,64 +136,63 @@ def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True):
ret = (model, clip, res)
return {"result": ret, "expand": r}
return io.NodeOutput(*ret, expand=r)
class PCLazyLoraLoaderAdvanced:
CACHE_KEY = cache_key_lora
class PCLazyLoraLoaderAdvanced(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="PCLazyLoraLoaderAdvanced",
display_name="PC: Schedule LoRAs (Advanced)",
enable_expand=True,
category="promptcontrol",
description="Returns a model and clip with LoRAs scheduled",
inputs=[
io.Model.Input("model", extra_dict={"rawLink": True}, optional=True),
io.Clip.Input("clip", extra_dict={"rawLink": True}, optional=True),
io.String.Input("text", multiline=True, default=""),
io.Boolean.Input("apply_hooks", default=True),
io.String.Input("tags", default=""),
io.Float.Input("start", min=0.0, max=1.0, default=0.0, step=0.01),
io.Float.Input("end", min=0.0, max=1.0, default=1.0, step=0.01),
io.Int.Input("num_steps", min=0, max=10000, default=0, step=1),
],
outputs=[io.Model.Output("model"), io.Clip.Output("clip"), io.Hooks.Output("hooks")],
)
@classmethod
def INPUT_TYPES(s):
return {
"optional": {
"model": ("MODEL", {"rawLink": True}),
"clip": ("CLIP", {"rawLink": True}),
"text": ("STRING", {"multiline": True, "default": ""}),
"apply_hooks": ("BOOLEAN", {"default": True}),
"tags": ("STRING", {"default": ""}),
"start": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 0.0, "step": 0.01}),
"end": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 1.0, "step": 0.01}),
"num_steps": ("INT", {"min": 0, "max": 10000, "default": 0, "step": 1}),
},
"hidden": {"unique_id": "UNIQUE_ID"},
}
RETURN_TYPES = ("MODEL", "CLIP", "HOOKS")
OUTPUT_TOOLTIPS = ("Returns a model and clip with LoRAs scheduled",)
CATEGORY = "promptcontrol"
FUNCTION = "apply"
def apply(
self, unique_id, model=None, clip=None, text="", apply_hooks=True, tags="", start=0.0, end=1.0, num_steps=0
):
def execute(cls, model=None, clip=None, text="", apply_hooks=True, tags="", start=0.0, end=1.0, num_steps=0):
schedule = parse_prompt_schedules(text, filters=tags, start=start, end=end, num_steps=num_steps)
graph = GraphBuilder()
r = build_lora_schedule(graph, schedule, model, clip, apply_hooks=apply_hooks)
return r
class PCLazyLoraLoader(PCLazyLoraLoaderAdvanced):
class PCLazyLoraLoader(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {
"optional": {
"model": ("MODEL", {"rawLink": True}),
"clip": ("CLIP", {"rawLink": True}),
"text": ("STRING", {"multiline": True, "default": ""}),
},
"hidden": {"unique_id": "UNIQUE_ID"},
}
def define_schema(cls):
return io.Schema(
node_id="PCLazyLoraLoader",
display_name="PC: Schedule LoRAs",
enable_expand=True,
category="promptcontrol",
description="Returns a model and clip with LoRAs scheduled",
inputs=[
io.Model.Input("model", extra_dict={"rawLink": True}, optional=True),
io.Clip.Input("clip", extra_dict={"rawLink": True}, optional=True),
io.String.Input("text", multiline=True, default=""),
],
outputs=[
io.Model.Output("model"),
io.Clip.Output("clip"),
],
)
RETURN_TYPES = (
"MODEL",
"CLIP",
)
CATEGORY = "promptcontrol"
def apply(self, *args, **kwargs):
r = super().apply(*args, **kwargs)
r["result"] = r["result"][:2]
return r
@classmethod
def execute(cls, model, clip, text):
no = PCLazyLoraLoaderAdvanced.execute(model, clip, text)
return io.NodeOutput(*no.args[:2], expand=no.expand)
def build_scheduled_prompts(graph, schedules, clip):
@@ -234,61 +224,62 @@ def build_scheduled_prompts(graph, schedules, clip):
g = graph.finalize()
log.debug("Built graph: %s", json.dumps(g))
return {"result": (node.out(0),), "expand": g}
return io.NodeOutput(node.out(0), expand=g)
def cache_key_from_inputs(cachekey, text, tags="", start=0.0, end=1.0, num_steps=0, **kwargs):
schedules = parse_prompt_schedules(text, filters=tags, start=start, end=end, num_steps=num_steps)
return [(pct, s[cachekey]) for pct, s in schedules]
class PCLazyTextEncodeAdvanced:
CACHE_KEY = cache_key_prompt
class PCLazyTextEncodeAdvanced(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="PCLazyTextEncodeAdvanced",
display_name="PC: Schedule prompt (Advanced)",
enable_expand=True,
category="promptcontrol",
inputs=[
io.Clip.Input("clip", extra_dict={"rawLink": True}),
io.String.Input("text", multiline=True, default=""),
io.String.Input("tags", default=""),
io.Float.Input("start", min=0.0, max=1.0, default=0.0, step=0.01),
io.Float.Input("end", min=0.0, max=1.0, default=1.0, step=0.01),
io.Int.Input("num_steps", min=0, max=10000, default=0, step=1),
],
outputs=[
io.Conditioning.Output("conditioning"),
],
)
@classmethod
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP", {"rawLink": True}), "text": ("STRING", {"multiline": True})},
"optional": {
"tags": ("STRING", {"default": ""}),
"start": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 0.0, "step": 0.01}),
"end": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 1.0, "step": 0.01}),
"num_steps": ("INT", {"min": 0, "max": 10000, "default": 0, "step": 1}),
},
"hidden": {"unique_id": "UNIQUE_ID"},
}
RETURN_TYPES = ("CONDITIONING",)
CATEGORY = "promptcontrol"
FUNCTION = "apply"
def apply(self, clip, text, unique_id, tags="", start=0.0, end=1.0, num_steps=0):
def execute(cls, clip, text, tags="", start=0.0, end=1.0, num_steps=0):
schedules = parse_prompt_schedules(text, filters=tags, start=start, end=end, num_steps=num_steps)
graph = GraphBuilder()
return build_scheduled_prompts(graph, schedules, clip)
class PCLazyTextEncode(PCLazyTextEncodeAdvanced):
class PCLazyTextEncode(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP", {"rawLink": True}), "text": ("STRING", {"multiline": True})},
"hidden": {"unique_id": "UNIQUE_ID"},
}
def define_schema(cls):
return io.Schema(
node_id="PCLazyTextEncode",
display_name="PC: Schedule prompt",
enable_expand=True,
category="promptcontrol",
inputs=[
io.Clip.Input("clip", extra_dict={"rawLink": True}),
io.String.Input("text", multiline=True, default=""),
],
outputs=[
io.Conditioning.Output("conditioning"),
],
)
CATEGORY = "promptcontrol"
@classmethod
def execute(cls, clip, text):
return PCLazyTextEncodeAdvanced.execute(clip, text)
NODE_CLASS_MAPPINGS = {
"PCLazyTextEncode": PCLazyTextEncode,
"PCLazyTextEncodeAdvanced": PCLazyTextEncodeAdvanced,
"PCLazyLoraLoader": PCLazyLoraLoader,
"PCLazyLoraLoaderAdvanced": PCLazyLoraLoaderAdvanced,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"PCLazyTextEncode": "PC: Schedule Prompt",
"PCLazyTextEncodeAdvanced": "PC: Schedule prompt (Advanced)",
"PCLazyLoraLoader": "PC: Schedule LoRAs",
"PCLazyLoraLoaderAdvanced": "PC: Schedule LoRAs (Advanced)",
}
NODES = [
PCLazyTextEncode,
PCLazyTextEncodeAdvanced,
PCLazyLoraLoader,
PCLazyLoraLoaderAdvanced,
]
+119 -171
View File
@@ -1,149 +1,106 @@
import logging
from .parser import parse_prompt_schedules, expand_macros
from .nodes_lazy import NODE_CLASS_MAPPINGS as LAZY_NODES
from .utils import expand_graph
import json
import folder_paths
from pathlib import Path
from comfy_api.latest import io
from .parser import expand_macros, parse_prompt_schedules
log = logging.getLogger("comfyui-prompt-control")
class PCSaveExpandedWorkflow:
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
class PCSetLogLevel(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"any": ("*", {}),
},
"hidden": {
"prompt": "PROMPT",
},
}
@classmethod
def VALIDATE_INPUTS(self, input_types):
return True
OUTPUT_NODE = True
RETURN_TYPES = ()
CATEGORY = "promptcontrol/tools"
DESCRIPTION = "Expands lazy prompt control nodes in the prompt and saves the expanded prompt into a JSON file"
FUNCTION = "apply"
def apply(self, any, prompt):
full_output_folder, filename, counter, subfolder, prefix = folder_paths.get_save_image_path(
"pc_workflow_debug", self.output_dir
def define_schema(cls):
return io.Schema(
node_id="PCSetLogLevel",
display_name="PC: Configure Logging (for debug)",
category="promptcontrol/tools",
description="A debug node to configure Prompt Control logging level. Pass a CLIP through it before you run any PC nodes",
inputs=[
io.Clip.Input("clip"),
io.Combo.Input("level", options=["INFO", "DEBUG", "WARNING", "ERROR"], default="INFO", optional=True),
],
outputs=[io.Clip.Output()],
)
expanded = expand_graph(LAZY_NODES, prompt)
file = f"{filename}_{counter:05}_.json"
full_path = Path(full_output_folder) / file
with open(full_path, "w") as f:
log.info(f"Saving workflow to {full_path}")
json.dump(expanded, f)
return ()
class PCSetLogLevel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"clip": ("CLIP",),
},
"optional": {
"level": (["INFO", "DEBUG", "WARNING", "ERROR"], {"default": "INFO"}),
},
}
def apply(self, clip, level="INFO"):
def execute(cls, clip, level="INFO") -> io.NodeOutput:
log.setLevel(getattr(logging, level))
log.info("Set logging level to %s", level)
return (clip,)
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/tools"
DESCRIPTION = (
"A debug node to configure Prompt Control logging level. Pass a CLIP through it before you run any PC nodes"
)
FUNCTION = "apply"
return io.NodeOutput(clip)
class PCAddMaskToCLIP:
class PCAddMaskToCLIP(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP",)},
"optional": {
"mask": ("MASK",),
},
}
def define_schema(cls):
return io.Schema(
node_id="PCAddMaskToCLIP",
display_name="PC: Attach Mask",
category="promptcontrol/tools",
description="Attaches a mask to a CLIP object so that they can be referred to in a prompt using IMASK(). Using this node multiple times adds more masks rather than replacing existing ones.",
inputs=[
io.Clip.Input("clip"),
io.Mask.Input("mask", optional=True),
],
outputs=[io.Clip.Output()],
)
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Attaches a mask to a CLIP object so that they can be referred to in a prompt using IMASK(). Using this node multiple times adds more masks rather than replacing existing ones."
def apply(self, clip, mask=None):
return PCAddMaskToCLIPMany().apply(clip, mask1=mask)
class PCAddMaskToCLIPMany:
@classmethod
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP",)},
"optional": {
"mask1": ("MASK",),
"mask2": ("MASK",),
"mask3": ("MASK",),
"mask4": ("MASK",),
},
}
def execute(cls, clip, mask=None) -> io.NodeOutput:
return PCAddMaskToCLIPMany.execute(clip, mask1=mask)
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Multi-input version of PCAddMaskToCLIP, for convenience"
def apply(self, clip, mask1=None, mask2=None, mask3=None, mask4=None):
class PCAddMaskToCLIPMany(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="PCAddMaskToCLIPMany",
display_name="PC: Attach Mask (multi)",
category="promptcontrol/tools",
description="Multi-input version of PCAddMaskToCLIP, for convenience",
inputs=[
io.Clip.Input("clip"),
io.Mask.Input("mask1", optional=True),
io.Mask.Input("mask2", optional=True),
io.Mask.Input("mask3", optional=True),
io.Mask.Input("mask4", optional=True),
],
outputs=[io.Clip.Output()],
)
@classmethod
def execute(cls, clip, mask1=None, mask2=None, mask3=None, mask4=None) -> io.NodeOutput:
clip = clip.clone()
current_masks = clip.patcher.model_options.get("x-promptcontrol.masks", [])
current_masks.extend(m for m in (mask1, mask2, mask3, mask4) if m is not None)
clip.patcher.model_options["x-promptcontrol.masks"] = current_masks
return (clip,)
return io.NodeOutput(clip)
class PCSetPCTextEncodeSettings:
class PCSetPCTextEncodeSettings(io.ComfyNode):
@classmethod
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP",)},
"optional": {
"mask_width": ("INT", {"default": 512, "min": 64, "max": 4096 * 4}),
"mask_height": ("INT", {"default": 512, "min": 64, "max": 4096 * 4}),
"sdxl_width": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
"sdxl_height": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
"sdxl_target_w": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
"sdxl_target_h": ("INT", {"default": 1024, "min": 0, "max": 4096 * 4}),
"sdxl_crop_w": ("INT", {"default": 0, "min": 0, "max": 4096 * 4}),
"sdxl_crop_h": ("INT", {"default": 0, "min": 0, "max": 4096 * 4}),
},
}
def define_schema(cls):
return io.Schema(
node_id="PCSetPCTextEncodeSettings",
display_name="PC: Configure PCTextEncode",
category="promptcontrol/tools",
description="Configures default values for PCTextEncode",
inputs=[
io.Clip.Input("clip"),
io.Int.Input("mask_width", default=512, min=64, max=4096 * 4, optional=True),
io.Int.Input("mask_height", default=512, min=64, max=4096 * 4, optional=True),
io.Int.Input("sdxl_width", default=1024, min=0, max=4096 * 4, optional=True),
io.Int.Input("sdxl_height", default=1024, min=0, max=4096 * 4, optional=True),
io.Int.Input("sdxl_target_w", default=1024, min=0, max=4096 * 4, optional=True),
io.Int.Input("sdxl_target_h", default=1024, min=0, max=4096 * 4, optional=True),
io.Int.Input("sdxl_crop_w", default=0, min=0, max=4096 * 4, optional=True),
io.Int.Input("sdxl_crop_h", default=0, min=0, max=4096 * 4, optional=True),
],
outputs=[io.Clip.Output()],
)
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Configures default values for PCTextEncode"
def apply(
self,
@classmethod
def execute(
cls,
clip,
mask_width=512,
mask_height=512,
@@ -153,7 +110,7 @@ class PCSetPCTextEncodeSettings:
sdxl_target_h=1024,
sdxl_crop_w=0,
sdxl_crop_h=0,
):
) -> io.NodeOutput:
settings = {
"mask_width": mask_width,
"mask_height": mask_height,
@@ -166,66 +123,57 @@ class PCSetPCTextEncodeSettings:
}
clip = clip.clone()
clip.patcher.model_options["x-promptcontrol.settings"] = settings
return (clip,)
return io.NodeOutput(clip)
class PCExtractScheduledPrompt:
class PCExtractScheduledPrompt(io.ComfyNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"multiline": True}),
"at": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 1.0, "step": 0.01}),
},
"optional": {"tags": ("STRING", {"default": ""})},
}
def define_schema(cls):
return io.Schema(
node_id="PCExtractScheduledPrompt",
display_name="PC: Extract Scheduled Prompt",
category="promptcontrol/tools",
description="Parses the input prompt and returns the prompt scheduled at the specified point",
inputs=[
io.String.Input("text", multiline=True),
io.Float.Input("at", min=0.0, max=1.0, default=1.0, step=0.01),
io.String.Input("tags", default="", optional=True),
],
outputs=[io.String.Output()],
)
RETURN_TYPES = ("STRING",)
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Parses the input prompt and returns the prompt scheduled at the specified point"
def apply(self, text, at, tags=""):
@classmethod
def execute(cls, text, at, tags="") -> io.NodeOutput:
schedule = parse_prompt_schedules(text, filters=tags)
_, entry = schedule.at_step(at, total_steps=1)
prompt_text = entry.get("prompt", "")
return (prompt_text,)
return io.NodeOutput(prompt_text)
class PCMacroExpand:
class PCMacroExpand(io.ComfyNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"multiline": True}),
},
}
def define_schema(cls):
return io.Schema(
node_id="PCMacroExpand",
display_name="PC: Expand Macros",
category="promptcontrol/tools",
description="Expands DEF macros in a string and returns the result",
inputs=[
io.String.Input("text", multiline=True),
],
outputs=[io.String.Output()],
)
RETURN_TYPES = ("STRING",)
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Expands DEF macros in a string and returns the result"
def apply(self, text):
return (expand_macros(text),)
@classmethod
def execute(cls, text) -> io.NodeOutput:
return io.NodeOutput(expand_macros(text))
NODE_CLASS_MAPPINGS = {
"PCSetPCTextEncodeSettings": PCSetPCTextEncodeSettings,
"PCAddMaskToCLIP": PCAddMaskToCLIP,
"PCAddMaskToCLIPMany": PCAddMaskToCLIPMany,
"PCSetLogLevel": PCSetLogLevel,
"PCExtractScheduledPrompt": PCExtractScheduledPrompt,
"PCSaveExpandedWorkflow": PCSaveExpandedWorkflow,
"PCMacroExpand": PCMacroExpand,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"PCSetPCTextEncodeSettings": "PC: Configure PCTextEncode",
"PCAddMaskToCLIP": "PC: Attach Mask",
"PCAddMaskToCLIPMany": "PC: Attach Mask (multi)",
"PCSetLogLevel": "PC: Configure Logging (for debug)",
"PCExtractScheduledPrompt": "PC: Extract Scheduled Prompt",
"PCSaveExpandedWorkflow": "PC: Save Expanded Workflow (for debug)",
"PCMacroExpand": "PC: Expand Macros",
}
NODES = [
PCSetPCTextEncodeSettings,
PCAddMaskToCLIP,
PCAddMaskToCLIPMany,
PCSetLogLevel,
PCExtractScheduledPrompt,
PCMacroExpand,
]
+16 -128
View File
@@ -1,21 +1,24 @@
# vim: sw=4 ts=4
import lark
from __future__ import annotations
import logging
from functools import lru_cache
from math import ceil
import lark
from .macros import expand_macros
logging.basicConfig()
log = logging.getLogger("comfyui-prompt-control")
import re
from functools import lru_cache
from .utils import get_function, find_closing_paren
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!",
"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",
]
@@ -31,19 +34,19 @@ ESCAPES = [
]
def escape_specials(string):
def escape_specials(string: str) -> str:
for ph, c in ESCAPES:
string = string.replace(rf"\{c}", ph)
return string
def restore_escaped(string):
def restore_escaped(string: str) -> str:
for ph, c in ESCAPES:
string = string.replace(ph, c)
return string
def remove_comments(string):
def remove_comments(string: str) -> str:
r = []
for line in string.split("\n"):
comment = line.find("#")
@@ -81,46 +84,6 @@ TAG: /[A-Z_]+/
)
cut_parser = lark.Lark(
r"""
!start: (prompt | /[][:()]/+)*
prompt: (cut | PLAIN | WHITESPACE)+
cut: "[CUT:" prompt ":" prompt [":" NUMBER [ ":" NUMBER [":" NUMBER [ ":" PLAIN ] ] ] ]"]"
WHITESPACE: /\s+/
PLAIN: /([^\[\]:])+/
%import common.SIGNED_NUMBER -> NUMBER
"""
)
class CutTransform(lark.Transformer):
def __default__(self, data, children, meta):
return children
def cut(self, args):
prompt, cutout, weight, strict_mask, start_from_masked, mask_token = args
return ("".join(flatten(prompt)), "".join(flatten(cutout)), weight, strict_mask, start_from_masked, mask_token)
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):
return args
def parse_cuts(text):
return CutTransform().transform(cut_parser.parse(text))
def flatten(x):
if type(x) in [str, tuple, int, type(None)] or isinstance(x, dict) and "type" in x:
yield x
@@ -173,7 +136,7 @@ def get_steps(tree, num_steps):
def sequence(self, tree):
steps = tree.children[1::2]
for i, steps in enumerate(steps):
for i, _ in enumerate(steps):
w = tostep(tree.children[i * 2 + 1])
tree.children[i * 2 + 1] = w
res.append(w)
@@ -230,7 +193,7 @@ def at_step(step, filters, tree):
previous_step = 0.0
prompts = args[::2]
steps = args[1::2]
for s, p in zip(steps, prompts):
for s, p in zip(steps, prompts, strict=False):
if s >= step and step >= previous_step:
previous_step = step
return p or ""
@@ -301,13 +264,12 @@ def at_step(step, filters, tree):
return name, params, lbw
def __default__(self, data, children, meta):
for child in children:
yield child
return children
return AtStep().transform(tree)
class PromptSchedule(object):
class PromptSchedule:
# 0 num_steps means unconfigured
def __init__(self, prompt, filters="", start=0.0, end=1.0, num_steps=0):
self.filters = filters
@@ -399,80 +361,6 @@ class PromptSchedule(object):
return len(self.parsed_prompt) - 1, self.parsed_prompt[-1]
def parse_search(search):
arg_start = search.find("(")
args = ""
name = search.strip()
if arg_start > 0:
arg_end = find_closing_paren(search, arg_start + 1)
if arg_end < 0:
arg_end = len(search)
name = search[:arg_start].strip()
args = search[arg_start + 1 : arg_end]
if not name:
return None
args = args.strip()
# If using the form DEF(F()=$1) then the default value of $1 is the empty string
if arg_start > 0:
args = [a.strip() for a in args.split(";")]
else:
args = []
return name, args
def expand_macros(text):
text, defs = get_function(text, "DEF", defaults=None)
res = text
prevres = text
replacements = []
for d in defs:
if not d.args:
continue
r = d.args[0].split("=", 1)
search = parse_search(r[0].strip())
if not search or len(r) != 2:
log.warning("Ignoring invalid DEF(%s)", d)
continue
replacements.append((search, r[1].strip()))
iterations = 0
while True:
iterations += 1
if iterations > 10:
raise ValueError("Unable to resolve DEFs, make sure there are no cycles!")
return text
for search, replace in replacements:
res = substitute_defcall(res, search, replace)
if res == prevres:
break
prevres = res
if res.strip() != text.strip():
res = res.strip()
log.info("DEFs expanded to: %s", res)
return res
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
paramvals = []
if parameters:
paramvals = [x.strip() for x in parameters[0].split(";")]
r = replace
for i, v in enumerate(paramvals):
r = re.sub(rf"\${i+1}\b", v, r)
for i, v in enumerate(default_args):
r = re.sub(rf"\${i+1}\b", v, r)
text = text.replace(ph, r)
return text
@lru_cache
def parse_prompt_schedules(prompt, **kwargs):
prompt = expand_macros(prompt)
+384
View File
@@ -0,0 +1,384 @@
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("\\[]:#")).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)
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)
+720
View File
@@ -0,0 +1,720 @@
# Vendored from https://github.com/python-parsy/parsy/blob/master/src/parsy/__init__.py
from __future__ import annotations
import enum
import operator
import re
from dataclasses import dataclass
from functools import wraps
from typing import Any, Callable, FrozenSet
__version__ = "2.2"
noop = lambda x: x
def line_info_at(stream, index):
if index > len(stream):
raise ValueError("invalid index")
line = stream.count("\n", 0, index)
last_nl = stream.rfind("\n", 0, index)
col = index - (last_nl + 1)
return (line, col)
class ParseError(RuntimeError):
def __init__(self, expected, stream, index):
self.expected = expected
self.stream = stream
self.index = index
def line_info(self) -> str:
try:
return "{}:{}".format(*line_info_at(self.stream, self.index))
except (TypeError, AttributeError): # not a str
return str(self.index)
def __str__(self):
expected_list = sorted(repr(e) for e in self.expected)
if len(expected_list) == 1:
return f"expected {expected_list[0]} at {self.line_info()}"
else:
return f"expected one of {', '.join(expected_list)} at {self.line_info()}"
@dataclass
class Result:
status: bool
index: int
value: Any
furthest: int
expected: FrozenSet[str]
@staticmethod
def success(index, value) -> Result:
return Result(True, index, value, -1, frozenset())
@staticmethod
def failure(index, expected) -> Result:
return Result(False, -1, None, index, frozenset([expected]))
# collect the furthest failure from self and other
def aggregate(self, other) -> Result:
if not other:
return self
if self.furthest > other.furthest:
return self
elif self.furthest == other.furthest:
# if we both have the same failure index, we combine the expected messages.
return Result(self.status, self.index, self.value, self.furthest, self.expected | other.expected)
else:
return Result(self.status, self.index, self.value, other.furthest, other.expected)
# Roughly, a stream is str|bytes|list, but in practice we are duck-typed
# and could accept other things.
# We should switch to this alias when all supported Python versions allow it:
# type Stream = str | bytes | list
class Parser:
"""
A Parser is an object that wraps a function whose arguments are
a string to be parsed and the index on which to begin parsing.
The function should return either Result.success(next_index, value),
where the next index is where to continue the parse and the value is
the yielded value, or Result.failure(index, expected), where expected
is a string indicating what was expected, and the index is the index
of the failure.
"""
def __init__(self, wrapped_fn: Callable[[str | bytes | list, int], Result]):
"""
Creates a new Parser from a function that takes a stream
and returns a Result.
"""
self.wrapped_fn = wrapped_fn
def __call__(self, stream: str | bytes | list, index: int) -> Any:
return self.wrapped_fn(stream, index)
def parse(self, stream: str | bytes | list) -> Any:
"""Parses a string or list of tokens and returns the result or raise a ParseError."""
(result, _) = (self << eof).parse_partial(stream)
return result
def parse_partial(self, stream: str | bytes | list) -> tuple[Any, str | bytes | list]:
"""
Parses the longest possible prefix of a given string.
Returns a tuple of the result and the unparsed remainder,
or raises ParseError
"""
result = self(stream, 0)
if result.status:
return (result.value, stream[result.index :])
else:
raise ParseError(result.expected, stream, result.furthest)
def bind(self, bind_fn: Callable[[Any], Parser]) -> Parser:
@Parser
def bound_parser(stream: str | bytes | list, index: int) -> Result:
result = self(stream, index)
if result.status:
next_parser = bind_fn(result.value)
return next_parser(stream, result.index).aggregate(result)
else:
return result
return bound_parser
def map(self, map_function: Callable) -> Parser:
"""
Returns a parser that transforms the produced value of the initial parser with map_function.
"""
return self.bind(lambda res: success(map_function(res)))
def combine(self, combine_fn: Callable) -> Parser:
"""
Returns a parser that transforms the produced values of the initial parser
with ``combine_fn``, passing the arguments using ``*args`` syntax.
The initial parser should return a list/sequence of parse results.
"""
return self.bind(lambda res: success(combine_fn(*res)))
def combine_dict(self, combine_fn: Callable) -> Parser:
"""
Returns a parser that transforms the value produced by the initial parser
using the supplied function/callable, passing the arguments using the
``**kwargs`` syntax.
The value produced by the initial parser must be a mapping/dictionary from
names to values, or a list of two-tuples, or something else that can be
passed to the ``dict`` constructor.
If ``None`` is present as a key in the dictionary it will be removed
before passing to ``fn``, as will all keys starting with ``_``.
"""
return self.bind(
lambda res: success(
combine_fn(
**{
k: v
for k, v in dict(res).items()
if k is not None and not (isinstance(k, str) and k.startswith("_"))
}
)
)
)
def concat(self) -> Parser:
"""
Returns a parser that concatenates together (as a string) the previously
produced values.
"""
return self.map("".join)
def then(self, other: Parser) -> Parser:
"""
Returns a parser which, if the initial parser succeeds, will
continue parsing with ``other``. This will produce the
value produced by ``other``.
"""
return seq(self, other).combine(lambda left, right: right)
def skip(self, other: Parser) -> Parser:
"""
Returns a parser which, if the initial parser succeeds, will
continue parsing with ``other``. It will produce the
value produced by the initial parser.
"""
return seq(self, other).combine(lambda left, right: left)
def result(self, value: Any) -> Parser:
"""
Returns a parser that, if the initial parser succeeds, always produces
the passed in ``value``.
"""
return self >> success(value)
def many(self) -> Parser:
"""
Returns a parser that expects the initial parser 0 or more times, and
produces a list of the results.
"""
return self.times(0, float("inf"))
def times(self, min: int, max: int = None) -> Parser:
"""
Returns a parser that expects the initial parser at least ``min`` times,
and at most ``max`` times, and produces a list of the results. If only one
argument is given, the parser is expected exactly that number of times.
"""
if max is None:
max = min
@Parser
def times_parser(stream: str | bytes | list, index: int) -> Result:
values = []
times = 0
result = None
while times < max:
result = self(stream, index).aggregate(result)
if result.status:
values.append(result.value)
index = result.index
times += 1
elif times >= min:
break
else:
return result
return Result.success(index, values).aggregate(result)
return times_parser
def at_most(self, n: int) -> Parser:
"""
Returns a parser that expects the initial parser at most ``n`` times, and
produces a list of the results.
"""
return self.times(0, n)
def at_least(self, n: int) -> Parser:
"""
Returns a parser that expects the initial parser at least ``n`` times, and
produces a list of the results.
"""
return self.times(n) + self.many()
def optional(self, default: Any = None) -> Parser:
"""
Returns a parser that expects the initial parser zero or once, and maps
the result to a given default value in the case of no match. If no default
value is given, ``None`` is used.
"""
return self.times(0, 1).map(lambda v: v[0] if v else default)
def until(self, other: Parser, min: int = 0, max: int = float("inf"), consume_other: bool = False) -> Parser:
"""
Returns a parser that expects the initial parser followed by ``other``.
The initial parser is expected at least ``min`` times and at most ``max`` times.
By default, it does not consume ``other`` and it produces a list of the
results excluding ``other``. If ``consume_other`` is ``True`` then
``other`` is consumed and its result is included in the list of results.
"""
@Parser
def until_parser(stream: str | bytes | list, index: int) -> Result:
values = []
times = 0
while True:
# try parser first
res = other(stream, index)
if res.status and times >= min:
if consume_other:
# consume other
values.append(res.value)
index = res.index
return Result.success(index, values)
# exceeded max?
if times >= max:
# return failure, it matched parser more than max times
return Result.failure(index, f"at most {max} items")
# failed, try parser
result = self(stream, index)
if result.status:
# consume
values.append(result.value)
index = result.index
times += 1
elif times >= min:
# return failure, parser is not followed by other
return Result.failure(index, "did not find other parser")
else:
# return failure, it did not match parser at least min times
return Result.failure(index, f"at least {min} items; got {times} item(s)")
return until_parser
def sep_by(self, sep: Parser, *, min: int = 0, max: int = float("inf")) -> Parser:
"""
Returns a new parser that repeats the initial parser and
collects the results in a list. Between each item, the ``sep`` parser
is run (and its return value is discarded). By default it
repeats with no limit, but minimum and maximum values can be supplied.
"""
zero_times = success([])
if max == 0:
return zero_times
res = self.times(1) + (sep >> self).times(min - 1, max - 1)
if min == 0:
res |= zero_times
return res
def desc(self, description: str) -> Parser:
"""
Returns a new parser with a description added, which is used in the error message
if parsing fails.
"""
@Parser
def desc_parser(stream: str | bytes | list, index: int) -> Result:
result = self(stream, index)
if result.status:
return result
else:
return Result.failure(index, description)
return desc_parser
def mark(self) -> Parser:
"""
Returns a parser that wraps the initial parser's result in a value
containing column and line information of the match, as well as the
original value. The new value is a 3-tuple:
((start_row, start_column),
original_value,
(end_row, end_column))
"""
@generate
def marked():
start = yield line_info
body = yield self
end = yield line_info
return (start, body, end)
return marked
def tag(self, name: str) -> Parser:
"""
Returns a parser that wraps the produced value of the initial parser in a
2 tuple containing ``(name, value)``. This provides a very simple way to
label parsed components
"""
return self.map(lambda v: (name, v))
def should_fail(self, description: str) -> Parser:
"""
Returns a parser that fails when the initial parser succeeds, and succeeds
when the initial parser fails (consuming no input). A description must
be passed which is used in parse failure messages.
This is essentially a negative lookahead
"""
@Parser
def fail_parser(stream: str | bytes | list, index: int) -> Result:
res = self(stream, index)
if res.status:
return Result.failure(index, description)
return Result.success(index, res)
return fail_parser
def __add__(self, other: Parser) -> Parser:
return seq(self, other).combine(operator.add)
def __mul__(self, other: int | range) -> Parser:
if isinstance(other, range):
return self.times(other.start, other.stop - 1)
return self.times(other)
def __or__(self, other: Parser) -> Parser:
return alt(self, other)
# haskelley operators, for fun #
# >>
def __rshift__(self, other: Parser) -> Parser:
return self.then(other)
# <<
def __lshift__(self, other: Parser) -> Parser:
return self.skip(other)
def alt(*parsers: Parser) -> Parser:
"""
Creates a parser from the passed in argument list of alternative
parsers, which are tried in order, moving to the next one if the
current one fails.
"""
if not parsers:
return fail("<empty alt>")
@Parser
def alt_parser(stream: str | bytes | list, index: int) -> Result:
result = None
for parser in parsers:
result = parser(stream, index).aggregate(result)
if result.status:
return result
return result
return alt_parser
def seq(*parsers: Parser, **kw_parsers: Parser) -> Parser:
"""
Takes a list of parsers, runs them in order,
and collects their individuals results in a list,
or in a dictionary if you pass them as keyword arguments.
"""
if not parsers and not kw_parsers:
return success([])
if parsers and kw_parsers:
raise ValueError("Use either positional arguments or keyword arguments with seq, not both")
if parsers:
@Parser
def seq_parser(stream: str | bytes | list, index: int) -> Result:
result = None
values = []
for parser in parsers:
result = parser(stream, index).aggregate(result)
if not result.status:
return result
index = result.index
values.append(result.value)
return Result.success(index, values).aggregate(result)
return seq_parser
else:
@Parser
def seq_kwarg_parser(stream: str | bytes | list, index: int) -> Result:
result = None
values = {}
for name, parser in kw_parsers.items():
result = parser(stream, index).aggregate(result)
if not result.status:
return result
index = result.index
values[name] = result.value
return Result.success(index, values).aggregate(result)
return seq_kwarg_parser
def generate(fn) -> Parser:
"""
Creates a parser from a generator function
"""
if isinstance(fn, str):
return lambda f: generate(f).desc(fn)
@Parser
@wraps(fn)
def generated(stream: str | bytes | list, index: int) -> Result:
# start up the generator
iterator = fn()
result = None
value = None
try:
while True:
next_parser = iterator.send(value)
result = next_parser(stream, index).aggregate(result)
if not result.status:
return result
value = result.value
index = result.index
except StopIteration as stop:
returnVal = stop.value
if isinstance(returnVal, Parser):
return returnVal(stream, index).aggregate(result)
return Result.success(index, returnVal).aggregate(result)
return generated
index = Parser(lambda _, index: Result.success(index, index))
line_info = Parser(lambda stream, index: Result.success(index, line_info_at(stream, index)))
def success(value: Any) -> Parser:
"""
Returns a parser that does not consume any of the stream, but
produces ``value``.
"""
return Parser(lambda _, index: Result.success(index, value))
def fail(expected: str) -> Parser:
"""
Returns a parser that always fails with the provided error message.
"""
return Parser(lambda _, index: Result.failure(index, expected))
def string(expected_string: str, transform: Callable[[str], str] = noop) -> Parser:
"""
Returns a parser that expects the ``expected_string`` and produces
that string value.
Optionally, a transform function can be passed, which will be used on both
the expected string and tested string.
"""
slen = len(expected_string)
transformed_s = transform(expected_string)
@Parser
def string_parser(stream: str, index: int) -> Result:
if transform(stream[index : index + slen]) == transformed_s:
return Result.success(index + slen, expected_string)
else:
return Result.failure(index, expected_string)
return string_parser
def regex(exp: str, flags=0, group: int | str | tuple = 0) -> Parser:
"""
Returns a parser that expects the given ``exp``, and produces the
matched string. ``exp`` can be a compiled regular expression, or a
string which will be compiled with the given ``flags``.
Optionally, accepts ``group``, which is passed to re.Match.group
https://docs.python.org/3/library/re.html#re.Match.group> to
return the text from a capturing group in the regex instead of the
entire match.
"""
if isinstance(exp, (str, bytes)):
exp = re.compile(exp, flags)
if isinstance(group, (str, int)):
group = (group,)
@Parser
def regex_parser(stream: str | bytes | list, index: int) -> Result:
match = exp.match(stream, index)
if match:
return Result.success(match.end(), match.group(*group))
else:
return Result.failure(index, exp.pattern)
return regex_parser
def test_item(func: Callable[..., bool], description: str) -> Parser:
"""
Returns a parser that tests a single item from the list of items being
consumed, using the callable ``func``. If ``func`` returns ``True``, the
parse succeeds, otherwise the parse fails with the description
``description``.
"""
@Parser
def test_item_parser(stream: str | bytes | list, index: int) -> Result:
if index < len(stream):
if isinstance(stream, bytes):
# Subscripting bytes with `[index]` instead of
# `[index:index + 1]` returns an int
item = stream[index : index + 1]
else:
item = stream[index]
if func(item):
return Result.success(index + 1, item)
return Result.failure(index, description)
return test_item_parser
def test_char(func: Callable[..., bool], description: str) -> Parser:
"""
Returns a parser that tests a single character with the callable
``func``. If ``func`` returns ``True``, the parse succeeds, otherwise
the parse fails with the description ``description``.
"""
# Implementation is identical to test_item
return test_item(func, description)
def match_item(item: Any, description: str = None) -> Parser:
"""
Returns a parser that tests the next item (or character) from the stream (or
string) for equality against the provided item. Optionally a string
description can be passed.
"""
if description is None:
description = str(item)
return test_item(lambda i: item == i, description)
def string_from(*strings: str, transform: Callable[[str], str] = noop):
"""
Accepts a sequence of strings as positional arguments, and returns a parser
that matches and returns one string from the list. The list is first sorted
in descending length order, so that overlapping strings are handled correctly
by checking the longest one first.
"""
# Sort longest first, so that overlapping options work correctly
return alt(*(string(s, transform) for s in sorted(strings, key=len, reverse=True)))
def char_from(string: str | bytes) -> Parser:
"""
Accepts a string and returns a parser that matches and returns one character
from the string.
"""
if isinstance(string, bytes):
return test_char(lambda c: c in string, b"[" + string + b"]")
else:
return test_char(lambda c: c in string, "[" + string + "]")
def peek(parser: Parser) -> Parser:
"""
Returns a lookahead parser that parses the input stream without consuming
chars.
"""
@Parser
def peek_parser(stream: str | bytes | list, index: int) -> Result:
result = parser(stream, index)
if result.status:
return Result.success(index, result.value)
else:
return result
return peek_parser
any_char = test_char(lambda c: True, "any character")
whitespace = regex(r"\s+")
letter = test_char(lambda c: c.isalpha(), "a letter")
digit = test_char(lambda c: c.isdigit(), "a digit")
decimal_digit = char_from("0123456789")
@Parser
def eof(stream: str | bytes | list, index: int) -> Result:
"""
A parser that only succeeds if the end of the stream has been reached.
"""
if index >= len(stream):
return Result.success(index, None)
else:
return Result.failure(index, "EOF")
def from_enum(enum_cls: type[enum.Enum], transform=noop) -> Parser:
"""
Given a class that is an enum.Enum class
https://docs.python.org/3/library/enum.html , returns a parser that
will parse the values (or the string representations of the values)
and return the corresponding enum item.
"""
items = sorted(
((str(enum_item.value), enum_item) for enum_item in enum_cls), key=lambda t: len(t[0]), reverse=True
)
return alt(*(string(value, transform=transform).result(enum_item) for value, enum_item in items))
class forward_declaration(Parser):
"""
An empty parser that can be used as a forward declaration,
especially for parsers that need to be defined recursively.
You must use `.become(parser)` before using.
"""
def __init__(self):
pass
def _raise_error(self, *args, **kwargs):
raise ValueError("You must use 'become' before attempting to call `parse` or `parse_partial`")
parse = _raise_error
parse_partial = _raise_error
def become(self, other: Parser):
"""
Take on the behavior of the given parser.
"""
self.__dict__ = other.__dict__
self.__class__ = other.__class__
+45 -41
View File
@@ -1,28 +1,31 @@
from __future__ import annotations
import logging
import re
import torch
import math
import re
from collections import defaultdict
from functools import partial
from typing import Any
import torch
from comfy_extras.nodes_mask import FeatherMask, MaskComposite
from nodes import ConditioningAverage
from .utils import (
safe_float,
get_function,
split_by_function,
parse_floats,
smarter_split,
call_node,
split_quotable,
FunctionSpec,
ComfyConditioning,
)
from .adv_encode import advanced_encode_from_tokens
from .cutoff import process_cuts
from .parser import parse_cuts
from .attention_couple_ppm import set_cond_attnmask
from .cutoff import process_cuts
from .cutoff_parser import parse_cuts
from .utils import (
ComfyConditioning,
FunctionSpec,
call_node,
get_function,
parse_floats,
safe_float,
smarter_split,
split_by_function,
split_quotable,
)
log = logging.getLogger("comfyui-prompt-control")
@@ -32,7 +35,7 @@ AVAILABLE_NORMALIZATIONS = ["none", "mean", "length", "length+mean"]
SHUFFLE_GEN = torch.Generator(device="cpu")
def get_sdxl(text, defaults):
def get_sdxl(text: str, defaults: dict[str, Any]) -> tuple[str, dict[str, int]]:
# Defaults fail to parse and get looked up from the defaults dict
text, sdxl = get_function(text, "SDXL", ["none", "none", "none"])
if not sdxl:
@@ -54,7 +57,7 @@ def get_sdxl(text, defaults):
return text, opts
def get_clipweights(text, existing_spec=None):
def get_clipweights(text: str, existing_spec: dict[str, float] | None = None) -> tuple[dict[str, float], str]:
text, spec = get_function(text, "TE_WEIGHT", defaults=None)
if not spec:
return existing_spec or {}, text
@@ -70,7 +73,7 @@ def get_clipweights(text, existing_spec=None):
return res, text
def get_style(text, default_style="comfy", default_normalization="none"):
def get_style(text: str, default_style="comfy", default_normalization="none") -> tuple[str, str, str]:
text, styles = get_function(text, "STYLE", [default_style, default_normalization])
if not styles:
return default_style, default_normalization, text
@@ -125,7 +128,8 @@ def shuffle_chunk(func_spec: FunctionSpec, c: str) -> str:
def fix_word_ids(tokens):
"""Fix word indexes. Tokenizing separately (when BREAKs exist) causes the indexes to restart which causes problems with some weighting algorithms that rely on them"""
"""Fix word indexes. Tokenizing separately (when BREAKs exist) causes the indexes
to restart which causes problems with some weighting algorithms that rely on them"""
for key in tokens:
max_idx = 0
for group in range(len(tokens[key])):
@@ -178,7 +182,7 @@ def tokenize(clip, text, can_break, empty_tokens):
need_word_ids = True
tokens = tokenize_chunks(clip, text, need_word_ids, can_break)
per_te_prompts = {}
per_te_prompts = defaultdict(list)
if l_prompts:
log.warning("Note: CLIP_L is deprecated. Use TE(l=prompt) instead")
per_te_prompts["l"] = [x.args for x in l_prompts]
@@ -198,9 +202,7 @@ def tokenize(clip, text, can_break, empty_tokens):
log.warning("Invalid TE call, no TE with key '%s', ignoring: %s", te)
log.info("Encoders available for TE: %s", ", ".join(tokens.keys()))
continue
l = per_te_prompts.get(te, [])
l.append(prompt)
per_te_prompts[te] = l
per_te_prompts[te].append(prompt)
if per_te_prompts:
for key in per_te_prompts:
@@ -395,7 +397,8 @@ def get_area(text):
area = (int(h) // 8, int(w) // 8, int(y) // 8, int(x) // 8)
else:
raise Exception(
f"AREA specified with invalid size {x} {w}, {h} {y}. They must either all be percentages between 0 and 1 or positive integer pixel values excluding 1"
f"AREA specified with invalid size {x} {w}, {h} {y}. They must either all"
" be percentages between 0 and 1 or positive integer pixel values excluding 1"
)
return text, (area, weight)
@@ -429,7 +432,8 @@ def make_mask(args, size, weight):
ys = int(y1), int(y2)
else:
raise Exception(
f"MASK specified with invalid size {x1} {x2}, {y1} {y2}. They must either all be percentages between 0 and 1 or positive integer pixel values excluding 1"
f"MASK specified with invalid size {x1} {x2}, {y1} {y2}. They must either all"
" be percentages between 0 and 1 or positive integer pixel values excluding 1"
)
mask = torch.full((h, w), 0, dtype=torch.float32, device="cpu")
@@ -450,9 +454,9 @@ def get_mask(text, size, input_masks):
return text, None, None
def feather(f, mask):
l, t, r, b, *_ = [int(x) for x in parse_floats(f[0], [0, 0, 0, 0], split_re="\\s+")]
mask = call_node(FeatherMask, mask, l, t, r, b)[0]
log.info("FeatherMask l=%s, t=%s, r=%s, b=%s", l, t, r, b)
left, top, right, bottom, *_ = [int(x) for x in parse_floats(f[0], [0, 0, 0, 0], split_re="\\s+")]
mask = call_node(FeatherMask, mask, left, top, right, bottom)[0]
log.info("FeatherMask l=%s, t=%s, r=%s, b=%s", left, top, right, bottom)
return mask
mask = None
@@ -478,22 +482,19 @@ def get_mask(text, size, input_masks):
idx = int(safe_float(idx, 0.0))
w = safe_float(w, 1.0)
if input_masks is None:
log.warn(
log.warning(
"IMASK requires you to attach custom masks to the CLIP object using PCAddMasksToClIP before using it"
)
input_masks = []
if len(input_masks) < idx + 1:
log.warn("IMASK index %s not found, ignoring...", idx)
log.warning("IMASK index %s not found, ignoring...", idx)
continue
nextmask = input_masks[idx] * w
if i < len(feathers):
nextmask = feather(feathers[i].args, nextmask)
i += 1
if mask is not None:
mask = call_node(MaskComposite, mask, nextmask, 0, 0, op)[0]
else:
mask = nextmask
mask = call_node(MaskComposite, mask, nextmask, 0, 0, op)[0] if mask is not None else nextmask
# apply leftover FEATHER() specs to the whole
for f in feathers[i:]:
@@ -556,7 +557,6 @@ def process_settings(prompt, defaults, masks, mask_size, sdxl_opts):
prompt = prompt.replace("FILL()", "")
settings["x-promptcontrol.fill"] = True
prompt, mask, mask_weight = get_mask(prompt, mask_size, masks)
prompt, noise_w, generator = get_noise(prompt)
prompt, area = get_area(prompt)
prompt, local_sdxl_opts = get_sdxl(prompt, defaults)
# Get weight last so other syntax doesn't interfere with it
@@ -604,6 +604,7 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
return f"MASK({args[0]})"
for prompt in prompts:
text, noise_w, generator = get_noise(text)
base_prompt, attn_couple_prompts = split_by_function(prompt, "COUPLE", defaults=None, require_args=False)
prompts = [base_prompt] + [couple_mask(f.args) + chunk for (chunk, f) in attn_couple_prompts]
@@ -617,16 +618,17 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
x = encode_prompt_segment(clip, p, settings, style, normalization)
encoded.append(x)
assert all(
len(c) == len(encoded[0]) for c in encoded
), "All encoded prompts didn't produce the same number of conds, I don't know what to do in this situation."
assert all(len(c) == len(encoded[0]) for c in encoded), (
"All encoded prompts didn't produce the same number of conds, I don't know what to do in this situation."
)
# each call to encode_prompt_segment can produce a number of conds based on any
# scheduled LoRA hooks on the clip model. Zip them together with coupled prompts
base_cond = []
for base_cond, *attention_couple in zip(*encoded):
for base_cond, *attention_couple in zip(*encoded, strict=False):
s = base_cond[1]
# If there are LoRAs on the CLIP, we need to fix start_percent and end_percent on the new conds for things to work properly.
# If there are LoRAs on the CLIP, we need to fix start_percent and
# end_percent on the new conds for things to work properly.
s["start_percent"] = s.get("clip_start_percent", s["start_percent"])
s["end_percent"] = s.get("clip_end_percent", s["end_percent"])
s.pop("clip_start_percent", None)
@@ -642,6 +644,8 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
[ensure_mask(c) for c in attention_couple],
fill=fill,
)
base_cond = [[apply_noise(c[0], noise_w, generator), c[1]] for c in base_cond]
conds.extend(base_cond)
return conds
-224
View File
@@ -1,224 +0,0 @@
import unittest
import unittest.mock as mock
import numpy.testing as npt
from os import environ
import nodes
import comfy_extras.nodes_mask
from .nodes_base import PCTextEncode
clips = []
import logging
logging.basicConfig()
def run(f, *args):
if hasattr(f, "execute"):
return f.execute(*args)
else:
return getattr(f, f.FUNCTION)(*args)
def compare_hookgroup_mask(h1, h2):
assert len(h1.hooks) == len(h2.hooks)
for a, b in zip(h1.hooks, h2.hooks):
assert (a.mask == b.mask).all()
@mock.patch("torch.cuda.current_device", lambda: "cpu")
class TestEncode(unittest.TestCase):
@classmethod
def setUpClass(cls):
global clips
print("Loading ComfyUI")
from comfy.sd import load_clip
from pathlib import Path
to_test = environ.get("TEST_TE", "clip_l").split()
model_dir = environ.get("COMFYUI_TE_DIR", ".")
te_root = Path(model_dir).resolve()
if "clip_l" in to_test:
clip_l = load_clip(
ckpt_paths=[str(te_root / "clip_l.safetensors")], clip_type="stable_diffusion", model_options={}
)
clips.append(("clip_l", clip_l))
if "t5" in to_test:
dual = load_clip(
[str(te_root / "clip_l.safetensors"), str(te_root / "t5xxl_fp16.safetensors")],
clip_type="flux",
model_options={},
)
clips.append(("clip_l+t5", dual))
print("Starting tests")
def tensorsEqual(self, t1, t2):
npt.assert_equal(t1.detach().numpy(), t2.detach().numpy())
def condEqual(self, c1, c2, key=None, key_assert=None):
self.assertEqual(len(c1), len(c2))
for i in range(len(c1)):
a, b = c1[i], c2[i]
if key:
(key_assert or self.assertEqual)(a[1].get(key), b[1].get(key))
else:
self.tensorsEqual(a[0], b[0])
def test_basic_encode(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
combine = nodes.ConditioningCombine()
average = nodes.ConditioningAverage()
concat = nodes.ConditioningConcat()
zeroout = nodes.ConditioningZeroOut()
for k, clip in clips:
with self.subTest(k):
with self.subTest("No exceptions"):
run(
pc,
clip,
"test AND test (test:1.2) BREAK test AND TE_WEIGHT(all=0) SDXL() AND AREA(,,) test CAT test",
)
with self.subTest("Basic"):
(c1,) = run(pc, clip, "test")
(c2,) = run(comfy, clip, "test")
c = c2 # Used in later tests
self.condEqual(c1, c2)
with self.subTest("Quotes"):
(c1,) = run(pc, clip, 'Text saying "DOG MASK AND CAT COUPLE MASK(X)"')
(c2,) = run(comfy, clip, 'Text saying "DOG MASK AND CAT COUPLE MASK(X)"')
self.condEqual(c1, c2)
with self.subTest("Function cornercase"):
(c1,) = run(pc, clip, "test SDXL function")
(c2,) = run(comfy, clip, "test SDXL function")
(c3,) = run(pc, clip, "test SDXL() function")
self.condEqual(c1, c2)
with self.subTest("Weights"):
(c1,) = run(pc, clip, "(test:1.2) (test:0.6)")
(c2,) = run(comfy, clip, "(test:1.2) (test:0.6)")
self.condEqual(c1, c2)
with self.subTest("Concat"):
(c1,) = run(pc, clip, "test CAT test")
(c2,) = run(concat, c, c)
self.condEqual(c1, c2)
with self.subTest("Combine"):
(c1,) = run(pc, clip, "test AND test")
(c2,) = run(combine, c, c)
self.condEqual(c1, c2)
with self.subTest("Zero out"):
(c1,) = run(pc, clip, "test TE_WEIGHT(all=0)")
(c2,) = run(zeroout, c)
self.condEqual(c1, c2)
with self.subTest("Average"):
(c1,) = run(comfy, clip, "test1")
(c2,) = run(comfy, clip, "test2")
(c3,) = run(pc, clip, "test1 AVG() test2")
(c4,) = run(pc, clip, "test1 AVG test2")
(avg,) = run(average, c1, c2, 0.5)
self.condEqual(avg, c3)
self.condEqual(avg, c4)
with self.subTest("Average multi"):
(c1,) = run(comfy, clip, "test1")
(c2,) = run(comfy, clip, "test2")
(c3,) = run(comfy, clip, "test3")
(c4,) = run(pc, clip, "test1 AVG() test2 AVG() test3")
(c5,) = run(pc, clip, "test1 AVG test2 AVG test3")
(avg1,) = run(average, c1, c2, 0.5)
(avg,) = run(average, avg1, c3, 0.5)
self.condEqual(avg, c4)
self.condEqual(avg, c5)
@unittest.expectedFailure
def test_failure(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
for k, clip in clips:
with self.subTest(k):
(c1,) = run(comfy, clip, "test SDXL function")
(c2,) = run(pc, clip, "test SDXL() function")
self.condEqual(c1, c2)
def test_weight(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
combine = nodes.ConditioningCombine()
strength = nodes.ConditioningSetAreaStrength()
for k, clip in clips:
(c,) = run(comfy, clip, "test")
(c2,) = run(strength, c, 0.5)
with self.subTest(f"Testing {k}"):
with self.subTest("Conditioning weights"):
(a,) = run(pc, clip, "test :0.5 AND test :0.5")
(b,) = run(combine, c2, c2)
self.condEqual(a, b)
self.condEqual(a, b, "strength")
with self.subTest("Weight == 0"):
(a,) = run(pc, clip, "test :0.5 AND test :0 AND test")
(b,) = run(combine, c2, c)
self.condEqual(a, b)
self.condEqual(a, b, "strength")
def test_attn_couple(self):
pc = PCTextEncode()
for k, clip in clips:
with self.subTest(f"Testing {k}"):
(c,) = run(pc, clip, "test COUPLE prompt1 AND test2 COUPLE prompt2")
(c2,) = run(pc, clip, "test COUPLE prompt1 COUPLE test2 COUPLE prompt2")
self.assertTrue(len(c) == 2)
self.assertTrue(len(c2) == 1)
with self.subTest(f"Testing {k} mask shortcut"):
(c,) = run(pc, clip, "test COUPLE() prompt1")
(c2,) = run(pc, clip, "test COUPLE MASK() prompt1")
self.condEqual(c, c2)
self.condEqual(c, c2, "hooks", compare_hookgroup_mask)
with self.subTest(f"Testing {k} mask shortcut 2"):
(c,) = run(pc, clip, "test COUPLE(0 0.2, 0.5) prompt1")
(c2,) = run(pc, clip, "test COUPLE MASK(0 0.2, 0.5) prompt1")
self.condEqual(c, c2)
self.condEqual(c, c2, "hooks", compare_hookgroup_mask)
def test_styles(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
for k, clip in clips:
(no_weights,) = run(comfy, clip, "this prompt has no weights")
for style in ["comfy", "A1111", "comfy++", "compel", "down_weight", "perp"]:
with self.subTest(f"TE {k} style {style} no weights equal comfy"):
(c,) = run(pc, clip, "this prompt has no weights")
self.condEqual(no_weights, c)
with self.subTest(f"TE {k} style {style} does not fail when encoding weights"):
for normalization in ["none", "mean", "length", "mean+length", "length+mean"]:
with self.subTest(f"TE {k} style {style} normalization {normalization}"):
(c,) = run(
pc,
clip,
f"STYLE({style}, {normalization}) (this prompt) (has weights:0.9), (a:1.2) (b:1.2)",
)
def test_masks(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
solidmask = comfy_extras.nodes_mask.SolidMask()
setMask = nodes.ConditioningSetMask()
for k, clip in clips:
(c1,) = run(pc, clip, "test MASK()")
(c2,) = run(comfy, clip, "test")
(c2,) = run(setMask, c2, run(solidmask, 1.0, 512, 512)[0], "default", 1.0)
self.condEqual(c1, c2)
self.condEqual(c1, c2, "mask", self.tensorsEqual)
if __name__ == "__main__":
unittest.main()
-244
View File
@@ -1,244 +0,0 @@
import unittest
import unittest.mock as mock
import logging
log = logging.getLogger("comfyui-prompt-control")
def reset_graphbuilder_state():
from comfy_execution.graph_utils import GraphBuilder
GraphBuilder.set_default_prefix("UID", 0, 0)
def find_file(name):
names = {"test": "test.safetensors", "other": "some/other.safetensors"}
return names.get(name)
def loraloader(text, adv=False, **kwargs):
from .nodes_lazy import PCLazyLoraLoader, PCLazyLoraLoaderAdvanced
reset_graphbuilder_state()
if adv:
cls = PCLazyLoraLoader
else:
cls = PCLazyLoraLoaderAdvanced
model = [0, 1]
clip = [0, 0]
return cls().apply(unique_id="UID", model=model, clip=clip, text=text, **kwargs)
def te(text, adv=False, **kwargs):
from .nodes_lazy import PCLazyTextEncode, PCLazyTextEncodeAdvanced
if adv:
cls = PCLazyTextEncode
else:
cls = PCLazyTextEncodeAdvanced
reset_graphbuilder_state()
clip = [0, 0]
return cls().apply(clip=clip, text=text, unique_id="UID", **kwargs)
@mock.patch("prompt_control.utils.lora_name_to_file", find_file)
@mock.patch("torch.cuda.current_device", lambda: "cpu")
class GraphTests(unittest.TestCase):
maxDiff = 4096
def test_textencode(self):
for p in ["test", "[test:0.2] test", "[test[test::0.5]]<lora:test:1>"]:
r1 = te(p)
r2 = te(p, adv=True)
with self.subTest(f"Expansion: {p}"):
self.assertEqual(r1, r2)
reset_graphbuilder_state()
with self.subTest("Expansion: LoRA"):
r = te("test<lora:test:1>")
self.assertEqual(
r,
{
"result": (["UID.0.0.2", 0],),
"expand": {
"UID.0.0.1": {"class_type": "PCTextEncode", "inputs": {"clip": [0, 0], "text": "test"}},
"UID.0.0.2": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.1", 0], "start": 0.0, "end": 1.0},
},
},
},
)
with self.subTest("Expansion: LoRA with schedule"):
r = te("simple [test:0.1,0.5] prompt<lora:test:1>")
self.assertEqual(
r,
{
"result": (["UID.0.0.8", 0],),
"expand": {
"UID.0.0.1": {
"class_type": "PCTextEncode",
"inputs": {"clip": [0, 0], "text": "simple prompt"},
},
"UID.0.0.2": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.1", 0], "start": 0.0, "end": 0.1},
},
"UID.0.0.3": {
"class_type": "PCTextEncode",
"inputs": {"clip": [0, 0], "text": "simple test prompt"},
},
"UID.0.0.4": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.3", 0], "start": 0.1, "end": 0.5},
},
"UID.0.0.5": {
"class_type": "PCTextEncode",
"inputs": {"clip": [0, 0], "text": "simple prompt"},
},
"UID.0.0.6": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.5", 0], "start": 0.5, "end": 1.0},
},
"UID.0.0.7": {
"class_type": "ConditioningCombine",
"inputs": {"conditioning_1": ["UID.0.0.2", 0], "conditioning_2": ["UID.0.0.4", 0]},
},
"UID.0.0.8": {
"class_type": "ConditioningCombine",
"inputs": {"conditioning_1": ["UID.0.0.7", 0], "conditioning_2": ["UID.0.0.6", 0]},
},
},
},
)
@mock.patch("prompt_control.utils.lora_name_to_file", find_file)
def test_loraloader(self):
with self.assertLogs(log, level="WARNING") as cm:
result = loraloader("prompt here <lora:nonexistent:1.0:0.5>")["expand"]
result_adv = loraloader("prompt here <lora:nonexistent:1.0:0.5>", adv=True)["expand"]
self.assertIn("LoRA 'nonexistent' not found", cm.output[0])
self.assertEqual(result, {})
self.assertEqual(result_adv, {})
result = loraloader("<lora:test:1>")["expand"]
result2 = loraloader("prompt here <lora:test:1.0:0.5><lora:test:0:0.5>")["expand"]
result3 = loraloader("prompt here <lora:test:1.0:0.5><lora:test:0:0.5>", adv=True)["expand"]
self.assertEqual(result, result2)
self.assertEqual(result2, result3)
self.assertEqual(
result,
{
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 1.0,
"lora_name": "test.safetensors",
},
}
},
)
result = loraloader("<lora:test:1><lora:other:0.5>")["expand"]
self.assertEqual(
result,
{
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 1.0,
"lora_name": "test.safetensors",
},
},
"UID.0.0.2": {
"class_type": "LoraLoader",
"inputs": {
"model": ["UID.0.0.1", 0],
"clip": ["UID.0.0.1", 1],
"strength_model": 0.5,
"strength_clip": 0.5,
"lora_name": "some/other.safetensors",
},
},
},
)
result = loraloader("prompt here <lora:test:1.0:0.5>")["expand"]
self.assertEqual(
result,
{
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 0.5,
"lora_name": "test.safetensors",
},
}
},
)
result = loraloader("prompt [<lora:test:0.5>:0.5]")["expand"]
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True)["expand"]
self.assertEqual(result, result2)
expected = {
"UID.0.0.1": {
"class_type": "CreateHookLora",
"inputs": {"lora_name": "test.safetensors", "strength_model": 0.5, "strength_clip": 0.5},
},
"UID.0.0.2": {
"class_type": "CreateHookKeyframe",
"inputs": {"strength_mult": 0.0, "start_percent": 0.0},
},
"UID.0.0.3": {
"class_type": "CreateHookKeyframe",
"inputs": {
"start_percent": 0.5,
"prev_hook_kf": ["UID.0.0.2", 0],
"strength_mult": 1.0,
},
},
"UID.0.0.4": {
"class_type": "SetHookKeyframes",
"inputs": {"hooks": ["UID.0.0.1", 0], "hook_kf": ["UID.0.0.3", 0]},
},
"UID.0.0.5": {
"class_type": "SetClipHooks",
"inputs": {
"clip": [0, 0],
"hooks": ["UID.0.0.4", 0],
"apply_to_conds": True,
"schedule_clip": True,
},
},
}
self.assertEqual(result, expected)
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True, start=0.6)["expand"]
self.assertEqual(
result2,
{
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 0.5,
"strength_clip": 0.5,
"lora_name": "test.safetensors",
},
}
},
)
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", end=0.5)["expand"]
self.assertEqual(result2, {})
if __name__ == "__main__":
unittest.main()
-247
View File
@@ -1,247 +0,0 @@
import unittest
from .parser import parse_prompt_schedules as parse, expand_macros
def prompt(until, text, *loras):
loras = {lora: {"weight": unet, "weight_clip": te} for lora, unet, te in loras}
return [until, {"prompt": text, "loras": loras}]
class TestParser(unittest.TestCase):
def assertPrompt(self, p, at, until, text, *loras):
self.assertEqual(p.at_step(at), prompt(until, text, *loras))
def test_no_scheduling(self):
p = parse("This is a (basic:0.6) (prompt) with [no scheduling] features")
expected = prompt(1.0, "This is a (basic:0.6) (prompt) with [no scheduling] features")
self.assertEqual(p.at_step(0), expected)
self.assertEqual(p.at_step(0.5), expected)
self.assertEqual(p.at_step(1), expected)
def test_quote(self):
p = parse('This is a text with a "QUOTED DEF(X=Y)"')
expected = prompt(1.0, 'This is a text with a "QUOTED DEF(X=Y)"')
self.assertEqual(p.at_step(0), expected)
self.assertEqual(p.at_step(0.5), expected)
self.assertEqual(p.at_step(1), expected)
def test_equivalences(self):
eqs = [
[parse(p) for p in ["[a:0.1]", "[:a:0.1]", "[:a:0,0.1]", "[:a::0.1,1.0]", "[:a::0.1]"]],
[parse(p) for p in ["[before:during:after:0.1]", "[before:during:after:0.1,1.0]", "[before:during:0.1]"]],
[parse(p) for p in ["[a:0.1,0.5]", "[[a:0.1]::0.5]", "[:a::0.1,0.5]", "[a::0.1,0.5]"]],
[parse(p) for p in ["[a:b:0.5]", "[a::b:0.5,0.5]"]],
[parse(p) for p in ["[a::0.5]", "[a:::0.5,0.5]"]],
]
for group in eqs:
for p in group[1:]:
with self.subTest(p):
self.assertEqual(group[0].parsed_prompt, p.parsed_prompt)
def test_basic(self):
p = parse(
"This is a (basic:0.6) (prompt) with (very [[simple]:(basic:0.6):0.5]:1.1) [features::0.8][ and this is ignored:1]"
)
self.assertPrompt(p, 0, 0.5, "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features")
self.assertPrompt(p, 0.5, 0.5, "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features")
self.assertPrompt(p, 0.7, 0.8, "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) features")
self.assertPrompt(p, 1.0, 1.0, "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) ")
def test_lora(self):
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)
)
self.assertEqual(p.at_step(0), expected)
self.assertEqual(p.at_step(0.5), expected)
self.assertEqual(p.at_step(1), expected)
def test_scheduled_lora(self):
p = parse(
"This is a (lora:0.6) (prompt) with [scheduling] features [<lora:foo:0.5>:<lora:bar:0.5:0.2>:0.3] <lora:bar:0.5:1.0>"
)
self.assertPrompt(
p,
0.1,
0.3,
"This is a (lora:0.6) (prompt) with [scheduling] features ",
("foo", 0.5, 0.5),
("bar", 0.5, 1.0),
)
self.assertPrompt(p, 0.5, 1.0, "This is a (lora:0.6) (prompt) with [scheduling] features ", ("bar", 1.0, 1.2))
def test_seq(self):
p = parse("This is a sequence of [SEQ:a:0.2::0.5:c:0.8][SEQ: and x:0.8]")
p2 = parse("This is a sequence of [[a:[c:0.5]:0.2]::0.8][ and x::0.8]")
prompts = {
0.2: "This is a sequence of a and x",
0.5: "This is a sequence of and x",
0.8: "This is a sequence of c and x",
1.0: "This is a sequence of ",
}
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
for k, v in prompts.items():
self.assertPrompt(p, k, k, v)
def test_shortcuts_scheduling(self):
p = parse("A schedule [a:0.1,0.7] b")
p2 = parse("A schedule [[a:0.1]::0.7] b")
p3 = parse("A schedule [a:b:0.5,0.8]")
p4 = parse("A schedule [[a:0.5]:b:0.8]")
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
self.assertEqual(p3.parsed_prompt, p4.parsed_prompt)
def test_range(self):
p = parse("test [excluded::excluded2:0.1,0.4] test")
self.assertPrompt(p, 0, 0.1, "test excluded test")
self.assertPrompt(p, 0.2, 0.4, "test test")
self.assertPrompt(p, 0.45, 1.0, "test excluded2 test")
p = parse("test [[:included::0.2,0.8]|[excluded::excluded2:0.4,0.9]:0.1] test")
self.assertPrompt(p, 0, 0.1, "test test")
self.assertPrompt(p, 0.25, 0.3, "test included test")
self.assertPrompt(p, 0.15, 0.2, "test excluded test")
self.assertPrompt(p, 0.25, 0.3, "test included test")
self.assertPrompt(p, 0.55, 0.6, "test test")
self.assertPrompt(p, 0.95, 1.0, "test excluded2 test")
def test_nested(self):
p = parse(
"This [prompt is [SEQ:[crazy:weird:0.2] stuff:0.5:<lora:cool:1>:0.7:nesting:1.0]:completely ignored with tags:HR]"
)
prompts = {
0.2: (0.2, "This prompt is crazy stuff"),
0.3: (0.5, "This prompt is weird stuff"),
0.5: (0.5, "This prompt is weird stuff"),
0.8: (1.0, "This prompt is nesting"),
}
for k in prompts:
self.assertEqual(p.at_step(k), [prompts[k][0], {"prompt": prompts[k][1], "loras": {}}])
self.assertPrompt(p, 0.6, 0.7, "This prompt is ", ("cool", 1.0, 1.0))
self.assertPrompt(p, 0.7, 0.7, "This prompt is ", ("cool", 1.0, 1.0))
p2 = p.with_filters(filters="hr, xyz")
self.assertEqual(p2.at_step(0), p2.at_step(1))
def test_def(self):
p = parse("DEF(X=0.5) [a:b:X] DEF(test = [c:X]) test test")
prompts = {
0.2: (0.5, "a "),
0.6: (1.0, "b c c"),
}
for k, v in prompts.items():
self.assertPrompt(p, k, v[0], v[1])
p = parse("DEF(X=[($1):($1:$2):$2])X(test;0.7)")
p2 = parse("[(test):(test:0.7):0.7]")
with self.subTest("parameters"):
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
p = parse("DEF(X=[($1):($1:$2):$2])DEF(Y=X(test;$1))Y(0.7) Y(0.5)")
p2 = parse("[(test):(test:0.7):0.7] [(test):(test:0.5):0.5]")
with self.subTest("two functions"):
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
p = expand_macros("DEF(X(a;b)=$1 $2 $3 d)X(A) X(A;B;C)")
with self.subTest("defaults"):
self.assertEqual(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)")
with self.subTest("Empty default for $1"):
self.assertEqual(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)]")
with self.subTest("defaults, DEF=X vs DEF=X()"):
self.assertEqual(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]")
with self.subTest("defaults, nested parens"):
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
with self.assertRaises(ValueError) as c:
expand_macros("DEF(X=recurse Y) DEF(Y=recurse X) X")
self.assertTrue("Unable to resolve DEFs" in str(c.exception))
def test_escapes(self):
p = parse(r"[a:\:a:0.5] :\[a:b:0.5]")
self.assertPrompt(p, 0, 0.5, r"a :\[a:b:0.5]")
self.assertPrompt(p, 0.55, 1, r":a :\[a:b:0.5]")
p = parse(r"[embedding\:a:embedding\:b:0.1,0.5]")
self.assertPrompt(p, 0.15, 0.5, r"embedding:a")
self.assertPrompt(p, 0.55, 1, r"embedding:b")
p = parse(r"[embedding\:a:embedding\:b:embedding\:c:0.1,0.5]")
self.assertPrompt(p, 0.0, 0.1, r"embedding:a")
self.assertPrompt(p, 0.15, 0.5, r"embedding:b")
self.assertPrompt(p, 0.55, 1, r"embedding:c")
p = parse(r"[a\:b\\:c:0.5]")
self.assertPrompt(p, 0.0, 0.5, "a:b\\")
self.assertPrompt(p, 0.55, 1, r"c")
p = parse(r"[a:\#b:0.5]")
self.assertPrompt(p, 0.0, 0.5, "a")
self.assertPrompt(p, 0.55, 1, "#b")
def test_comments(self):
p = parse("this is a # comment")
self.assertPrompt(p, 0, 1.0, "this is a ")
p = parse("this is a [comment#:scheduled:0.6]")
self.assertPrompt(p, 0, 1.0, "this is a [comment")
p = parse(r"this is a [comment\#:scheduled:0.6]")
self.assertPrompt(p, 0, 0.6, "this is a comment#")
self.assertPrompt(p, 0.65, 1.0, "this is a scheduled")
p = parse("#this is a comment\nthis is a prompt")
self.assertPrompt(p, 0, 1.0, "\nthis is a prompt")
def test_misc(self):
p = parse("[[a:c:0.5]:0.7]")
p2 = parse("[:[a:c:0.5]:0.7]")
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
p = parse("test [[a:[b<lora:test:0.5>:0.6]:0.5]:HR]")
p2 = parse("test [:[a:[:b<lora:test:0.5>:0.6]:0.5]:HR]")
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
pf = p.with_filters(filters="hr")
self.assertEqual(pf.parsed_prompt, p2.with_filters(filters="hr").parsed_prompt)
self.assertPrompt(pf, 0, 0.5, "test a")
self.assertPrompt(pf, 0.55, 0.6, "test ")
self.assertPrompt(pf, 0.8, 1.0, "test b", ("test", 0.5, 0.5))
p = parse("[:[<lora:test:1>:c:0.5]:0.3]")
self.assertPrompt(p, 0, 0.3, "")
self.assertPrompt(p, 0.4, 0.5, "", ("test", 1.0, 1.0))
self.assertPrompt(p, 1.0, 1.0, "c")
p = parse("an [<emb:foo>:<emb:bar>:0.5]")
prompts = {
0.2: (0.5, "an embedding:foo"),
0.8: (1.0, "an embedding:bar"),
}
for k, v in prompts.items():
self.assertPrompt(p, k, v[0], v[1])
def test_alternating(self):
p = parse("[cat|dog|tiger]")
p2 = parse("[cat|dog|tiger:0.1]")
p3 = parse("[cat|[dog|wolf]|tiger]")
p4 = parse("[cat|[dog:wolf<lora:canine:1>:0.5]:0.2]")
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
for i, x in enumerate(["cat", "wolf", "tiger", "cat", "dog", "tiger", "cat", "wolf", "tiger", "cat"]):
step = round((i * 0.1) + 0.1, 2)
with self.subTest(step):
self.assertPrompt(p3, step, step, x)
for i, x in enumerate([["cat"], ["dog"], ["cat"], ["wolf", ("canine", 1.0, 1.0)], ["cat"]]):
step = round((i * 0.2) + 0.2, 2)
with self.subTest(step):
self.assertPrompt(p4, step, step, *x)
self.assertPrompt(p4, 0.7, 0.8, "wolf", ("canine", 1.0, 1.0))
if __name__ == "__main__":
unittest.main()
+11 -11
View File
@@ -1,11 +1,12 @@
from __future__ import annotations
from pathlib import Path
import re
import logging
import copy
import copy
import logging
import re
from collections.abc import Iterator
from dataclasses import dataclass
from typing import Any, TypeAlias, Iterator, TypeVar, TYPE_CHECKING
from pathlib import Path
from typing import TYPE_CHECKING, Any, TypeAlias, TypeVar
if TYPE_CHECKING:
import torch # flakes8: noqa
@@ -119,10 +120,8 @@ def find_closing_paren(text: str, start: int) -> int:
def find_function_spans(
text: str, func: str, require_args: bool, defaults: FunctionArgs | None
) -> Iterator[tuple[int, int, str, FunctionArgs]]:
if require_args:
rex = re.compile(rf"\b{func}\(", re.MULTILINE)
else:
rex = re.compile(rf"\b{func}\b", re.MULTILINE)
e = r"\(" if require_args else r"\b"
rex = re.compile(rf"\b{func}{e}", re.MULTILINE)
idx = 0
match = rex.search(text)
@@ -191,7 +190,8 @@ def split_by_function(
text: str, func: str, defaults: list[str] | None = None, require_args: bool = True
) -> tuple[str, list[tuple[str, FunctionSpec]]]:
"""
Splits a string by function calls, returning the leftover text along with a list of functions with their associated text chunk.
Splits a string by function calls, returning the leftover text
along with a list of functions with their associated text chunk.
"""
text, functions = get_function(text, func, defaults, require_args=require_args)
chunks = []
@@ -291,7 +291,7 @@ def expand_graph(node_mappings, graph):
node = node_mappings[data["class_type"]]()
inputs = map_inputs(input_map, data["inputs"].copy())
inputs["unique_id"] = k
fn = getattr(node, getattr(node, "FUNCTION"))
fn = getattr(node, node.FUNCTION)
expansion = fn(**inputs)
for i, v in enumerate(expansion["result"]):
input_map[(k, i)] = v
+35
View File
@@ -6,6 +6,8 @@ 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"
[project.urls]
Repository = "https://github.com/asagi4/comfyui-prompt-control"
@@ -13,3 +15,36 @@ Repository = "https://github.com/asagi4/comfyui-prompt-control"
PublisherId = "asagi4"
DisplayName = "ComfyUI Prompt Control"
Icon = ""
[tool.pyright]
extraPaths = ["../../"]
exclude = ["prompt_control/*test*"]
[tool.ty.src]
exclude = ["tests/*.py", "prompt_control/*test*.py", "prompt_control/parsy.py"]
[tool.ty.environment]
extra-paths = ["../.."]
[tool.ty.rules]
# ComfyUI executes give this...
invalid-method-override = "ignore"
[tool.ruff]
exclude = ["prompt_control/parsy.py"]
line-length = 120
[tool.ruff.lint]
# Ignore line length and let the formatter handle it
ignore = ["E501"]
select = [
"E",
"F",
"UP",
"B",
"SIM",
"I",
]
[tool.pytest.ini_options]
testpaths = ["tests"]
View File
+5
View File
@@ -0,0 +1,5 @@
import logging
def pytest_runtest_setup(item):
logging.getLogger("comfyui-prompt-control").setLevel(logging.CRITICAL)
+26
View File
@@ -0,0 +1,26 @@
import pytest
from prompt_control.cutoff_parser import parse_cuts
@pytest.fixture(scope="module", autouse=True)
def parser():
return parse_cuts
def test_parse_no_cuts(parser):
prompt, cutouts = parse_cuts("a b c")
assert prompt == "a b c"
assert cutouts == []
def test_parse_cuts(parser):
prompt, cutouts = parse_cuts("a [CUT:b:d:0] c")
assert prompt == "a b c"
assert cutouts == [("b", "d", 0.0, None, None, None)]
def test_parse_cuts_multiple(parser):
prompt, cutouts = parse_cuts("a [CUT:b:d:0] [CUT:c:e:1.0:0.5:0.9:-]")
assert prompt == "a b c"
assert cutouts == [("b", "d", 0.0, None, None, None), ("c", "e", 1.0, 0.5, 0.9, "-")]
+243
View File
@@ -0,0 +1,243 @@
import numpy.testing as npt
import pytest
def run(f, *args):
if hasattr(f, "execute"):
return f.execute(*args)
else:
return getattr(f, f.FUNCTION)(*args)
def compare_hookgroup_mask(h1, h2):
assert len(h1.hooks) == len(h2.hooks)
for a, b in zip(h1.hooks, h2.hooks, strict=True):
assert (a.mask == b.mask).all()
@pytest.fixture(scope="module")
def text_encoder_clips():
import os
from pathlib import Path
from comfy.sd import load_clip
clips = []
to_test = os.environ.get("TEST_TE", "clip_l").split()
model_dir = os.environ.get("COMFYUI_TE_DIR", ".")
te_root = Path(model_dir).resolve()
if "clip_l" in to_test:
clip_l = load_clip(
ckpt_paths=[str(te_root / "clip_l.safetensors")], clip_type="stable_diffusion", model_options={}
)
clips.append(("clip_l", clip_l))
if "t5" in to_test:
dual = load_clip(
[str(te_root / "clip_l.safetensors"), str(te_root / "t5xxl_fp16.safetensors")],
clip_type="flux",
model_options={},
)
clips.append(("clip_l+t5", dual))
return clips
@pytest.fixture
def pc_text_encode():
from prompt_control.nodes_base import PCTextEncode
return PCTextEncode()
@pytest.fixture
def node_class_objs():
import comfy_extras.nodes_mask
import nodes
# Return all used node class objects
return {
"comfy": nodes.CLIPTextEncode(),
"combine": nodes.ConditioningCombine(),
"average": nodes.ConditioningAverage(),
"concat": nodes.ConditioningConcat(),
"zeroout": nodes.ConditioningZeroOut(),
"strength": nodes.ConditioningSetAreaStrength(),
"solidmask": comfy_extras.nodes_mask.SolidMask(),
"setmask": nodes.ConditioningSetMask(),
}
def tensors_equal(t1, t2):
npt.assert_equal(t1.detach().numpy(), t2.detach().numpy())
def cond_equal(c1, c2, key=None, key_assert=None):
assert len(c1) == len(c2)
for i in range(len(c1)):
a, b = c1[i], c2[i]
if key:
(key_assert or assert_equal)(a[1].get(key), b[1].get(key))
else:
tensors_equal(a[0], b[0])
def assert_equal(a, b):
assert a == b
@pytest.mark.usefixtures("text_encoder_clips", "pc_text_encode", "node_class_objs")
class TestPCTextEncode:
def test_basic_encode(self, text_encoder_clips, pc_text_encode, node_class_objs):
comfy = node_class_objs["comfy"]
combine = node_class_objs["combine"]
average = node_class_objs["average"]
concat = node_class_objs["concat"]
zeroout = node_class_objs["zeroout"]
for _k, clip in text_encoder_clips:
# No exceptions
run(
pc_text_encode,
clip,
"test AND test (test:1.2) BREAK test AND TE_WEIGHT(all=0) SDXL() AND AREA(,,) test CAT test",
)
# Basic
(c1,) = run(pc_text_encode, clip, "test")
(c2,) = run(comfy, clip, "test")
c = c2 # Used in later tests
cond_equal(c1, c2)
# Quotes
(c1,) = run(pc_text_encode, clip, 'Text saying "DOG MASK AND CAT COUPLE MASK(X)"')
(c2,) = run(comfy, clip, 'Text saying "DOG MASK AND CAT COUPLE MASK(X)"')
cond_equal(c1, c2)
# Function cornercase
(c1,) = run(pc_text_encode, clip, "test SDXL function")
(c2,) = run(comfy, clip, "test SDXL function")
(c3,) = run(pc_text_encode, clip, "test SDXL() function")
cond_equal(c1, c2)
# Weights
(c1,) = run(pc_text_encode, clip, "(test:1.2) (test:0.6)")
(c2,) = run(comfy, clip, "(test:1.2) (test:0.6)")
cond_equal(c1, c2)
# Concat
(c1,) = run(pc_text_encode, clip, "test CAT test")
(c2,) = run(concat, c, c)
cond_equal(c1, c2)
# Combine
(c1,) = run(pc_text_encode, clip, "test AND test")
(c2,) = run(combine, c, c)
cond_equal(c1, c2)
# Zero out
(c1,) = run(pc_text_encode, clip, "test TE_WEIGHT(all=0)")
(c2,) = run(zeroout, c)
cond_equal(c1, c2)
# Average
(c1,) = run(comfy, clip, "test1")
(c2,) = run(comfy, clip, "test2")
(c3,) = run(pc_text_encode, clip, "test1 AVG() test2")
(c4,) = run(pc_text_encode, clip, "test1 AVG test2")
(avg,) = run(average, c1, c2, 0.5)
cond_equal(avg, c3)
cond_equal(avg, c4)
def test_avg(self, text_encoder_clips, pc_text_encode, node_class_objs):
comfy = node_class_objs["comfy"]
average = node_class_objs["average"]
for _k, clip in text_encoder_clips:
(c1,) = run(comfy, clip, "test1")
(c2,) = run(comfy, clip, "test2")
(c3,) = run(comfy, clip, "test3")
(c4,) = run(pc_text_encode, clip, "test1 AVG() test2 AVG() test3")
(c5,) = run(pc_text_encode, clip, "test1 AVG test2 AVG test3")
(avg1,) = run(average, c1, c2, 0.5)
(avg,) = run(average, avg1, c3, 0.5)
cond_equal(avg, c4)
cond_equal(avg, c5)
@pytest.mark.xfail
def test_failure(self, text_encoder_clips, pc_text_encode, node_class_objs):
comfy = node_class_objs["comfy"]
for _k, clip in text_encoder_clips:
(c1,) = run(comfy, clip, "test SDXL function")
(c2,) = run(pc_text_encode, clip, "test SDXL() function")
cond_equal(c1, c2)
def test_weight(self, text_encoder_clips, pc_text_encode, node_class_objs):
comfy = node_class_objs["comfy"]
combine = node_class_objs["combine"]
strength = node_class_objs["strength"]
for _k, clip in text_encoder_clips:
(c,) = run(comfy, clip, "test")
(c2,) = run(strength, c, 0.5)
# Conditioning weights
(a,) = run(pc_text_encode, clip, "test :0.5 AND test :0.5")
(b,) = run(combine, c2, c2)
cond_equal(a, b)
cond_equal(a, b, "strength")
# Weight == 0
(a,) = run(pc_text_encode, clip, "test :0.5 AND test :0 AND test")
(b,) = run(combine, c2, c)
cond_equal(a, b)
cond_equal(a, b, "strength")
def test_attn_couple(self, text_encoder_clips, pc_text_encode):
for _k, clip in text_encoder_clips:
(c,) = run(pc_text_encode, clip, "test COUPLE prompt1 AND test2 COUPLE prompt2")
(c2,) = run(pc_text_encode, clip, "test COUPLE prompt1 COUPLE test2 COUPLE prompt2")
assert len(c) == 2
assert len(c2) == 1
def test_styles(self, text_encoder_clips, pc_text_encode, node_class_objs):
comfy = node_class_objs["comfy"]
for _k, clip in text_encoder_clips:
(no_weights,) = run(comfy, clip, "this prompt has no weights")
for style in ["comfy", "A1111", "comfy++", "compel", "down_weight", "perp"]:
# no weights equal comfy
(c,) = run(pc_text_encode, clip, "this prompt has no weights")
cond_equal(no_weights, c)
# does not fail when encoding weights
for normalization in ["none", "mean", "length", "mean+length", "length+mean"]:
run(
pc_text_encode,
clip,
f"STYLE({style}, {normalization}) (this prompt) (has weights:0.9), (a:1.2) (b:1.2)",
)
# Just checking for exceptions
def test_masks(self, text_encoder_clips, pc_text_encode, node_class_objs):
comfy = node_class_objs["comfy"]
solidmask = node_class_objs["solidmask"]
setmask = node_class_objs["setmask"]
for _k, clip in text_encoder_clips:
(c1,) = run(pc_text_encode, clip, "test MASK()")
(c2,) = run(comfy, clip, "test")
(c2,) = run(setmask, c2, run(solidmask, 1.0, 512, 512)[0], "default", 1.0)
cond_equal(c1, c2)
cond_equal(c1, c2, "mask", tensors_equal)
def test_cutoff_nofail(self, text_encoder_clips, pc_text_encode, node_class_objs):
for _k, clip in text_encoder_clips:
(c1,) = run(pc_text_encode, clip, "test [CUT:a:b:0.5]")
def test_couple_mask_shortcut(self, text_encoder_clips, pc_text_encode, node_class_objs):
for _k, clip in text_encoder_clips:
(c,) = run(pc_text_encode, clip, "test COUPLE() prompt1")
(c2,) = run(pc_text_encode, clip, "test COUPLE MASK() prompt1")
cond_equal(c, c2)
cond_equal(c, c2, "hooks", compare_hookgroup_mask)
(c,) = run(pc_text_encode, clip, "test COUPLE(0 0.2, 0.5) prompt1")
(c2,) = run(pc_text_encode, clip, "test COUPLE MASK(0 0.2, 0.5) prompt1")
cond_equal(c, c2)
cond_equal(c, c2, "hooks", compare_hookgroup_mask)
+584
View File
@@ -0,0 +1,584 @@
import logging
import pytest
from comfy_execution.graph_utils import GraphBuilder
from prompt_control.nodes_lazy import (
PCLazyLoraLoader,
PCLazyLoraLoaderAdvanced,
PCLazyTextEncode,
PCLazyTextEncodeAdvanced,
)
log = logging.getLogger("comfyui-prompt-control")
def reset_graphbuilder_state():
GraphBuilder.set_default_prefix("UID", 0, 0)
def find_file(name):
names = {"test": "test.safetensors", "other": "some/other.safetensors"}
return names.get(name)
def as_dict(out):
return {"result": out.result, "expand": out.expand}
def loraloader(text, adv=False, **kwargs):
reset_graphbuilder_state()
cls = PCLazyLoraLoaderAdvanced if adv else PCLazyLoraLoader
model = [0, 1]
clip = [0, 0]
return as_dict(cls.execute(model=model, clip=clip, text=text, **kwargs))
def te(text, adv=False, **kwargs):
cls = PCLazyTextEncode if adv else PCLazyTextEncodeAdvanced
reset_graphbuilder_state()
clip = [0, 0]
return as_dict(cls.execute(clip=clip, text=text, **kwargs))
@pytest.fixture(autouse=True)
def patch_lora_name_to_file(monkeypatch):
import prompt_control.utils
monkeypatch.setattr(prompt_control.utils, "lora_name_to_file", find_file)
@pytest.fixture(autouse=True)
def patch_torch_cuda_current_device(monkeypatch):
import torch.cuda
monkeypatch.setattr(torch.cuda, "current_device", lambda: "cpu")
def test_textencode_expansion():
for p in ["test", "[test:0.2] test", "[test[test::0.5]]<lora:test:1>"]:
r1 = te(p)
r2 = te(p, adv=True)
assert r1 == r2
def test_textencode_alternating():
r = te("[a|b]")
expected_result = {
"expand": {
"UID.0.0.1": {
"class_type": "PCTextEncode",
"inputs": {
"clip": [
0,
0,
],
"text": "a",
},
},
"UID.0.0.10": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {
"conditioning": [
"UID.0.0.9",
0,
],
"end": 0.5,
"start": 0.4,
},
},
"UID.0.0.11": {
"class_type": "PCTextEncode",
"inputs": {
"clip": [
0,
0,
],
"text": "b",
},
},
"UID.0.0.12": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {
"conditioning": [
"UID.0.0.11",
0,
],
"end": 0.6,
"start": 0.5,
},
},
"UID.0.0.13": {
"class_type": "PCTextEncode",
"inputs": {
"clip": [
0,
0,
],
"text": "a",
},
},
"UID.0.0.14": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {
"conditioning": [
"UID.0.0.13",
0,
],
"end": 0.7,
"start": 0.6,
},
},
"UID.0.0.15": {
"class_type": "PCTextEncode",
"inputs": {
"clip": [
0,
0,
],
"text": "b",
},
},
"UID.0.0.16": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {
"conditioning": [
"UID.0.0.15",
0,
],
"end": 0.8,
"start": 0.7,
},
},
"UID.0.0.17": {
"class_type": "PCTextEncode",
"inputs": {
"clip": [
0,
0,
],
"text": "a",
},
},
"UID.0.0.18": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {
"conditioning": [
"UID.0.0.17",
0,
],
"end": 0.9,
"start": 0.8,
},
},
"UID.0.0.19": {
"class_type": "PCTextEncode",
"inputs": {
"clip": [
0,
0,
],
"text": "b",
},
},
"UID.0.0.2": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {
"conditioning": [
"UID.0.0.1",
0,
],
"end": 0.1,
"start": 0.0,
},
},
"UID.0.0.20": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {
"conditioning": [
"UID.0.0.19",
0,
],
"end": 1.0,
"start": 0.9,
},
},
"UID.0.0.21": {
"class_type": "ConditioningCombine",
"inputs": {
"conditioning_1": [
"UID.0.0.2",
0,
],
"conditioning_2": [
"UID.0.0.4",
0,
],
},
},
"UID.0.0.22": {
"class_type": "ConditioningCombine",
"inputs": {
"conditioning_1": [
"UID.0.0.21",
0,
],
"conditioning_2": [
"UID.0.0.6",
0,
],
},
},
"UID.0.0.23": {
"class_type": "ConditioningCombine",
"inputs": {
"conditioning_1": [
"UID.0.0.22",
0,
],
"conditioning_2": [
"UID.0.0.8",
0,
],
},
},
"UID.0.0.24": {
"class_type": "ConditioningCombine",
"inputs": {
"conditioning_1": [
"UID.0.0.23",
0,
],
"conditioning_2": [
"UID.0.0.10",
0,
],
},
},
"UID.0.0.25": {
"class_type": "ConditioningCombine",
"inputs": {
"conditioning_1": [
"UID.0.0.24",
0,
],
"conditioning_2": [
"UID.0.0.12",
0,
],
},
},
"UID.0.0.26": {
"class_type": "ConditioningCombine",
"inputs": {
"conditioning_1": [
"UID.0.0.25",
0,
],
"conditioning_2": [
"UID.0.0.14",
0,
],
},
},
"UID.0.0.27": {
"class_type": "ConditioningCombine",
"inputs": {
"conditioning_1": [
"UID.0.0.26",
0,
],
"conditioning_2": [
"UID.0.0.16",
0,
],
},
},
"UID.0.0.28": {
"class_type": "ConditioningCombine",
"inputs": {
"conditioning_1": [
"UID.0.0.27",
0,
],
"conditioning_2": [
"UID.0.0.18",
0,
],
},
},
"UID.0.0.29": {
"class_type": "ConditioningCombine",
"inputs": {
"conditioning_1": [
"UID.0.0.28",
0,
],
"conditioning_2": [
"UID.0.0.20",
0,
],
},
},
"UID.0.0.3": {
"class_type": "PCTextEncode",
"inputs": {
"clip": [
0,
0,
],
"text": "b",
},
},
"UID.0.0.4": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {
"conditioning": [
"UID.0.0.3",
0,
],
"end": 0.2,
"start": 0.1,
},
},
"UID.0.0.5": {
"class_type": "PCTextEncode",
"inputs": {
"clip": [
0,
0,
],
"text": "a",
},
},
"UID.0.0.6": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {
"conditioning": [
"UID.0.0.5",
0,
],
"end": 0.3,
"start": 0.2,
},
},
"UID.0.0.7": {
"class_type": "PCTextEncode",
"inputs": {
"clip": [
0,
0,
],
"text": "b",
},
},
"UID.0.0.8": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {
"conditioning": [
"UID.0.0.7",
0,
],
"end": 0.4,
"start": 0.3,
},
},
"UID.0.0.9": {
"class_type": "PCTextEncode",
"inputs": {
"clip": [
0,
0,
],
"text": "a",
},
},
},
"result": (
[
"UID.0.0.29",
0,
],
),
}
assert r == expected_result
def test_textencode_lora():
reset_graphbuilder_state()
r = te("test<lora:test:1>")
assert r == {
"result": (["UID.0.0.2", 0],),
"expand": {
"UID.0.0.1": {"class_type": "PCTextEncode", "inputs": {"clip": [0, 0], "text": "test"}},
"UID.0.0.2": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.1", 0], "start": 0.0, "end": 1.0},
},
},
}
def test_textencode_lora_with_schedule():
r = te("simple [test:0.1,0.5] prompt<lora:test:1>")
assert r == {
"result": (["UID.0.0.8", 0],),
"expand": {
"UID.0.0.1": {
"class_type": "PCTextEncode",
"inputs": {"clip": [0, 0], "text": "simple prompt"},
},
"UID.0.0.2": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.1", 0], "start": 0.0, "end": 0.1},
},
"UID.0.0.3": {
"class_type": "PCTextEncode",
"inputs": {"clip": [0, 0], "text": "simple test prompt"},
},
"UID.0.0.4": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.3", 0], "start": 0.1, "end": 0.5},
},
"UID.0.0.5": {
"class_type": "PCTextEncode",
"inputs": {"clip": [0, 0], "text": "simple prompt"},
},
"UID.0.0.6": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.5", 0], "start": 0.5, "end": 1.0},
},
"UID.0.0.7": {
"class_type": "ConditioningCombine",
"inputs": {"conditioning_1": ["UID.0.0.2", 0], "conditioning_2": ["UID.0.0.4", 0]},
},
"UID.0.0.8": {
"class_type": "ConditioningCombine",
"inputs": {"conditioning_1": ["UID.0.0.7", 0], "conditioning_2": ["UID.0.0.6", 0]},
},
},
}
def test_loraloader_empty(monkeypatch, caplog):
result = loraloader("prompt here <lora:nonexistent:1.0:0.5>")["expand"]
result_adv = loraloader("prompt here <lora:nonexistent:1.0:0.5>", adv=True)["expand"]
assert result == {}
assert result_adv == {}
def test_loraloader_duplicate_results():
result = loraloader("<lora:test:1>")["expand"]
result2 = loraloader("prompt here <lora:test:1.0:0.5><lora:test:0:0.5>")["expand"]
result3 = loraloader("prompt here <lora:test:1.0:0.5><lora:test:0:0.5>", adv=True)["expand"]
assert result == result2
assert result2 == result3
assert result == {
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 1.0,
"lora_name": "test.safetensors",
},
}
}
def test_loraloader_multiple_loras():
result = loraloader("<lora:test:1><lora:other:0.5>")["expand"]
assert result == {
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 1.0,
"lora_name": "test.safetensors",
},
},
"UID.0.0.2": {
"class_type": "LoraLoader",
"inputs": {
"model": ["UID.0.0.1", 0],
"clip": ["UID.0.0.1", 1],
"strength_model": 0.5,
"strength_clip": 0.5,
"lora_name": "some/other.safetensors",
},
},
}
def test_loraloader_strength_clip():
result = loraloader("prompt here <lora:test:1.0:0.5>")["expand"]
assert result == {
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 0.5,
"lora_name": "test.safetensors",
},
}
}
def test_loraloader_scheduled_compare():
result = loraloader("prompt [<lora:test:0.5>:0.5]")["expand"]
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True)["expand"]
assert result == result2
expected = {
"UID.0.0.1": {
"class_type": "CreateHookLora",
"inputs": {"lora_name": "test.safetensors", "strength_model": 0.5, "strength_clip": 0.5},
},
"UID.0.0.2": {
"class_type": "CreateHookKeyframe",
"inputs": {"strength_mult": 0.0, "start_percent": 0.0},
},
"UID.0.0.3": {
"class_type": "CreateHookKeyframe",
"inputs": {"start_percent": 0.5, "prev_hook_kf": ["UID.0.0.2", 0], "strength_mult": 1.0},
},
"UID.0.0.4": {
"class_type": "SetHookKeyframes",
"inputs": {"hooks": ["UID.0.0.1", 0], "hook_kf": ["UID.0.0.3", 0]},
},
"UID.0.0.5": {
"class_type": "SetClipHooks",
"inputs": {
"clip": [0, 0],
"hooks": ["UID.0.0.4", 0],
"apply_to_conds": True,
"schedule_clip": True,
},
},
}
assert result == expected
def test_loraloader_adv_start():
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True, start=0.6)["expand"]
assert result2 == {
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 0.5,
"strength_clip": 0.5,
"lora_name": "test.safetensors",
},
}
}
def test_loraloader_end_zero():
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True, end=0.5)["expand"]
assert result2 == {}
+376
View File
@@ -0,0 +1,376 @@
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}
def prompt(until, text, *loras):
return (until, {"prompt": text, "loras": lora_dict(*loras)})
def prompts_match(a, b):
return list(a) == list(b)
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:
params.append(old_parse)
if "new" in parsers_to_test:
params.append(new_parse)
@pytest.fixture(scope="module", autouse=True, params=params)
def parse(request):
return request.param
@pytest.mark.parametrize("step", [0, 0.5, 1])
def test_no_scheduling(step, parse):
p = parse(r"This is a (basic:0.6) (prompt) with [no scheduling] features and \(escaped parens\)")
expected = prompt(1.0, r"This is a (basic:0.6) (prompt) with [no scheduling] features and \(escaped parens\)")
assert prompts_match(p.at_step(step), expected)
def test_integer_steps(parse):
p = parse("[a:b:25]", num_steps=50)
assert prompts_match(p.at_step(0), prompt(0.5, "a"))
assert prompts_match(p.at_step(25), prompt(0.5, "a"))
assert prompts_match(p.at_step(0.5), prompt(0.5, "a"))
assert prompts_match(p.at_step(0.51), prompt(1.0, "b"))
assert prompts_match(p.at_step(30), prompt(1.0, "b"))
def test_mixed_steps(parse):
p = parse("[a:b:25] [c:d:0.25]", num_steps=50)
assert prompts_match(p.at_step(0), prompt(0.25, "a c"))
assert prompts_match(p.at_step(25), prompt(0.5, "a d"))
assert prompts_match(p.at_step(0.5), prompt(0.5, "a d"))
assert prompts_match(p.at_step(0.51), prompt(1.0, "b d"))
assert prompts_match(p.at_step(30), prompt(1.0, "b d"))
@pytest.mark.parametrize("step", [0, 0.5, 1])
def test_quote(step, parse):
p = parse('This is a text with a "QUOTED DEF(X=Y)"')
expected = prompt(1.0, 'This is a text with a "QUOTED DEF(X=Y)"')
assert prompts_match(p.at_step(step), expected)
@pytest.mark.parametrize(
"group",
[
["[a:0.1]", "[:a:0.1]", "[:a::0.1,1.0]", "[:a::0.1,1.0]", "[:a::0.1]"],
["[before:during:after:0.1]", "[before:during:after:0.1,1.0]", "[before:during:0.1]"],
["[a:0.1,0.5]", "[[a:0.1]::0.5]", "[:a::0.1,0.5]", "[a::0.1,0.5]"],
["[a:b:0.5]", "[a::b:0.5,0.5]"],
["[a::0.5]", "[a:::0.5,0.5]"],
],
)
def test_equivalences(group, parse):
objects = [parse(g) for g in group]
first = objects[0].parsed_prompt
for obj in objects[1:]:
assert obj.parsed_prompt == first
def test_basic(parse):
p = parse(
"This is a (basic:0.6) (prompt) with (very [[simple]:(basic:0.6):0.5]:1.1) [features::0.8][ and this is ignored:1]"
)
assert_prompt(p, 0.5, 0.5, "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features")
@pytest.mark.parametrize("step", [0, 0.5, 1])
def test_basic_cornercase(parse, step):
p = parse("This contains[ an ignored segment in:1] the prompt")
assert_prompt(p, step, 1.0, "This contains the prompt")
def test_basic_ok(parse):
p = parse(
"This is a (basic:0.6) (prompt) with (very [[simple]:(basic:0.6):0.5]:1.1) [features::0.8][ and this is ignored:1]"
)
assert_prompt(p, 0, 0.5, "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features")
assert_prompt(p, 0.7, 0.8, "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) features")
assert_prompt(p, 1.0, 1.0, "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) ")
@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>")
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)
)
assert prompts_match(p.at_step(step), expected)
def test_scheduled_lora(parse):
p = parse(
"This is a (lora:0.6) (prompt) with [scheduling] features [<lora:foo:0.5>:<lora:bar:0.5:0.2>:0.3] <lora:bar:0.5:1.0>"
)
assert_prompt(
p, 0.1, 0.3, "This is a (lora:0.6) (prompt) with [scheduling] features ", ("foo", 0.5, 0.5), ("bar", 0.5, 1.0)
)
assert_prompt(p, 0.5, 1.0, "This is a (lora:0.6) (prompt) with [scheduling] features ", ("bar", 1.0, 1.2))
@pytest.mark.parametrize(
"text",
[
"This is a sequence of [SEQ:a:0.2::0.5:c:0.8][SEQ: and x:0.8]",
"This is a sequence of [[a:[c:0.5]:0.2]::0.8][ and x::0.8]",
],
)
def test_seq(parse, text):
p = parse(text)
prompts = {
0.2: "This is a sequence of a and x",
0.5: "This is a sequence of and x",
0.8: "This is a sequence of c and x",
1.0: "This is a sequence of ",
}
for k, v in prompts.items():
assert_prompt(p, k, k, v)
def test_shortcuts_scheduling(parse):
p = parse("A schedule [a:0.1,0.7] b")
p2 = parse("A schedule [[a:0.1]::0.7] b")
p3 = parse("A schedule [a:b:0.5,0.8]")
p4 = parse("A schedule [[a:0.5]:b:0.8]")
assert p.parsed_prompt == p2.parsed_prompt
assert p3.parsed_prompt == p4.parsed_prompt
@pytest.mark.parametrize(
"step,until,text",
[
(0, 0.1, "test excluded test"),
(0.2, 0.4, "test test"),
(0.45, 1.0, "test excluded2 test"),
],
)
def test_range_1(step, until, text, parse):
p = parse("test [excluded::excluded2:0.1,0.4] test")
assert_prompt(p, step, until, text)
@pytest.mark.parametrize(
"step,until,text",
[
(0, 0.1, "test test"),
(0.25, 0.3, "test included test"),
(0.15, 0.2, "test excluded test"),
(0.55, 0.6, "test test"),
(0.95, 1.0, "test excluded2 test"),
],
)
def test_range_2(step, until, text, parse):
p = parse("test [[:included::0.2,0.8]|[excluded::excluded2:0.4,0.9]:0.1] test")
assert_prompt(p, step, until, text)
def test_nested(parse):
p = parse(
"This [prompt is [SEQ:[crazy:weird:0.2] stuff:0.5:<lora:cool:1>:0.7:nesting:1.0]:completely ignored with tags:HR]"
)
prompts = {
0.2: (0.2, "This prompt is crazy stuff"),
0.3: (0.5, "This prompt is weird stuff"),
0.5: (0.5, "This prompt is weird stuff"),
0.8: (1.0, "This prompt is nesting"),
}
for k in prompts:
exp = [prompts[k][0], {"prompt": prompts[k][1], "loras": {}}]
assert prompts_match(p.at_step(k), exp)
assert_prompt(p, 0.6, 0.7, "This prompt is ", ("cool", 1.0, 1.0))
assert_prompt(p, 0.7, 0.7, "This prompt is ", ("cool", 1.0, 1.0))
p2 = p.with_filters(filters="hr, xyz")
assert prompts_match(p2.at_step(0), p2.at_step(1))
def test_def(parse):
p = parse("DEF(X=0.5) [a:b:X] DEF(test = [c:X]) test test")
cases = [
(0.2, 0.5, "a "),
(0.6, 1.0, "b c c"),
]
for k, until, text in cases:
assert_prompt(p, k, until, text)
p = parse("DEF(X=[($1):($1:$2):$2])X(test;0.7)")
p2 = parse("[(test):(test:0.7):0.7]")
assert p.parsed_prompt == p2.parsed_prompt
p = parse("DEF(X=[($1):($1:$2):$2])DEF(Y=X(test;$1))Y(0.7) Y(0.5)")
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",
[
(r"[embedding\:a:embedding\:b:0.1,0.5]", [(0.15, 0.5, r"embedding:a"), (0.55, 1, r"embedding:b")]),
(
r"[embedding\:a:embedding\:b:embedding\:c:0.1,0.5]",
[(0.0, 0.1, r"embedding:a"), (0.15, 0.5, r"embedding:b"), (0.55, 1, r"embedding:c")],
),
(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")]),
],
)
def test_escapes(text, cases, parse):
p = parse(text)
for step, until, val in cases:
assert_prompt(p, step, until, val)
# I think these were wrong in the old parser too
@pytest.mark.xfail(reason="Old parser behaviour, possibly buggy")
@pytest.mark.parametrize(
"text, cases",
[
(r"[a:\:a:0.5] :\[a:b:0.5]", [(0, 0.5, r"a :\[a:b:0.5]"), (0.55, 1, r":a :\[a:b:0.5]")]),
],
)
def test_escapes_fail(text, cases, parse):
p = parse(text)
for step, until, val in cases:
assert_prompt(p, step, until, val)
def test_comments(parse):
p = parse("this is a # comment")
assert_prompt(p, 0, 1.0, "this is a ")
p = parse("this is a [comment#:scheduled:0.6]")
assert_prompt(p, 0, 1.0, "this is a [comment")
p = parse(r"this is a [comment\#:scheduled:0.6]")
assert_prompt(p, 0, 0.6, "this is a comment#")
assert_prompt(p, 0.65, 1.0, "this is a scheduled")
p = parse("#this is a comment\nthis is a prompt")
assert_prompt(p, 0, 1.0, "\nthis is a prompt")
def test_misc(parse):
p = parse("[[a:c:0.5]:0.7]")
p2 = parse("[:[a:c:0.5]:0.7]")
assert p.parsed_prompt == p2.parsed_prompt
p = parse("test [[a:[b<lora:test:0.5>:0.6]:0.5]:HR]")
p2 = parse("test [:[a:[:b<lora:test:0.5>:0.6]:0.5]:HR]")
assert p.parsed_prompt == p2.parsed_prompt
def test_filters(parse):
p = parse("test [[a:[b<lora:test:0.5>:0.6]:0.5]:HR]")
p2 = parse("test [:[a:[:b<lora:test:0.5>:0.6]:0.5]:HR]")
assert p.parsed_prompt == p2.parsed_prompt
pf = p.with_filters(filters="hr")
assert pf.parsed_prompt == p2.with_filters(filters="hr").parsed_prompt
assert_prompt(pf, 0, 0.5, "test a")
assert_prompt(pf, 0.55, 0.6, "test ")
assert_prompt(pf, 0.8, 1.0, "test b", ("test", 0.5, 0.5))
p = parse("[:[<lora:test:1>:c:0.5]:0.3]")
assert_prompt(p, 0, 0.3, "")
assert_prompt(p, 0.4, 0.5, "", ("test", 1.0, 1.0))
assert_prompt(p, 1.0, 1.0, "c")
def test_emb(parse):
p = parse("an [<emb:foo>:<emb:bar>:0.5]")
prompts = {
0.2: (0.5, "an embedding:foo"),
0.8: (1.0, "an embedding:bar"),
}
for k, (until, val) in prompts.items():
assert_prompt(p, k, until, val)
def test_alternating_defaultstep(parse):
p = parse("[cat|dog|tiger]")
p2 = parse("[cat|dog|tiger:0.1]")
assert p.parsed_prompt == p2.parsed_prompt
def test_alternating_basic(parse):
p = parse("[cat|dog|tiger]")
p2 = parse("[cat|dog|tiger:0.1]")
assert p.parsed_prompt == p2.parsed_prompt
@pytest.mark.parametrize(
"equivalent",
[
"[cat::0.1][dog:0.1,0.2][tiger:0.2,0.3][cat:0.3,0.4][dog:0.4,0.5][tiger:0.5,0.6][cat:0.6,0.7][dog:0.7,0.8][tiger:0.8,0.9][cat:0.9,1.0]"
],
)
def test_alternating_equivalences(parse, equivalent):
p = parse("[cat|dog|tiger]")
p2 = parse(equivalent)
assert p.parsed_prompt == p2.parsed_prompt
@pytest.mark.xfail(reason="Old parser behaviour")
def test_cornercase_failure(parse):
"""p1 returns a prompt entry until 0 at the start"""
p = parse("[cat:0,0.1]")
p2 = parse("[cat::0.1]")
assert p.parsed_prompt == p2.parsed_prompt
def test_cornercase_corrected(parse):
p = parse("[cat:0,0.1]")
p2 = parse("[cat::0.1]")
assert p.parsed_prompt[0][0] == 0.0
assert p.parsed_prompt[1:] == 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(
[(["cat"],), (["dog"],), (["cat"],), (["wolf", ("canine", 1.0, 1.0)],), (["cat"],)]
):
step = round((i * 0.2) + 0.2, 2)
assert_prompt(p4, step, step, *text)
assert_prompt(p4, 0.7, 0.8, "wolf", ("canine", 1.0, 1.0))
def test_alternating_nested(parse):
p3 = parse("[cat|[dog|wolf]|tiger]")
catdogtigers = ["cat", "wolf", "tiger", "cat", "dog", "tiger", "cat", "wolf", "tiger", "cat"]
for i, x in enumerate(catdogtigers):
step = round((i * 0.1) + 0.1, 2)
assert_prompt(p3, step, step, x)
-13
View File
@@ -1,13 +0,0 @@
#!/usr/bin/env python3
from prompt_control.utils import expand_graph
from prompt_control.nodes_lazy import NODE_CLASS_MAPPINGS as LN
import json
import sys
# Needs ComfyUI in Python path
# Usage: PYTHONPATH=../..:. python tools/expand_graph < graph_in_api_format.json > out.json
if __name__ == "__main__":
graph = json.load(sys.stdin)
new = expand_graph(LN, graph)
print(json.dumps(new))