Compare commits

..
8 Commits
Author SHA1 Message Date
asagi4 8aec0d8f46 I can't run tests if this exists... 2025-12-18 23:51:43 +02:00
asagi4 013586904b Appease typing 2025-12-18 23:51:43 +02:00
asagi4 85ff8466ad Appease typing 2025-12-18 23:51:43 +02:00
asagi4 97d17d2bfd Minor refactor to appease typing 2025-12-18 23:51:43 +02:00
asagi4 c0e671b2de More typing 2025-12-18 23:51:43 +02:00
asagi4 7ff55c6717 Make split_quotable an iterator 2025-12-16 18:08:35 +02:00
asagi4 524738f21d Add some more typing 2025-12-16 18:01:50 +02:00
asagi4 5c7d507e91 Refactor get_function to makes its use consistent
Add some typing, just for fun
2025-12-16 17:47:10 +02:00
44 changed files with 3701 additions and 5888 deletions
+2 -7
View File
@@ -11,13 +11,8 @@ 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 pytest typing-extensions
- run: PYTHONPATH=ComfyUI pytest tests/test_parser.py
- run: pip install -r requirements.txt
- run: python -m prompt_control.test_parser
+4 -2
View File
@@ -31,10 +31,12 @@ jobs:
- name: install-torch
run: pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu
- name: install ComfyUI
run: pip install pytest typing-extensions -r ComfyUI/requirements.txt
run: pip install -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 pytest tests/test_graph.py tests/test_encode.py
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
+6 -15
View File
@@ -1,30 +1,21 @@
ARGS=
all: format check test
@echo "Done"
check:
ty check && ruff check
fix:
ruff check --fix
find . -name "*.py" | xargs pyflakes
format:
ruff format
find . -name "*.py" | xargs black -l 120
test:
PYTHONPATH=../../ pytest tests/test_parser.py tests/test_cutout.py tests/test_macros.py $(ARGS)
python -m prompt_control.test_parser
test_graph:
PYTHONPATH=../../ pytest tests/test_graph.py $(ARGS)
PYTHONPATH=../../ python -m prompt_control.test_graph
test_encode:
PYTHONPATH=../../ pytest tests/test_encode.py $(ARGS)
test_workflow:
PYTHONPATH=../../ pytest tests/test_workflow.py $(ARGS)
PYTHONPATH=../../ python -m prompt_control.test_encode --verbose
test_encode_both:
TEST_TE="clip_l t5" PYTHONPATH=../../ pytest tests/test_encode.py $(ARGS)
TEST_TE="clip_l t5" PYTHONPATH=../../ python -m prompt_control.test_encode --verbose
test_heavy: test_graph test_encode_both
+22 -17
View File
@@ -1,32 +1,24 @@
# ComfyUI prompt control
Control LoRA and prompt scheduling, advanced text encoding, regional prompting, and much more, through your text prompt. Prompt Control generates dynamic graphs that are literally identical to handcrafted noodle soup, condensing complicated workflows with dozens of nodes into simple text prompts.
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.
Prompt Control comes with `PCTextEncode`, which provides advanced text encoding with many additional features compared to ComfyUI's base `CLIPTextEncode`.
A `Basic Text to Image` template is included with the extension, and can be loaded from ComfyUI's template library.
> [!NOTE]
> v3.0.0 is backwards compatible with existing workflows, but requires at least ComfyUI v0.8.0
> The parser was rewritten using parsy. It is intended to have the same behaviour as the old parser, but is **significantly** faster.
> Please report any bugs or incompatibilities you find.
## Notable changes
- `PC: Schedule Prompt` now strips surrounding whitespace by default, which may change some prompts. Add `NOSTRIP()` to your prompt to restore previous behaviour.
## What can it do?
You can use text prompts to control the following:
- A1111-style prompt scheduling and filtering without noodle soup.
- LoRA loading and [scheduling](/doc/schedules.md) using ComfyUI's built-in hook system.
- LoRA loading and [scheduling](/doc/schedules.md) via the prompt, using ComfyUI's hook system
- Masking, composition and area control ([regional prompting](/doc/regional_prompts.md)) with an implementation of [Attention Couple](/doc/attention_couple.md), also fully schedulable.
- [Advanced prompt encoding](/doc/basic.md)
- Per-encoder prompts for models with multiple text encoders, such as SDXL and Flux.
- Per-encoder prompts for models with multiple text encoders, such as SDXL and Flux
- Prompt combinators like `BREAK`, as well as `CAT`, `AVG()` and `AND` corresponding to ComfyUI's `ConditioningConcat`, `ConditioningAverage` and `ConditioningCombine` nodes.
- Different weight interpretation types (ComfyUI, A1111, compel, etc.)
- Prompt masking with an implementation of [cutoff](https://github.com/BlenderNeko/ComfyUI_Cutoff).
- Organize complicated prompts with [segments and prompt macros](/doc/macros.md).
- [Schedule your own encoder nodes](/doc/node_function.md), allowing prompt control of eg. video or audio models with non-text inputs.
- Prompt masking with an implementation of [cutoff](https://github.com/BlenderNeko/ComfyUI_Cutoff)
- Simple [prompt macros](/doc/macros.md) with `DEF`
All features are fully schedulable unless otherwise stated. See the [scheduling syntax documentation](doc/schedules.md) to get started.
@@ -42,9 +34,16 @@ If you encounter issues as a user or if you're a node developer and Prompt Contr
## Requirements
The v3 node schema uses features that require at least ComfyUI v0.8.0
For LoRA scheduling to work, you'll need at least version 0.3.7 of ComfyUI (0.3.36 of ComfyUI desktop).
If you run into problems, update ComfyUI first.
You need to have `lark` installed in your Python environment for parsing to work (If you reuse A1111's venv, it'll already be there).
If you use the portable version of ComfyUI on Windows with its embedded Python, you must open a terminal in the ComfyUI installation directory and run the command:
```
.\python_embeded\python.exe -m pip install lark
```
Then restart ComfyUI afterwards.
# Core nodes
@@ -56,6 +55,10 @@ If you run into problems, update ComfyUI first.
for example, if you first encode `[cat:dog:0.1]` and later change that to `[cat:dog:0.5]`, no re-encoding takes place.
for added fun, put `NODE(NodeClassName, textinputname)` in a prompt to generate a graph using **any other node** that's compatible. The node can't have required parameters besides a single CLIP parameter (which must be named `clip`) and the text prompt, and it must return a `CONDITIONING` as its first return value. The "default" values are `PCTextEncode` and `text`.
For example, if you for some reason do not want the advanced features of `PCTextEncode`, use `NODE(CLIPTextEncode)` in the prompt and you'll still get scheduling with ComfyUI's regular TE node.
The advanced node enables filtering the prompt for multi-pass workflows.
## PCLazyLoraLoader and PCLazyLoraLoaderAdvanced
@@ -80,4 +83,6 @@ This node configures `PCTextEncode` default values for some functions by attachi
# Known issues
- ComfyUI's caching mechanism has an issue that makes it unnecessarily invalidate caches for certain inputs; you'll still get some benefit from the lazy nodes, but changing inputs that shouldn't affect downstream nodes (especially if using filtering) will still cause them to be recomputed because ComfyUI doesn't realize the inputs haven't changed.
- Cutoff does not work with models that use non-CLIP text encoders, like Flux. This might be fixable, but it's uncertain if cutoff even makes sense for those models.
+16 -25
View File
@@ -5,41 +5,32 @@
@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"
v1_modules = []
v3_modules = []
# Importing things here breaks pytest for whatever reason...
if "PYTEST_CURRENT_TEST" not in os.environ:
import importlib
nodes = ["base", "lazy", "tools", "hooks"]
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", "anima"]:
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()
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)
+8 -20
View File
@@ -1,5 +1,7 @@
# Attention Couple
NOTE: This is still considered an experimental feature, so the syntax may change.
Attention Couple is an attention-based implementation of regional prompting. it is faster and often more flexible than latent-based masking.
The implementation is based on the one by [pamparamm](https://github.com/pamparamm/ComfyUI-ppm.git), modified to use ComfyUI's hook system. This enables it to work with prompt scheduling.
@@ -10,11 +12,6 @@ As a consequence of this, however, you can also use `COUPLE` in your negative pr
To enable batching negative prompts, run your positive and negative prompt through the `PPCAttentionCoupleBatchNegative` node. This will make the outputs identical to pamparamm's implementation and will also improve performance. It will fall back to the default behaviour in cases where batching can't be done, so it should always be safe to use.
## Anima
There is a **very experimental** port of pamparamm's Anima support for Attention Couple in Prompt Control. Because ComfyUI lacks the built-in schedulable hooks required, you must first patch your model with `PC: Anima Attention Couple Model Patch` in addition to using `COUPLE` as usual.
The code was hacked together with minimal thought, so expect bugs and misbehaviour. The port is also currently *not* compatible with NegPIP.
## Syntax
@@ -24,13 +21,6 @@ See also the [regional prompting documentation](/doc/regional_prompts.md) for in
You can use `COUPLE` to attach attention-coupled prompts to a base prompt:
For example:
```
dog FILL() COUPLE(0.5 1) cat
```
The full syntax looks as follows (to use `IMASK` you need to attach a custom mask)
`base_prompt COUPLE MASK(0 0.5) coupled prompt 1 with mask COUPLE IMASK(0) coupled prompt 2 with custom mask`
as a shortcut, `COUPLE(maskparams)` is expanded to `COUPLE MASK(maskparams)`, so the above prompt can also be written as:
@@ -38,17 +28,15 @@ as a shortcut, `COUPLE(maskparams)` is expanded to `COUPLE MASK(maskparams)`, so
`base_prompt COUPLE(0 0.5) coupled prompt 1 with mask COUPLE IMASK(0) coupled prompt 2 with custom mask`
Behaviour:
- If no mask is specified, an implicit `MASK()` is assumed, meaning that the prompt affects the entire image.
- If no mask is specified, an implicit `MASK()` is assumed.
- For the base prompt, you can use `FILL()` to automatically mask all parts not masked by other coupled prompts
- For the base prompt, you can also use `FILL()` to automatically mask all parts not masked by coupled prompts
- If the base prompt has weight set to zero (ie. ´:0` at the end), then the first coupled prompt with non-zero weight becomes the base prompt:
- If the base prompt has weight set to zero (ie. ´:0` at the end), then the first coupled prompt with non-zero weight becomes the base prompt.
For example:
```
disabled prompt :0 COUPLE new base prompt COUPLE coupled prompt
dog FILL() COUPLE(0.5 1) cat
```
You can also schedule the weight normally: `prompt :[1:0:0.35]`
> ![NOTE]
> Note that because the generation still sees and diffuses the full latent, attention coupling is not guaranteed to perfectly limit the effect of your prompt to the masked area.
Note that because the generation still sees and diffuses the full latent, attention coupling is not guaranteed to perfectly limit the effect of your prompt to the masked area.
-1
View File
@@ -31,7 +31,6 @@ Prompt operators are processed in the following order, meaning that all features
- DEF macros are expanded
- Scheduling is expanded, and for each scheduled prompt:
- SEGs are processed and the template is expanded
- The prompt is split by AND, and for each:
- Prompts are split by COUPLE. and for each:
- Most functions (like MASK) and cutoffs are evaluated
-39
View File
@@ -58,42 +58,3 @@ a "$1" b "$2"
a "" b "$2"
a "A" b "$2"
```
## SEG: Split your prompt into named segments
Syntax: `SEG(segment_name)`
To help with organizing prompts, you can use the `SEG` function. For example:
```
This is a comic
Top panel: $CAT. $SEG3
Bottom panel: $DOG
SEG(DOG)
A dog chasing its
tail in a living room.
SEG(CAT)
a sleeping cat
SEG
The cat has orange fur with white stripes
```
This produces:
```
This is a comic
Top panel: a sleeping cat. The cat has orange fur with white stripes
Bottom panel: A dog chasing its
tail in a living room.
```
> [!NOTE]
> Unlike macros, SEGs are processed *after* scheduling syntax has been expanded, except in the LoRA loader (this may change later, but requires a bit of refactoring)
In this case, the first section before any `SEG` becomes the *template* and any text after a `SEG` call becomes part of that segment. Whitespace is stripped from the start and end of segments and the template.
In the template, you can refer to segments by either their index (starting from 1) or the given name, prefixed with a `$SEG`, so in this example, `$SEG1` is the same as `$PANEL1`
Segments can also refer to each other. Recursion will terminate, but produces weird outputs.
Naming segments is optional, in which case you will have to refer to it by its index.
-41
View File
@@ -1,41 +0,0 @@
# The NODE function
The `NODE` function allows you to use any other text encoding node within `PC: Schedule Prompt`, replacing the default `PCTextEncode` and allowing for example video model scheduling.
> [!NOTE]
> When using NODE, you lose access to *all* special syntax provided by `PCTextEncode`. Only SEGs, macros and scheduling will continue to work since those are processed at graph expansion time before the text prompt is passed into the node.
## Basic usage
Use `NODE(NodeClassName, textinputname)` in a prompt to generate a graph using any node that's compatible. The requirements are as follows:
- The node must have a CLIP parameter (which must be named `clip`)
- It must have a text field
- It must return a `CONDITIONING` as its first return value.
For example, if you for some reason do not want the advanced features of `PCTextEncode`, use `NODE(CLIPTextEncode)` in the prompt and you'll still get scheduling with ComfyUI's regular TE node.
The default parameters are `PCTextEncode` and `text`.
## Advanced Usage with arbitrary parameters
Advanced usage of `NODE` can be complicated. For an example, see [The H3 workflow](/example_workflows/Prompt%20Control%20with%20MiniMax%20H3.json?raw=1). You can also find it in the template library.
The full synopsis of the function is `NODE(NodeClassName, textinputname, arg_spec)` where `arg_spec` is a semicolon-separated list of `parameter_name json_value` pairs. In raw form, it looks like this:
```
NODE(MiniMaxH3ImageToVideo, prompt, vae ["1", 0]; width 1024; height 1024; first_frame ["2", 0])
```
The names and inputs must match the ComfyUI API format which **may differ from frontend names**. You can export your workflow in API format and inspect it to see how inputs are passed in to nodes.
The arrays are literal ComfyUI node links, meaning the `vae` parameter is taken from node ID "1" first output and `first_frame` from node ID "2" first output.
The values are arbitrary JSON literals, meaning that you can also pass in constant values. To pass in literal strings for example, you need to use quotes `"like this"`.
This is intended to be used with the helper node `PC: NODE Input Helper`, which can be used to pass arbitrary parameters (named `$a` to `$n`) to the encoder. The recommended pattern is to put something like the following:
```
SEG(node)
NODE(MiniMaxH3ImageToVideo, prompt, vae $a; width $b; height $c; length $d; first_frame $e)
```
to the helper and then concatenate it at the end of your prompt (use whitespace as a separator). You can then trigger the node with `$node` in your prompt. (see [documentation](/doc/macros.md) for `SEG`)
The helper will replace the parameters with the correct ComfyUI link values.
-3
View File
@@ -14,9 +14,6 @@ Besides the syntax documented below, the [basic syntax](/doc/basic.md) and [prom
a [large::0.1] [cat|dog:0.05] [<lora:somelora:0.5:0.6>::0.5]
[in a park:in space:0.4]
```
## Note on whitespace
`PC: Schedule Prompt` will strip leading and following whitespace from the prompt automatically. If you really want whitespace in your prompt, include `NOSTRIP()` in your prompt.
## Comments and escaping
In schedules, any text on a line following a `#` is considered a comment and removed, including the `#` character.
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+9 -20
View File
@@ -1,9 +1,8 @@
import itertools
import logging
from math import copysign
import numpy as np
import torch
import numpy as np
from math import copysign
import logging
import itertools
log = logging.getLogger("comfyui-prompt-control")
@@ -42,18 +41,12 @@ 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, strict=False)]
for x, y in zip(weights, word_ids, strict=False)
]
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)]
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, strict=False)]
for x, y in zip(tokens, word_ids, strict=False)
]
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)]
mask = np.array(word_ids) == target_id
return (new_tokens, mask)
@@ -167,7 +160,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), strict=False))
sums = dict(zip(*np.unique(word_ids, return_counts=True)))
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
@@ -176,9 +169,7 @@ 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, strict=False) for w, id in zip(x, y, strict=False) if id != 0]
)
delta = 1 - np.mean([w for x, y in zip(weights, word_ids) for w, id in zip(x, y) 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
@@ -314,9 +305,7 @@ 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], strict=False) if w != 1.0
)
weight_dict = dict((id, w) for id, w in zip(wids, np.array(weights).reshape(-1)[inds]) 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
-167
View File
@@ -1,167 +0,0 @@
# Adapted from https://github.com/pamparamm/ComfyUI-ppm
import itertools
from collections.abc import Callable
from functools import partial
from math import lcm
import torch
import torch.nn.functional as F
from comfy.ldm.anima.model import Anima as AnimaDIT
from comfy.ldm.cosmos.predict2 import Attention as CosmosAttention
from comfy.patcher_extension import WrapperExecutor
from comfy.sampler_helpers import convert_cond
from comfy.samplers import process_conds
COND = 0
UNCOND = 1
def reshape_mask(mask: torch.Tensor, size: tuple[int, int], bs: int, num_tokens: int) -> torch.Tensor:
num_conds = mask.shape[0]
mask_downsample = F.interpolate(mask, size=size, mode="nearest")
mask_downsample_reshaped = mask_downsample.view(num_conds, num_tokens, 1).repeat_interleave(bs, dim=0)
return mask_downsample_reshaped
def wrap_forwards(anima_model):
backups = {}
for block_name, b in (
(n, b) for n, b in anima_model.named_modules() if "cross_attn" in n and isinstance(b, CosmosAttention)
):
backups[block_name] = b.forward
b.forward = partial(cosmos_attention_forward_couple, b.forward)
return backups
def unwrap_forwards(anima_model, backups):
for block_name, b in (
(n, b) for n, b in anima_model.named_modules() if "cross_attn" in n and isinstance(b, CosmosAttention)
):
b.forward = backups[block_name]
def anima_sample_wrapper(executor, *args, **kwargs):
guider, _, extra_options, _, noise, latent_image, denoise_mask, *_ = args
seed = extra_options["seed"]
device = "cuda" # TODO: fix
def pc_process_conds(pc_conds):
conds = [convert_cond([c])[0] for c in pc_conds]
conds = process_conds(
guider.inner_model,
noise,
{"positive": conds},
device,
latent_image,
denoise_mask,
seed,
latent_shapes=[latent_image.shape],
)
return [
c["model_conds"]["c_crossattn"].cond * pc_conds[i][1].get("strength", 1.0)
for i, c in enumerate(conds["positive"])
]
extra_options["model_options"]["transformer_options"]["pc_process_conds"] = pc_process_conds
return executor(*args, **kwargs)
def anima_forward_wrapper(executor: WrapperExecutor, *args, **kwargs):
"""Model wrapper does something with activation shapes?"""
anima_model: AnimaDIT = executor.class_obj # type: ignore
x: torch.Tensor = args[0]
transformer_options: dict = kwargs.get("transformer_options", {}).copy()
pc = transformer_options.get("pc_couple")
if pc and "processed_conds" not in pc:
pc["processed_conds"] = transformer_options["pc_process_conds"](pc["conds"])
patch_spatial = anima_model.patch_spatial
activations_shape = list(x.shape)
activations_shape[-2] = activations_shape[-2] // patch_spatial
activations_shape[-1] = activations_shape[-1] // patch_spatial
transformer_options["activations_shape"] = activations_shape
kwargs["transformer_options"] = transformer_options
b = {}
if pc:
b = wrap_forwards(anima_model)
r = executor(*args, **kwargs)
if pc:
unwrap_forwards(anima_model, b)
return r
def cosmos_attention_forward_couple(_forward: Callable, x, context, rope_emb, transformer_options):
"""attention block wrapper"""
if "pc_couple" not in transformer_options:
return _forward(x, context, rope_emb, transformer_options)
c: torch.Tensor = context
# FIXME: base cond weight
# c = args["processed_conds"][0]
args = transformer_options["pc_couple"]
mask = args["mask"]
conds = args["processed_conds"][1:]
num_conds = len(conds) + 1
num_tokens_c: list[int] = [c.shape[1] for c in conds]
cond_or_uncond = transformer_options["cond_or_uncond"]
cond_or_uncond_couple = []
num_chunks = len(cond_or_uncond)
bs = x.shape[0] // num_chunks
x_chunks = x.chunk(num_chunks, dim=0)
c_chunks = c.chunk(num_chunks, dim=0)
lcm_tokens_c = lcm(c.shape[1], *num_tokens_c)
conds_c_tensor = torch.cat(
[cond.repeat(bs, lcm_tokens_c // num_tokens_c[i], 1) for i, cond in enumerate(conds)],
dim=0,
)
xs, cs = [], []
for i, cond_type in enumerate(cond_or_uncond):
x_target = x_chunks[i]
c_target = c_chunks[i].repeat(1, lcm_tokens_c // c.shape[1], 1)
if cond_type == UNCOND:
xs.append(x_target)
cs.append(c_target)
cond_or_uncond_couple.append(UNCOND)
else:
xs.append(x_target.repeat(num_conds, 1, 1))
cs.append(torch.cat([c_target, conds_c_tensor], dim=0))
cond_or_uncond_couple.extend(itertools.repeat(COND, num_conds))
xs = torch.cat(xs, dim=0)
cs = torch.cat(cs, dim=0)
out = _forward(xs, cs, rope_emb, transformer_options)
size = tuple(transformer_options["activations_shape"][-2:])
num_tokens = out.shape[1]
mask_downsample = reshape_mask(mask, size, bs, num_tokens)
outputs = []
cond_outputs = []
i_cond = 0
for i, cond_type in enumerate(cond_or_uncond_couple):
pos, next_pos = i * bs, (i + 1) * bs
if cond_type == UNCOND:
outputs.append(out[pos:next_pos])
else:
pos_cond, next_pos_cond = i_cond * bs, (i_cond + 1) * bs
masked_output = out[pos:next_pos] * mask_downsample[pos_cond:next_pos_cond]
cond_outputs.append(masked_output)
i_cond += 1
if len(cond_outputs) > 0:
cond_output = torch.stack(cond_outputs).sum(0)
outputs.append(cond_output)
return torch.cat(outputs, dim=0)
+6 -13
View File
@@ -9,6 +9,7 @@ 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
@@ -52,7 +53,7 @@ class Proxy:
return self
def __call__(self, *args, **kwargs):
return self.function(*args, **kwargs)
return self.function(*args, *kwargs)
class AttentionCoupleHook(TransformerOptionsHook):
@@ -63,23 +64,20 @@ class AttentionCoupleHook(TransformerOptionsHook):
def __init__(self):
super().__init__(hook_scope=EnumHookScope.HookedOnly)
self.transformers_dict: dict[str, Any] = {
self.transformers_dict = {
"patches": {
"attn2_output_patch": [Proxy(self.attn2_output_patch)],
"attn2_patch": [Proxy(self.attn2_patch)],
},
"pc_couple": {},
}
}
self.has_negpip = False
# The list will be calculated later. All clones must refer to the same kv dict
self.kv: dict[str, list] = {"k": None, "v": None} # type: ignore
# calculate later. All clones must refer to the same kv dict
self.kv = {}
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: list[float] = [cond[1].get("strength", 1.0) for cond in conds]
self.comfy_conds = [base_cond] + 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]
@@ -119,11 +117,6 @@ class AttentionCoupleHook(TransformerOptionsHook):
self.mask = mask / mask.sum(dim=0, keepdim=True)
def on_apply_hooks(self, model: ModelPatcher, transformer_options: dict[str, Any]):
self.transformers_dict["pc_couple"] = {
"conds": self.comfy_conds,
"num_conds": self.num_conds,
"mask": self.mask,
}
if self.kv["k"] is None:
self.has_negpip = model.model_options.get("ppm_negpip", False)
log.debug("AttentionCouple has_negpip=%s", self.has_negpip)
+8 -7
View File
@@ -1,9 +1,9 @@
import torch
import copy
import logging
import re
import numpy as np
import torch
import logging
log = logging.getLogger("comfyui-prompt-control")
@@ -69,7 +69,10 @@ 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]
weight = 1.0 if weight is None else float(weight)
if weight is None:
weight = 1.0
else:
weight = float(weight)
region_text = region_text.strip()
target_text = target_text.strip()
@@ -136,7 +139,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))), strict=False))
mask_ids = list(zip(*np.nonzero(mask.reshape((len(weighted_tokens), -1)))))
new_prompt = copy.deepcopy(weighted_tokens)
for x, y in mask_ids:
new_prompt[x][y] = (mask_token,) + new_prompt[x][y][1:]
@@ -197,9 +200,7 @@ 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"], strict=False
):
for region, target, weight in zip(clip_regions["regions"], clip_regions["targets"], clip_regions["weights"]):
region_masking = torch.tensor(
regions_normalized * region * weight, dtype=base_embedding_full.dtype, device=base_embedding_full.device
).unsqueeze(-1)
-25
View File
@@ -1,25 +0,0 @@
import re
from .utils import parse_args
CUTOFF_RE = re.compile(r"\[CUT:((.*?):(.*?))\]")
def noop(x):
return x
def parse_cuts(string):
text = CUTOFF_RE.sub(r"\2", string)
cutoffs = CUTOFF_RE.findall(string)
cs = []
for x, *_ in cutoffs:
p = x.split(":")
args = parse_args(
p, [(str, ""), (str, ""), (float, 0), (float, None), (float, None), (noop, None)], strip=False
)
args = tuple(args)
if not args[0] or not args[1] or (args[5] is not None and not args[5].strip()):
raise ValueError(f"Invalid CUT spec: [CUT:{x}]")
cs.append(args)
return text, cs
-131
View File
@@ -1,131 +0,0 @@
# vim: sw=4 ts=4
from __future__ import annotations
import logging
import re
from .utils import find_closing_paren, get_function, split_by_function
log = logging.getLogger("comfyui-prompt-control")
def substitute_template(template, segments, do_subs):
def _substitute(template, segments, stack):
name = ""
if "$" in template:
for name, value in sorted(segments):
value = substitute_var(value, name, "")
if name not in stack:
stack.add(name)
value = _substitute(value, segments, stack)
stack.remove(name)
template = substitute_var(template, name, value)
if do_subs and name not in stack:
template = expand_subs(template)
return template
return _substitute(template, segments, set())
def expand_segs(text, do_subs=True):
template, segments = split_by_function(text, "SEG", defaults=[""], require_args=True)
named_segs = [(f.args[0].strip() or f"SEG{i + 1}", c.strip()) for i, (c, f) in enumerate(segments)]
new_text = substitute_template(template, named_segs, do_subs).strip()
if new_text != text.strip():
log.debug("Template expanded to: %s", new_text)
return new_text
def expand_subs(text):
text, subs = get_function(text, "SUB", defaults=None)
subs = [spec.strip() for f in subs for spec in f.args[0].split(";")]
for spec in subs:
if len(spec) <= 3 or spec[0] != "s":
log.warning("Invalid SUB spec ignored: '%s'", spec)
continue
splitchar = spec[1]
search, replace, *_ = spec[2:].split(splitchar)
text = re.sub(search, replace, text)
return text
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, defs=None):
silent = False
if defs is None:
text, defs = get_function(text, "DEF", defaults=None)
else:
silent = True
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!")
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()
if not silent:
log.debug("DEFs expanded to: %s", res)
return res
def substitute_var(text, name, replace, boundary=r"\b"):
if f"${name}" not in text:
return text
name = re.escape(str(name))
return re.sub(rf"\${name}{boundary}", replace, text)
def substitute_defcall(text, search, replace):
name, default_args = search
def run_macro(*parameters):
paramvals = []
if parameters:
paramvals = [x.strip() for x in parameters[0].split(";")]
r = replace
end_re = r"(?![0-9])"
for i, v in enumerate(paramvals):
r = substitute_var(r, i + 1, v, boundary=end_re)
for i, v in enumerate(default_args):
r = substitute_var(r, i + 1, v, boundary=end_re)
return r
text, _ = get_function(text, name, defaults=None, processor=run_macro, require_args=False)
return text
+46
View File
@@ -0,0 +1,46 @@
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
-52
View File
@@ -1,52 +0,0 @@
# Adapted from ComfyUI-ppm into hook form
import comfy.model_management
import comfy.patcher_extension
from comfy.model_base import Anima
from comfy.model_patcher import ModelPatcher
from comfy_api.latest import io
from .anima_couple import (
anima_forward_wrapper,
anima_sample_wrapper,
)
class PCAnimaAttnCouplePatch(io.ComfyNode):
@classmethod
def define_schema(cls) -> io.Schema:
return io.Schema(
node_id="PCAnimaAttnCouplePatch",
display_name="PC: Anima attention Couple Model Patch",
category="promptcontrol/experimental",
inputs=[
io.Model.Input("model"),
],
outputs=[
io.Model.Output(),
],
)
@classmethod
def execute(cls, model: ModelPatcher) -> io.NodeOutput:
model_type = type(model.model)
m = model
if issubclass(model_type, Anima):
m = model.clone()
m.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL,
cls.__name__,
anima_forward_wrapper,
)
m.add_wrapper_with_key(
comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE,
cls.__name__,
anima_sample_wrapper,
)
return io.NodeOutput(m)
NODES = [PCAnimaAttnCouplePatch]
+35 -69
View File
@@ -1,86 +1,52 @@
# pyright: reportSelfClsParameterName=false
import logging
from comfy_api.latest import io
from .macros import expand_segs
from .prompts import encode_prompt, hook_te
from .prompts import encode_prompt
log = logging.getLogger("comfyui-prompt-control")
class PCTextEncodeWithRange(io.ComfyNode):
class PCTextEncodeWithRange:
@classmethod
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()],
)
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}),
},
}
@classmethod
def execute(cls, clip, text, start=0.0, end=1.0) -> io.NodeOutput:
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):
log.debug("PCTextEncode: Encoding '%s'", text)
defaults = clip.patcher.model_options.get("x-promptcontrol.defaults", {})
masks = clip.patcher.model_options.get("x-promptcontrol.masks", None)
text = expand_segs(text)
out = encode_prompt(clip, text, start, end, defaults, masks)
return io.NodeOutput(out)
return (encode_prompt(clip, text, start, end, defaults, masks),)
class PCTextEncode(io.ComfyNode):
class PCTextEncode:
@classmethod
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()],
)
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP",), "text": ("STRING", {"multiline": True})},
}
@classmethod
def execute(cls, clip, text) -> io.NodeOutput:
# Use the WithRange node for the range 0.0, 1.0
return PCTextEncodeWithRange.execute(clip, text, 0.0, 1.0)
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(clip, text, 0.0, 1.0)
class PCHookEncoderModsInternal(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="PCHookTextEncoderModsInternal",
display_name="PC: Apply Text Encoder Mods",
category="promptcontrol",
description="Apply TE modifications (internal)",
is_experimental=True,
is_dev_only=True,
inputs=[
io.Clip.Input("clip"),
io.String.Input("te_names"),
io.String.Input("style"),
io.String.Input("normalization"),
io.Custom("PC_EXTRA_DATA").Input("extra", optional=True),
],
outputs=[io.Clip.Output()],
)
NODE_CLASS_MAPPINGS = {"PCTextEncode": PCTextEncode, "PCTextEncodeWithRange": PCTextEncodeWithRange}
@classmethod
def execute(cls, clip, te_names, style, normalization, extra) -> io.NodeOutput:
te_names = [x.strip() for x in te_names.split(",")]
clip = hook_te(clip, te_names, style, normalization, extra)
return io.NodeOutput(clip)
NODES = [PCTextEncodeWithRange, PCTextEncode, PCHookEncoderModsInternal]
NODE_DISPLAY_NAME_MAPPINGS = {
"PCTextEncode": "PC: Text Encode (no scheduling)",
"PCTextEncodeWithRange": "PC: Text Encode with Range (no scheduling)",
}
+49 -48
View File
@@ -1,10 +1,10 @@
# pyright: reportSelfClsParameterName=false
import logging
import comfy.hooks
import comfy.utils
import folder_paths
from comfy_api.latest import io
from typing_extensions import override
from comfy.comfy_types.node_typing import IO, ComfyNodeABC, InputTypeDict
from .attention_couple_ppm import AttentionCoupleHook
from .parser import parse_prompt_schedules
@@ -13,27 +13,24 @@ from .utils import consolidate_schedule
log = logging.getLogger("comfyui-prompt-control")
class PCLoraHooksFromText(io.ComfyNode):
class PCLoraHooksFromText:
@classmethod
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()],
)
def INPUT_TYPES(s):
return {
"required": {"text": ("STRING",)},
}
@classmethod
def execute(cls, text) -> io.NodeOutput:
RETURN_TYPES = ("HOOKS",)
OUTPUT_TOOLTIPS = ("set of hooks created from the prompt schedule",)
CATEGORY = "promptcontrol/v2"
FUNCTION = "apply"
EXPERIMENTAL = True
def apply(self, text):
prompt_schedule = parse_prompt_schedules(text)
consolidated = consolidate_schedule(prompt_schedule)
hooks = lora_hooks_from_schedule(consolidated, {})
return io.NodeOutput(hooks)
return (hooks,)
def lora_hooks_from_schedule(schedules, non_scheduled):
@@ -41,7 +38,7 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
lora_cache = {}
all_hooks = []
def create_hook(loras, start_pct, end_pct, non_scheduled):
def create_hook(loraspec, start_pct, end_pct, non_scheduled):
hooks = []
hook_kf = comfy.hooks.HookKeyframeGroup()
for path, info in loras.items():
@@ -55,8 +52,9 @@ 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
ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
new_hook.hooks[0].hook_ref = ref
new_hook.hooks[0].hook_ref = ref # pyright: ignore[reportAttributeAccessIssue]
hooks.append(new_hook)
if start_pct > 0.0:
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=0.0)
@@ -78,42 +76,40 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
start_pct = end_pct
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(io.ComfyNode):
class PCAttentionCoupleBatchNegative(ComfyNodeABC):
@classmethod
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"),
],
)
def INPUT_TYPES(s) -> InputTypeDict:
return {
"required": {
"positive": (IO.CONDITIONING, {}),
"negative": (IO.CONDITIONING, {}),
},
}
@classmethod
@override
def execute(cls, positive, negative) -> io.NodeOutput:
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):
if len(negative) != 1:
log.warning("Batching scheduled negatives is not supported yet")
return io.NodeOutput(positive, negative)
return (positive, negative)
negative_batch = []
for p in positive:
n = [negative[0][0], negative[0][1].copy()]
n_hook_group = n[1].get("hooks", comfy.hooks.HookGroup()).clone()
p_hook_group = p[1].get("hooks", comfy.hooks.HookGroup())
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())
attn_couple = [hook for hook in p_hook_group.hooks if isinstance(hook, AttentionCoupleHook)]
for hook in attn_couple:
n_hook_group.add(hook)
@@ -122,10 +118,15 @@ class PCAttentionCoupleBatchNegative(io.ComfyNode):
n[1]["end_percent"] = p[1].get("end_percent", 1.0)
negative_batch.append(n)
return io.NodeOutput(positive, negative_batch)
return (positive, negative_batch)
NODES = [
PCLoraHooksFromText,
PCAttentionCoupleBatchNegative,
]
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)",
}
+126 -233
View File
@@ -1,19 +1,34 @@
# pyright: reportSelfClsParameterName=false
from __future__ import annotations
import json
import logging
from comfy_api.latest import io
from comfy_execution.graph import ExecutionBlocker
from comfy_execution.graph_utils import GraphBuilder
from .macros import expand_macros, expand_segs
from .parser import parse_prompt_schedules
from .utils import consolidate_schedule, find_nonscheduled_loras, get_function, split_by_function
from comfy_execution.graph_utils import GraphBuilder, is_link
from comfy_execution.graph import ExecutionBlocker
from .utils import 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)
def create_lora_loader_nodes(graph, model, clip, loras):
for path, info in loras.items():
@@ -132,155 +147,64 @@ def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True):
ret = (model, clip, res)
return io.NodeOutput(*ret, expand=r)
return {"result": ret, "expand": r}
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")],
)
class PCLazyLoraLoaderAdvanced:
CACHE_KEY = cache_key_lora
@classmethod
def execute(cls, model=None, clip=None, text="", apply_hooks=True, tags="", start=0.0, end=1.0, num_steps=0):
text = expand_segs(expand_macros(text))
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: tuple[str, ...] = ("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
):
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(io.ComfyNode):
class PCLazyLoraLoader(PCLazyLoraLoaderAdvanced):
@classmethod
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"),
],
)
def INPUT_TYPES(s):
return {
"optional": {
"model": ("MODEL", {"rawLink": True}),
"clip": ("CLIP", {"rawLink": True}),
"text": ("STRING", {"multiline": True, "default": ""}),
},
"hidden": {"unique_id": "UNIQUE_ID"},
}
@classmethod
def execute(cls, model, clip, text):
no = PCLazyLoraLoaderAdvanced.execute(model, clip, text)
return io.NodeOutput(*no.args[:2], expand=no.expand)
RETURN_TYPES = (
"MODEL",
"CLIP",
)
CATEGORY = "promptcontrol"
def parse_extra_inputs(args, defaults):
params = {}
if not args.strip():
return defaults + [{}]
defaults = defaults[:]
defaults.append("")
for i, v in enumerate(args.split(",", maxsplit=len(defaults) - 1)):
defaults[i] = v
# We should strip extra whitespace so that people don't have to worry about functions.
magic_spec = defaults[-1]
magic_spec.replace(r"\;", "__ESCAPED_SEMICOLON__")
extra_inputs = magic_spec.split(";") if magic_spec.strip() else []
for e in extra_inputs:
e = e.strip()
if not e:
continue
e = e.replace("__ESCAPED_SEMICOLON__", ";")
name, jsondata = e.split(maxsplit=1)
jsondata = jsondata.strip()
if not jsondata.strip():
continue
# From helper node:
if jsondata == "__EMPTY__":
continue
try:
params[name.strip()] = json.loads(jsondata.strip())
except ValueError as e:
raise ValueError(f"Invalid JSON input: '{jsondata}'") from e
return [x.strip() for x in defaults[:-1]] + [params]
def make_node(graph, p, clip, strip):
p, classnames = get_function(p, "NODE", defaults=None)
p, filters = get_function(p, "FILTER", defaults=None)
args = ""
if len(classnames) > 1:
log.warning("You have more than one NODE call in your prompt. Only the first one will be used")
if classnames:
args = classnames[0].args[0]
if not args.strip():
raise ValueError("NODE can't be empty!")
classname, paramname, extras = parse_extra_inputs(args, ["PCTextEncode", "text"])
# We should strip extra whitespace so that people don't have to worry about functions.
node = graph.node(classname.strip())
node.set_input("clip", clip)
node.set_input(paramname.strip(), p.strip() if strip else p)
for e, v in extras.items():
node.set_input(e, v)
for f in filters:
classname, paramname, extras = parse_extra_inputs(f.args[0], ["", "conditioning"])
if not classname:
raise ValueError("FILTER requires a Node class name")
extras[paramname] = node.out(0)
node = graph.node(classname)
for e, v in extras.items():
node.set_input(e, v)
return node
def build_prompt(graph, prompt, clip, start=None, end=None):
p = prompt
strip = "NOSTRIP()" not in p
p = p.replace("NOSTRIP()", "")
# Need to explicitly expand SEGs here *before* NODE is processed
p = expand_segs(p)
p, combines = split_by_function(p, "COMBINE")
current_cond = make_node(graph, p, clip, strip)
for text, f in combines:
classname, param1, param2, extra = parse_extra_inputs(f.args[0], ["", "conditioning_1", "conditioning_2"])
if classname.strip() == "":
raise ValueError("Can't use COMBINE without a class name")
combiner = graph.node(classname.strip())
c2 = make_node(graph, text, clip, strip)
extra[param1] = current_cond.out(0)
extra[param2] = c2.out(0)
for e, v in extra.items():
combiner.set_input(e, v)
current_cond = combiner
node = current_cond
if start is not None and end is not None:
node = graph.node("ConditioningSetTimestepRange")
node.set_input("conditioning", current_cond.out(0))
node.set_input("start", start)
node.set_input("end", end)
return node
def apply(self, *args, **kwargs):
r = super().apply(*args, **kwargs)
r["result"] = r["result"][:2]
return r
def build_scheduled_prompts(graph, schedules, clip):
@@ -288,10 +212,20 @@ def build_scheduled_prompts(graph, schedules, clip):
start_pct = 0.0
for end_pct, c in schedules:
p = c["prompt"]
node = build_prompt(graph, p, clip, start_pct, end_pct)
nodes.append(node)
p, classnames = get_function(p, "NODE", ["PCTextEncode", "text"])
classname = "PCTextEncode"
paramname = "text"
if classnames:
classname, paramname = classnames[0].args
node = graph.node(classname)
node.set_input("clip", clip)
node.set_input(paramname, p)
timestep = graph.node("ConditioningSetTimestepRange")
timestep.set_input("conditioning", node.out(0))
timestep.set_input("start", start_pct)
timestep.set_input("end", end_pct)
nodes.append(timestep)
start_pct = end_pct
node = nodes[0]
for othernode in nodes[1:]:
combiner = graph.node("ConditioningCombine")
@@ -302,102 +236,61 @@ def build_scheduled_prompts(graph, schedules, clip):
g = graph.finalize()
log.debug("Built graph: %s", json.dumps(g))
return io.NodeOutput(node.out(0), expand=g)
return {"result": (node.out(0),), "expand": g}
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"),
],
)
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
@classmethod
def execute(cls, clip, text, tags="", start=0.0, end=1.0, num_steps=0):
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):
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(io.ComfyNode):
class PCLazyTextEncode(PCLazyTextEncodeAdvanced):
@classmethod
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"),
],
)
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP", {"rawLink": True}), "text": ("STRING", {"multiline": True})},
"hidden": {"unique_id": "UNIQUE_ID"},
}
@classmethod
def execute(cls, clip, text):
return PCLazyTextEncodeAdvanced.execute(clip, text)
CATEGORY = "promptcontrol"
predefined_macros = get_function(
"""
DEF(AND=COMBINE(ConditioningCombine, conditioning_1, conditioning_2))
DEF(CAT=COMBINE(ConditioningConcat, conditioning_to, conditioning_from))
DEF(AVG(0.5)=COMBINE(ConditioningAverage, conditioning_from, conditioning_to, conditioning_to_strength $1))
""",
"DEF",
defaults=None,
)
NODE_CLASS_MAPPINGS = {
"PCLazyTextEncode": PCLazyTextEncode,
"PCLazyTextEncodeAdvanced": PCLazyTextEncodeAdvanced,
"PCLazyLoraLoader": PCLazyLoraLoader,
"PCLazyLoraLoaderAdvanced": PCLazyLoraLoaderAdvanced,
}
class PCLazyTextEncodeSingle(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="PCLazyTextEncodeSingle",
display_name="PC: Prompt (without scheduling)",
is_experimental=True,
is_dev_only=True,
enable_expand=True,
category="promptcontrol",
inputs=[
io.Clip.Input("clip", raw_link=True),
io.String.Input("text", multiline=True, default=""),
],
outputs=[
io.Conditioning.Output("conditioning"),
],
)
@classmethod
def execute(cls, clip, text):
graph = GraphBuilder()
text = expand_macros(text, predefined_macros)
node = build_prompt(graph, text, clip)
g = graph.finalize()
return io.NodeOutput(node.out(0), expand=g)
NODES = [
PCLazyTextEncode,
PCLazyTextEncodeAdvanced,
PCLazyTextEncodeSingle,
PCLazyLoraLoader,
PCLazyLoraLoaderAdvanced,
]
NODE_DISPLAY_NAME_MAPPINGS = {
"PCLazyTextEncode": "PC: Schedule Prompt",
"PCLazyTextEncodeAdvanced": "PC: Schedule prompt (Advanced)",
"PCLazyLoraLoader": "PC: Schedule LoRAs",
"PCLazyLoraLoaderAdvanced": "PC: Schedule LoRAs (Advanced)",
}
+174 -201
View File
@@ -1,111 +1,150 @@
import json
# pyright: reportSelfClsParameterName=false
import logging
from comfy_api.latest import io
from .macros import expand_macros as macroexpand
from .macros import expand_segs as segexpand
from .macros import expand_subs as subexpand
from .macros import substitute_var
from .parser import parse_prompt_schedules
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
log = logging.getLogger("comfyui-prompt-control")
class PCSetLogLevel(io.ComfyNode):
@classmethod
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()],
)
class PCSaveExpandedWorkflow:
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
@classmethod
def execute(cls, clip, level="INFO") -> io.NodeOutput:
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
)
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"):
log.setLevel(getattr(logging, level))
log.info("Set logging level to %s", level)
return io.NodeOutput(clip)
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"
class PCAddMaskToCLIP(io.ComfyNode):
class PCAddMaskToCLIP:
@classmethod
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()],
)
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP",)},
"optional": {
"mask": ("MASK",),
},
}
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 execute(cls, clip, mask=None) -> io.NodeOutput:
return PCAddMaskToCLIPMany.execute(clip, mask1=mask)
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP",)},
"optional": {
"mask1": ("MASK",),
"mask2": ("MASK",),
"mask3": ("MASK",),
"mask4": ("MASK",),
},
}
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Multi-input version of PCAddMaskToCLIP, for convenience"
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:
def apply(self, clip, mask1=None, mask2=None, mask3=None, mask4=None):
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 io.NodeOutput(clip)
return (clip,)
class PCSetPCTextEncodeSettings(io.ComfyNode):
class PCSetPCTextEncodeSettings:
@classmethod
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()],
)
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}),
},
}
@classmethod
def execute(
cls,
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Configures default values for PCTextEncode"
def apply(
self,
clip,
mask_width=512,
mask_height=512,
@@ -115,7 +154,7 @@ class PCSetPCTextEncodeSettings(io.ComfyNode):
sdxl_target_h=1024,
sdxl_crop_w=0,
sdxl_crop_h=0,
) -> io.NodeOutput:
):
settings = {
"mask_width": mask_width,
"mask_height": mask_height,
@@ -128,132 +167,66 @@ class PCSetPCTextEncodeSettings(io.ComfyNode):
}
clip = clip.clone()
clip.patcher.model_options["x-promptcontrol.settings"] = settings
return io.NodeOutput(clip)
return (clip,)
class PCExtractScheduledPrompt(io.ComfyNode):
class PCExtractScheduledPrompt:
@classmethod
def define_schema(cls):
return io.Schema(
node_id="PCExtractScheduledPrompt",
display_name="PC: Show 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),
io.Boolean.Input("expand_segs", default=False, optional=True),
io.Boolean.Input("expand_subs", default=False, optional=True),
io.Boolean.Input("expand_macros", default=False, optional=True),
],
outputs=[io.String.Output()],
search_aliases=["extract scheduled prompt"],
)
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": ""})},
}
@classmethod
def execute(cls, text, at, tags="", expand_segs=False, expand_subs=False, expand_macros=False) -> io.NodeOutput:
if expand_macros:
text = macroexpand(text)
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=""):
schedule = parse_prompt_schedules(text, filters=tags)
_, entry = schedule.at_step(at)
_, entry = schedule.at_step(at, total_steps=1)
prompt_text = entry.get("prompt", "")
if expand_segs:
prompt_text = segexpand(prompt_text, do_subs=expand_subs)
if expand_subs:
prompt_text = subexpand(prompt_text)
return io.NodeOutput(prompt_text)
return (prompt_text,)
class PCMacroExpand(io.ComfyNode):
class PCMacroExpand:
@classmethod
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()],
)
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"multiline": True}),
},
}
@classmethod
def execute(cls, text) -> io.NodeOutput:
return io.NodeOutput(macroexpand(text))
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),)
class PCLinkHelper(io.ComfyNode):
# a-z
NAMES = [chr(97 + i) for i in range(26)]
NODE_CLASS_MAPPINGS = {
"PCSetPCTextEncodeSettings": PCSetPCTextEncodeSettings,
"PCAddMaskToCLIP": PCAddMaskToCLIP,
"PCAddMaskToCLIPMany": PCAddMaskToCLIPMany,
"PCSetLogLevel": PCSetLogLevel,
"PCExtractScheduledPrompt": PCExtractScheduledPrompt,
"PCSaveExpandedWorkflow": PCSaveExpandedWorkflow,
"PCMacroExpand": PCMacroExpand,
}
@classmethod
def define_schema(cls):
t1 = io.Autogrow.TemplateNames(io.AnyType.Input("link", raw_link=True), min=0, names=cls.NAMES)
t2 = io.Autogrow.TemplateNames(
io.AnyType.Input("value", lazy=True), min=0, names=[f"var{i + 1}" for i in range(50)]
)
return io.Schema(
node_id="PCNODELinkHelper",
display_name="PC: Extra argument helper for NODE",
category="promptcontrol/tools",
description="Takes in arbitrary inputs and renders them as NODE-compatible values, replacing $a -> $z with JSON link values.",
is_experimental=True,
inputs=[
io.Autogrow.Input("links", template=t1),
io.Autogrow.Input(
"vars",
template=t2,
),
io.String.Input(
"template",
tooltip="The variables $a to $z will be replaced in this text with their corresponding input's JSON link value",
placeholder="In this text you can refer to the input links as $a, $b etc. and the var inputs as either $var1 or $json1 etc. (the latter will be rendered through Python's json.dumps function which will cause strings to be quoted)",
multiline=True,
),
],
outputs=[io.String.Output()],
)
# This requires https://github.com/Comfy-Org/ComfyUI/pull/15103 to work properly
# Without that PR, all inputs will be evaluated non-lazily
@classmethod
def check_lazy_status(cls, template, links, vars):
r = []
for name, (v, input_name) in vars.items():
if v is None and f"${name}" in template or v is None and f"$json{name[3:]}" in template:
r.append(input_name)
return r
@classmethod
def execute(cls, template, links, vars) -> io.NodeOutput:
text = template
for k in cls.NAMES:
v = "__EMPTY__"
if k in links:
# Replace : with \: to avoid breaking scheduling syntax when linking subgraphs. Any function that consumes this should replace \: with :
v = json.dumps(links[k]).replace(":", r"\:")
text = substitute_var(text, k, v)
for i in range(50):
v = "__EMPTY__"
k = f"var{i + 1}"
if k in vars:
v = vars[k]
text = substitute_var(text, k, str(v))
if f"$json{i + 1}" in text:
v = v if v == "__EMPTY__" else json.dumps(v)
text = substitute_var(text, f"json{i + 1}", v)
return io.NodeOutput(text)
NODES = [
PCSetPCTextEncodeSettings,
PCAddMaskToCLIP,
PCAddMaskToCLIPMany,
PCSetLogLevel,
PCExtractScheduledPrompt,
PCMacroExpand,
PCLinkHelper,
]
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",
}
+446 -356
View File
@@ -1,397 +1,487 @@
# vim: sw=4 ts=4
from __future__ import annotations
import itertools as it
from dataclasses import dataclass
import lark
import logging
from math import ceil
from typing import Any, TypeAlias
from typing_extensions import override
logging.basicConfig()
log = logging.getLogger("comfyui-prompt-control")
import re
from .macros import expand_macros
from .parsy import any_char, char_from, digit, eof, forward_declaration, generate, regex, seq, string, success
from functools import lru_cache
from .utils import get_function, find_closing_paren
FOREVER = float("inf")
if lark.__version__ == "0.12.0":
from sys import executable
EvalResult: TypeAlias = tuple[float, str, list["LoRA"]]
x = "\n".join(
[
"Your lark package reports an ancient version (0.12.0) and will not work. If you have the 'lark-parser' package in your Python environment, remove that and *reinstall* lark!",
f"{executable} -m pip uninstall lark-parser lark",
f"{executable} -m pip install lark",
]
)
log.error(x)
raise ImportError(x)
def merge_until(i: EvalResult, minimum: float):
until, p, loras = i
until = min(until, minimum)
return until, p, loras
ESCAPES = [
("XxPCBackslashESCAPExX", "\\"),
("XxPCColonESCAPExX", ":"),
("XxPCCommentESCAPExX", "#"),
]
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
def escape_specials(string):
for ph, c in ESCAPES:
string = string.replace(rf"\{c}", ph)
return string
EvalResult: TypeAlias = tuple[float, str, list["LoRA"]]
def restore_escaped(string):
for ph, c in ESCAPES:
string = string.replace(ph, c)
return string
class Expression:
def eval(self, step: float, tags: list[str]) -> EvalResult:
return (FOREVER, "", [])
def required_steps(self, max_steps: float) -> set[float]:
return set()
def remove_comments(string):
r = []
for line in string.split("\n"):
comment = line.find("#")
if comment >= 0:
r.append(line[:comment])
else:
r.append(line)
return "\n".join(r)
@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, []
prompt_parser = lark.Lark(
r"""
!start: (prompt | /[][():|]/+)*
prompt: (emphasized | embedding | scheduled | alternate | sequence | loraspec | PLAIN | | /\\:/ | /</ | />/ | WHITESPACE)+
!emphasized: "(" prompt? ")"
| "(" prompt ":" prompt ")"
| "[" prompt "]"
promptlist: ([prompt] ":")~1..3
scheduled: "[" promptlist _WS? NUMBER ["," NUMBER] "]"
| "[" promptlist _WS? TAG "]"
sequence.5: "[SEQ" ":" [prompt] ":" NUMBER (":" [prompt] ":" NUMBER)* "]"
alternate: "[" [prompt] ("|" [prompt])+ [":" NUMBER] "]"
loraspec.99: "<lora:" FILENAME lora_weights [lora_block_weights] ">"
lora_weights.1: (":" _WS? NUMBER)~1..2
lora_block_weights.-1: ":" PLAIN
embedding.100: "<emb:" FILENAME ">"
WHITESPACE: /\s+/
_WS: WHITESPACE
PLAIN: /([^<>\\\[\]():|]|\\.)+/
FILENAME: /[^<>:]+/
TAG: /[A-Z_]+/
%import common.SIGNED_NUMBER -> NUMBER
""",
lexer="dynamic",
)
@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
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
"""
)
@dataclass
class Sequence(Expression):
prompts: list[tuple[Expression, float]]
class CutTransform(lark.Transformer):
def __default__(self, data, children, meta):
return children
@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
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(flatten(prompt)), # pyright: ignore
"".join(flatten(cutout)), # pyright: ignore
weight,
strict_mask,
start_from_masked,
mask_token,
) # pyright: ignore
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
else:
for g in x:
yield from flatten(g)
def clamp(a, b, c):
"""clamp b between a and c"""
return min(max(a, b), c)
def get_steps(tree, num_steps):
res = [num_steps or 100]
def tostep(s):
steps = num_steps or 100
if "." in str(s) or not num_steps:
w = float(s)
value = w * steps
else:
w = int(s)
value = w
if w > 1 and not num_steps:
log.warning(
"You haven't configured the number of steps for Prompt Control to use, %s will be clipped to 1.0", w
)
value = steps
return int(clamp(0, value, steps))
class CollectSteps(lark.Visitor):
def scheduled(self, tree):
i = tree.children[-1]
if i and i.type == "TAG":
return
for i in [-1, -2]:
if tree.children[i] is not None:
tree.children[i] = tostep(tree.children[i])
res.append(tree.children[i])
def interp_steps(self, tree):
tree.children[-1] = tostep(tree.children[-1] or 0.1)
for i, _ in enumerate(tree.children[:-1]):
tree.children[i] = tostep(tree.children[i])
res.extend(tree.children[:-1])
def sequence(self, tree):
steps = tree.children[1::2]
for i, steps in enumerate(steps):
w = tostep(tree.children[i * 2 + 1])
tree.children[i * 2 + 1] = w
res.append(w)
def alternate(self, tree):
step_size = tostep(round(float(tree.children[-1] or 0.1), 2))
tree.children[-1] = step_size
res.extend([x for x in range(step_size, num_steps or 100, step_size)])
CollectSteps().visit(tree)
return sorted(set(res))
def at_step(step, filters, tree):
class AtStep(lark.Transformer):
def scheduled(self, args):
before = None
during = None
after = None
when_end = None
pl, when, *rest = args
if rest:
when_end = rest[0]
pl = list(pl)
if len(pl) == 1:
(during,) = pl # [after:0.5] == [::after:0.5,0.5]
if when_end is None:
when_end = when
after = during
elif len(pl) == 2:
during, after = pl # [during:after:0.5] = [before::after:0.5,0.5]
if when_end is None:
when_end = when
before = during
else:
before, during, after = pl # [before:during:after:0.5,0.8]
if isinstance(when, str):
return before or "" if when not in filters else after or ""
if when_end is None:
when_end = 1000_000
if step <= when:
return before or ""
if when < step <= when_end:
return during or ""
else:
return after or ""
def sequence(self, args):
previous_step = 0.0
prompts = args[::2]
steps = args[1::2]
for s, p in zip(steps, prompts):
if s >= step and step >= previous_step:
previous_step = step
return p or ""
else:
previous_step = s
return ""
def alternate(self, args):
step_size = args[-1]
idx = ceil(step / step_size)
return args[(idx - 1) % (len(args) - 1)] or ""
def start(self, args):
prompt = []
loraspecs = {}
args = flatten(args)
for a in args:
if isinstance(a, str):
prompt.append(a)
elif isinstance(a, tuple):
# sum identical specs together
n = a[0]
# if clip weight is not provided, use unet weight
w, w_clip = a[1][0], a[1][1 % len(a[1])]
e = loraspecs.get(n, {})
loraspecs[n] = {
"weight": round(e.get("weight", 0.0) + w, 2),
"weight_clip": round(e.get("weight_clip", 0.0) + w_clip, 2),
}
lbw = a[2]
if lbw:
loraspecs[n]["lbw"] = lbw
if loraspecs[n]["weight"] == 0 and loraspecs[n]["weight_clip"] == 0 and not lbw:
del loraspecs[n]
else:
pass
p = "".join(prompt)
return {"prompt": p, "loras": loraspecs}
def PLAIN(self, args):
return restore_escaped(args)
def FILENAME(self, value):
return str(value)
def embedding(self, args):
return "embedding:" + str(args[0])
def lora_weights(self, args):
return [float(str(a)) for a in args]
def lora_block_weights(self, args):
vals = args[0].split(";")
r = {}
for v in vals:
x = v.split("=", 2)
if len(x) != 2:
continue
k, v = x[0].strip().upper(), x[1].strip()
r[k] = v
return r
def loraspec(self, args):
name = args[0]
params = args[1]
lbw = args[2]
return name, params, lbw
def __default__(self, data, children, meta):
return children
return AtStep().transform(tree)
class PromptSchedule(object):
# 0 num_steps means unconfigured
def __init__(self, prompt, filters="", start=0.0, end=1.0, num_steps=0):
self.filters = filters
self.start = start
self.end = end
self.num_steps = num_steps
# placeholder is restored on parse
self.prompt = remove_comments(escape_specials(prompt.strip()))
self.defaults = {}
self.loaded_loras = {}
self.parsed_prompt = self._parse(num_steps)
def __iter__(self):
# Filter out zero, it's only useful for interpolation
return (x for x in self.parsed_prompt if x[0] != 0)
def _parse(self, num_steps):
filters = [x.strip() for x in self.filters.upper().split(",")]
try:
parsed = []
tree = prompt_parser.parse(self.prompt)
steps = get_steps(tree, num_steps=num_steps)
def f(x):
return round(x / (num_steps or 100), 2)
for t in steps:
p = at_step(t, filters, tree)
parsed.append([f(t), p])
except lark.exceptions.LarkError as e:
log.error("Prompt editing parse error: %s", e)
parsed = [[1.0, {"prompt": self.prompt, "loras": {}}]]
raise
# Tag filtering may return redundant prompts, so filter them out here
res = []
prev_end = -1
for end_at, p in parsed:
if end_at < self.start:
continue
elif end_at <= self.end:
res.append([end_at, p])
prev_end = end_at
elif end_at > self.end and prev_end < self.end:
res.append([end_at, p])
break
return merge_until(item.eval(step, tags), found_step)
# Always use the last prompt if everything was filtered
if len(res) == 0:
res = [[1.0, parsed[-1][1]]]
@override
def required_steps(self, max_steps: float):
return set(step for _, step in self.prompts if step <= max_steps)
final = [res[0]]
@dataclass
class Schedule(Expression):
before: Prompt
during: Prompt
after: Prompt
start: float
end: float
tag: str | None
def tag_matches(self, tags: list[str]):
return self.tag in tags
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
if self.tag is not None and not self.tag_matches(tags):
return self.before.eval(step, tags)
if self.tag_matches(tags):
return self.during.eval(step, tags)
if step <= self.start:
return merge_until(self.before.eval(step, tags), self.start)
if self.start < step <= self.end:
return merge_until(self.during.eval(step, tags), self.end)
if step > self.end:
return self.after.eval(step, tags)
raise AssertionError("How are you here?")
@override
def required_steps(self, max_steps: float):
r = set()
if self.start < max_steps:
r.add(self.start)
if self.end < max_steps:
r.add(self.end)
r.update(self.before.required_steps(max_steps))
r.update(self.during.required_steps(max_steps))
r.update(self.after.required_steps(max_steps))
return r
@dataclass
class Prompt(Expression):
data: list[Expression]
@override
def eval(self, step: float, tags: list[str]) -> EvalResult:
evals = [x.eval(step, tags) for x in self.data]
text = "".join(x[1] for x in evals)
untils = [x[0] for x in evals]
loras = []
for x in evals:
loras.extend(x[2])
until = FOREVER if not untils else min(untils)
return until, text, loras
@override
def required_steps(self, max_steps):
r = set()
for x in self.data:
r.update(x.required_steps(max_steps))
return r
@dataclass
class LoRA(Expression):
filename: str
w_model: float = 1.0
w_te: float = 1.0
def eval(self, step: float, tags: list[str]) -> EvalResult:
return FOREVER, "", [self]
def find_weight_at(weights: list[tuple[float, float]], step: float, until: float):
res_w = 0
for this, next in zip(weights, it.chain(weights[1:], [(0, FOREVER)]), strict=False):
w, start = this
_, next_start = next
if start > step or next_start < step:
until = min(until, start)
continue
res_w = w
return until, res_w
@dataclass
class LoRACTL(Expression):
filename: str
w_model: list[tuple[float, float]]
w_te: list[tuple[float, float]]
def eval(self, step: float, tags: list[str]) -> EvalResult:
until, w1 = find_weight_at(self.w_model, step, FOREVER)
until, w2 = find_weight_at(self.w_te, step, until)
lora = []
if w1 != 0 or w1 != 0:
lora = [LoRA(self.filename, w1, w2)]
return until, "", lora
def required_steps(self, max_steps):
r = set(x[1] for x in self.w_model)
r.update(set(x[1] for x in self.w_te))
return r
def combine_arglist(prompts, start_end) -> Schedule:
a, b, c = prompts
start_or_tag, end = start_end
empty = Prompt([])
start = start_or_tag
# Handle [a:b:TAG]
if isinstance(start_or_tag, str):
if b is None:
before = empty
during = a # [a:TAG] produces a when tag is active
else:
before, during = a, b # [a:b:TAG] changes from a to b when tag is active
return Schedule(before, during, empty, start=0.0, end=FOREVER, tag=start_or_tag)
during = before = after = empty
if end is not None:
if b is None: # [a:0,0.5] == [:a:0,0.5]
during = a
before = after = empty
elif c is None: # [a:b:0,0.5]
before = empty
during = a
after = b
else:
before, during, after = a, b, c
else:
end = FOREVER
if b is None: # [a:0.5] == [::a:0.5,0.5]
before = empty
during = a
after = a
else:
before = a
during = b
after = b
# c always gets ignored
start = float(start) # for typechecking
return Schedule(before, during, after, start, end, tag=None)
def token(s: str):
return string(s).map(Text)
def combine_prompt(*prompts):
p = prompts
if len(p) == 1:
p = p[0]
if isinstance(p, Prompt):
p = p.data[0] if len(p.data) == 1 else combine_prompt(*p.data)
if isinstance(p, Expression):
return p
p = [combine_prompt(x) for x in p]
return Prompt(p)
@dataclass
class PromptSchedule:
parse_tree: Expression
filters: list[str]
start: float
end: float
num_steps: int
def at_step(self, step: float) -> tuple[float, dict[str, Any]]:
max_step = self.num_steps or 1.0
if max_step > 1 and step < 1:
step = step * max_step
until, p, lora_list = self.parse_tree.eval(step, self.filters)
loras = {}
for lora in lora_list:
d = loras.get(lora.filename, {})
d["weight"] = d.get("weight", 0) + lora.w_model
d["weight_clip"] = d.get("weight_clip", 0) + lora.w_te
loras[lora.filename] = d
if max_step > 0 and until > 1:
# TODO: better logic for this?
until = min(until / max_step, 1.0)
return (min(max_step, round(until, 2)), {"prompt": p, "loras": loras})
def with_filters(self, filters: str | None = None, start: float | None = None, end: float | None = None):
return PromptSchedule(
self.parse_tree,
self.filters if filters is None else parse_filters(filters),
self.start if start is None else start,
self.end if end is None else end,
self.num_steps,
)
# Clean up duplicates
for p in res[1:]:
if p[1] != final[-1][1]:
final.append(p)
else:
final[-1][0] = p[0]
return final
def clone(self):
return self.with_filters()
def __iter__(self):
return (x for x in self.parsed_prompt if x[0] != 0)
def with_filters(self, filters=None, start=None, end=None, defaults=None):
def ifspecified(x, defval):
return x if x is not None else defval
@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})
p = PromptSchedule(
self.prompt,
filters=ifspecified(filters, self.filters),
start=ifspecified(start, self.start),
end=ifspecified(end, self.end),
num_steps=self.num_steps,
)
return p
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
def at_step(self, step, total_steps=1):
_, x = self.at_step_idx(step, total_steps)
return x
if len(res) == 0:
res = [[1.0], prompts[-1][1]]
return res
def at_step_idx(self, step, total_steps=1):
for i, x in enumerate(self.parsed_prompt):
if x[0] * total_steps >= step:
return i, x
return len(self.parsed_prompt) - 1, self.parsed_prompt[-1]
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]
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]
return parser.desc("lora_weights")
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
prompt = forward_declaration()
empty = Text("")
comma = token(",")
col = token(":")
lsq = token("[")
rsq = token("]")
lpar = token("(")
rpar = token(")")
tag = regex(r"[A-Z_]+")
non_special = regex(r"[^:\[\]()|\\<>#]+").map(Text)
filename = regex(r"[^:<>]+")
comment = string("#") >> any_char.until(eof | char_from("\n")) >> success(empty)
escape = (string("\\") >> char_from("\\[]:#") | string(r"\(") | string(r"\)")).map(Text)
emphasis = seq(lpar, (prompt | col).at_least(0), rpar)
sign = string("+") | string("-")
number = (
(sign.optional("") + (digit.many() + string(".") * 1 + digit.many() | digit.at_least(1)).concat())
.concat()
.map(float)
)
opt_prompt = prompt.optional(empty)
step_range = seq(number | tag, (comma >> number).optional())
arglist = seq((opt_prompt << col).optional() * 3, step_range)
schedule = lsq >> arglist.combine(combine_arglist) << rsq
alternate = (lsq >> seq(prompt.sep_by(string("|"), min=1), (col >> number).optional(0.1)) << rsq).combine(Alternate)
sequence = (lsq >> string("SEQ") >> seq(col >> opt_prompt << col, number).at_least(1) << rsq).map(Sequence)
bracketed = seq(lsq, prompt.at_least(0), rsq) | sequence | schedule | alternate
lora = (string("<lora:") >> filename * 1 + lora_weights(number) << string(">")).combine(LoRA)
ctlweight = seq(number, (string("@") >> number).optional(0)).sep_by(comma, min=1)
loractl = (string("<loractl:") >> filename * 1 + lora_weights(ctlweight) << string(">")).combine(LoRACTL)
emb = (string("<emb:") >> filename << string(">")).map(lambda f: Text(f"embedding:{f}"))
expr = (
escape
| comment
| non_special
| bracketed
| emphasis.combine(combine_prompt)
| lora
| loractl
| emb
| char_from("<>").map(Text)
)
prompt_ = expr.at_least(1).combine(combine_prompt)
prompt.become(prompt_)
# Treat any character that isn't valid prompt syntax as just text
all = (prompt | any_char.map(Text)).at_least(0).combine(combine_prompt)
def 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 parse_filters(filters: str):
return [x.strip().upper() for x in filters.split(",") if x.strip()]
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
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)
@lru_cache
def parse_prompt_schedules(prompt, **kwargs):
prompt = expand_macros(prompt)
return PromptSchedule(prompt, **kwargs)
-720
View File
@@ -1,720 +0,0 @@
# 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__
+41 -46
View File
@@ -1,31 +1,28 @@
from __future__ import annotations
import logging
import math
import re
from collections import defaultdict
from functools import partial
from typing import Any
import torch
import math
from functools import partial
from comfy_extras.nodes_mask import FeatherMask, MaskComposite
from nodes import ConditioningAverage
from .adv_encode import advanced_encode_from_tokens
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,
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
log = logging.getLogger("comfyui-prompt-control")
@@ -35,7 +32,7 @@ AVAILABLE_NORMALIZATIONS = ["none", "mean", "length", "length+mean"]
SHUFFLE_GEN = torch.Generator(device="cpu")
def get_sdxl(text: str, defaults: dict[str, Any]) -> tuple[str, dict[str, int]]:
def get_sdxl(text, defaults):
# Defaults fail to parse and get looked up from the defaults dict
text, sdxl = get_function(text, "SDXL", ["none", "none", "none"])
if not sdxl:
@@ -57,7 +54,7 @@ def get_sdxl(text: str, defaults: dict[str, Any]) -> tuple[str, dict[str, int]]:
return text, opts
def get_clipweights(text: str, existing_spec: dict[str, float] | None = None) -> tuple[dict[str, float], str]:
def get_clipweights(text, existing_spec=None):
text, spec = get_function(text, "TE_WEIGHT", defaults=None)
if not spec:
return existing_spec or {}, text
@@ -73,7 +70,7 @@ def get_clipweights(text: str, existing_spec: dict[str, float] | None = None) ->
return res, text
def get_style(text: str, default_style="comfy", default_normalization="none") -> tuple[str, str, str]:
def get_style(text, default_style="comfy", default_normalization="none"):
text, styles = get_function(text, "STYLE", [default_style, default_normalization])
if not styles:
return default_style, default_normalization, text
@@ -128,8 +125,7 @@ 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])):
@@ -182,7 +178,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 = defaultdict(list)
per_te_prompts = {}
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]
@@ -202,7 +198,9 @@ 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
per_te_prompts[te].append(prompt)
l = per_te_prompts.get(te, [])
l.append(prompt)
per_te_prompts[te] = l
if per_te_prompts:
for key in per_te_prompts:
@@ -241,7 +239,7 @@ def encode_prompt_segment(
can_break = {}
for k in empty:
tokenizer = getattr(clip.tokenizer, f"clip_{k}", getattr(clip.tokenizer, k, None))
can_break[k] = tokenizer and getattr(tokenizer, "pad_to_max_length", False)
can_break[k] = tokenizer and tokenizer.pad_to_max_length
clip = hook_te(clip, empty.keys(), style, normalization, extra)
@@ -397,8 +395,7 @@ 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)
@@ -432,8 +429,7 @@ 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")
@@ -454,9 +450,9 @@ def get_mask(text, size, input_masks):
return text, None, None
def feather(f, mask):
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)
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)
return mask
mask = None
@@ -482,19 +478,22 @@ 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.warning(
log.warn(
"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.warning("IMASK index %s not found, ignoring...", idx)
log.warn("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
mask = call_node(MaskComposite, mask, nextmask, 0, 0, op)[0] if mask is not None else nextmask
if mask is not None:
mask = call_node(MaskComposite, mask, nextmask, 0, 0, op)[0]
else:
mask = nextmask
# apply leftover FEATHER() specs to the whole
for f in feathers[i:]:
@@ -557,6 +556,7 @@ 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
@@ -598,13 +598,11 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
return c
def couple_mask(args):
assert len(args) <= 1, "Argument parsing failure. This is a bug in Prompt Control"
if not args:
if args is None:
return ""
return f"MASK({args[0]})"
return f"MASK({args})"
for prompt in prompts:
prompt, noise_w, generator = get_noise(prompt)
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]
@@ -618,17 +616,16 @@ 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, strict=False):
for base_cond, *attention_couple in zip(*encoded):
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)
@@ -644,8 +641,6 @@ 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
+195
View File
@@ -0,0 +1,195 @@
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)
class TestEncode(unittest.TestCase):
@classmethod
def setUpClass(cls):
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)
@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)
def test_styles(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
for k, clip in clips:
(no_weights,) = run(comfy, clip, "this prompt has no weights")
for style in ["comfy", "A1111", "comfy++", "compel", "down_weight", "perp"]:
with self.subTest(f"TE {k} style {style} no weights equal comfy"):
(c,) = run(pc, clip, "this prompt has no weights")
self.condEqual(no_weights, c)
with self.subTest(f"TE {k} style {style} does not fail when encoding weights"):
for normalization in ["none", "mean", "length", "mean+length", "length+mean"]:
with self.subTest(f"TE {k} style {style} normalization {normalization}"):
(c,) = run(
pc,
clip,
f"STYLE({style}, {normalization}) (this prompt) (has weights:0.9), (a:1.2) (b:1.2)",
)
def test_masks(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
solidmask = comfy_extras.nodes_mask.SolidMask()
setMask = nodes.ConditioningSetMask()
for k, clip in clips:
(c1,) = run(pc, clip, "test MASK()")
(c2,) = run(comfy, clip, "test")
(c2,) = run(setMask, c2, run(solidmask, 1.0, 512, 512)[0], "default", 1.0)
self.condEqual(c1, c2)
self.condEqual(c1, c2, "mask", self.tensorsEqual)
if __name__ == "__main__":
unittest.main()
+244
View File
@@ -0,0 +1,244 @@
import unittest
import unittest.mock as mock
import logging
log = logging.getLogger("comfyui-prompt-control")
def reset_graphbuilder_state():
from comfy_execution.graph_utils import GraphBuilder
GraphBuilder.set_default_prefix("UID", 0, 0)
def find_file(name):
names = {"test": "test.safetensors", "other": "some/other.safetensors"}
return names.get(name)
def loraloader(text, adv=False, **kwargs):
from .nodes_lazy import PCLazyLoraLoader, PCLazyLoraLoaderAdvanced
reset_graphbuilder_state()
if adv:
cls = PCLazyLoraLoader
else:
cls = PCLazyLoraLoaderAdvanced
model = [0, 1]
clip = [0, 0]
return cls().apply(unique_id="UID", model=model, clip=clip, text=text, **kwargs)
def te(text, adv=False, **kwargs):
from .nodes_lazy import PCLazyTextEncode, PCLazyTextEncodeAdvanced
if adv:
cls = PCLazyTextEncode
else:
cls = PCLazyTextEncodeAdvanced
reset_graphbuilder_state()
clip = [0, 0]
return cls().apply(clip=clip, text=text, unique_id="UID", **kwargs)
@mock.patch("prompt_control.utils.lora_name_to_file", find_file)
@mock.patch("torch.cuda.current_device", lambda: "cpu")
class GraphTests(unittest.TestCase):
maxDiff = 4096
def test_textencode(self):
for p in ["test", "[test:0.2] test", "[test[test::0.5]]<lora:test:1>"]:
r1 = te(p)
r2 = te(p, adv=True)
with self.subTest(f"Expansion: {p}"):
self.assertEqual(r1, r2)
reset_graphbuilder_state()
with self.subTest("Expansion: LoRA"):
r = te("test<lora:test:1>")
self.assertEqual(
r,
{
"result": (["UID.0.0.2", 0],),
"expand": {
"UID.0.0.1": {"class_type": "PCTextEncode", "inputs": {"clip": [0, 0], "text": "test"}},
"UID.0.0.2": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.1", 0], "start": 0.0, "end": 1.0},
},
},
},
)
with self.subTest("Expansion: LoRA with schedule"):
r = te("simple [test:0.1,0.5] prompt<lora:test:1>")
self.assertEqual(
r,
{
"result": (["UID.0.0.8", 0],),
"expand": {
"UID.0.0.1": {
"class_type": "PCTextEncode",
"inputs": {"clip": [0, 0], "text": "simple prompt"},
},
"UID.0.0.2": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.1", 0], "start": 0.0, "end": 0.1},
},
"UID.0.0.3": {
"class_type": "PCTextEncode",
"inputs": {"clip": [0, 0], "text": "simple test prompt"},
},
"UID.0.0.4": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.3", 0], "start": 0.1, "end": 0.5},
},
"UID.0.0.5": {
"class_type": "PCTextEncode",
"inputs": {"clip": [0, 0], "text": "simple prompt"},
},
"UID.0.0.6": {
"class_type": "ConditioningSetTimestepRange",
"inputs": {"conditioning": ["UID.0.0.5", 0], "start": 0.5, "end": 1.0},
},
"UID.0.0.7": {
"class_type": "ConditioningCombine",
"inputs": {"conditioning_1": ["UID.0.0.2", 0], "conditioning_2": ["UID.0.0.4", 0]},
},
"UID.0.0.8": {
"class_type": "ConditioningCombine",
"inputs": {"conditioning_1": ["UID.0.0.7", 0], "conditioning_2": ["UID.0.0.6", 0]},
},
},
},
)
@mock.patch("prompt_control.utils.lora_name_to_file", find_file)
def test_loraloader(self):
with self.assertLogs(log, level="WARNING") as cm:
result = loraloader("prompt here <lora:nonexistent:1.0:0.5>")["expand"]
result_adv = loraloader("prompt here <lora:nonexistent:1.0:0.5>", adv=True)["expand"]
self.assertIn("LoRA 'nonexistent' not found", cm.output[0])
self.assertEqual(result, {})
self.assertEqual(result_adv, {})
result = loraloader("<lora:test:1>")["expand"]
result2 = loraloader("prompt here <lora:test:1.0:0.5><lora:test:0:0.5>")["expand"]
result3 = loraloader("prompt here <lora:test:1.0:0.5><lora:test:0:0.5>", adv=True)["expand"]
self.assertEqual(result, result2)
self.assertEqual(result2, result3)
self.assertEqual(
result,
{
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 1.0,
"lora_name": "test.safetensors",
},
}
},
)
result = loraloader("<lora:test:1><lora:other:0.5>")["expand"]
self.assertEqual(
result,
{
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 1.0,
"lora_name": "test.safetensors",
},
},
"UID.0.0.2": {
"class_type": "LoraLoader",
"inputs": {
"model": ["UID.0.0.1", 0],
"clip": ["UID.0.0.1", 1],
"strength_model": 0.5,
"strength_clip": 0.5,
"lora_name": "some/other.safetensors",
},
},
},
)
result = loraloader("prompt here <lora:test:1.0:0.5>")["expand"]
self.assertEqual(
result,
{
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 1.0,
"strength_clip": 0.5,
"lora_name": "test.safetensors",
},
}
},
)
result = loraloader("prompt [<lora:test:0.5>:0.5]")["expand"]
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True)["expand"]
self.assertEqual(result, result2)
expected = {
"UID.0.0.1": {
"class_type": "CreateHookLora",
"inputs": {"lora_name": "test.safetensors", "strength_model": 0.5, "strength_clip": 0.5},
},
"UID.0.0.2": {
"class_type": "CreateHookKeyframe",
"inputs": {"strength_mult": 0.0, "start_percent": 0.0},
},
"UID.0.0.3": {
"class_type": "CreateHookKeyframe",
"inputs": {
"start_percent": 0.5,
"prev_hook_kf": ["UID.0.0.2", 0],
"strength_mult": 1.0,
},
},
"UID.0.0.4": {
"class_type": "SetHookKeyframes",
"inputs": {"hooks": ["UID.0.0.1", 0], "hook_kf": ["UID.0.0.3", 0]},
},
"UID.0.0.5": {
"class_type": "SetClipHooks",
"inputs": {
"clip": [0, 0],
"hooks": ["UID.0.0.4", 0],
"apply_to_conds": True,
"schedule_clip": True,
},
},
}
self.assertEqual(result, expected)
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True, start=0.6)["expand"]
self.assertEqual(
result2,
{
"UID.0.0.1": {
"class_type": "LoraLoader",
"inputs": {
"model": [0, 1],
"clip": [0, 0],
"strength_model": 0.5,
"strength_clip": 0.5,
"lora_name": "test.safetensors",
},
}
},
)
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", end=0.5)["expand"]
self.assertEqual(result2, {})
if __name__ == "__main__":
unittest.main()
+247
View File
@@ -0,0 +1,247 @@
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()
+19 -39
View File
@@ -1,12 +1,11 @@
from __future__ import annotations
import copy
import logging
import re
from collections.abc import Callable, Iterator
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING, Any, TypeAlias, TypeVar
import re
import logging
import copy
from dataclasses import dataclass
from typing import Any, TypeAlias, Iterator, TypeVar, TYPE_CHECKING
if TYPE_CHECKING:
import torch # flakes8: noqa
@@ -20,6 +19,7 @@ class FunctionSpec:
name: str
args: FunctionArgs
position: int
placeholder: str | None
# Allow testing
@@ -34,14 +34,6 @@ except ImportError:
log = logging.getLogger("comfyui-prompt-control")
def flatten(x):
if type(x) in [str, tuple, int, type(None)] or isinstance(x, dict) and "type" in x:
yield x
else:
for g in x:
yield from flatten(g)
def call_node(cls, *args, **kwargs):
if hasattr(cls, "execute"):
# v3 node
@@ -127,8 +119,10 @@ 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]]:
e = r"\(" if require_args else r"\b"
rex = re.compile(rf"\b{func}{e}", re.MULTILINE)
if require_args:
rex = re.compile(rf"\b{func}\(", re.MULTILINE)
else:
rex = re.compile(rf"\b{func}\b", re.MULTILINE)
idx = 0
match = rex.search(text)
@@ -141,10 +135,6 @@ def find_function_spans(
if text[at_paren:after_first_paren] == "(":
end = find_closing_paren(text, after_first_paren)
if end < 0:
# Unclosed paren: skip past this match so the loop terminates
idx += match.end()
text = text[match.end() :]
match = rex.search(text)
continue
args = parse_strings(text[after_first_paren:end], defaults)
end += 1
@@ -158,26 +148,21 @@ def find_function_spans(
def get_function(
text: str,
func: str,
defaults: list[str] | None,
processor: Callable[..., str] | None = None,
require_args: bool = True,
text: str, func: str, defaults: list[str] | None, placeholder: str = "", require_args: bool = True
) -> tuple[str, list[FunctionSpec]]:
spans = [x.span() for x in re.finditer(r'".+?"', text)]
instances = []
count = 0
chunks = []
current = 0
skipped = 0
for start, end, funcname, args in find_function_spans(text, func, require_args, defaults):
ph = None
if spans_include(spans, start, end):
continue
instances.append(FunctionSpec(funcname, args, start - skipped))
skipped += end - start
chunks.append(text[current:start])
if processor:
chunks.append(processor(*args))
if placeholder:
ph = f"\0{placeholder}{count}\0"
instances.append(FunctionSpec(funcname, args, start, ph))
chunks.append(text[current:start] + (ph or ""))
current = end
count += 1
chunks.append(text[current:])
@@ -204,8 +189,7 @@ 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 = []
@@ -279,10 +263,6 @@ def lora_name_to_file(name: str) -> str | None:
search = [f for f in filenames if all(p in f for p in parts)]
if len(search) == 1:
return search[0]
elif len(search) > 1:
if len(search) > 4:
search[4] = "..."
log.warning("Ignored LoRA search 's%'; matched more than one file: %s", name, ", ".join(search[:5]))
return None
@@ -309,7 +289,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, node.FUNCTION)
fn = getattr(node, getattr(node, "FUNCTION"))
expansion = fn(**inputs)
for i, v in enumerate(expansion["result"]):
input_map[(k, i)] = v
+4 -33
View File
@@ -1,10 +1,10 @@
[project]
name = "comfyui-prompt-control"
description = "Nodes for prompt editing and LoRA scheduling, advanced regional prompting (including attention masking) and advanced prompt encoding, all controlled through your text prompt. Feature keywords: comfyui-prompt-control, schedule, macros, attention couple, loractl, A1111"
version = "3.0.0-beta.10"
description = "Provides nodes for prompt editing and LoRA scheduling, advanced regional prompting (including attention masking) and more, all controlled through your text prompt"
version = "2.1.1"
license = { file = "LICENSE" }
requires-python = ">= 3.10"
# 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"]
[project.urls]
Repository = "https://github.com/asagi4/comfyui-prompt-control"
@@ -17,32 +17,3 @@ 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"]
+2
View File
@@ -0,0 +1,2 @@
# 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
lark >= 1.1.9
View File
-5
View File
@@ -1,5 +0,0 @@
import logging
def pytest_runtest_setup(item):
logging.getLogger("comfyui-prompt-control").setLevel(logging.CRITICAL)
-26
View File
@@ -1,26 +0,0 @@
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, "-")]
-265
View File
@@ -1,265 +0,0 @@
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_neq(c1, c2, key=None, key_assert=None):
ok = False
try:
cond_equal(c1, c2, key=key, key_assert=key_assert)
except AssertionError:
ok = True
if not ok:
raise ValueError("Tensors should not be equal")
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)
def test_noise_weight0(self, text_encoder_clips, pc_text_encode, node_class_objs):
for _k, clip in text_encoder_clips:
(c1,) = run(pc_text_encode, clip, "test")
(c2,) = run(pc_text_encode, clip, "test NOISE(0, 0)")
cond_equal(c1, c2)
def test_noise(self, text_encoder_clips, pc_text_encode, node_class_objs):
for _k, clip in text_encoder_clips:
(c1,) = run(pc_text_encode, clip, "test")
(c2,) = run(pc_text_encode, clip, "test NOISE(1, 0)")
cond_neq(c1, c2)
-686
View File
@@ -1,686 +0,0 @@
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_textencode_custom():
r = te("NODE(CLIPTextEncode)simple [test:0.1,0.5] $p SEG(p) prompt")
assert r == {
"result": (["UID.0.0.8", 0],),
"expand": {
"UID.0.0.1": {
"class_type": "CLIPTextEncode",
"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": "CLIPTextEncode",
"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": "CLIPTextEncode",
"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_textencode_custom_extra():
r = te(
'NODE(CustomTextEncode, prompt, image ["1\:1", 0]; option "test"; float [10.0:__EMPTY__:0.5])simple [test:0.1,0.5] prompt'
)
assert r == {
"result": (["UID.0.0.8", 0],),
"expand": {
"UID.0.0.1": {
"class_type": "CustomTextEncode",
"inputs": {
"clip": [0, 0],
"prompt": "simple prompt",
"image": ["1:1", 0],
"option": "test",
"float": 10.0,
},
},
"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": "CustomTextEncode",
"inputs": {
"clip": [0, 0],
"prompt": "simple test prompt",
"image": ["1:1", 0],
"option": "test",
"float": 10.0,
},
},
"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": "CustomTextEncode",
"inputs": {"clip": [0, 0], "prompt": "simple prompt", "image": ["1:1", 0], "option": "test"},
},
"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 == {}
def test_loraloader_segs():
result = loraloader("prompt [<lora:test:0.5>:0.5]")["expand"]
result2 = loraloader("prompt [$lora:0.5]\nSEG(lora)<lora:test:0.5>\nSEG(lora2)<lora:ignored:1>")["expand"]
assert result == result2
-77
View File
@@ -1,77 +0,0 @@
from textwrap import dedent
import pytest
from prompt_control.macros import expand_macros, expand_segs
@pytest.mark.parametrize(
"text, result",
[
("DEF(X(a;b)=$1 $2 $3 d)X(A) X(A;B;C)", "A b $3 d A B C d"),
(
"DEF(MACRO()=[empty:$1:$2])MACRO MACRO(;) MACRO(;0.5) MACRO(a;0.5)",
"[empty::$2] [empty::] [empty::0.5] [empty:a:0.5]",
),
("DEF(X=$1)DEF(Y()=$1)[X Y][X() Y()][X(1) Y(1)]", "[$1 ][ ][1 1]"),
(
"DEF(C_ANIMAL=cat)DEF(D_ANIMAL=dog)DEF(IT=It is a $1_ANIMAL $10_ANIMAL)IT(D) IT(C)",
"It is a dog $10_ANIMAL It is a cat $10_ANIMAL",
),
],
)
def test_basic_macro(text, result):
assert expand_macros(text) == result
def test_macro_recursion():
with pytest.raises(ValueError) as c:
expand_macros("DEF(X=recurse Y) DEF(Y=recurse X) X")
assert "Unable to resolve DEFs" in str(c.value)
def test_parsing_cornercase():
r = expand_macros("This should not get stuck DEF(")
assert r == "This should not get stuck DEF("
@pytest.mark.parametrize(
"input, output",
[
(
"""\
A red $b and
a blue $a
SEG(a)
cat
SEG(b)
dog
SEG(c)""",
"A red dog and\na blue cat",
),
(
"""\
$a and $b
SEG(a)
cat, $b
SEG(b)
dog, $c
SEG(c)
tiger
""",
"cat, dog, tiger and dog, tiger",
),
(
"""\
$a
SEG(a)
a $b
SEG(b)
b $a""",
"a b a b $a",
),
],
)
def test_segments(input, output):
assert expand_segs(dedent(input)) == output
-369
View File
@@ -1,369 +0,0 @@
import pytest
from prompt_control.parser import parse_prompt_schedules as parse
def lora_dict(*loras):
return {lora: {"weight": unet, "weight_clip": te} for lora, unet, te in loras}
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))
params = []
params.append(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
@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")]),
(r"[a:b \(test\):0.2]", [(0, 0.2, r"a"), (0.25, 1, r"b \(test\)")]),
],
)
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_ltgt_in_schedule(parse):
p = parse("This should [<parse> correctly:be <Picture 1>:0.1]<lora:test:1>")
assert_prompt(p, 0.1, 0.1, "This should <parse> correctly", ("test", 1.0, 1.0))
assert_prompt(p, 0.15, 1.0, "This should be <Picture 1>", ("test", 1.0, 1.0))
def test_floats(parse):
p = parse("[a:b:0.5] [c:d:e:0.2,0.7] <lora:test:-0.3>")
p2 = parse("[a:b:.5] [c:d:e:.2,.7] <lora:test:-.3>")
assert p.parsed_prompt == p2.parsed_prompt
def test_alternating_lora(parse):
p4 = parse("[cat|[dog:wolf<lora:canine:1>:0.5]:0.2]")
for i, (text, *_loras) in enumerate(
[(["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)
def test_alternating_with_tags(parse):
p1 = parse("[[a|b]:HR]", filters="HR")
p2 = parse("[a|b]")
assert p1.parsed_prompt == p2.parsed_prompt
-7
View File
@@ -1,7 +0,0 @@
from prompt_control import utils
def test_smart_split():
assert utils.smarter_split(",", "foo,bar") == ["foo", "bar"]
assert utils.smarter_split(",", "(foo,bar),zonk") == ["(foo,bar)", "zonk"]
assert utils.smarter_split(",", r"\(foo,bar),zonk") == [r"\(foo", "bar)", "zonk"]
-139
View File
@@ -1,139 +0,0 @@
{
"1": {
"inputs": {
"text": "positive prompt",
"clip": [
"4",
1
]
},
"class_type": "PCLazyTextEncode",
"_meta": {
"title": "PC: Schedule prompt"
}
},
"2": {
"inputs": {
"ckpt_name": "$TEST_CHECKPOINT"
},
"class_type": "CheckpointLoaderSimple",
"_meta": {
"title": "Load Checkpoint"
}
},
"3": {
"inputs": {
"seed": 0,
"steps": 8,
"cfg": 3,
"sampler_name": "euler",
"scheduler": "simple",
"denoise": 1,
"model": [
"4",
0
],
"positive": [
"9",
0
],
"negative": [
"9",
1
],
"latent_image": [
"5",
0
]
},
"class_type": "KSampler",
"_meta": {
"title": "KSampler"
}
},
"4": {
"inputs": {
"text": "<lora:$TEST_LORA:1>",
"model": [
"2",
0
],
"clip": [
"2",
1
]
},
"class_type": "PCLazyLoraLoader",
"_meta": {
"title": "PC: Schedule LoRAs"
}
},
"5": {
"inputs": {
"width": 1024,
"height": 1024,
"batch_size": 1
},
"class_type": "EmptyLatentImage",
"_meta": {
"title": "Empty Latent Image"
}
},
"6": {
"inputs": {
"samples": [
"3",
0
],
"vae": [
"2",
2
]
},
"class_type": "VAEDecode",
"_meta": {
"title": "VAE Decode"
}
},
"7": {
"inputs": {
"images": [
"6",
0
]
},
"class_type": "PreviewImage",
"_meta": {
"title": "Preview Image"
}
},
"8": {
"inputs": {
"text": "worst quality,",
"clip": [
"4",
1
]
},
"class_type": "CLIPTextEncode",
"_meta": {
"title": "CLIP Text Encode (Prompt)"
}
},
"9": {
"inputs": {
"positive": [
"1",
0
],
"negative": [
"8",
0
]
},
"class_type": "PCAttentionCoupleBatchNegative",
"_meta": {
"title": "PC: Attention Couple (batch negative)"
}
}
}
-40
View File
@@ -1,40 +0,0 @@
import json
import os
import uuid
from time import sleep
import pytest
import requests
@pytest.fixture(scope="module", autouse=True)
def workflow(request):
with open(str(request.path).replace(".py", ".json")) as f:
data = f.read()
data = data.replace("$TEST_CHECKPOINT", os.environ["PC_TEST_CHECKPOINT"])
data = data.replace("$TEST_LORA", os.environ["PC_TEST_LORA"])
return json.loads(data)
def assert_prompt(url, p):
timeout = 60
r = requests.post(f"{url}/prompt", json={"prompt": p, "client_id": str(uuid.uuid4())}).json()
prompt_id = r["prompt_id"]
r = {"status": "pending"}
while r["status"] in ["pending", "in_progress"]:
sleep(1)
assert timeout > 0
timeout -= 1
r = requests.get(f"{url}/api/jobs/{prompt_id}").json()
assert r["status"] == "completed"
@pytest.fixture
def comfyui():
return os.environ.get("PC_TEST_COMFYUI", "http://localhost:8188")
def test_workflow(workflow, comfyui):
prompt = "DEF(blue=green)a blue dog and a cat sitting [COUPLE(0 0.5, 0 1) red (cat,:1.3) COUPLE(0.5 1, 0 1) (blue:1.2) dog,:0.1]"
workflow["1"]["inputs"]["text"] = prompt
assert_prompt(comfyui, workflow)
+13
View File
@@ -0,0 +1,13 @@
#!/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))