Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8aec0d8f46 | ||
|
|
013586904b | ||
|
|
85ff8466ad | ||
|
|
97d17d2bfd | ||
|
|
c0e671b2de | ||
|
|
7ff55c6717 | ||
|
|
524738f21d | ||
|
|
5c7d507e91 |
@@ -1 +1,2 @@
|
||||
__pycache__
|
||||
.pyre
|
||||
|
||||
@@ -197,6 +197,7 @@ class AdvancedEncoder:
|
||||
def add_encoder(cls, name, fn):
|
||||
cls.STYLES[name] = fn
|
||||
|
||||
@classmethod
|
||||
def add_normalization_op(cls, name, fn):
|
||||
cls.NORMALIZATION_OPS[name] = fn
|
||||
|
||||
@@ -296,7 +297,7 @@ class AdvancedEncoder:
|
||||
w_mix = np.diff([0] + w.tolist())
|
||||
w_mix = torch.tensor(w_mix, dtype=embs.dtype, device=embs.device).reshape((-1, 1, 1))
|
||||
|
||||
weighted_emb = (w_mix * embs).sum(axis=0, keepdim=True)
|
||||
weighted_emb = (w_mix * embs).sum(dim=0, keepdim=True)
|
||||
pooled = pooled_base
|
||||
if pooled is not None and self.max_length:
|
||||
pooled = weighted_emb[0, self.max_length - 1 : self.max_length, :]
|
||||
@@ -328,12 +329,13 @@ class AdvancedEncoder:
|
||||
masks = torch.cat(masks)
|
||||
|
||||
embs = base_emb.expand(embs.shape) - embs
|
||||
pooled = None
|
||||
if pooled_base is not None and self.max_length:
|
||||
pooled = embs[0, self.max_length - 1 : self.max_length, :]
|
||||
pooled_start = pooled_base.expand(len(ws), -1)
|
||||
ws = torch.tensor(ws).reshape(-1, 1).expand(pooled_start.shape)
|
||||
pooled = (pooled - pooled_start) * (ws - 1)
|
||||
pooled = pooled.mean(axis=0, keepdim=True)
|
||||
pooled = pooled.mean(dim=0, keepdim=True)
|
||||
pooled = pooled_base + pooled
|
||||
|
||||
if embs.shape[0] != masks.shape[0]:
|
||||
|
||||
@@ -71,14 +71,13 @@ class AttentionCoupleHook(TransformerOptionsHook):
|
||||
}
|
||||
}
|
||||
self.has_negpip = False
|
||||
|
||||
# calculate later. All clones must refer to the same kv dict
|
||||
self.kv = {"k": None, "v": None}
|
||||
self.kv = {}
|
||||
|
||||
def initialize_regions(self, base_cond, conds, fill):
|
||||
self.num_conds = len(conds) + 1
|
||||
self.base_strength = base_cond[1].get("strength", 1.0)
|
||||
self.strengths = [cond[1].get("strength", 1.0) for cond in conds]
|
||||
self.strengths: list[float] = [cond[1].get("strength", 1.0) for cond in conds]
|
||||
self.conds: list[torch.Tensor] = [base_cond[0]] + [cond[0] for cond in conds]
|
||||
base_mask = base_cond[1].get("mask", None)
|
||||
masks = [cond[1].get("mask") * cond[1].get("mask_strength") for cond in conds]
|
||||
@@ -215,6 +214,7 @@ class AttentionCoupleHook(TransformerOptionsHook):
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
assert self.num_conds is not None, "this is a bug"
|
||||
cond_or_uncond_couple.extend(itertools.repeat(self.COND, self.num_conds))
|
||||
|
||||
q = torch.cat(qs, dim=0)
|
||||
|
||||
@@ -1,47 +0,0 @@
|
||||
import comfy_execution.caching
|
||||
from comfy_execution.graph_utils import is_link
|
||||
import nodes
|
||||
from os import environ
|
||||
|
||||
import logging
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
include_unique_id_in_input = comfy_execution.caching.include_unique_id_in_input
|
||||
|
||||
|
||||
def promptcontrol_get_immediate_node_signature(self, dynprompt, node_id, ancestor_order_mapping):
|
||||
if not dynprompt.has_node(node_id):
|
||||
# This node doesn't exist -- we can't cache it.
|
||||
return [float("NaN")]
|
||||
node = dynprompt.get_node(node_id)
|
||||
class_type = node["class_type"]
|
||||
class_def = nodes.NODE_CLASS_MAPPINGS[class_type]
|
||||
inputs = node["inputs"]
|
||||
if hasattr(class_def, "CACHE_KEY"):
|
||||
inputs = getattr(class_def, "CACHE_KEY")(inputs)
|
||||
signature = [class_type, self.is_changed_cache.get(node_id)]
|
||||
if (
|
||||
self.include_node_id_in_input()
|
||||
or (hasattr(class_def, "NOT_IDEMPOTENT") and class_def.NOT_IDEMPOTENT)
|
||||
or include_unique_id_in_input(class_type)
|
||||
):
|
||||
signature.append(node_id)
|
||||
for key in sorted(inputs.keys()):
|
||||
if is_link(inputs[key]):
|
||||
(ancestor_id, ancestor_socket) = inputs[key]
|
||||
ancestor_index = ancestor_order_mapping[ancestor_id]
|
||||
signature.append((key, ("ANCESTOR", ancestor_index, ancestor_socket)))
|
||||
else:
|
||||
signature.append((key, inputs[key]))
|
||||
return signature
|
||||
|
||||
|
||||
def init():
|
||||
if environ.get("PROMPTCONTROL_ENABLE_CACHE_HACK") != "1":
|
||||
return
|
||||
log.warning("Enabling Prompt Control cache hack")
|
||||
comfy_execution.caching.CacheKeySetInputSignature.get_immediate_node_signature = (
|
||||
promptcontrol_get_immediate_node_signature
|
||||
)
|
||||
@@ -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(axis=0)
|
||||
region_embeddings = torch.stack(region_embeddings).sum(dim=0)
|
||||
|
||||
embeddings_final_mask = torch.tensor(
|
||||
global_region_mask, dtype=base_embedding_full.dtype, device=base_embedding_full.device
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
import logging
|
||||
from .prompts import encode_prompt
|
||||
|
||||
@@ -40,7 +41,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(self, clip, text, 0.0, 1.0)
|
||||
return PCTextEncodeWithRange().apply(clip, text, 0.0, 1.0)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"PCTextEncode": PCTextEncode, "PCTextEncodeWithRange": PCTextEncodeWithRange}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
import logging
|
||||
|
||||
import comfy.hooks
|
||||
@@ -38,7 +39,6 @@ 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,7 +53,8 @@ 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
|
||||
new_hook.hooks[0].hook_ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
|
||||
ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
|
||||
new_hook.hooks[0].hook_ref = ref # pyright: ignore[reportAttributeAccessIssue]
|
||||
hooks.append(new_hook)
|
||||
if start_pct > 0.0:
|
||||
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=0.0)
|
||||
@@ -74,8 +75,6 @@ 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:
|
||||
@@ -85,7 +84,7 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
|
||||
class PCAttentionCoupleBatchNegative(ComfyNodeABC):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> InputTypeDict:
|
||||
def INPUT_TYPES(s) -> InputTypeDict:
|
||||
return {
|
||||
"required": {
|
||||
"positive": (IO.CONDITIONING, {}),
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
from __future__ import annotations
|
||||
import logging
|
||||
from .parser import parse_prompt_schedules
|
||||
from comfy_execution.graph_utils import GraphBuilder, is_link
|
||||
@@ -167,7 +169,7 @@ class PCLazyLoraLoaderAdvanced:
|
||||
"hidden": {"unique_id": "UNIQUE_ID"},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CLIP", "HOOKS")
|
||||
RETURN_TYPES: tuple[str, ...] = ("MODEL", "CLIP", "HOOKS")
|
||||
OUTPUT_TOOLTIPS = ("Returns a model and clip with LoRAs scheduled",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
import logging
|
||||
from .parser import parse_prompt_schedules, expand_macros
|
||||
from .nodes_lazy import NODE_CLASS_MAPPINGS as LAZY_NODES
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# vim: sw=4 ts=4
|
||||
from __future__ import annotations
|
||||
import lark
|
||||
import logging
|
||||
from math import ceil
|
||||
@@ -100,7 +101,15 @@ class CutTransform(lark.Transformer):
|
||||
def cut(self, args):
|
||||
prompt, cutout, weight, strict_mask, start_from_masked, mask_token = args
|
||||
|
||||
return ("".join(flatten(prompt)), "".join(flatten(cutout)), weight, strict_mask, start_from_masked, mask_token)
|
||||
# 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 = []
|
||||
@@ -301,8 +310,7 @@ def at_step(step, filters, tree):
|
||||
return name, params, lbw
|
||||
|
||||
def __default__(self, data, children, meta):
|
||||
for child in children:
|
||||
yield child
|
||||
return children
|
||||
|
||||
return AtStep().transform(tree)
|
||||
|
||||
|
||||
@@ -598,10 +598,9 @@ 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:
|
||||
base_prompt, attn_couple_prompts = split_by_function(prompt, "COUPLE", defaults=None, require_args=False)
|
||||
|
||||
@@ -20,17 +20,9 @@ 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
|
||||
@@ -129,17 +121,6 @@ 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()
|
||||
@@ -178,16 +159,6 @@ 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()
|
||||
|
||||
@@ -155,15 +155,13 @@ 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 - skipped, ph))
|
||||
skipped += end - start
|
||||
instances.append(FunctionSpec(funcname, args, start, ph))
|
||||
chunks.append(text[current:start] + (ph or ""))
|
||||
current = end
|
||||
count += 1
|
||||
|
||||
+5
-1
@@ -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.3"
|
||||
version = "2.1.1"
|
||||
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,3 +13,7 @@ Repository = "https://github.com/asagi4/comfyui-prompt-control"
|
||||
PublisherId = "asagi4"
|
||||
DisplayName = "ComfyUI Prompt Control"
|
||||
Icon = ""
|
||||
|
||||
[tool.pyright]
|
||||
extraPaths = ["../../"]
|
||||
exclude = ["prompt_control/*test*"]
|
||||
|
||||
Reference in New Issue
Block a user