Compare commits

...
22 Commits
Author SHA1 Message Date
asagi4 68cda3663e v2.0.0-rc.6 2025-06-08 01:02:16 +03:00
asagi4 3d46f705b6 Test downweighting too, and normalizations 2025-06-08 01:01:57 +03:00
asagi4 eec4bc4da9 Split old code into its own file for easy removal later 2025-06-08 00:55:43 +03:00
asagi4 3d3218e831 Tests for verifying refactor 2025-06-08 00:55:43 +03:00
asagi4 f4e57ec514 Switch on new implementation by default 2025-06-08 00:55:40 +03:00
asagi4 74f65c1b31 Use old from_masked batching for now
I can't figure out why from_masked works differently from down_weight
which also used that batching function but could be replaced with a simple
torch.cat
2025-06-08 00:50:26 +03:00
asagi4 4d94cca88f Restore adv_encoding to original implementation to compare them 2025-06-08 00:44:56 +03:00
asagi4 4569ecccf9 Manual testing... 2025-06-08 00:44:31 +03:00
asagi4 192e6d30d4 refactor adv_encode 2025-06-07 21:41:13 +03:00
asagi4 67112f11e0 T5 makes STYLE(perp) return NaNs. Just replace them with 0 2025-06-07 21:41:13 +03:00
asagi4 f761dfac86 Fix comfy++ with more than one weight, see #115
I'm not sure if this is correct, but it at least doesn't fail.
2025-06-07 14:09:32 +03:00
asagi4 278e733835 Fix normalization validity check 2025-06-07 14:09:32 +03:00
asagi4 2bf65720eb Fix STYLE(perp) exception, see #115 2025-06-07 14:09:32 +03:00
asagi4 99f3af92b7 Add tests for weighting and a way to run manual testing 2025-06-07 14:09:14 +03:00
asagi4 2437cd4daf Merge pull request #116 from pamparamm/pooled_none_check
Partially resolve #115
2025-06-07 13:30:01 +03:00
asagi4 d262a7dc7a Remove an extra clone. 2025-06-07 13:23:06 +03:00
Pam 6c319ad5b4 Partially resolve #115 2025-06-07 09:38:13 +05:00
asagi4 34056cac19 v2.0.0-rc.5 2025-06-06 19:52:02 +03:00
asagi4 8ae436abf1 Merge pull request #114 from pamparamm/negpip_option
Use ppm_negpip option to detect NegPiP
2025-06-06 16:49:20 +03:00
Pam c2ce2ce023 Use ppm_negpip option to detect NegPiP 2025-06-06 16:24:15 +05:00
asagi4 7a9e69ec31 Fix errors found in testing
Who knew tests could be useful, too?
2025-06-05 23:49:02 +03:00
asagi4 aaff8dc7da Encoding tests 2025-06-05 23:49:02 +03:00
10 changed files with 741 additions and 204 deletions
+3
View File
@@ -14,4 +14,7 @@ test_graph:
test_encode:
PYTHONPATH=../../ python -m prompt_control.test_encode
manual_test:
PYTHONPATH=../../ python -im prompt_control.manual_test
.PHONY: check format all
+314 -177
View File
@@ -1,6 +1,11 @@
import torch
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")
def _norm_mag(w, n):
@@ -9,23 +14,31 @@ def _norm_mag(w, 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)
@@ -39,47 +52,6 @@ def mask_word_id(tokens, word_ids, target_id, mask_token):
return (new_tokens, mask)
def from_masked(tokens, weights, word_ids, base_emb, pooled_base, max_length, encode_func, m_token):
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), torch.zeros_like(pooled_base) if pooled_base is not None else None
weight_tensor = weights_like(weights, base_emb)
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)
masks.append(weights_like(m, base_emb))
ws.append(w)
embs, pooled = encode_func(tokens)
masks = torch.cat(masks)
embs = base_emb.expand(embs.shape) - embs
if pooled is not None and max_length:
pooled = embs[0, max_length - 1 : 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_base + pooled
embs *= masks
embs = embs.sum(axis=0, keepdim=True)
return ((weight_tensor - 1) * embs), pooled
def mask_inds(tokens, inds, mask_token):
clip_len = len(tokens[0])
inds_set = set(inds)
@@ -89,39 +61,6 @@ def mask_inds(tokens, inds, mask_token):
return new_tokens
def down_weight(tokens, weights, word_ids, base_emb, pooled_base, max_length, encode_func, m_token):
w, w_inv = np.unique(weights, return_inverse=True)
if np.sum(w < 1) == 0:
return (
base_emb,
tokens,
base_emb[0, max_length - 1 : max_length, :] if (pooled_base is not None and max_length) else None,
)
m_token = (m_token, 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, pooled = encode_func(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)
if pooled and max_length:
pooled = weighted_emb[0, max_length - 1 : max_length, :]
return weighted_emb, masked_current, pooled
def scale_emb_to_mag(base_emb, weighted_emb):
norm_base = torch.linalg.norm(base_emb)
norm_weighted = torch.linalg.norm(weighted_emb)
@@ -129,12 +68,6 @@ def scale_emb_to_mag(base_emb, weighted_emb):
return embeddings_final
def recover_dist(base_emb, weighted_emb):
fixed_std = (base_emb.std() / weighted_emb.std()) * (weighted_emb - weighted_emb.mean())
embeddings_final = fixed_std + (base_emb.mean() - fixed_std.mean())
return embeddings_final
def perp_weight(weights, unweighted_embs, empty_embs):
unweighted, unweighted_pooled = unweighted_embs
zero, zero_pooled = empty_embs
@@ -153,9 +86,281 @@ def perp_weight(weights, unweighted_embs, empty_embs):
result[~over1] = (unweighted - (1 - weights) * perp)[~over1]
result[weights == 0.0] = zero[weights == 0.0]
# Not sure if this is an implementation bug or if this just doesn't make sense with T5
nans = result.isnan()
if nans.any():
log.warning("perp weight returned NaNs (known to happen with T5), replacing with 0")
result[nans] = 0.0
return result, unweighted_pooled
def style_comfy(encoder, tokens, **kwargs):
tokens = encoder.without_word_ids(tokens)
return encoder.encode_fn(tokens)
def style_a1111(encoder, tokens, **kwargs):
base_emb, pooled = encoder.base_emb(tokens)
weighted_emb = base_emb * weights_like(encoder.weights(tokens), base_emb)
weighted_emb = (base_emb.mean() / weighted_emb.mean()) * weighted_emb # renormalize
return weighted_emb, pooled
def style_compel(encoder, tokens, **kwargs):
pos_tokens = encoder.weighted_with(tokens, lambda w: w if w > 1.0 else 1.0)
weighted_emb, pooled = encoder.encode_fn(pos_tokens)
weighted_emb, _, pooled = encoder.down_weight(
pos_tokens, encoder.weights(tokens), encoder.word_ids(tokens), weighted_emb, pooled
)
return weighted_emb, pooled
def style_comfypp(encoder, tokens, **kwargs):
unweighted_tokens = encoder.unweighted(tokens)
base_emb, pooled_base = encoder.base_emb(tokens)
weighted_emb, tokens_down, _ = encoder.down_weight(
unweighted_tokens, encoder.weights(tokens), encoder.word_ids(tokens), base_emb, pooled_base
)
weights = encoder.weights(encoder.weighted_with(tokens, lambda w: w if w > 1.0 else 1.0))
embs, pooled = encoder.from_masked(
unweighted_tokens,
weights,
encoder.word_ids(tokens),
base_emb,
pooled_base,
)
weighted_emb += embs
return weighted_emb, pooled
def style_downweight(encoder, tokens, **kwargs):
weights = scale_to_norm(encoder.weights(tokens), encoder.word_ids(tokens), encoder.w_max)
base_emb, pooled_base = encoder.base_emb(tokens)
weighted_emb, _, pooled = encoder.down_weight(
encoder.unweighted(tokens), weights, encoder.word_ids(tokens), base_emb, pooled_base
)
return weighted_emb, pooled
def style_perp(encoder, tokens, **kwargs):
zero_emb, zero_pooled = encoder.encode_fn(encoder.tokenizer.tokenize_with_weights(""))
base_emb, pooled = encoder.base_emb(tokens)
return perp_weight(encoder.weights(tokens), (base_emb, pooled), (zero_emb, zero_pooled))
def apply_negpip(encoder, emb, pooled, **kwargs):
original_tokens = kwargs["original_tokens"]
emb_negpip = torch.empty_like(emb).repeat(1, 2, 1)
emb_negpip[:, 0::2, :] = emb
emb_negpip[:, 1::2, :] = emb * weights_like(encoder.signs(original_tokens), emb)
return emb_negpip, pooled
def norm_length(encoder, tokens, **kwargs):
word_ids = encoder.word_ids(tokens)
sums = dict(zip(*np.unique(word_ids, return_counts=True)))
sums[0] = 1
tokens = [[(t, _norm_mag(w, sums[id]) if id != 0 else 1.0, id) for (t, w, id) in x] for x in tokens]
return tokens
def norm_mean(encoder, tokens, **kwargs):
weights = encoder.weights(tokens)
word_ids = encoder.word_ids(tokens)
delta = 1 - np.mean([w for x, y in zip(weights, word_ids) for w, id in zip(x, y) if id != 0])
tokens = [[(t, w if id == 0 else w + delta, id) for (t, w, id) in x] for x in tokens]
return tokens
def norm_none(encoder, tokens, **kwargs):
return tokens
class AdvancedEncoder:
STYLES = {
"A1111": style_a1111,
"comfy": style_comfy,
"comfy++": style_comfypp,
"compel": style_compel,
"down_weight": style_downweight,
"perp": style_perp,
}
NORMALIZATION_OPS = {
"none": norm_none,
"length": norm_length,
"mean": norm_mean,
}
@classmethod
def add_encoder(cls, name, fn):
cls.STYLES[name] = fn
def add_normalization_op(cls, name, fn):
cls.NORMALIZATION_OPS[name] = fn
@classmethod
def weighted_with(cls, tokens, fn=id, word_ids=True):
w = ([(t, fn(w), id) for t, w, id in x] for x in tokens)
if not word_ids:
w = cls.without_word_ids(w)
return list(w)
@classmethod
def unweighted(cls, tokens, word_ids=False):
return cls.weighted_with(tokens, fn=lambda w: 1.0, word_ids=word_ids)
@classmethod
def tokens_only(cls, tokens):
return list([t[0] for t in x] for x in tokens)
@classmethod
def weights(cls, tokens):
return list([t[1] for t in x] for x in tokens)
@classmethod
def word_ids(cls, tokens):
return list([t[2] for t in x] for x in tokens)
@classmethod
def signs(cls, tokens):
return list([copysign(1, t[1]) for t in x] for x in tokens)
@classmethod
def without_word_ids(cls, tokens):
return list([(t, w) for t, w, _ in x] for x in tokens)
def __init__(self, encode_fn, style, normalization, tokenizer, m_token="+", w_max=1.0, **extra_args):
self.encode_fn = encode_fn
self.preprocessors = []
self.postprocessors = []
self.tokenizer = tokenizer
self.extra_args = extra_args
self.m_token = tokenizer.tokenize_with_weights(m_token)[0][tokenizer.tokens_start]
self.max_length = tokenizer.max_length if tokenizer.pad_to_max_length else None
self.w_max = w_max
if style == "comfy++" and not self.max_length:
log.warning("comfy++ does not work with tokenizer %s, using default weighting", tokenizer)
style = "comfy"
norms = normalization.split("+")
assert style in self.STYLES, f"Invalid weight interpretation: {style}"
self.weight_fn = self.STYLES[style]
for n in norms:
n = n.strip()
assert n in self.NORMALIZATION_OPS, f"Invalid normalization: {normalization}"
self.preprocessors.append(self.NORMALIZATION_OPS[n])
negpip = extra_args.get("has_negpip")
if negpip:
def _encode(t):
emb, pooled = encode_fn(t)
return emb[:, 0::2, :], pooled
self.encode_fn = _encode
self.preprocessors.insert(lambda encoder, tokens, **kwargs: encoder.weighted_with(tokens, abs))
self.postprocessors.insert(0, apply_negpip)
def base_emb(self, tokens):
unweighted = self.unweighted(tokens)
return self.encode_fn(unweighted)
def down_weight(self, tokens, weights, word_ids, base_emb, pooled_base):
w, w_inv = np.unique(weights, return_inverse=True)
if np.sum(w < 1) == 0:
return (
base_emb,
tokens,
(
base_emb[0, self.max_length - 1 : self.max_length, :]
if (pooled_base is not None and self.max_length)
else None
),
)
masked_current = tokens
emblist = [base_emb]
for i in range(len(w)):
if w[i] >= 1:
continue
masked_current = mask_inds(masked_current, np.where(w_inv == i)[0], self.m_token)
masked, _ = self.encode_fn(masked_current)
emblist.append(masked)
embs = torch.cat(emblist)
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)
pooled = pooled_base
if pooled is not None and self.max_length:
pooled = weighted_emb[0, self.max_length - 1 : self.max_length, :]
return weighted_emb, masked_current, pooled
def from_masked(self, tokens, weights, word_ids, base_emb, pooled_base):
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), torch.zeros_like(pooled_base) if pooled_base is not None else None
weight_tensor = weights_like(weights, base_emb)
ws = []
masked_tokens = []
masks = []
# create prompts
for id, w in weight_dict.items():
masked, m = mask_word_id(tokens, word_ids, id, self.m_token)
masks.append(weights_like(m, base_emb))
masked_tokens.extend(masked)
ws.append(w)
# TODO: figure out how to get rid of this
embs = batched_clip_encode(masked_tokens, self.max_length, self.encode_fn, len(tokens))
masks = torch.cat(masks)
embs = base_emb.expand(embs.shape) - embs
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_base + pooled
if embs.shape[0] != masks.shape[0]:
embs = embs.repeat(masks.shape[0], 1, 1)
embs *= masks
embs = embs.sum(axis=0, keepdim=True)
return ((weight_tensor - 1) * embs), pooled
def __call__(self, tokens, apply_to_pooled=False, return_pooled=False):
normalized_tokens = tokens
for op in self.preprocessors:
normalized_tokens = op(self, normalized_tokens)
emb, pooled = self.weight_fn(self, normalized_tokens, original_tokens=tokens)
for fn in self.postprocessors:
emb, pooled = fn(self, emb, pooled, tokens=tokens, original_tokens=tokens)
if return_pooled:
if not apply_to_pooled:
_, pooled = self.base_emb(tokens)
return emb, pooled
return emb, None
def advanced_encode_from_tokens(
tokenized,
token_normalization,
@@ -166,90 +371,22 @@ def advanced_encode_from_tokens(
return_pooled=False,
apply_to_pooled=False,
tokenizer=None,
**extra_args
**extra_args,
):
negpip = extra_args.get("has_negpip")
if negpip:
weights_sign = [[copysign(1, w) for _, w, _ in x] for x in tokenized]
tokenized = [[(t, abs(w), p) for t, w, p in x] for x in tokenized]
orig_encode = encode_func
def _encode(t):
emb, pooled = orig_encode(t)
return emb[:, 0::2, :], pooled
encode_func = _encode
assert tokenizer, "Must pass tokenizer"
max_length = None
if tokenizer.pad_to_max_length:
max_length = tokenizer.max_length
m_token = tokenizer.tokenize_with_weights(m_token)[0][tokenizer.tokens_start]
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]
for op in token_normalization.split("+"):
op = op.strip()
if op == "length":
# distribute down/up weights over word lengths
weights = divide_length(word_ids, weights)
if op == "mean":
weights = shift_mean_weight(word_ids, weights)
pooled = None
if weight_interpretation == "comfy":
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
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:
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 = base_emb * weights_like(weights, base_emb) # from_zero
weighted_emb = (base_emb.mean() / weighted_emb.mean()) * weighted_emb # renormalize
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, pooled = encode_func(pos_tokens)
weighted_emb, _, pooled = down_weight(
pos_tokens, weights, word_ids, weighted_emb, pooled, max_length, encode_func, m_token
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,
)
if weight_interpretation == "comfy++":
weighted_emb, tokens_down, _ = down_weight(
unweighted_tokens, weights, word_ids, base_emb, pooled_base, max_length, encode_func, m_token
)
weights = [[w if w > 1.0 else 1.0 for w in x] for x in weights]
embs, pooled = from_masked(
unweighted_tokens, weights, word_ids, base_emb, pooled_base, max_length, encode_func, m_token
)
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, pooled_base, max_length, encode_func, m_token
)
if weight_interpretation == "perp":
weighted_emb, pooled = perp_weight(
weights, (base_emb, pooled_base), encode_func(extra_args["tokenizer"].tokenize_with_weights(""))
)
if negpip:
emb_negpip = torch.empty_like(weighted_emb).repeat(1, 2, 1)
emb_negpip[:, 0::2, :] = weighted_emb
emb_negpip[:, 1::2, :] = weighted_emb * weights_like(weights_sign, weighted_emb)
weighted_emb = emb_negpip
if return_pooled:
if apply_to_pooled:
return weighted_emb, pooled
else:
return weighted_emb, pooled_base
return weighted_emb, None
+235
View File
@@ -0,0 +1,235 @@
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
+3 -3
View File
@@ -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
+46
View File
@@ -0,0 +1,46 @@
import main
import nodes
import prompt_control.adv_encode
(l,) = nodes.CLIPLoader.load_clip(None, "clip_l.safetensors")
(t5,) = nodes.CLIPLoader.load_clip(None, "t5base.safetensors")
id(main) # get rid of warning
def adv(t, text, style="A1111", norm="none", new=True, **kwargs):
c = t.tokenize(text, return_word_ids=True)
if new:
style = "new+" + style
if t is t5:
te = t.patcher.model.t5base.encode_token_weights
token = t.tokenizer.clip_t5base
tok = c["t5base"]
else:
te = t.patcher.model.clip_l.encode_token_weights
token = t.tokenizer.clip_l
tok = c["l"]
return prompt_control.adv_encode.advanced_encode_from_tokens(tok, norm, style, te, tokenizer=token)
def adv_all(t, text, styles=[], **kwargs):
r = []
for s in styles or prompt_control.adv_encode.AdvancedEncoder.STYLES:
print("Testing", s, kwargs)
r.append([s, adv(t, text, style=s, **kwargs)])
return r
def replacenan(t):
t[t.isnan()] = 42.123321
return t
def adv_equal(t, text, **kwargs):
old = adv_all(t, text, new=False, **kwargs)
new = adv_all(t, text, new=True, **kwargs)
r = {}
for i, o in enumerate(old):
n = new[i]
r[n[0]] = (replacenan(n[1][0]) == replacenan(o[1][0])).all()
return r
-1
View File
@@ -112,7 +112,6 @@ class PCAttentionCoupleBatchNegative(ComfyNodeABC):
n_hook_group: comfy.hooks.HookGroup = n[1].get("hooks", comfy.hooks.HookGroup()).clone()
p_hook_group: comfy.hooks.HookGroup = p[1].get("hooks", comfy.hooks.HookGroup())
attn_couple = [hook for hook in p_hook_group.hooks if isinstance(hook, AttentionCoupleHook)]
n_hook_group = n_hook_group.clone()
for hook in attn_couple:
n_hook_group.add(hook)
n[1]["hooks"] = p_hook_group if n_hook_group.hooks == p_hook_group.hooks else n_hook_group
+12 -14
View File
@@ -65,13 +65,14 @@ def get_style(text, default_style="comfy", default_normalization="none"):
style, normalization = styles[0]
style = style.strip()
normalization = normalization.strip()
if style not in AVAILABLE_STYLES:
if style.replace("old+", "") not in AVAILABLE_STYLES:
log.warning("Unrecognized prompt style: %s. Using %s", style, default_style)
style = default_style
if normalization not in AVAILABLE_NORMALIZATIONS:
log.warning("Unrecognized prompt normalization: %s. Using %s", normalization, default_normalization)
normalization = default_normalization
for part in normalization.split("+"):
if part not in AVAILABLE_NORMALIZATIONS:
log.warning("Unrecognized prompt normalization: %s. Using %s", normalization, default_normalization)
normalization = default_normalization
break
return style, normalization, text
@@ -294,7 +295,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 +324,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 +336,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(
+71 -8
View File
@@ -1,4 +1,5 @@
import unittest
import numpy.testing as npt
clip_l = None
dual = None
@@ -9,24 +10,85 @@ 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_styles(self):
pc = PCTextEncode()
comfy = nodes.CLIPTextEncode()
for k, clip in [("l", clip_l), ("dual", dual)]:
(no_weights,) = run(comfy, clip, "this prompt has no weights")
for style in ["comfy", "A1111", "comfy++", "compel", "down_weight", "perp"]:
with self.subTest(f"TE {k} style {style} no weights equal comfy"):
(c,) = run(pc, clip, "this prompt has no weights")
self.condEqual(no_weights, c)
with self.subTest(f"TE {k} style {style} does not fail when encoding weights"):
for normalization in ["none", "mean", "length", "mean+length", "length+mean"]:
with self.subTest(f"TE {k} style {style} normalization {normalization}"):
(c,) = run(
pc,
clip,
f"STYLE({style}, {normalization}) (this prompt) (has weights:0.9), (a:1.2) (b:1.2)",
)
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 +97,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")
+56
View File
@@ -0,0 +1,56 @@
import unittest
import numpy.testing as npt
clip_l = None
dual = None
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 [("l", clip_l), ("dual", dual)]:
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")
import main
id(main) # get rid of flake warning
import nodes
from .nodes_base import PCTextEncode
(clip_l,) = nodes.CLIPLoader().load_clip("clip_l.safetensors")
(dual,) = nodes.DualCLIPLoader().load_clip("clip_l.safetensors", "clip_g.safetensors", "sdxl")
print("Starting tests")
unittest.main()
+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-rc.4"
version = "2.0.0-rc.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"]