Update for Comfy changes

Updated for recent Comfy changes to CLIP stuff
Updated how weights between 0 and 1 and negative weights are handled
This commit is contained in:
BVH
2023-12-08 13:22:23 +05:30
committed by GitHub
parent 6c9df5dad6
commit e7c24c9fa1
+30 -12
View File
@@ -15,7 +15,7 @@ class CLIPTextEncodePerpWeight:
def encode(self, clip, text):
empty_tokens = clip.tokenize("")
sdxl_flag = isinstance(empty_tokens, dict)
sdxl_flag = "g" in empty_tokens.keys()
if sdxl_flag:
empty_cond, empty_cond_pooled = clip.encode_from_tokens(empty_tokens, return_pooled=True)
@@ -36,9 +36,13 @@ class CLIPTextEncodePerpWeight:
if weight_l > 1.0:
cond[i][j][:768] = token_vector_l + (weight_l * perp_l)
elif (weight_l > 0.0) and (weight_l < 1.0):
cond[i][j][:768] = token_vector_l - ((1 - weight_l) * perp_l)
elif weight_l < 0.0:
cond[i][j][:768] = token_vector_l + (weight_l * perp_l)
cond[i][j][:768] = zero_vector_l + (weight_l * (token_vector_l - zero_vector_l))
elif (weight_l > -1.0) and (weight_l < 0.0):
cond[i][j][:768] = -1 * (zero_vector_l + (abs(weight_l) * (token_vector_l - zero_vector_l)))
elif weight_l == -1.0:
cond[i][j][:768] = -1 * token_vector_l
elif weight_l < -1.0:
cond[i][j][:768] = -1 * (token_vector_l + (abs(weight_l) * perp_l))
elif weight_l == 0.0:
cond[i][j][:768] = empty_cond[0][(j%77)][:768]
@@ -50,21 +54,31 @@ class CLIPTextEncodePerpWeight:
if weight_g > 1.0:
cond[i][j][768:] = token_vector_g + (weight_g * perp_g)
elif (weight_g > 0.0) and (weight_g < 1.0):
cond[i][j][768:] = token_vector_g - ((1 - weight_g) * perp_g)
elif (weight_g < 0.0):
cond[i][j][768:] = token_vector_g + (weight_g * perp_g)
cond[i][j][768:] = weight_g * token_vector_g
elif (weight_g > -1.0) and (weight_g < 0.0):
cond[i][j][768:] = -1 * (zero_vector_g + (abs(weight_g) * (token_vector_g - zero_vector_g)))
elif weight_g == -1.0:
cond[i][j][768:] = -1 * token_vector_g
elif weight_g < -1.0:
cond[i][j][768:] = -1 * (token_vector_g + (abs(weight_g) * perp_g))
elif weight_g == 0.0:
cond[i][j][768:] = empty_cond[0][(j%77)][768:]
else:
empty_cond, empty_cond_pooled = clip.encode_from_tokens(empty_tokens, return_pooled=True)
tokens = clip.tokenize(text)
unweighted_tokens = [[(t, 1.0) for t,_ in x] for x in tokens]
unweighted_tokens = {}
if "h" in empty_tokens.keys():
unweighted_tokens["h"] = [[(t, 1.0) for t,_ in x] for x in tokens["h"]]
clip_name = "h"
else:
unweighted_tokens["l"] = [[(t, 1.0) for t,_ in x] for x in tokens["l"]]
clip_name = "l"
unweighted_cond, unweighted_pooled = clip.encode_from_tokens(unweighted_tokens, return_pooled=True)
cond = torch.clone(unweighted_cond)
for i in range(unweighted_cond.shape[0]):
for j in range(unweighted_cond.shape[1]):
weight = tokens[(j//77)][(j%77)][1]
weight = tokens[clip_name][(j//77)][(j%77)][1]
if weight != 1.0:
token_vector = unweighted_cond[i][j]
zero_vector = empty_cond[0][(j%77)]
@@ -72,9 +86,13 @@ class CLIPTextEncodePerpWeight:
if weight > 1.0:
cond[i][j] = token_vector + (weight * perp)
elif (weight > 0.0) and (weight < 1.0):
cond[i][j] = token_vector - ((1 - weight) * perp)
elif (weight < 0.0):
cond[i][j] = token_vector + (weight * perp)
cond[i][j] = weight * token_vector
elif (weight > -1.0) and (weight_g < 0.0):
cond[i][j] = -1 * (zero_vector + (abs(weight) * (token_vector - zero_vector)))
elif weight == -1.0:
cond[i][j] = -1 * token_vector
elif weight < -1.0:
cond[i][j] = -1 * (token_vector + (abs(weight) * perp))
elif weight == 0.0:
cond[i][j] = empty_cond[0][(j%77)]