Proper fix for 999999 issue: Don't pad T5.

This commit is contained in:
zer0int
2024-08-28 15:57:58 +02:00
committed by GitHub
parent 7d57f1c71b
commit cb0f859b76
+10 -31
View File
@@ -211,51 +211,34 @@ class SDLongTokenizer:
pad_token = self.end_token
else:
pad_token = 0
from comfy.sd1_clip import token_weights,escape_important,unescape_important
text = escape_important(text)
parsed_weights = token_weights(text, 1.0)
tokens = []
total_token_count = 0
for weighted_segment, weight in parsed_weights:
to_tokenize = unescape_important(weighted_segment).replace("\n", " ").split(' ')
to_tokenize = [x for x in to_tokenize if x != ""]
for word in to_tokenize:
if total_token_count >= self.max_length:
print(f"Reached max length of {self.max_length} tokens.")
break
#if we find an embedding, deal with the embedding
if word.startswith(self.embedding_identifier) and self.embedding_directory is not None:
embedding_name = word[len(self.embedding_identifier):].strip('\n')
embed, leftover = self._try_get_embedding(embedding_name)
if embed is None:
print(f"Warning: embedding '{embedding_name}' does not exist, ignoring.")
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
else:
if len(embed.shape) == 1:
tokens.append([(embed, weight)])
else:
tokens.append([(embed[x], weight) for x in range(embed.shape[0])])
total_token_count += len(tokens[-1]) # Add count of the newly added tokens
#if we accidentally have leftover text, continue parsing using leftover, else move on to next word
if leftover != "":
word = leftover
else:
continue
# Tokenize word and add to tokens
word_tokens = [(t, weight) for t in self.tokenizer(word)[0][self.tokens_start:-1]]
if total_token_count + len(word_tokens) > self.max_length:
# Truncate the current group of tokens to fit within the max length
word_tokens = word_tokens[:self.max_length - total_token_count]
tokens.append(word_tokens)
total_token_count += len(word_tokens)
# Debug: Print the final token structure and counts
#print(f"Final token structure: {tokens}")
#print(f"Total token count: {total_token_count}")
#parse word
tokens.append([(t, weight) for t in self.tokenizer(word)[0][self.tokens_start:-1]])
#reshape token array to CLIP input size
batched_tokens = []
@@ -297,9 +280,6 @@ class SDLongTokenizer:
if not return_word_ids:
batched_tokens = [[(t, w) for t, w,_ in x] for x in batched_tokens]
if len(tokens) > self.max_length:
print("Warning: Token sequence exceeds max length") # This should not happen due to truncation.
return batched_tokens
@@ -401,7 +381,7 @@ class LongCLIPFluxModel(torch.nn.Module):
# Encode using T5XXL
t5_out, t5_pooled = self.t5xxl.encode_token_weights(token_weight_pairs_t5)
return t5_out, l_pooled # Adjust as per the output requirement of the diffusion model
return t5_out, l_pooled
def load_sd(self, sd):
if "text_model.encoder.layers.1.mlp.fc1.weight" in sd:
@@ -412,7 +392,7 @@ class LongCLIPFluxModel(torch.nn.Module):
class LongCLIPFluxTokenizer:
def __init__(self):
self.clip_l = None
self.clip_g = None
self.t5xxl = None
def tokenize_with_weights(self, text: str, return_word_ids=False):
# Tokenize with both Long-CLIP and T5XXL
@@ -420,14 +400,13 @@ class LongCLIPFluxTokenizer:
out["l"] = self.clip_l.tokenize_with_weights(text, return_word_ids) # Long-CLIP tokenization
out["t5xxl"] = self.t5xxl.tokenize_with_weights(text, return_word_ids) # T5XXL tokenization
# Check the number of tokens and pad the shorter sequence if necessary
# Check the number of tokens
l_tokens = token_num(out["l"])
t5_tokens = token_num(out["t5xxl"])
# Leaving this here as a reminder: Do NOT pad T5XXL!
if l_tokens > t5_tokens:
out["t5xxl"] = pad_tokens(out["t5xxl"], self.t5xxl, l_tokens - t5_tokens)
elif t5_tokens > l_tokens:
out["l"] = pad_tokens(out["l"], self.clip_l, t5_tokens - l_tokens)
pass # Do not pad T5XXL
return out