Compare commits
36
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9857c662af | ||
|
|
32f18ef338 | ||
|
|
d86d9f902c | ||
|
|
f600b1ef90 | ||
|
|
72484d1034 | ||
|
|
e7b444316f | ||
|
|
b9db1c7c4c | ||
|
|
3ee001a3fa | ||
|
|
3fcd708dc6 | ||
|
|
cb5ef8d97a | ||
|
|
24fe09ba4e | ||
|
|
1c7149a811 | ||
|
|
b01bdc3d50 | ||
|
|
5ffd94d7d2 | ||
|
|
a433b246a7 | ||
|
|
bdd0c62665 | ||
|
|
1c763f743c | ||
|
|
c223603100 | ||
|
|
c9efa87291 | ||
|
|
e15d4d693c | ||
|
|
26d5bfc961 | ||
|
|
cdb6c418eb | ||
|
|
dc08fa93b7 | ||
|
|
e29f3ca2d0 | ||
|
|
ee186a73e5 | ||
|
|
7440b309fc | ||
|
|
2ccc697311 | ||
|
|
4559060535 | ||
|
|
60ceddbd73 | ||
|
|
6fd28cfca7 | ||
|
|
11b1313ed4 | ||
|
|
9be7133e48 | ||
|
|
22e78f08eb | ||
|
|
39befe2a34 | ||
|
|
ea9017862f | ||
|
|
6ae6cf0d65 |
@@ -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
|
||||
|
||||
@@ -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 +1,2 @@
|
||||
__pycache__
|
||||
.pyre
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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,
|
||||
]
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
import logging
|
||||
|
||||
|
||||
def pytest_runtest_setup(item):
|
||||
logging.getLogger("comfyui-prompt-control").setLevel(logging.CRITICAL)
|
||||
@@ -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, "-")]
|
||||
@@ -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)
|
||||
@@ -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 == {}
|
||||
@@ -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)
|
||||
@@ -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))
|
||||
Reference in New Issue
Block a user