Compare commits

...
37 Commits
Author SHA1 Message Date
asagi4 ce0c1cd698 Document TE_WEIGHT as experimental 2025-01-12 15:56:25 +02:00
asagi4 3745bc6879 doc and reformat 2025-01-12 15:56:25 +02:00
asagi4 99c815a3b6 Mechanism for using attention masks via PCTextEncode and a Hook node
This is experimental and may still change
2025-01-12 15:56:01 +02:00
asagi4 9f6c9c11e6 Don't throw an exception when text input is None 2025-01-07 18:13:00 +02:00
asagi4 15127d2466 Experiment: DEF
use DEF(x=whatever goes = here) to define a macro. Any mention of x in the prompt
will be replaced with "whatever, goes = here" (using \bx\b as the regexp)

whitespace is stripped from the ends

DEF is expanded *before* scheduling

Will expand all defined macros in a loop until no changes occur or until a limit of 10
iterations is reached.
2024-12-30 02:33:08 +02:00
asagi4 3fbae90478 v2.0.0-beta.4 2024-12-29 22:38:59 +02:00
asagi4 2f12069821 Adjust logging a bit 2024-12-29 22:38:59 +02:00
asagi4 3baeabb8ee Add note about debug logging in issue template 2024-12-29 22:27:34 +02:00
asagi4 c85134b31e Clarify docs a bit 2024-12-29 22:18:46 +02:00
asagi4 26f7e1ff24 Add node to configure PC logging and change categories a bit 2024-12-29 22:13:17 +02:00
asagi4 e9e8b75d7f Make SHUFFLE and SHIFT a bit smarter when emphasis is used 2024-12-29 21:54:09 +02:00
asagi4 01526e3923 🤦
See #83
2024-12-29 19:50:39 +02:00
asagi4 e888238625 More debug logging 2024-12-29 19:32:37 +02:00
asagi4 53a6d48cb1 Fix cache hack for LazyLoRALoader 2024-12-29 19:32:24 +02:00
asagi4 a0df992741 Just remove prompt caching altogether, ComfyUI's own caching should take care of it 2024-12-29 18:58:23 +02:00
asagi4 be51e0dfc4 Fix broken PCTextEncode
Mistake was hidden because it's usually not used directly
2024-12-29 18:58:23 +02:00
asagi4 c023956b4c Revert "Slightly optimize prompt encoding in some cases"
This reverts commit 3a2d08fcf7.

See #82

