Compare commits
12
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d1cc60b00a | ||
|
|
efe8939250 | ||
|
|
39b353b916 | ||
|
|
485bc7f2ab | ||
|
|
1d03ded9dd | ||
|
|
94a4d076e0 | ||
|
|
228dc4b22b | ||
|
|
db523e1f16 | ||
|
|
6b08c7a90e | ||
|
|
76142c4b7e | ||
|
|
1d84fdaf9e | ||
|
|
1c50ae5297 |
@@ -12,10 +12,10 @@ test_graph:
|
||||
PYTHONPATH=../../ python -m prompt_control.test_graph
|
||||
|
||||
test_encode:
|
||||
PYTHONPATH=../../ python -m prompt_control.test_encode
|
||||
PYTHONPATH=../../ python -m prompt_control.test_encode --verbose
|
||||
|
||||
test_encode_both:
|
||||
TEST_TE="clip_l t5" PYTHONPATH=../../ python -m prompt_control.test_encode
|
||||
TEST_TE="clip_l t5" PYTHONPATH=../../ python -m prompt_control.test_encode --verbose
|
||||
|
||||
test_heavy: test_graph test_encode_both
|
||||
|
||||
|
||||
@@ -32,21 +32,6 @@ Prompt Control uses graph generation, and tries to delegate functionality to co
|
||||
|
||||
If you encounter issues as a user or if you're a node developer and Prompt Control somehow breaks something, feel free to file a bug report.
|
||||
|
||||
## 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. 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
|
||||
|
||||
### Removed features
|
||||
|
||||
- Prompt interpolation syntax; it was too cumbersome to maintain
|
||||
- LoRA block weight integration; ditto, for now.
|
||||
|
||||
### Everything broke, where are the old nodes?
|
||||
|
||||
If you really need them, you can install the [legacy nodes](https://github.com/asagi4/comfyui-prompt-control-legacy). However, I will not fix bugs in those nodes, and I strongly recommend just migrating your workflows to the new nodes.
|
||||
|
||||
You can have both installed at the same time; none of the nodes conflict.
|
||||
|
||||
## Requirements
|
||||
|
||||
For LoRA scheduling to work, you'll need at least version 0.3.7 of ComfyUI (0.3.36 of ComfyUI desktop).
|
||||
@@ -100,8 +85,4 @@ This node configures `PCTextEncode` default values for some functions by attachi
|
||||
|
||||
- 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.
|
||||
|
||||
- Cutoff does not work with models that use non-CLIP text encoders, like Flux. This might be fixable, but it's uncertain if cutoff even makes sense for those models.
|
||||
|
||||
@@ -23,9 +23,6 @@ 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 = {}
|
||||
|
||||
|
||||
+8
-1
@@ -89,10 +89,17 @@ You can refer to LoRAs by using the filename without extension and subdirectorie
|
||||
|
||||
Alternatively, the name can include the full directory path relative to ComfyUI's search paths, without extension: `<lora:XL/sdxllora:0.5>`. In this case, the *full* path must match.
|
||||
|
||||
You can also give the exact path (including the extension) as shown in `LoRALoader`.
|
||||
|
||||
If no match is found, the node will try to replace spaces with underscores and search again. That is, `<lora:cats and dogs:1>` will find `cats_and_dogs.safetensors`. This helps with some autocompletion scripts that replace underscores with spaces.
|
||||
|
||||
Finally, you can give the exact path (including the extension) as shown in `LoRALoader`.
|
||||
Finally, if none of the above produce a match, the search term will be split by whitespace and files that contain all of the parts in any order will be considered. If this returns only a single match, it will be loaded. For example, consider LoRAs:
|
||||
|
||||
- `xl/red_cats.safetensors`
|
||||
- `flux/blue_cats.safetensors`
|
||||
- `flux/red_cats.safetensors`
|
||||
|
||||
Then `<lora:cats xl:1>` would match the red cats LoRA, but `cats flux` would be ambiguous and not match.
|
||||
|
||||
## Alternating
|
||||
|
||||
|
||||
@@ -3,7 +3,6 @@ import numpy as np
|
||||
from math import copysign
|
||||
import logging
|
||||
import itertools
|
||||
from .adv_encode_old import old_advanced_encode_from_tokens
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
@@ -373,20 +372,7 @@ def advanced_encode_from_tokens(
|
||||
tokenizer=None,
|
||||
**extra_args,
|
||||
):
|
||||
if "old+" not in weight_interpretation:
|
||||
enc = AdvancedEncoder(
|
||||
encode_func, weight_interpretation, token_normalization, tokenizer, m_token, w_max, **extra_args
|
||||
)
|
||||
return enc(tokenized, return_pooled=return_pooled, apply_to_pooled=apply_to_pooled)
|
||||
else:
|
||||
weight_interpretation = weight_interpretation.replace("old+", "")
|
||||
log.warning("Using old implementation of %s", weight_interpretation)
|
||||
return old_advanced_encode_from_tokens(
|
||||
tokenized,
|
||||
token_normalization,
|
||||
weight_interpretation,
|
||||
encode_func,
|
||||
266,
|
||||
return_pooled=return_pooled,
|
||||
apply_to_pooled=apply_to_pooled,
|
||||
)
|
||||
enc = AdvancedEncoder(
|
||||
encode_func, weight_interpretation, token_normalization, tokenizer, m_token, w_max, **extra_args
|
||||
)
|
||||
return enc(tokenized, return_pooled=return_pooled, apply_to_pooled=apply_to_pooled)
|
||||
|
||||
@@ -1,235 +0,0 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
import logging
|
||||
import itertools
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def _norm_mag(w, n):
|
||||
d = w - 1
|
||||
return 1 + np.sign(d) * np.sqrt(np.abs(d) ** 2 / n)
|
||||
# return np.sign(w) * np.sqrt(np.abs(w)**2 / n)
|
||||
|
||||
|
||||
def _grouper(n, iterable):
|
||||
it = iter(iterable)
|
||||
while True:
|
||||
chunk = list(itertools.islice(it, n))
|
||||
if not chunk:
|
||||
return
|
||||
yield chunk
|
||||
|
||||
|
||||
def batched_clip_encode(tokens, length, encode_func, num_chunks):
|
||||
embs = []
|
||||
for e in _grouper(32, tokens):
|
||||
enc, pooled = encode_func(e)
|
||||
enc = enc.reshape((len(e), length, -1))
|
||||
embs.append(enc)
|
||||
|
||||
embs = torch.cat(embs)
|
||||
embs = embs.reshape((len(tokens) // num_chunks, length * num_chunks, -1))
|
||||
return embs
|
||||
|
||||
|
||||
def weights_like(weights, emb):
|
||||
return torch.tensor(weights, dtype=emb.dtype, device=emb.device).reshape(1, -1, 1).expand(emb.shape)
|
||||
|
||||
|
||||
def divide_length(word_ids, weights):
|
||||
sums = dict(zip(*np.unique(word_ids, return_counts=True)))
|
||||
sums[0] = 1
|
||||
weights = [[_norm_mag(w, sums[id]) if id != 0 else 1.0 for w, id in zip(x, y)] for x, y in zip(weights, word_ids)]
|
||||
return weights
|
||||
|
||||
|
||||
def shift_mean_weight(word_ids, weights):
|
||||
delta = 1 - np.mean([w for x, y in zip(weights, word_ids) for w, id in zip(x, y) if id != 0])
|
||||
weights = [[w if id == 0 else w + delta for w, id in zip(x, y)] for x, y in zip(weights, word_ids)]
|
||||
return weights
|
||||
|
||||
|
||||
def scale_to_norm(weights, word_ids, w_max):
|
||||
top = np.max(weights)
|
||||
w_max = min(top, w_max)
|
||||
weights = [[w_max if id == 0 else (w / top) * w_max for w, id in zip(x, y)] for x, y in zip(weights, word_ids)]
|
||||
return weights
|
||||
|
||||
|
||||
def mask_word_id(tokens, word_ids, target_id, mask_token):
|
||||
new_tokens = [[mask_token if wid == target_id else t for t, wid in zip(x, y)] for x, y in zip(tokens, word_ids)]
|
||||
mask = np.array(word_ids) == target_id
|
||||
return (new_tokens, mask)
|
||||
|
||||
|
||||
def from_masked(tokens, weights, word_ids, base_emb, length, encode_func, m_token=266):
|
||||
pooled_base = base_emb[0, length - 1 : length, :]
|
||||
wids, inds = np.unique(np.array(word_ids).reshape(-1), return_index=True)
|
||||
weight_dict = dict((id, w) for id, w in zip(wids, np.array(weights).reshape(-1)[inds]) if w != 1.0)
|
||||
|
||||
if len(weight_dict) == 0:
|
||||
return torch.zeros_like(base_emb), base_emb[0, length - 1 : length, :]
|
||||
|
||||
weight_tensor = torch.tensor(weights, dtype=base_emb.dtype, device=base_emb.device)
|
||||
weight_tensor = weight_tensor.reshape(1, -1, 1).expand(base_emb.shape)
|
||||
|
||||
# m_token = (clip.tokenizer.end_token, 1.0) if clip.tokenizer.pad_with_end else (0,1.0)
|
||||
# TODO: find most suitable masking token here
|
||||
m_token = (m_token, 1.0)
|
||||
|
||||
ws = []
|
||||
masked_tokens = []
|
||||
masks = []
|
||||
|
||||
# create prompts
|
||||
for id, w in weight_dict.items():
|
||||
masked, m = mask_word_id(tokens, word_ids, id, m_token)
|
||||
masked_tokens.extend(masked)
|
||||
|
||||
m = torch.tensor(m, dtype=base_emb.dtype, device=base_emb.device)
|
||||
m = m.reshape(1, -1, 1).expand(base_emb.shape)
|
||||
masks.append(m)
|
||||
|
||||
ws.append(w)
|
||||
|
||||
# batch process prompts
|
||||
embs = batched_clip_encode(masked_tokens, length, encode_func, len(tokens))
|
||||
masks = torch.cat(masks)
|
||||
|
||||
embs = base_emb.expand(embs.shape) - embs
|
||||
pooled = embs[0, length - 1 : length, :]
|
||||
|
||||
embs *= masks
|
||||
embs = embs.sum(axis=0, keepdim=True)
|
||||
|
||||
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)
|
||||
|
||||
return ((weight_tensor - 1) * embs), pooled_base + pooled
|
||||
|
||||
|
||||
def mask_inds(tokens, inds, mask_token):
|
||||
clip_len = len(tokens[0])
|
||||
inds_set = set(inds)
|
||||
new_tokens = [
|
||||
[mask_token if i * clip_len + j in inds_set else t for j, t in enumerate(x)] for i, x in enumerate(tokens)
|
||||
]
|
||||
return new_tokens
|
||||
|
||||
|
||||
def down_weight(tokens, weights, word_ids, base_emb, length, encode_func):
|
||||
w, w_inv = np.unique(weights, return_inverse=True)
|
||||
|
||||
if np.sum(w < 1) == 0:
|
||||
return base_emb, tokens, base_emb[0, length - 1 : length, :]
|
||||
# m_token = (clip.tokenizer.end_token, 1.0) if clip.tokenizer.pad_with_end else (0,1.0)
|
||||
# using the comma token as a masking token seems to work better than aos tokens for SD 1.x
|
||||
m_token = (266, 1.0)
|
||||
|
||||
masked_tokens = []
|
||||
|
||||
masked_current = tokens
|
||||
for i in range(len(w)):
|
||||
if w[i] >= 1:
|
||||
continue
|
||||
masked_current = mask_inds(masked_current, np.where(w_inv == i)[0], m_token)
|
||||
masked_tokens.extend(masked_current)
|
||||
|
||||
embs = batched_clip_encode(masked_tokens, length, encode_func, len(tokens))
|
||||
embs = torch.cat([base_emb, embs])
|
||||
w = w[w <= 1.0]
|
||||
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)
|
||||
return weighted_emb, masked_current, weighted_emb[0, length - 1 : length, :]
|
||||
|
||||
|
||||
def scale_emb_to_mag(base_emb, weighted_emb):
|
||||
norm_base = torch.linalg.norm(base_emb)
|
||||
norm_weighted = torch.linalg.norm(weighted_emb)
|
||||
embeddings_final = (norm_base / norm_weighted) * weighted_emb
|
||||
return embeddings_final
|
||||
|
||||
|
||||
# For verification
|
||||
def A1111_renorm(base_emb, weighted_emb):
|
||||
embeddings_final = (base_emb.mean() / weighted_emb.mean()) * weighted_emb
|
||||
return embeddings_final
|
||||
|
||||
|
||||
def from_zero(weights, base_emb):
|
||||
weight_tensor = torch.tensor(weights, dtype=base_emb.dtype, device=base_emb.device)
|
||||
weight_tensor = weight_tensor.reshape(1, -1, 1).expand(base_emb.shape)
|
||||
return base_emb * weight_tensor
|
||||
|
||||
|
||||
def old_advanced_encode_from_tokens(
|
||||
tokenized,
|
||||
token_normalization,
|
||||
weight_interpretation,
|
||||
encode_func,
|
||||
m_token=266,
|
||||
w_max=1.0,
|
||||
return_pooled=False,
|
||||
apply_to_pooled=False,
|
||||
**extra_args,
|
||||
):
|
||||
length = 77
|
||||
tokens = [[t for t, _, _ in x] for x in tokenized]
|
||||
weights = [[w for _, w, _ in x] for x in tokenized]
|
||||
word_ids = [[wid for _, _, wid in x] for x in tokenized]
|
||||
|
||||
# weight normalization
|
||||
# ====================
|
||||
|
||||
# distribute down/up weights over word lengths
|
||||
if token_normalization.startswith("length"):
|
||||
weights = divide_length(word_ids, weights)
|
||||
|
||||
# make mean of word tokens 1
|
||||
if token_normalization.endswith("mean"):
|
||||
weights = shift_mean_weight(word_ids, weights)
|
||||
|
||||
# weight interpretation
|
||||
# =====================
|
||||
pooled = None
|
||||
|
||||
if weight_interpretation in ["comfy", "perp"]:
|
||||
weighted_tokens = [[(t, w) for t, w in zip(x, y)] for x, y in zip(tokens, weights)]
|
||||
weighted_emb, pooled_base = encode_func(weighted_tokens)
|
||||
pooled = pooled_base
|
||||
else:
|
||||
unweighted_tokens = [[(t, 1.0) for t, _, _ in x] for x in tokenized]
|
||||
base_emb, pooled_base = encode_func(unweighted_tokens)
|
||||
|
||||
if weight_interpretation == "A1111":
|
||||
weighted_emb = from_zero(weights, base_emb)
|
||||
weighted_emb = A1111_renorm(base_emb, weighted_emb)
|
||||
pooled = pooled_base
|
||||
|
||||
if weight_interpretation == "compel":
|
||||
pos_tokens = [[(t, w) if w >= 1.0 else (t, 1.0) for t, w in zip(x, y)] for x, y in zip(tokens, weights)]
|
||||
weighted_emb, _ = encode_func(pos_tokens)
|
||||
weighted_emb, _, pooled = down_weight(pos_tokens, weights, word_ids, weighted_emb, length, encode_func)
|
||||
|
||||
if weight_interpretation == "comfy++":
|
||||
weighted_emb, tokens_down, _ = down_weight(unweighted_tokens, weights, word_ids, base_emb, length, encode_func)
|
||||
weights = [[w if w > 1.0 else 1.0 for w in x] for x in weights]
|
||||
# unweighted_tokens = [[(t,1.0) for t, _,_ in x] for x in tokens_down]
|
||||
embs, pooled = from_masked(unweighted_tokens, weights, word_ids, base_emb, length, encode_func)
|
||||
weighted_emb += embs
|
||||
|
||||
if weight_interpretation == "down_weight":
|
||||
weights = scale_to_norm(weights, word_ids, w_max)
|
||||
weighted_emb, _, pooled = down_weight(unweighted_tokens, weights, word_ids, base_emb, length, encode_func)
|
||||
|
||||
if return_pooled:
|
||||
if apply_to_pooled:
|
||||
return weighted_emb, pooled
|
||||
else:
|
||||
return weighted_emb, pooled_base
|
||||
return weighted_emb, None
|
||||
@@ -76,10 +76,6 @@ class AttentionCoupleHook(TransformerOptionsHook):
|
||||
self.kv = {"k": None, "v": None}
|
||||
|
||||
def initialize_regions(self, base_cond, conds, fill):
|
||||
self._base_cond = base_cond
|
||||
self._conds = conds
|
||||
self._fill = 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]
|
||||
|
||||
@@ -448,16 +448,9 @@ def expand_macros(text):
|
||||
return res
|
||||
|
||||
|
||||
def substitute_def(text, search, replace):
|
||||
search, default_args = search
|
||||
for i, v in enumerate(default_args):
|
||||
replace = re.sub(rf"\${i+1}\b", v, replace)
|
||||
return re.sub(rf"\b{re.escape(search)}\b", replace, text)
|
||||
|
||||
|
||||
def substitute_defcall(text, search, replace):
|
||||
name, default_args = search
|
||||
text, defns = get_function(text, name, defaults=None, placeholder=f"DEFNCALL{name}")
|
||||
text, defns = get_function(text, name, defaults=None, placeholder=f"DEFNCALL{name}", require_args=False)
|
||||
for i, parameters in enumerate(defns):
|
||||
ph = f"\0DEFNCALL{name}{i}\0"
|
||||
paramvals = []
|
||||
|
||||
+20
-10
@@ -1,11 +1,12 @@
|
||||
import logging
|
||||
import re
|
||||
import torch
|
||||
import math
|
||||
from functools import partial
|
||||
from comfy_extras.nodes_mask import FeatherMask, MaskComposite
|
||||
from nodes import ConditioningAverage
|
||||
|
||||
from .utils import safe_float, get_function, split_by_function, parse_floats, smarter_split
|
||||
from .utils import safe_float, get_function, split_by_function, parse_floats, smarter_split, call_node
|
||||
from .adv_encode import advanced_encode_from_tokens
|
||||
from .cutoff import process_cuts
|
||||
from .parser import parse_cuts
|
||||
@@ -231,7 +232,7 @@ def encode_prompt_segment(
|
||||
|
||||
# Chunks to ConditioningAverage:
|
||||
|
||||
text, averages = split_by_function(text, "AVG", ["0.5"])
|
||||
text, averages = split_by_function(text, "AVG", ["0.5"], require_args=False)
|
||||
prompts_to_avg = []
|
||||
for avg in averages:
|
||||
w = safe_float(avg["args"][0], 0.5)
|
||||
@@ -264,13 +265,22 @@ def encode_prompt_segment(
|
||||
w = next_w
|
||||
continue
|
||||
for i in range(len(base)):
|
||||
(cond,) = ConditioningAverage.addWeighted(None, [base[i]], [cond[i]], w)
|
||||
(cond,) = call_node(ConditioningAverage, [base[i]], [cond[i]], w)
|
||||
base[i] = cond[0]
|
||||
w = next_w
|
||||
|
||||
return base
|
||||
|
||||
|
||||
def calc_w(tensor, w):
|
||||
if math.isclose(w, 0):
|
||||
return torch.zeros_like(tensor)
|
||||
elif math.isclose(w, 1.0):
|
||||
return tensor
|
||||
else:
|
||||
return tensor * w
|
||||
|
||||
|
||||
def apply_weights(output, te_name, spec):
|
||||
"""Applies weights to TE outputs"""
|
||||
if not spec:
|
||||
@@ -292,16 +302,16 @@ def apply_weights(output, te_name, spec):
|
||||
if pooled_w is None:
|
||||
pooled_w = 1.0
|
||||
log.info("Weighting %s output by %s, pooled by %s", te_name, w, pooled_w)
|
||||
out = out * w
|
||||
out = calc_w(out, w)
|
||||
if pooled is not None:
|
||||
pooled = pooled * pooled_w
|
||||
pooled = calc_w(pooled, pooled_w)
|
||||
|
||||
return out, pooled
|
||||
else:
|
||||
if te_name in spec or default is not None:
|
||||
w = spec.get(te_name, default)
|
||||
log.info("Weighting %s output by %s", te_name, w)
|
||||
output = output * w
|
||||
output = calc_w(output, w)
|
||||
return output
|
||||
|
||||
|
||||
@@ -429,7 +439,7 @@ def get_mask(text, size, input_masks):
|
||||
|
||||
def feather(f, mask):
|
||||
l, t, r, b, *_ = [int(x) for x in parse_floats(f[0], [0, 0, 0, 0], split_re="\\s+")]
|
||||
mask = FeatherMask().feather(mask, l, t, r, b)[0]
|
||||
mask = call_node(FeatherMask, mask, l, t, r, b)[0]
|
||||
log.info("FeatherMask l=%s, t=%s, r=%s, b=%s", l, t, r, b)
|
||||
return mask
|
||||
|
||||
@@ -447,7 +457,7 @@ def get_mask(text, size, input_masks):
|
||||
i += 1
|
||||
if mask is not None:
|
||||
log.info("MaskComposite op=%s", op)
|
||||
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
|
||||
mask = call_node(MaskComposite, mask, nextmask, 0, 0, op)[0]
|
||||
else:
|
||||
mask = nextmask
|
||||
|
||||
@@ -462,7 +472,7 @@ def get_mask(text, size, input_masks):
|
||||
nextmask = feather(feathers[i], nextmask)
|
||||
i += 1
|
||||
if mask is not None:
|
||||
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
|
||||
mask = call_node(MaskComposite, mask, nextmask, 0, 0, op)[0]
|
||||
else:
|
||||
mask = nextmask
|
||||
|
||||
@@ -573,7 +583,7 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
|
||||
return f"MASK({args})"
|
||||
|
||||
for prompt in prompts:
|
||||
base_prompt, attn_couple_prompts = split_by_function(prompt, "COUPLE", defaults=None)
|
||||
base_prompt, attn_couple_prompts = split_by_function(prompt, "COUPLE", defaults=None, require_args=False)
|
||||
|
||||
prompts = [base_prompt] + [couple_mask(p["args"]) + p["text"] for p in attn_couple_prompts]
|
||||
encoded = []
|
||||
|
||||
@@ -14,7 +14,10 @@ logging.basicConfig()
|
||||
|
||||
|
||||
def run(f, *args):
|
||||
return getattr(f, f.FUNCTION)(*args)
|
||||
if hasattr(f, "execute"):
|
||||
return f.execute(*args)
|
||||
else:
|
||||
return getattr(f, f.FUNCTION)(*args)
|
||||
|
||||
|
||||
@mock.patch("torch.cuda.current_device", lambda: "cpu")
|
||||
@@ -80,6 +83,12 @@ class TestEncode(unittest.TestCase):
|
||||
c = c2 # Used in later tests
|
||||
self.condEqual(c1, c2)
|
||||
|
||||
with self.subTest("Function cornercase"):
|
||||
(c1,) = run(pc, clip, "test SDXL function")
|
||||
(c2,) = run(comfy, clip, "test SDXL function")
|
||||
(c3,) = run(pc, clip, "test SDXL() function")
|
||||
self.condEqual(c1, c2)
|
||||
|
||||
with self.subTest("Weights"):
|
||||
(c1,) = run(pc, clip, "(test:1.2) (test:0.6)")
|
||||
(c2,) = run(comfy, clip, "(test:1.2) (test:0.6)")
|
||||
@@ -104,8 +113,20 @@ class TestEncode(unittest.TestCase):
|
||||
(c1,) = run(comfy, clip, "test1")
|
||||
(c2,) = run(comfy, clip, "test2")
|
||||
(c3,) = run(pc, clip, "test1 AVG() test2")
|
||||
(c4,) = run(pc, clip, "test1 AVG test2")
|
||||
(avg,) = run(average, c1, c2, 0.5)
|
||||
self.condEqual(avg, c3)
|
||||
self.condEqual(avg, c4)
|
||||
|
||||
@unittest.expectedFailure
|
||||
def test_failure(self):
|
||||
pc = PCTextEncode()
|
||||
comfy = nodes.CLIPTextEncode()
|
||||
for k, clip in clips:
|
||||
with self.subTest(k):
|
||||
(c1,) = run(comfy, clip, "test SDXL function")
|
||||
(c2,) = run(pc, clip, "test SDXL() function")
|
||||
self.condEqual(c1, c2)
|
||||
|
||||
def test_weight(self):
|
||||
pc = PCTextEncode()
|
||||
|
||||
@@ -1,71 +0,0 @@
|
||||
import unittest
|
||||
import numpy.testing as npt
|
||||
from os import environ
|
||||
|
||||
clips = []
|
||||
|
||||
|
||||
def run(f, *args):
|
||||
return getattr(f, f.FUNCTION)(*args)
|
||||
|
||||
|
||||
class TestEncode(unittest.TestCase):
|
||||
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)):
|
||||
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_styles(self):
|
||||
pc = PCTextEncode()
|
||||
for k, clip in clips:
|
||||
for style in ["comfy++", "A1111", "comfy++", "compel", "down_weight"]:
|
||||
with self.subTest(f"TE {k} style {style} does not fail when encoding weights"):
|
||||
for normalization in ["none", "mean", "length", "length+mean"]:
|
||||
with self.subTest(f"TE {k} style {style} normalization {normalization}"):
|
||||
(c,) = run(
|
||||
pc,
|
||||
clip,
|
||||
f"STYLE(old+{style}, {normalization}) this prompt has weights, (a:1.2) (b:1.2)",
|
||||
)
|
||||
(c2,) = run(
|
||||
pc,
|
||||
clip,
|
||||
f"STYLE({style}, {normalization}) this prompt has weights, (a:1.2) (b:1.2)",
|
||||
)
|
||||
self.condEqual(c, c2)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("Loading ComfyUI")
|
||||
from comfy.sd import load_clip
|
||||
from .nodes_base import PCTextEncode
|
||||
from pathlib import Path
|
||||
|
||||
to_test = environ.get("TEST_TE", "clip_l").split()
|
||||
model_path = environ.get("COMFYUI_MODEL_ROOT", ".")
|
||||
|
||||
te_root = (Path(model_path) / "text_encoders").resolve()
|
||||
|
||||
if "clip_l" in to_test:
|
||||
clip_l = load_clip(
|
||||
ckpt_paths=[str(te_root / "clip_l.safetensors")], clip_type="stable_diffusion", model_options={}
|
||||
)
|
||||
clips.append(("clip_l", clip_l))
|
||||
|
||||
if "t5" in to_test:
|
||||
dual = load_clip(
|
||||
[str(te_root / "clip_l.safetensors"), str(te_root / "t5xxl_fp16.safetensors")],
|
||||
clip_type="flux",
|
||||
model_options={},
|
||||
)
|
||||
clips.append(("clip_l+t5", dual))
|
||||
|
||||
print("Starting tests")
|
||||
unittest.main()
|
||||
+25
-5
@@ -15,6 +15,15 @@ except ImportError:
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def call_node(cls, *args, **kwargs):
|
||||
if hasattr(cls, "execute"):
|
||||
# v3 node
|
||||
return cls.execute(*args, **kwargs)
|
||||
else:
|
||||
func = getattr(cls(), cls.FUNCTION)
|
||||
return func(*args, **kwargs)
|
||||
|
||||
|
||||
def consolidate_schedule(prompt_schedule):
|
||||
prev_loras = {}
|
||||
not_found = []
|
||||
@@ -88,14 +97,19 @@ def find_closing_paren(text, start):
|
||||
return len(text)
|
||||
|
||||
|
||||
def get_function(text, func, defaults, return_func_name=False, placeholder="", return_dict=False):
|
||||
rex = re.compile(rf"\b{func}\b", re.MULTILINE)
|
||||
def get_function(text, func, defaults, return_func_name=False, placeholder="", return_dict=False, require_args=True):
|
||||
if require_args:
|
||||
rex = re.compile(rf"\b{func}\(", re.MULTILINE)
|
||||
else:
|
||||
rex = re.compile(rf"\b{func}\b", re.MULTILINE)
|
||||
instances = []
|
||||
match = rex.search(text)
|
||||
count = 0
|
||||
while match:
|
||||
# Match start, content start
|
||||
start, at_paren = match.span()
|
||||
if require_args:
|
||||
at_paren = at_paren - 1
|
||||
funcname = text[start:at_paren]
|
||||
after_first_paren = at_paren + 1
|
||||
if text[at_paren:after_first_paren] == "(":
|
||||
@@ -104,7 +118,7 @@ def get_function(text, func, defaults, return_func_name=False, placeholder="", r
|
||||
end += 1
|
||||
else:
|
||||
end = at_paren
|
||||
args = None
|
||||
args = defaults
|
||||
ph = None
|
||||
if placeholder:
|
||||
ph = f"\0{placeholder}{count}\0"
|
||||
@@ -131,11 +145,11 @@ def get_function(text, func, defaults, return_func_name=False, placeholder="", r
|
||||
return text, instances
|
||||
|
||||
|
||||
def split_by_function(text, func, defaults=None):
|
||||
def split_by_function(text, func, defaults=None, require_args=True):
|
||||
"""
|
||||
Splits a string by function calls, returning the text preceding the first call and a list of dictionaries with a "text" key with the prompt before the next split or until hthe end of the text.
|
||||
"""
|
||||
text, functions = get_function(text, func, defaults, return_dict=True)
|
||||
text, functions = get_function(text, func, defaults, return_dict=True, require_args=require_args)
|
||||
chunks = []
|
||||
prev = 0
|
||||
for f in functions:
|
||||
@@ -195,6 +209,12 @@ def lora_name_to_file(name):
|
||||
p = Path(f).with_suffix("")
|
||||
if p.name == n or str(p) == n:
|
||||
return f
|
||||
# Finally, try to find unique match from parts
|
||||
parts = name.split()
|
||||
search = [f for f in filenames if all(p in f for p in parts)]
|
||||
if len(search) == 1:
|
||||
return search[0]
|
||||
|
||||
return None
|
||||
|
||||
|
||||
|
||||
+1
-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.0.0-rc.9"
|
||||
version = "2.1.0"
|
||||
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