Compare commits

..
7 Commits
Author SHA1 Message Date
asagi4 68766215f2 v2.1.3 2026-02-04 16:32:26 +02:00
asagi4 287f554a68 Test COUPLE mask shortcut 2026-01-16 18:20:46 +02:00
asagi4 71465f914c Properly parse COUPLE(), see #134 2026-01-16 18:05:30 +02:00
asagi4 d7be7bc29e v2.1.2 2026-01-13 21:50:09 +02:00
asagi4 329d4cf95f Add a test for #133 2026-01-13 21:48:34 +02:00
asagi4 c52ace71aa #133 properly set function locations in get_function 2026-01-13 20:45:10 +02:00
asagi4 4806cf5959 refactor get_function to make it more consistent 2026-01-13 20:45:10 +02:00
14 changed files with 99 additions and 38 deletions
-1
View File
@@ -1,2 +1 @@
__pycache__
.pyre
+2 -4
View File
@@ -197,7 +197,6 @@ class AdvancedEncoder:
def add_encoder(cls, name, fn):
cls.STYLES[name] = fn
@classmethod
def add_normalization_op(cls, name, fn):
cls.NORMALIZATION_OPS[name] = fn
@@ -297,7 +296,7 @@ class AdvancedEncoder:
w_mix = np.diff([0] + w.tolist())
w_mix = torch.tensor(w_mix, dtype=embs.dtype, device=embs.device).reshape((-1, 1, 1))
weighted_emb = (w_mix * embs).sum(dim=0, keepdim=True)
weighted_emb = (w_mix * embs).sum(axis=0, keepdim=True)
pooled = pooled_base
if pooled is not None and self.max_length:
pooled = weighted_emb[0, self.max_length - 1 : self.max_length, :]
@@ -329,13 +328,12 @@ class AdvancedEncoder:
masks = torch.cat(masks)
embs = base_emb.expand(embs.shape) - embs
pooled = None
if pooled_base is not None and self.max_length:
pooled = embs[0, self.max_length - 1 : self.max_length, :]
pooled_start = pooled_base.expand(len(ws), -1)
ws = torch.tensor(ws).reshape(-1, 1).expand(pooled_start.shape)
pooled = (pooled - pooled_start) * (ws - 1)
pooled = pooled.mean(dim=0, keepdim=True)
pooled = pooled.mean(axis=0, keepdim=True)
pooled = pooled_base + pooled
if embs.shape[0] != masks.shape[0]:
+3 -3
View File
@@ -71,13 +71,14 @@ class AttentionCoupleHook(TransformerOptionsHook):
}
}
self.has_negpip = False
# calculate later. All clones must refer to the same kv dict
self.kv = {}
self.kv = {"k": None, "v": None}
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.strengths = [cond[1].get("strength", 1.0) for cond in conds]
self.conds: list[torch.Tensor] = [base_cond[0]] + [cond[0] for cond in conds]
base_mask = base_cond[1].get("mask", None)
masks = [cond[1].get("mask") * cond[1].get("mask_strength") for cond in conds]
@@ -214,7 +215,6 @@ class AttentionCoupleHook(TransformerOptionsHook):
dim=0,
)
)
assert self.num_conds is not None, "this is a bug"
cond_or_uncond_couple.extend(itertools.repeat(self.COND, self.num_conds))
q = torch.cat(qs, dim=0)
+47
View File
@@ -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
)
+1 -1
View File
@@ -215,7 +215,7 @@ def encode_regions(clip_regions, encode, tokenizer):
region_emb *= region_masking
region_embeddings.append(region_emb)
region_embeddings = torch.stack(region_embeddings).sum(dim=0)
region_embeddings = torch.stack(region_embeddings).sum(axis=0)
embeddings_final_mask = torch.tensor(
global_region_mask, dtype=base_embedding_full.dtype, device=base_embedding_full.device
+1 -2
View File
@@ -1,4 +1,3 @@
# pyright: reportSelfClsParameterName=false
import logging
from .prompts import encode_prompt
@@ -41,7 +40,7 @@ class PCTextEncode:
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)
return PCTextEncodeWithRange.apply(self, clip, text, 0.0, 1.0)
NODE_CLASS_MAPPINGS = {"PCTextEncode": PCTextEncode, "PCTextEncodeWithRange": PCTextEncodeWithRange}
+5 -4
View File
@@ -1,4 +1,3 @@
# pyright: reportSelfClsParameterName=false
import logging
import comfy.hooks
@@ -39,6 +38,7 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
all_hooks = []
def create_hook(loraspec, start_pct, end_pct, non_scheduled):
nonlocal lora_cache
hooks = []
hook_kf = comfy.hooks.HookKeyframeGroup()
for path, info in loras.items():
@@ -53,8 +53,7 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
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 # pyright: ignore[reportAttributeAccessIssue]
new_hook.hooks[0].hook_ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
hooks.append(new_hook)
if start_pct > 0.0:
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=0.0)
@@ -75,6 +74,8 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
all_hooks.append(hook)
start_pct = end_pct
del lora_cache
all_hooks = [x for x in all_hooks if x]
if all_hooks:
@@ -84,7 +85,7 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
class PCAttentionCoupleBatchNegative(ComfyNodeABC):
@classmethod
def INPUT_TYPES(s) -> InputTypeDict:
def INPUT_TYPES(cls) -> InputTypeDict:
return {
"required": {
"positive": (IO.CONDITIONING, {}),
+1 -3
View File
@@ -1,5 +1,3 @@
# pyright: reportSelfClsParameterName=false
from __future__ import annotations
import logging
from .parser import parse_prompt_schedules
from comfy_execution.graph_utils import GraphBuilder, is_link
@@ -169,7 +167,7 @@ class PCLazyLoraLoaderAdvanced:
"hidden": {"unique_id": "UNIQUE_ID"},
}
RETURN_TYPES: tuple[str, ...] = ("MODEL", "CLIP", "HOOKS")
RETURN_TYPES = ("MODEL", "CLIP", "HOOKS")
OUTPUT_TOOLTIPS = ("Returns a model and clip with LoRAs scheduled",)
CATEGORY = "promptcontrol"
FUNCTION = "apply"
-1
View File
@@ -1,4 +1,3 @@
# pyright: reportSelfClsParameterName=false
import logging
from .parser import parse_prompt_schedules, expand_macros
from .nodes_lazy import NODE_CLASS_MAPPINGS as LAZY_NODES
+3 -11
View File
@@ -1,5 +1,4 @@
# vim: sw=4 ts=4
from __future__ import annotations
import lark
import logging
from math import ceil
@@ -101,15 +100,7 @@ class CutTransform(lark.Transformer):
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
return ("".join(flatten(prompt)), "".join(flatten(cutout)), weight, strict_mask, start_from_masked, mask_token)
def start(self, args):
prompt = []
@@ -310,7 +301,8 @@ def at_step(step, filters, tree):
return name, params, lbw
def __default__(self, data, children, meta):
return children
for child in children:
yield child
return AtStep().transform(tree)
+3 -2
View File
@@ -598,9 +598,10 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
return c
def couple_mask(args):
if args is None:
assert len(args) <= 1, "Argument parsing failure. This is a bug in Prompt Control"
if not args:
return ""
return f"MASK({args})"
return f"MASK({args[0]})"
for prompt in prompts:
base_prompt, attn_couple_prompts = split_by_function(prompt, "COUPLE", defaults=None, require_args=False)
+29
View File
@@ -20,9 +20,17 @@ def run(f, *args):
return getattr(f, f.FUNCTION)(*args)
def compare_hookgroup_mask(h1, h2):
assert len(h1.hooks) == len(h2.hooks)
for a, b in zip(h1.hooks, h2.hooks):
assert (a.mask == b.mask).all()
@mock.patch("torch.cuda.current_device", lambda: "cpu")
class TestEncode(unittest.TestCase):
@classmethod
def setUpClass(cls):
global clips
print("Loading ComfyUI")
from comfy.sd import load_clip
from pathlib import Path
@@ -121,6 +129,17 @@ class TestEncode(unittest.TestCase):
self.condEqual(avg, c3)
self.condEqual(avg, c4)
with self.subTest("Average multi"):
(c1,) = run(comfy, clip, "test1")
(c2,) = run(comfy, clip, "test2")
(c3,) = run(comfy, clip, "test3")
(c4,) = run(pc, clip, "test1 AVG() test2 AVG() test3")
(c5,) = run(pc, clip, "test1 AVG test2 AVG test3")
(avg1,) = run(average, c1, c2, 0.5)
(avg,) = run(average, avg1, c3, 0.5)
self.condEqual(avg, c4)
self.condEqual(avg, c5)
@unittest.expectedFailure
def test_failure(self):
pc = PCTextEncode()
@@ -159,6 +178,16 @@ class TestEncode(unittest.TestCase):
(c2,) = run(pc, clip, "test COUPLE prompt1 COUPLE test2 COUPLE prompt2")
self.assertTrue(len(c) == 2)
self.assertTrue(len(c2) == 1)
with self.subTest(f"Testing {k} mask shortcut"):
(c,) = run(pc, clip, "test COUPLE() prompt1")
(c2,) = run(pc, clip, "test COUPLE MASK() prompt1")
self.condEqual(c, c2)
self.condEqual(c, c2, "hooks", compare_hookgroup_mask)
with self.subTest(f"Testing {k} mask shortcut 2"):
(c,) = run(pc, clip, "test COUPLE(0 0.2, 0.5) prompt1")
(c2,) = run(pc, clip, "test COUPLE MASK(0 0.2, 0.5) prompt1")
self.condEqual(c, c2)
self.condEqual(c, c2, "hooks", compare_hookgroup_mask)
def test_styles(self):
pc = PCTextEncode()
+3 -1
View File
@@ -155,13 +155,15 @@ def get_function(
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
if placeholder:
ph = f"\0{placeholder}{count}\0"
instances.append(FunctionSpec(funcname, args, start, ph))
instances.append(FunctionSpec(funcname, args, start - skipped, ph))
skipped += end - start
chunks.append(text[current:start] + (ph or ""))
current = end
count += 1
+1 -5
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-prompt-control"
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"
version = "2.1.3"
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"]
@@ -13,7 +13,3 @@ Repository = "https://github.com/asagi4/comfyui-prompt-control"
PublisherId = "asagi4"
DisplayName = "ComfyUI Prompt Control"
Icon = ""
[tool.pyright]
extraPaths = ["../../"]
exclude = ["prompt_control/*test*"]