The sharing of outputs from this node is broken, to be fixed later
2024-12-29 18:58:10 +02:00
asagi4 b33f24e0cb Dump generated graphs in debug mode 2024-12-29 18:25:28 +02:00
asagi4 a637321356 Make error message with the broken lark package even more obvious 2024-12-22 14:13:04 +02:00
asagi4 3a2d08fcf7 Slightly optimize prompt encoding in some cases 2024-12-17 23:38:06 +02:00
asagi4 99d966d74a Add PCTextEncodeWithRange 2024-12-17 23:32:15 +02:00
asagi4 724488d20b Fix debug logging a bit 2024-12-17 01:28:14 +02:00
asagi4 bec28affbe v2.0.0-beta.3 2024-12-15 15:39:47 +02:00
asagi4 08844019cc Cache hack causes lots of parser calls, memoize parser 2024-12-15 15:34:32 +02:00
asagi4 f7d78e54d5 Change logging format 2024-12-15 15:34:32 +02:00
asagi4 25990cb17e Cache hack for performance
Set PROMPTCONTROL_ENABLE_CACHE_HACK=1 in your environment to enable
2024-12-15 15:34:27 +02:00
asagi4 de1ad8512a Note about cache problem 2024-12-15 00:13:16 +02:00
asagi4 09925976d4 Remove print 2024-12-14 21:25:16 +02:00
asagi4 00061e18f6 Eh, why is caching now not working again? 2024-12-13 23:07:31 +02:00
asagi4 06f1291727 apply_hooks isn't actually required 2024-12-13 23:07:31 +02:00
asagi4 ee914b2920 Merge pull request #77 from DrJKL/patch-2
Add declaration for prev_keyframe
2024-12-13 23:06:38 +02:00
Alexander Brown b8d002facc Add declaration for prev_keyframe
Otherwise you can hit
```
UnboundLocalError: local variable 'prev_keyframe' referenced before assignment
```
2024-12-13 12:22:38 -08:00
asagi4 2f5d62b46b Tag a non-broken release 2024-12-13 20:24:00 +02:00
asagi4 48c0286f09 Fix extra parameter 2024-12-13 20:15:22 +02:00
asagi4 e64a71fc6a Make the description a bit less terse. 2024-12-13 19:38:34 +02:00
asagi4 dbd5a0e6d6 Get rid of dead code 2024-12-13 18:19:12 +02:00
asagi4 f3a4b12bc0 Clarifications 2024-12-13 18:06:37 +02:00
14 changed files with 402 additions and 144 deletions
+2
View File
@@ -20,5 +20,7 @@ A clear and concise description of what the bug is.
Information needed to trigger the problem.
If possible, attach a workflow to reproduce the problem
If a workflow works, but isn't producing the correct output, please enable debug logging with the `PCSetLogLevel` node (from `promptcontrol/tools`) and run your workflow with debug logging enabled, and copy the outputs here.
**Expected behavior**
A description of what you expected to happen.
+14 -11
View File
@@ -1,13 +1,10 @@
# ComfyUI prompt control
Nodes for LoRA and prompt scheduling that make basic operations in ComfyUI completely prompt-controllable.
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 v2
Prompt control has been almost completely rewritten. It now uses ComfyUI's lazy execution to build graphs from the text prompt at runtime. This has some advantages:
- ComfyUI will not re-run unchanged parts of generated graphs. This is especially useful for two-pass workflows where previously you'd be forced to re-run the first sampling pass even with filtering. That is no longer the case and it does the right thing.
- The generated graph is often exactly equivalent to a manually built workflow using native ComfyUI nodes. There are no more weird sampling hooks that could cause problems with other nodes
Prompt control has been almost completely rewritten. It now uses ComfyUI's lazy execution to build graphs from the text prompt at runtime. The generated graph is often exactly equivalent to a manually built workflow using native ComfyUI nodes. There are no more weird sampling hooks that could cause problems with other nodes
Prompt Control also comes with `PCTextEncode`, which provides advanced text encoding with many additional features compared to ComfyUI's base `CLIPTextEncode`.
@@ -65,23 +62,25 @@ Then restart ComfyUI afterwards.
# Core nodes
**Note**: The documentation refers to the nodes with their internal names for consistency. The display name may change, but ComfyUI's search will always find the nodes with the internal name. `PCLazyTextEncode` and `PCLazyLoraLoader` are the main ones you'll want to use, also known as `PC: Schedule Prompt` and `PC: Schedule LoRas`.
## PCLazyTextEncode and PCLazyTextEncodeAdvanced
`PCLazyTextEncode` uses ComfyUI's lazy graph execution mechanism to generate a graph of `PCTextEncode` and `SetConditioningTimestepRange` nodes from a prompt with schedules. This has the advantage that if a part of the schedule doesn't change, ComfyUI's caching mechanism allows you to avoid re-encoding the non-changed part.
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, paramname)` 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.
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
This node reads LoRA expressions from the scheduled prompt and constructs a graph of `LoraLoader`s and `CreateHookLora`s as necessary to provide the necessary LoRA scheduling.
This node reads LoRA expressions from the scheduled prompt and constructs a graph of `LoraLoader`s and `CreateHookLora`s as necessary to provide the necessary LoRA scheduling. Just use it in place of a `LoRALoader` and use the output normally.
If you have `apply_hooks` set to true, you **do not** need to apply the `HOOKS` output to a CLIP model separately; it's provided in case you want to use it elsewhere.
The advanced node enables filtering the prompt for multi-pass workflows.
The Advanced node gives you access to the generated hooks. If you have `apply_hooks` set to true, you **do not** need to apply the `HOOKS` output to a CLIP model separately; it's provided in case you want to use it elsewhere. The advanced node also enables filtering the prompt for multi-pass workflows.
## PCTextEncode
@@ -168,4 +167,8 @@ The parameters affect how the masked and unmasked prompts are combined to produc
# Known issues
- None at the moment
- 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.
If you want to enable a hack to fix this, set `PROMPTCONTROL_ENABLE_CACHE_HACK=1` in your environment. Unset it to disable.
It's a purely optional performance optimization that allows Prompt Control nodes to override their cache keys in a way that should not interfere with other nodes. Note that the optimization only works if the text input to the lazy nodes is a constant (so either directly on the node or from a primitive); outputs from other nodes can't be optimized.
+21 -5
View File
@@ -1,3 +1,10 @@
"""
@author: asagi4
@title: ComfyUI Prompt Control
@nickname: ComfyUI Prompt Control
@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 os
import sys
import logging
@@ -8,7 +15,7 @@ log = logging.getLogger("comfyui-prompt-control")
log.propagate = False
if not log.handlers:
h = logging.StreamHandler(sys.stdout)
h.setFormatter(logging.Formatter("[%(levelname)s] PromptControl: %(message)s"))
h.setFormatter(logging.Formatter("[PromptControl] %(levelname)s: %(message)s"))
log.addHandler(h)
if os.environ.get("PROMPTCONTROL_DEBUG"):
@@ -16,19 +23,28 @@ if os.environ.get("PROMPTCONTROL_DEBUG"):
else:
log.setLevel(logging.INFO)
cache_hack = importlib.import_module(".prompt_control.cache_hack", package=__name__)
cache_hack.init()
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
nodes = ["base", "lazy", "tools"]
optional_nodes = ["attnmask"]
if importlib.util.find_spec("comfy.hooks"):
nodes.append("hooks")
nodes.extend(["hooks"])
else:
log.warning(
"Your ComfyUI version is too old, can't import comfy.hooks for PCEncodeSchedule and PCLoraHooksFromSchedule. Update your installation."
)
log.error("Your ComfyUI version is too old, can't import comfy.hooks. Update your installation.")
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)
for node in optional_nodes:
try:
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)
except ImportError:
log.info(f"Could not import optional nodes: {node}; continuing anyway")
+13 -1
View File
@@ -135,7 +135,7 @@ These functions are applied to each prompt chunk **after** `BREAK`, `AND` etc. h
Multiple instances of these functions are applied in the order they appear in the prompt.
**NOTE:** These functions are *not* smart about syntax and will break emphasis if the separator occurs inside parentheses. I might fix this at some point, but for now, keep this in mind.
**NOTE** To avoid breaking emphasis syntax, the functions ignore any separators inside parentheses
For example:
- `SHIFT(1) cat, dog, tiger, mouse` does a shift and results in `dog, tiger, mouse, cat`. (whitespace may vary)
@@ -195,3 +195,15 @@ The order of the `FEATHER` and `MASK` calls doesn't matter; you can have `FEATHE
## Miscellaneous
- `<emb:xyz>` is alternative syntax for `embedding:xyz` to work around a syntax conflict with `[embedding:xyz:0.5]` which is parsed as a schedule that switches from `embedding` to `xyz`.
# Experimental features
Experimental features are unstable and may disappear or break without warning.
## Attention masking
Use `ATTN()` in combination with `MASK()` or `IMASK()` to enable attention masking. Currently, it's pretty slow and only works with SDXL. You need to have a recent enough version of ComfyUI for this to work.
## TE_WEIGHT
For models using multiple text encoders, you can set weights per TE using the syntax `TE_WEIGHT(clipname=weight, clipname2=weight2, ...)` where `clipname` is one of `g`, `l`, or `t5xxl`. For example with SDXL, try `TE_WEIGHT(g=0.25, l=0.75`). The weights are applied as a multiplier to the TE output.
+47
View File
@@ -0,0 +1,47 @@
import comfy_execution.caching
from comfy_execution.graph_utils import is_link
import nodes
from os import environ
import logging
log = logging.getLogger("comfyui-prompt-control")
include_unique_id_in_input = comfy_execution.caching.include_unique_id_in_input
def promptcontrol_get_immediate_node_signature(self, dynprompt, node_id, ancestor_order_mapping):
if not dynprompt.has_node(node_id):
# This node doesn't exist -- we can't cache it.
return [float("NaN")]
node = dynprompt.get_node(node_id)
class_type = node["class_type"]
class_def = nodes.NODE_CLASS_MAPPINGS[class_type]
inputs = node["inputs"]
if hasattr(class_def, "CACHE_KEY"):
inputs = getattr(class_def, "CACHE_KEY")(inputs)
signature = [class_type, self.is_changed_cache.get(node_id)]
if (
self.include_node_id_in_input()
or (hasattr(class_def, "NOT_IDEMPOTENT") and class_def.NOT_IDEMPOTENT)
or include_unique_id_in_input(class_type)
):
signature.append(node_id)
for key in sorted(inputs.keys()):
if is_link(inputs[key]):
(ancestor_id, ancestor_socket) = inputs[key]
ancestor_index = ancestor_order_mapping[ancestor_id]
signature.append((key, ("ANCESTOR", ancestor_index, ancestor_socket)))
else:
signature.append((key, inputs[key]))
return signature
def init():
if environ.get("PROMPTCONTROL_ENABLE_CACHE_HACK") != "1":
return
log.warning("Enabling Prompt Control cache hack")
comfy_execution.caching.CacheKeySetInputSignature.get_immediate_node_signature = (
promptcontrol_get_immediate_node_signature
)
+79
View File
@@ -0,0 +1,79 @@
import logging
log = logging.getLogger("comfyui-prompt-control")
from comfy.hooks import TransformerOptionsHook, HookGroup, EnumHookScope
from comfy.ldm.modules.attention import optimized_attention
import torch.nn.functional as F
import torch
from math import sqrt
class MaskedAttn2:
def __init__(self, mask):
self.mask = mask
def __call__(self, q, k, v, extra_options):
mask = self.mask
orig_shape = extra_options["original_shape"]
_, _, oh, ow = orig_shape
seq_len = q.shape[1]
mask_h = oh / sqrt(oh * ow / seq_len)
mask_h = int(mask_h) + int((seq_len % int(mask_h)) != 0)
mask_w = seq_len // mask_h
r = optimized_attention(q, k, v, extra_options["n_heads"])
mask = F.interpolate(mask.unsqueeze(1), size=(mask_h, mask_w), mode="nearest").squeeze(1)
mask = mask.view(mask.shape[0], -1, 1).repeat(1, 1, r.shape[2])
return mask * r
def create_attention_hook(mask):
attn_replacements = {}
mask = mask.detach().to(device="cuda", dtype=torch.float16)
masked_attention = MaskedAttn2(mask)
for id in [4, 5, 7, 8]: # id of input_blocks that have cross attention
block_indices = range(2) if id in [4, 5] else range(10) # transformer_depth
for index in block_indices:
k = ("input", id, index)
attn_replacements[k] = masked_attention
for id in range(6): # id of output_blocks that have cross attention
block_indices = range(2) if id in [3, 4, 5] else range(10) # transformer_depth
for index in block_indices:
k = ("output", id, index)
attn_replacements[k] = masked_attention
for index in range(10):
k = ("middle", 1, index)
attn_replacements[k] = masked_attention
hook = TransformerOptionsHook(
transformers_dict={"patches_replace": {"attn2": attn_replacements}}, hook_scope=EnumHookScope.HookedOnly
)
group = HookGroup()
group.add(hook)
return group
class AttentionMaskHookExperimental:
@classmethod
def INPUT_TYPES(s):
return {
"required": {"mask": ("MASK",)},
}
RETURN_TYPES = ("HOOKS",)
CATEGORY = "promptcontrol/_testing"
FUNCTION = "apply"
EXPERIMENTAL = True
DESCRIPTION = "Experimental attention masking hook. For testing only"
def apply(self, mask):
return (create_attention_hook(mask),)
NODE_CLASS_MAPPINGS = {"AttentionMaskHookExperimental": AttentionMaskHookExperimental}
NODE_DISPLAY_NAME_MAPPINGS = {}
+28 -7
View File
@@ -4,6 +4,29 @@ from .prompts import encode_prompt
log = logging.getLogger("comfyui-prompt-control")
class PCTextEncodeWithRange:
@classmethod
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP",), "text": ("STRING", {"multiline": True})},
"optional": {
"start": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 0.0, "step": 0.01}),
"end": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 1.0, "step": 0.01}),
},
}
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)
return (encode_prompt(clip, text, start, end, defaults, masks),)
class PCTextEncode:
@classmethod
def INPUT_TYPES(s):
@@ -14,17 +37,15 @@ class PCTextEncode:
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):
defaults = clip.patcher.model_options.get("x-promptcontrol.defaults", {})
masks = clip.patcher.model_options.get("x-promptcontrol.masks", None)
return (encode_prompt(clip, text, 0, 1.0, defaults, masks),)
return PCTextEncodeWithRange.apply(self, clip, text, 0.0, 1.0)
NODE_CLASS_MAPPINGS = {
"PCTextEncode": PCTextEncode,
}
NODE_CLASS_MAPPINGS = {"PCTextEncode": PCTextEncode, "PCTextEncodeWithRange": PCTextEncodeWithRange}
NODE_DISPLAY_NAME_MAPPINGS = {
"PCTextEncode": "PC Text Encode (no scheduling)",
"PCTextEncode": "PC: Text Encode (no scheduling)",
"PCTextEncodeWithRange": "PC: Text Encode with Range (no scheduling)",
}
+1 -1
View File
@@ -84,5 +84,5 @@ NODE_CLASS_MAPPINGS = {
}
NODE_DISPLAY_NAME_MAPPINGS = {
"PCLoraHooksFromText": "PC LoRA Hooks From Text (non-lazy)",
"PCLoraHooksFromText": "PC: LoRA Hooks From Text (non-lazy)",
}
+57 -28
View File
@@ -1,12 +1,29 @@
import logging
from .parser import parse_prompt_schedules
from comfy_execution.graph_utils import GraphBuilder
from comfy_execution.graph_utils import GraphBuilder, is_link
from .prompts 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):
@@ -24,9 +41,10 @@ def create_lora_loader_nodes(graph, model, clip, loras):
def create_hook_nodes_for_lora(graph, path, info, existing_node, start_pct, end_pct):
prev_keyframe = None
next_keyframe = None
if not existing_node:
log.debug("Creating hook for %s", path)
log.debug("Creating hook for %s, weight=%s, weight_clip=%s", path, info["weight"], info["weight_clip"])
hook_node = graph.node("CreateHookLora")
hook_node.set_input("lora_name", path)
hook_node.set_input("strength_model", info["weight"])
@@ -59,7 +77,7 @@ def create_hook_nodes_for_lora(graph, path, info, existing_node, start_pct, end_
next_keyframe.set_input("strength_mult", 1.0)
prev_hook_kf = next_keyframe.out(0)
if end_pct < 1.0:
log.debug("Creating end keyframe for %s, start=%s", path, start_pct)
log.debug("Creating end keyframe for %s, start=%s", path, end_pct)
next_keyframe = graph.node("CreateHookKeyframe")
next_keyframe.set_input("strength_mult", 0.0)
next_keyframe.set_input("start_percent", end_pct)
@@ -100,21 +118,22 @@ def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True, return_h
# Finally, combine all hooks and optionally apply
if len(hooks) > 0:
res = hooks[0]
for h in hooks[:1]:
for h in hooks[1:]:
n = graph.node("CombineHooks2")
n.set_input("hooks_A", res.out(0))
n.set_input("hooks_B", h.out(0))
res = n
res = res.out(0)
if apply_hooks:
n = graph.node("SetClipHooks")
n.set_input("clip", clip)
n.set_input("hooks", res)
n.set_input("apply_to_conds", True)
n.set_input("schedule_clip", True)
clip = n.out(0)
if apply_hooks:
n = graph.node("SetClipHooks")
n.set_input("clip", clip)
n.set_input("hooks", res)
n.set_input("apply_to_conds", True)
n.set_input("schedule_clip", True)
clip = n.out(0)
r = graph.finalize()
log.debug("LazyLoraLoader built graph: %s", json.dumps(r))
if return_hooks:
ret = (model, clip, res)
@@ -125,6 +144,8 @@ def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True, return_h
class PCLazyLoraLoaderAdvanced:
CACHE_KEY = cache_key_lora
@classmethod
def INPUT_TYPES(s):
return {
@@ -132,9 +153,9 @@ class PCLazyLoraLoaderAdvanced:
"text": ("STRING", {"multiline": True}),
"model": ("MODEL", {"rawLink": True}),
"clip": ("CLIP", {"rawLink": True}),
"apply_hooks": ("BOOLEAN", {"default": True}),
},
"optional": {
"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}),
@@ -147,13 +168,15 @@ class PCLazyLoraLoaderAdvanced:
CATEGORY = "promptcontrol"
FUNCTION = "apply"
def apply(self, model, clip, text, apply_hooks, unique_id, tags="", start=0.0, end=1.0):
schedule = parse_prompt_schedules(text).with_filters(filters=tags, start=start, end=end)
def apply(self, model, clip, text, unique_id, apply_hooks=True, tags="", start=0.0, end=1.0):
schedule = parse_prompt_schedules(text, filters=tags, start=start, end=end)
graph = GraphBuilder(f"PCLazyLoraLoaderAdvanced-{unique_id}")
return build_lora_schedule(graph, schedule, model, clip, apply_hooks=apply_hooks, return_hooks=True)
class PCLazyLoraLoader:
CACHE_KEY = cache_key_lora
@classmethod
def INPUT_TYPES(s):
return {
@@ -173,7 +196,7 @@ class PCLazyLoraLoader:
CATEGORY = "promptcontrol"
FUNCTION = "apply"
def apply(self, model, clip, text, apply_hooks, unique_id):
def apply(self, model, clip, text, unique_id):
graph = GraphBuilder(f"PCLazyLoraLoader-{unique_id}")
schedule = parse_prompt_schedules(text)
return build_lora_schedule(graph, schedule, model, clip, apply_hooks=True, return_hooks=False)
@@ -182,7 +205,6 @@ class PCLazyLoraLoader:
def build_scheduled_prompts(graph, schedules, clip):
nodes = []
start_pct = 0.0
prompt_cache = {}
for end_pct, c in schedules:
p = c["prompt"]
p, classnames = get_function(p, "NODE", ["PCTextEncode", "text"])
@@ -191,12 +213,9 @@ def build_scheduled_prompts(graph, schedules, clip):
if classnames:
classname = classnames[0][0]
paramname = classnames[0][1]
node = prompt_cache.get((p, classname, paramname))
if not node:
node = graph.node(classname)
node.set_input("clip", clip)
node.set_input(paramname, p)
prompt_cache[(p, classname, paramname)] = node
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)
@@ -210,16 +229,24 @@ def build_scheduled_prompts(graph, schedules, clip):
combiner.set_input("conditioning_2", othernode.out(0))
node = combiner
return {"result": (node.out(0),), "expand": graph.finalize()}
g = graph.finalize()
log.debug("Built graph: %s", json.dumps(g))
return {"result": (node.out(0),), "expand": g}
def cache_key_from_inputs(cachekey, text, tags="", start=0.0, end=1.0, **kwargs):
schedules = parse_prompt_schedules(text, filters=tags, start=start, end=end)
return [(pct, s[cachekey]) for pct, s in schedules]
class PCLazyTextEncode:
CACHE_KEY = cache_key_prompt
@classmethod
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP", {"rawLink": True}), "text": ("STRING", {"multiline": True})},
# "optional": {"defaults": ("SCHEDULE_DEFAULTS",)},
"hidden": {"unique_id": "UNIQUE_ID"},
}
RETURN_TYPES = ("CONDITIONING",)
@@ -227,13 +254,15 @@ class PCLazyTextEncode:
CATEGORY = "promptcontrol"
FUNCTION = "apply"
def apply(self, clip, text, unique_id):
def apply(self, clip, text):
schedules = parse_prompt_schedules(text)
graph = GraphBuilder(f"PCEncodeLazy-{unique_id}")
graph = GraphBuilder()
return build_scheduled_prompts(graph, schedules, clip)
class PCLazyTextEncodeAdvanced:
CACHE_KEY = cache_key_prompt
@classmethod
def INPUT_TYPES(s):
return {
@@ -251,7 +280,7 @@ class PCLazyTextEncodeAdvanced:
FUNCTION = "apply"
def apply(self, clip, text, unique_id, tags="", start=0.1, end=1.0):
schedules = parse_prompt_schedules(text).with_filters(start=start, end=end, filters=tags)
schedules = parse_prompt_schedules(text, filters=tags, start=start, end=end)
graph = GraphBuilder(f"PCLazyTextEncodeAdvanced-{unique_id}")
return build_scheduled_prompts(graph, schedules, clip)
+32 -4
View File
@@ -3,6 +3,32 @@ import logging
log = logging.getLogger("comfyui-prompt-control")
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 (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:
@classmethod
def INPUT_TYPES(s):
@@ -14,7 +40,7 @@ class PCAddMaskToCLIP:
}
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/v2"
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
def apply(self, clip, mask=None):
@@ -35,7 +61,7 @@ class PCAddMaskToCLIPMany:
}
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/v2"
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
def apply(self, clip, mask1=None, mask2=None, mask3=None, mask4=None):
@@ -65,7 +91,7 @@ class PCSetPCTextEncodeSettings:
}
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/v2"
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
def apply(
@@ -101,10 +127,12 @@ NODE_CLASS_MAPPINGS = {
"PCSetPCTextEncodeSettings": PCSetPCTextEncodeSettings,
"PCAddMaskToCLIP": PCAddMaskToCLIP,
"PCAddMaskToCLIPMany": PCAddMaskToCLIPMany,
"PCSetLogLevel": PCSetLogLevel,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"PCSetTextEncodeSettings": "PC: Configure PCTextEncode",
"PCSetPCTextEncodeSettings": "PC: Configure PCTextEncode",
"PCAddMaskToCLIP": "PC: Attach Mask",
"PCAddMaskToCLIPMany": "PC: Attach Mask (multi)",
"PCSetLogLevel": "PC: Configure Logging (for debug)",
}
+49 -80
View File
@@ -4,9 +4,21 @@ from math import ceil
logging.basicConfig()
log = logging.getLogger("comfyui-prompt-control")
import re
from functools import lru_cache
from .utils import get_function
if lark.__version__ == "0.12.0":
x = "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!"
from sys import executable
x = "\n".join(
[
"Your lark package reports an ancient version (0.12.0) and will not work. If you have the 'lark-parser' package in your Python environment, remove that and *reinstall* lark!",
f"{executable} -m pip uninstall lark-parser lark",
f"{executable} -m pip install lark",
]
)
log.error(x)
raise ImportError(x)
@@ -14,16 +26,13 @@ if lark.__version__ == "0.12.0":
prompt_parser = lark.Lark(
r"""
!start: (prompt | /[][():|]/+)*
prompt: (emphasized | embedding | scheduled | alternate | sequence | interpolate | loraspec | PLAIN | /</ | />/ | WHITESPACE)+
prompt: (emphasized | embedding | scheduled | alternate | sequence | loraspec | PLAIN | /</ | />/ | WHITESPACE)+
!emphasized: "(" prompt? ")"
| "(" prompt ":" prompt ")"
| "[" prompt "]"
scheduled: "[" [prompt ":"] [prompt] ":" _WS? NUMBER ["," NUMBER] "]"
| "[" [prompt ":"] [prompt] ":" _WS? TAG "]"
sequence: "[SEQ" ":" [prompt] ":" NUMBER (":" [prompt] ":" NUMBER)+ "]"
interpolate.100: "[INT" ":" interp_prompts ":" interp_steps "]"
interp_prompts: prompt (":" [prompt])+
interp_steps: NUMBER ("," NUMBER)+ [":" NUMBER]
alternate: "[" [prompt] ("|" [prompt])+ [":" NUMBER] "]"
loraspec.99: "<lora:" FILENAME lora_weights [lora_block_weights] ">"
lora_weights.1: (":" _WS? NUMBER)~1..2
@@ -39,6 +48,7 @@ TAG: /[A-Z_]+/
lexer="dynamic",
)
cut_parser = lark.Lark(
r"""
!start: (prompt | /[][:()]/+)*
@@ -94,7 +104,6 @@ def clamp(a, b, c):
def get_steps(tree):
res = [100]
interpolation_steps = []
def tostep(s):
w = float(s) * 100
@@ -116,7 +125,6 @@ def get_steps(tree):
for i, _ in enumerate(tree.children[:-1]):
tree.children[i] = tostep(tree.children[i])
interpolation_steps.append((tuple(tree.children[:-1]), tree.children[-1]))
res.extend(tree.children[:-1])
def sequence(self, tree):
@@ -134,7 +142,7 @@ def get_steps(tree):
CollectSteps().visit(tree)
return sorted(set(interpolation_steps)), sorted(set(res))
return sorted(set(res))
def at_step(step, filters, tree):
@@ -180,24 +188,6 @@ def at_step(step, filters, tree):
previous_step = s
return ""
def interpolate(self, args):
prompts, starts = args
starts = starts[:-1]
prev_prompt = None
if step < starts[0]:
return prompts[0]
for i, x in enumerate(starts):
prev_prompt = prompts[i]
if x >= step:
break
return prev_prompt
def interp_steps(self, args):
return list(args)
def interp_prompts(self, args):
return ["".join(flatten(a or [])) for a in args]
def alternate(self, args):
step_size = args[-1]
idx = ceil(step / step_size)
@@ -268,22 +258,15 @@ def at_step(step, filters, tree):
class PromptSchedule(object):
def __init__(self, prompt, filters="", start=0.0, end=1.0, defaults=None, masks=None):
def __init__(self, prompt, filters="", start=0.0, end=1.0):
self.filters = filters
self.start = start
self.end = end
self.prompt = prompt.strip()
self.defaults = {}
if defaults:
self.defaults = defaults
self.loaded_loras = {}
self.interpolations = None
self.parsed_prompt = None
self.interpolations, self.parsed_prompt = self._parse()
self.masks = masks
if masks is None:
self.masks = []
self.parsed_prompt = self._parse()
def __iter__(self):
# Filter out zero, it's only useful for interpolation
@@ -293,27 +276,14 @@ class PromptSchedule(object):
filters = [x.strip() for x in self.filters.upper().split(",")]
try:
parsed = []
interpolations = set()
tree = prompt_parser.parse(self.prompt)
interpolation_steps, steps = get_steps(tree)
log.debug("Interpolation steps: %s", interpolation_steps)
steps = get_steps(tree)
def f(x):
return round(x / 100, 2)
for t in steps:
p = at_step(t, filters, tree)
for control_points, step in interpolation_steps:
interp_start = None
interp_end = None
if t == control_points[-1]:
interp_start = max(control_points[0], int(self.start * 100))
interp_end = min(control_points[-1], int(self.end * 100))
control_points = tuple(
sorted(set(f(c) for c in control_points if c >= interp_start or c <= interp_end))
)
if interp_start is not None and interp_end is not None and interp_end > interp_start:
interpolations.add((control_points, f(step)))
parsed.append([f(t), p])
except lark.exceptions.LarkError as e:
@@ -322,14 +292,9 @@ class PromptSchedule(object):
# Tag filtering may return redundant prompts, so filter them out here
res = []
prev_p = None
prev_end = -1
for end_at, p in parsed:
# Preserve prompt if it ends at the start of an interpolation, otherwise bump its end time
if p == prev_p and res[-1][0] not in [x[0][0] for x in interpolations]:
res[-1][0] = end_at
continue
if end_at < self.start:
continue
elif end_at <= self.end:
@@ -338,18 +303,12 @@ class PromptSchedule(object):
elif end_at > self.end and prev_end < self.end:
res.append([end_at, p])
break
prev_p = p
# Always use the last prompt if everything was filtered
if len(res) == 0:
res = [[1.0, parsed[-1][1]]]
return interpolations, res
def add_masks(self, *masks):
for mask in masks:
if mask is not None:
self.masks.append(mask)
return res
def clone(self):
return self.with_filters()
@@ -363,8 +322,6 @@ class PromptSchedule(object):
filters=ifspecified(filters, self.filters),
start=ifspecified(start, self.start),
end=ifspecified(end, self.end),
defaults=ifspecified(defaults, self.defaults),
masks=self.masks[:],
)
return p
@@ -378,23 +335,35 @@ class PromptSchedule(object):
return i, x
return len(self.parsed_prompt) - 1, self.parsed_prompt[-1]
def interpolation_at(self, step, total_steps=1):
i, x = self.at_step_idx(step, total_steps)
for y in self.parsed_prompt[i:]:
step = min(y[0], 1.0)
if x[1]["prompt"] != y[1]["prompt"]:
return step, y
return 1.0, self.parsed_prompt[-1]
def load_loras(self, lora_cache=None):
from .utils import Timer, load_loras_from_schedule
if lora_cache is not None:
self.loaded_loras = lora_cache
with Timer("PromptSchedule.load_loras()"):
self.loaded_loras = load_loras_from_schedule(self.parsed_prompt, self.loaded_loras)
return self.loaded_loras
def replace_defs(text):
text, defs = get_function(text, "DEF", defaults=None)
res = text
prevres = text
replacements = []
for d in defs:
r = d.split("=", 1)
if len(r) != 2 or not r[0].strip():
log.warning("Ignoring invalid DEF(%s)", d)
continue
replacements.append(r[0].strip(), r[1].strip())
iterations = 0
while True:
iterations += 1
if iterations > 10:
log.error("Unable to resolve DEFs, make sure there are no cycles!")
return text
for search, replace in replacements:
res = re.sub(rf"\b{re.escape(search)}\b", replace, res)
if res == prevres:
break
prevres = res
if res != text:
log.info("DEFs expanded to: %s", res)
return res
def parse_prompt_schedules(prompt):
return PromptSchedule(prompt)
@lru_cache
def parse_prompt_schedules(prompt, **kwargs):
prompt = replace_defs(prompt)
return PromptSchedule(prompt, **kwargs)
+29 -3
View File
@@ -4,11 +4,26 @@ import torch
from functools import partial
from comfy_extras.nodes_mask import FeatherMask, MaskComposite
from .utils import safe_float, get_function, parse_floats
from .utils import safe_float, get_function, parse_floats, smarter_split
from .adv_encode import advanced_encode_from_tokens
from .cutoff import process_cuts
from .parser import parse_cuts
try:
from .nodes_attnmask import create_attention_hook
from comfy.hooks import set_hooks_for_conditioning
def set_cond_attnmask(cond, mask):
hook = create_attention_hook(mask)
return set_hooks_for_conditioning(cond, hooks=hook)
except ImportError:
def set_cond_attnmask(cond, mask):
log.info("Attention masking is not available")
return cond
log = logging.getLogger("comfyui-prompt-control")
AVAILABLE_STYLES = ["comfy", "perp", "A1111", "compel", "comfy++", "down_weight"]
@@ -88,8 +103,9 @@ def shuffle_chunk(shuffle, c):
"separator": separator,
}.get(joiner, joiner)
log.info("%s arg=%s sep=%s join=%s", func, shuffle_count, separator, joiner)
separated = c.split(separator)
log.debug("%s arg=%s sep=%s join=%s", func, shuffle_count, separator, joiner)
separated = smarter_split(separator, c)
log.debug("Prompt split into %s", separated)
if func == "SHIFT":
shuffle_count = shuffle_count % len(separated)
permutation = separated[shuffle_count:] + separated[:shuffle_count]
@@ -429,6 +445,11 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
# TODO: is this still needed?
# scale = sum(abs(weight(p)[0]) for p in prompts if not ("AREA(" in p or "MASK(" in p))
for prompt in prompts:
attn = False
if "ATTN()" in prompt:
prompt = prompt.replace("ATTN()", "")
attn = True
log.info("Using attention masking for prompt segment")
prompt, mask, mask_weight = get_mask(prompt, mask_size, masks)
w, opts, prompt = weight(prompt)
text, noise_w, generator = get_noise(text)
@@ -451,6 +472,11 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
settings["start_percent"] = start_pct
settings["end_percent"] = end_pct
x = encode_prompt_segment(clip, prompt, settings, style, normalization)
if attn and mask is not None:
mask = settings.pop("mask")
strength = settings.pop("mask_strength")
x = set_cond_attnmask(x, mask * strength)
conds.extend(x)
return conds
+29 -3
View File
@@ -2,7 +2,14 @@ from pathlib import Path
import re
import logging
import folder_paths
# Allow testing
try:
from folder_paths import get_filename_list
except ImportError:
def get_filename_list(x):
raise NotImplementedError("How did you get here?")
log = logging.getLogger("comfyui-prompt-control")
@@ -35,7 +42,6 @@ def find_nonscheduled_loras(consolidated_schedule):
if not consolidated_schedule:
return {}
last_end, candidate_loras = consolidated_schedule[0]
print(candidate_loras)
to_remove = set()
for candidate, weights in candidate_loras.items():
for end, loras in consolidated_schedule[1:]:
@@ -48,6 +54,26 @@ def find_nonscheduled_loras(consolidated_schedule):
return {k: v for (k, v) in candidate_loras.items() if k not in to_remove}
def smarter_split(separator, string):
"""Does not break () when splitting"""
splits = []
prev = 0
stack = 0
escape = False
for idx, x in enumerate(string):
if x == "(" and not escape:
stack += 1
elif x == ")" and not escape:
stack = max(0, stack - 1)
elif x == separator and stack == 0:
splits.append(string[prev:idx])
prev = idx + 1
escape = x == "\\"
splits.append(string[prev : idx + 1])
return splits
def find_closing_paren(text, start):
stack = 1
for i, char in enumerate(text[start:]):
@@ -119,7 +145,7 @@ def safe_float(f, default):
def lora_name_to_file(name):
filenames = folder_paths.get_filename_list("loras")
filenames = get_filename_list("loras")
# Return exact matches as is
if name in filenames:
return name
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-prompt-control"
description = "Nodes for convenient prompt editing, making many common operations prompt-controllable"
version = "2.0.0-beta.1"
version = "2.0.0-beta.4"
license = { file = "LICENSE" }
# some lark versions older than 1.1.9 apparently have a bug that breaks things, see https://github.com/asagi4/comfyui-prompt-control/issues/35
dependencies = ["lark >= 1.1.9"]