Compare commits

..
Author SHA1 Message Date
asagi4 91ba4c881f v2.0.0-beta.6 2025-02-27 20:03:53 +02:00
asagi4 106ebe49aa Add some tests to ensure that prompts don't break 2025-02-27 20:03:04 +02:00
asagi4 549b4347fd Make [:xyz:N] work
Fixes #91
2025-02-27 20:00:56 +02:00
asagi4 4cbce5df06 Parse [SEQ:a:N] properly
See #93
2025-02-27 10:09:40 +02:00
asagi4 21208bd733 v2.0.0-beta.5 2025-02-23 19:45:46 +02:00
asagi4 a4065415e7 Add PCExtractScheduledPrompt, fixes #90 2025-02-23 19:44:56 +02:00
asagi4 0d546f1a08 Fix DEF 2025-02-23 19:13:08 +02:00
asagi4 390e1ec6b8 Merge pull request #88 from asagi4/attn_mask
Experimental attention masking
2025-01-12 15:58:42 +02:00
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
15 changed files with 465 additions and 46 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.
+3
View File
@@ -5,4 +5,7 @@ check:
format:
find . -name "*.py" | xargs black -l 120
test:
python -m prompt_control.test
.PHONY: check format all
+2
View File
@@ -62,6 +62,8 @@ 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.
+11 -4
View File
@@ -30,14 +30,21 @@ 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.
+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)",
}
+27 -19
View File
@@ -7,15 +7,25 @@ from .prompts import get_function
log = logging.getLogger("comfyui-prompt-control")
from .utils import consolidate_schedule, find_nonscheduled_loras
import json
def cache_key_hack(inputs):
def _cache_key(cachekey, inputs):
out = inputs.copy()
if not is_link(inputs["text"]):
out["text"] = cache_key_from_inputs(**inputs)
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():
log.info("Creating LoraLoader for %s", path)
@@ -34,7 +44,7 @@ def create_hook_nodes_for_lora(graph, path, info, existing_node, start_pct, end_
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"])
@@ -67,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)
@@ -108,7 +118,7 @@ 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))
@@ -123,6 +133,7 @@ def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True, return_h
clip = n.out(0)
r = graph.finalize()
log.debug("LazyLoraLoader built graph: %s", json.dumps(r))
if return_hooks:
ret = (model, clip, res)
@@ -133,7 +144,7 @@ def build_lora_schedule(graph, schedule, model, clip, apply_hooks=True, return_h
class PCLazyLoraLoaderAdvanced:
CACHE_KEY = cache_key_hack
CACHE_KEY = cache_key_lora
@classmethod
def INPUT_TYPES(s):
@@ -164,7 +175,7 @@ class PCLazyLoraLoaderAdvanced:
class PCLazyLoraLoader:
CACHE_KEY = cache_key_hack
CACHE_KEY = cache_key_lora
@classmethod
def INPUT_TYPES(s):
@@ -194,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"])
@@ -203,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)
@@ -223,17 +230,18 @@ def build_scheduled_prompts(graph, schedules, clip):
node = combiner
g = graph.finalize()
log.debug("Built graph: %s", json.dumps(g))
return {"result": (node.out(0),), "expand": g}
def cache_key_from_inputs(text, tags="", start=0.0, end=1.0, **kwargs):
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["prompt"]) for pct, s in schedules]
return [(pct, s[cachekey]) for pct, s in schedules]
class PCLazyTextEncode:
CACHE_KEY = cache_key_hack
CACHE_KEY = cache_key_prompt
@classmethod
def INPUT_TYPES(s):
@@ -253,7 +261,7 @@ class PCLazyTextEncode:
class PCLazyTextEncodeAdvanced:
CACHE_KEY = cache_key_hack
CACHE_KEY = cache_key_prompt
@classmethod
def INPUT_TYPES(s):
+61 -4
View File
@@ -1,8 +1,35 @@
import logging
from .parser import parse_prompt_schedules
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,8 +41,9 @@ class PCAddMaskToCLIP:
}
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/v2"
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)
@@ -35,8 +63,9 @@ class PCAddMaskToCLIPMany:
}
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/v2"
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Multi-input version of PCAddMaskToCLIP, for convenience"
def apply(self, clip, mask1=None, mask2=None, mask3=None, mask4=None):
clip = clip.clone()
@@ -65,8 +94,9 @@ class PCSetPCTextEncodeSettings:
}
RETURN_TYPES = ("CLIP",)
CATEGORY = "promptcontrol/v2"
CATEGORY = "promptcontrol/tools"
FUNCTION = "apply"
DESCRIPTION = "Configures default values for PCTextEncode"
def apply(
self,
@@ -97,14 +127,41 @@ class PCSetPCTextEncodeSettings:
return (clip,)
class PCExtractScheduledPrompt:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"multiline": True}),
"at": ("FLOAT", {"min": 0.0, "max": 1.0, "default": 1.0, "step": 0.01}),
},
"optional": {"tags": ("STRING", {"default": ""})},
}
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, total_steps=1)
prompt_text = entry.get("prompt", "")
return (prompt_text,)
NODE_CLASS_MAPPINGS = {
"PCSetPCTextEncodeSettings": PCSetPCTextEncodeSettings,
"PCAddMaskToCLIP": PCAddMaskToCLIP,
"PCAddMaskToCLIPMany": PCAddMaskToCLIPMany,
"PCSetLogLevel": PCSetLogLevel,
"PCExtractScheduledPrompt": PCExtractScheduledPrompt,
}
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)",
"PCExtractScheduledPrompt": "PC: Extract Scheduled Prompt",
}
+43 -4
View File
@@ -4,11 +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)
@@ -20,9 +30,9 @@ prompt: (emphasized | embedding | scheduled | alternate | sequence | loraspec |
!emphasized: "(" prompt? ")"
| "(" prompt ":" prompt ")"
| "[" prompt "]"
scheduled: "[" [prompt ":"] [prompt] ":" _WS? NUMBER ["," NUMBER] "]"
| "[" [prompt ":"] [prompt] ":" _WS? TAG "]"
sequence: "[SEQ" ":" [prompt] ":" NUMBER (":" [prompt] ":" NUMBER)+ "]"
scheduled: "[" [[prompt] ":"] [prompt] ":" _WS? NUMBER ["," NUMBER] "]"
| "[" [[prompt] ":"] [prompt] ":" _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
@@ -38,6 +48,7 @@ TAG: /[A-Z_]+/
lexer="dynamic",
)
cut_parser = lark.Lark(
r"""
!start: (prompt | /[][:()]/+)*
@@ -325,6 +336,34 @@ class PromptSchedule(object):
return len(self.parsed_prompt) - 1, self.parsed_prompt[-1]
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
@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
+136
View File
@@ -0,0 +1,136 @@
import unittest
from .parser import parse_prompt_schedules as parse
class TestParser(unittest.TestCase):
def test_no_scheduling(self):
p = parse("This is a (basic:0.6) (prompt) with [no scheduling] features")
expected = [1.0, {"prompt": "This is a (basic:0.6) (prompt) with [no scheduling] features", "loras": {}}]
self.assertEqual(p.at_step(0), expected)
self.assertEqual(p.at_step(0.5), expected)
self.assertEqual(p.at_step(1), expected)
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]"
)
expected = [0.5, {"prompt": "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features", "loras": {}}]
expected2 = [
0.8,
{"prompt": "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) features", "loras": {}},
]
expected3 = [1.0, {"prompt": "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) ", "loras": {}}]
self.assertEqual(p.at_step(0), expected)
self.assertEqual(p.at_step(0.5), expected)
self.assertEqual(p.at_step(0.7), expected2)
self.assertEqual(p.at_step(1), expected3)
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 = [
1.0,
{
"prompt": "This is a (lora:0.6) (prompt) with [no scheduling] features ",
"loras": {"foo": {"weight": 0.5, "weight_clip": 0.5}, "bar": {"weight": 0.5, "weight_clip": 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>"
)
expected = [
0.3,
{
"prompt": "This is a (lora:0.6) (prompt) with [scheduling] features ",
"loras": {"foo": {"weight": 0.5, "weight_clip": 0.5}, "bar": {"weight": 0.5, "weight_clip": 1.0}},
},
]
expected2 = [
1.0,
{
"prompt": "This is a (lora:0.6) (prompt) with [scheduling] features ",
"loras": {"bar": {"weight": 1.0, "weight_clip": 1.2}},
},
]
self.assertEqual(p.at_step(0.1), expected)
self.assertEqual(p.at_step(1), expected2)
def test_seq(self):
p = parse("This is a sequence of [SEQ:a:0.2::0.5:c:0.8]")
p2 = parse("This is a sequence of [[a:[c:0.5]:0.2]::0.8]")
prompts = {
0.2: "This is a sequence of a",
0.5: "This is a sequence of ",
0.8: "This is a sequence of c",
1.0: "This is a sequence of ",
}
self.assertEqual(p.parsed_prompt, p2.parsed_prompt)
for k in prompts:
self.assertEqual(p.at_step(k), [k, {"prompt": prompts[k], "loras": {}}])
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_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.assertEqual(
p.at_step(0.6), [0.7, {"prompt": "This prompt is ", "loras": {"cool": {"weight": 1.0, "weight_clip": 1.0}}}]
)
self.assertEqual(
p.at_step(0.7), [0.7, {"prompt": "This prompt is ", "loras": {"cool": {"weight": 1.0, "weight_clip": 1.0}}}]
)
p2 = p.with_filters(filters="hr, xyz")
# TODO: for some reason, this does not deduplicate
# self.assertEqual(p2.at_step(0), p2.at_step(1))
self.assertEqual(p2.at_step(0)[1]["prompt"], p2.at_step(1)[1]["prompt"])
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 in prompts:
self.assertEqual(p.at_step(k), [prompts[k][0], {"prompt": prompts[k][1], "loras": {}}])
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("[:[<lora:test:1>:c:0.5]:0.3]")
self.assertEqual(p.at_step(0), [0.3, {"prompt": "", "loras": {}}])
self.assertEqual(p.at_step(0.4), [0.5, {"prompt": "", "loras": {"test": {"weight": 1.0, "weight_clip": 1.0}}}])
self.assertEqual(p.at_step(1.0), [1.0, {"prompt": "c", "loras": {}}])
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 in prompts:
self.assertEqual(p.at_step(k), [prompts[k][0], {"prompt": prompts[k][1], "loras": {}}])
if __name__ == "__main__":
unittest.main()
+29 -2
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")
@@ -47,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:]):
@@ -118,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.3"
version = "2.0.0-beta.6"
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"]