Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
34056cac19 | ||
|
|
8ae436abf1 | ||
|
|
c2ce2ce023 | ||
|
|
7a9e69ec31 | ||
|
|
aaff8dc7da |
@@ -5,6 +5,7 @@
|
||||
import itertools
|
||||
import logging
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -122,10 +123,9 @@ class AttentionCoupleHook(TransformerOptionsHook):
|
||||
|
||||
self.mask = mask / mask.sum(dim=0, keepdim=True)
|
||||
|
||||
def on_apply_hooks(self, model: ModelPatcher, transformer_options: dict[str]):
|
||||
def on_apply_hooks(self, model: ModelPatcher, transformer_options: dict[str, Any]):
|
||||
if self.conds_k is None:
|
||||
attn_patches = model.model_options["transformer_options"].get("patches", {}).get("attn2_patch", [])
|
||||
self.has_negpip = any("negpip_attn" in i.__name__ for i in attn_patches)
|
||||
self.has_negpip = model.model_options.get("ppm_negpip", False)
|
||||
log.debug("AttentionCouple has_negpip=%s", self.has_negpip)
|
||||
|
||||
# Skip the base cond here, which is always first
|
||||
|
||||
@@ -294,7 +294,8 @@ def apply_weights(output, te_name, spec):
|
||||
pooled_w = 1.0
|
||||
log.info("Weighting %s output by %s, pooled by %s", te_name, w, pooled_w)
|
||||
out = out * w
|
||||
pooled = pooled * pooled_w
|
||||
if pooled is not None:
|
||||
pooled = pooled * pooled_w
|
||||
|
||||
return out, pooled
|
||||
else:
|
||||
@@ -322,9 +323,10 @@ def hook_te(clip, te_names, style, normalization, extra):
|
||||
return clip
|
||||
newclip = clip.clone()
|
||||
for te_name in te_names:
|
||||
if hasattr(clip.tokenizer, "clip_" + te_name):
|
||||
tokenizer = getattr(clip.tokenizer, f"clip_{te_name}", getattr(clip.tokenizer, te_name, None))
|
||||
if tokenizer:
|
||||
x = extra.copy()
|
||||
x["tokenizer"] = getattr(clip.tokenizer, "clip_" + te_name)
|
||||
x["tokenizer"] = tokenizer
|
||||
if not hasattr(clip.patcher.model, te_name):
|
||||
te_name = "clip_" + te_name
|
||||
if not hasattr(clip.patcher.model, te_name):
|
||||
@@ -333,12 +335,7 @@ def hook_te(clip, te_names, style, normalization, extra):
|
||||
|
||||
log.debug("Hooked into te=%s with style=%s, normalization=%s", te_name, style, normalization)
|
||||
encode = clip.patcher.get_model_object(f"{te_name}.encode_token_weights")
|
||||
# A better way to do this would be nice. negpip uses a partial function
|
||||
if "negpip" in getattr(getattr(encode, "func", None), "__name__", "no_func"):
|
||||
if "negpip" in make_patch.__name__:
|
||||
log.info("Detected active NegPiP monkeypatch, disabling native support")
|
||||
else:
|
||||
x["has_negpip"] = True
|
||||
x["has_negpip"] = clip.patcher.model_options.get("ppm_negpip", False)
|
||||
newclip.patcher.add_object_patch(
|
||||
f"{te_name}.encode_token_weights",
|
||||
make_patch(
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import unittest
|
||||
import numpy.testing as npt
|
||||
|
||||
clip_l = None
|
||||
dual = None
|
||||
@@ -9,24 +10,67 @@ def run(f, *args):
|
||||
|
||||
|
||||
class TestEncode(unittest.TestCase):
|
||||
def condEqual(self, c1, c2):
|
||||
def tensorsEqual(self, t1, t2):
|
||||
npt.assert_equal(t1.detach().numpy(), t2.detach().numpy())
|
||||
|
||||
def condEqual(self, c1, c2, key=None, key_assert=None):
|
||||
self.assertEqual(len(c1), len(c2))
|
||||
for i in range(len(c1)):
|
||||
self.assertTrue((c1[i][0] == c2[i][0]).all())
|
||||
a, b = c1[i], c2[i]
|
||||
if key:
|
||||
(key_assert or self.assertEqual)(a[1][key], b[1][key])
|
||||
else:
|
||||
self.tensorsEqual(a[0], b[0])
|
||||
|
||||
def test_basic_encode(self):
|
||||
pc = PCTextEncode()
|
||||
comfy = nodes.CLIPTextEncode()
|
||||
combine = nodes.ConditioningCombine()
|
||||
concat = nodes.ConditioningConcat()
|
||||
zeroout = nodes.ConditioningZeroOut()
|
||||
for k, clip in [("l", clip_l), ("dual", dual)]:
|
||||
with self.subTest(k):
|
||||
(c1,) = run(pc, clip, "test")
|
||||
(c2,) = run(comfy, clip, "test")
|
||||
self.condEqual(c1, c2)
|
||||
with self.subTest("No exceptions"):
|
||||
run(
|
||||
pc,
|
||||
clip,
|
||||
"test AND test (test:1.2) BREAK test AND TE_WEIGHT(all=0) SDXL() AND AREA(,,) test CAT test",
|
||||
)
|
||||
with self.subTest("Basic"):
|
||||
(c1,) = run(pc, clip, "test")
|
||||
(c2,) = run(comfy, clip, "test")
|
||||
c = c2 # Used in later tests
|
||||
self.condEqual(c1, c2)
|
||||
|
||||
(c3,) = run(pc, clip, "test CAT test")
|
||||
(c4,) = run(concat, c2, c2)
|
||||
self.condEqual(c3, c4)
|
||||
(c1,) = run(pc, clip, "(test:1.2)")
|
||||
(c2,) = run(comfy, clip, "(test:1.2)")
|
||||
|
||||
with self.subTest("Concat"):
|
||||
(c1,) = run(pc, clip, "test CAT test")
|
||||
(c2,) = run(concat, c, c)
|
||||
self.condEqual(c1, c2)
|
||||
|
||||
with self.subTest("Combine"):
|
||||
(c1,) = run(pc, clip, "test AND test")
|
||||
(c2,) = run(combine, c, c)
|
||||
self.condEqual(c1, c2)
|
||||
|
||||
with self.subTest("Zero out"):
|
||||
(c1,) = run(pc, clip, "test TE_WEIGHT(all=0)")
|
||||
(c2,) = run(zeroout, c)
|
||||
self.condEqual(c1, c2)
|
||||
|
||||
def test_masks(self):
|
||||
pc = PCTextEncode()
|
||||
comfy = nodes.CLIPTextEncode()
|
||||
solidmask = comfy_extras.nodes_mask.SolidMask()
|
||||
setMask = nodes.ConditioningSetMask()
|
||||
for k, clip in [("l", clip_l), ("dual", dual)]:
|
||||
(c1,) = run(pc, clip, "test MASK()")
|
||||
(c2,) = run(comfy, clip, "test")
|
||||
(c2,) = run(setMask, c2, run(solidmask, 1.0, 512, 512)[0], "default", 1.0)
|
||||
self.condEqual(c1, c2)
|
||||
self.condEqual(c1, c2, "mask", self.tensorsEqual)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
@@ -35,6 +79,7 @@ if __name__ == "__main__":
|
||||
|
||||
id(main) # get rid of flake warning
|
||||
import nodes
|
||||
import comfy_extras.nodes_mask
|
||||
from .nodes_base import PCTextEncode
|
||||
|
||||
(clip_l,) = nodes.CLIPLoader().load_clip("clip_l.safetensors")
|
||||
|
||||
+1
-1
@@ -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-rc.4"
|
||||
version = "2.0.0-rc.5"
|
||||
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"]
|
||||
|
||||
Reference in New Issue
Block a user