one more sdxl embeddings fix

This commit is contained in:
Vladimir Mandic
2023-09-28 16:03:18 -04:00
parent a785f3da2e
commit e830b41dfa
3 changed files with 27 additions and 20 deletions
+1 -1
View File
@@ -201,7 +201,7 @@ def process_diffusers(p: StableDiffusionProcessing, seeds, prompts, negative_pro
prompt_embed, pooled, negative_embed, negative_pooled = prompt_parser_diffusers.compel_encode_prompts(model, prompts, negative_prompts, prompts_2, negative_prompts_2, is_refiner, kwargs.pop("clip_skip", None))
parser = shared.opts.prompt_attention
except Exception as e:
shared.log.error(f'Prompt parser: {e}')
shared.log.error(f'Prompt parser encode: {e}')
if 'prompt' in possible:
if hasattr(model, 'text_encoder') and 'prompt_embeds' in possible and prompt_embed is not None:
if type(pooled) == list:
+7 -8
View File
@@ -3,8 +3,7 @@ import typing
import torch
from compel import Compel, ReturnedEmbeddingsType
from compel.embeddings_provider import BaseTextualInversionManager
import modules.shared as shared
import modules.prompt_parser as prompt_parser
from modules import shared, devices, prompt_parser
debug_output = os.environ.get('SD_PROMPT_DEBUG', None)
@@ -142,20 +141,20 @@ def compel_encode_prompt(
prompt_2 = convert_to_compel(prompt_2)
negative_prompt_2 = convert_to_compel(negative_prompt_2)
textual_inversion_manager = DiffusersTextualInversionManager(pipeline)
textual_inversion_manager_te1 = DiffusersTextualInversionManager(pipeline)
compel_te1 = Compel(
tokenizer=pipeline.tokenizer,
text_encoder=pipeline.text_encoder,
returned_embeddings_type=embedding_type,
requires_pooled=False,
# truncate_long_prompts=False,
device=shared.device,
textual_inversion_manager=textual_inversion_manager
device=devices.device,
textual_inversion_manager=textual_inversion_manager_te1
)
if 'XL' in pipeline.__class__.__name__ and not is_refiner:
compel_te2 = Compel(tokenizer=pipeline.tokenizer_2, text_encoder=pipeline.text_encoder_2, returned_embeddings_type=embedding_type, requires_pooled=True, device=shared.device, textual_inversion_manager=textual_inversion_manager)
# TODO textual_inversion_manager=textual_inversion_manager_te2 - DiffusersTextualInversionManager needs to use tokenizer_2
compel_te2 = Compel(tokenizer=pipeline.tokenizer_2, text_encoder=pipeline.text_encoder_2, returned_embeddings_type=embedding_type, requires_pooled=True, device=devices.device)
positive_te1 = compel_te1(prompt)
positive_te2, positive_pooled = compel_te2(prompt_2)
positive = torch.cat((positive_te1, positive_te2), dim=-1)
@@ -169,7 +168,7 @@ def compel_encode_prompt(
return prompt_embed, positive_pooled, negative_embed, negative_pooled
elif 'XL' in pipeline.__class__.__name__ and is_refiner:
compel_te2 = Compel(tokenizer=pipeline.tokenizer_2, text_encoder=pipeline.text_encoder_2, returned_embeddings_type=embedding_type, requires_pooled=True, device=shared.device)
compel_te2 = Compel(tokenizer=pipeline.tokenizer_2, text_encoder=pipeline.text_encoder_2, returned_embeddings_type=embedding_type, requires_pooled=True, device=devices.device)
positive, positive_pooled = compel_te2(prompt)
negative, negative_pooled = compel_te2(negative_prompt)
+19 -11
View File
@@ -107,7 +107,12 @@ class EmbeddingDatabase:
def register_embedding(self, embedding, model):
self.word_embeddings[embedding.name] = embedding
ids = model.cond_stage_model.tokenize([embedding.name])[0]
if hasattr(model, 'cond_stage_model'):
ids = model.cond_stage_model.tokenize([embedding.name])[0]
elif hasattr(model, 'tokenizer'):
ids = model.tokenizer.convert_tokens_to_ids(embedding.name)
if type(ids) != list:
ids = [ids]
first_id = ids[0]
if first_id not in self.ids_lookup:
self.ids_lookup[first_id] = []
@@ -141,7 +146,10 @@ class EmbeddingDatabase:
self.register_embedding(embedding, shared.sd_model)
except Exception:
pass
is_loaded = pipe.tokenizer.convert_tokens_to_ids(name) > 49407
is_loaded = pipe.tokenizer.convert_tokens_to_ids(name)
if type(is_loaded) != list:
is_loaded = [is_loaded]
is_loaded = is_loaded[0] > 49407
if is_loaded:
self.register_embedding(embedding, shared.sd_model)
else:
@@ -166,28 +174,28 @@ class EmbeddingDatabase:
for vec in tokens:
embeddings_dict['clip_l'].append(vec)
"""
clip_l = pipe.text_encoder.get_input_embeddings().weight if hasattr(pipe, 'text_encoder') and hasattr(pipe.text_encoder, "resize_token_embeddings") else None
clip_g = pipe.text_encoder_2.get_input_embeddings().weight if hasattr(pipe, 'text_encoder_2') and hasattr(pipe.text_encoder_2, "resize_token_embeddings") else None
clip_l = pipe.text_encoder if hasattr(pipe, 'text_encoder') else None
clip_g = pipe.text_encoder_2 if hasattr(pipe, 'text_encoder_2') else None
is_sd = clip_l is not None and 'clip_l' in embeddings_dict and clip_g is None and 'clip_g' not in embeddings_dict
is_xl = clip_l is not None and 'clip_l' in embeddings_dict and clip_g is not None and 'clip_g' in embeddings_dict
tokens = []
for i in range(len(embeddings_dict["clip_l"])):
if (is_sd or is_xl) and (len(clip_l.data[0]) == len(embeddings_dict["clip_l"][i])):
if (is_sd or is_xl) and (len(clip_l.get_input_embeddings().weight.data[0]) == len(embeddings_dict["clip_l"][i])):
tokens.append(name if i == 0 else f"{name}_{i}")
num_added = pipe.tokenizer.add_tokens(tokens)
if num_added > 0:
token_ids = pipe.tokenizer.convert_tokens_to_ids(tokens)
if is_sd: # only used for sd15 if load_textual_inversion failed and format is safetensors
pipe.text_encoder.resize_token_embeddings(len(pipe.tokenizer))
clip_l.resize_token_embeddings(len(pipe.tokenizer))
for i in range(len(token_ids)):
clip_l.data[token_ids[i]] = embeddings_dict["clip_l"][i]
clip_l.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_l"][i]
elif is_xl:
pipe.tokenizer_2.add_tokens(tokens)
pipe.text_encoder.resize_token_embeddings(len(pipe.tokenizer))
pipe.text_encoder_2.resize_token_embeddings(len(pipe.tokenizer))
clip_l.resize_token_embeddings(len(pipe.tokenizer))
clip_g.resize_token_embeddings(len(pipe.tokenizer))
for i in range(len(token_ids)):
pipe.text_encoder.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_l"][i]
pipe.text_encoder_2.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_g"][i]
clip_l.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_l"][i]
clip_g.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_g"][i]
self.register_embedding(embedding, shared.sd_model)
else:
raise NotImplementedError