fix embeddings refresh

This commit is contained in:
Vladimir Mandic
2024-02-03 12:06:10 -05:00
parent 4fa23eed29
commit 2d75e37e1b
2 changed files with 97 additions and 120 deletions
+3 -3
View File
@@ -41,7 +41,7 @@ def control_run(units: List[unit.Unit], inputs, inits, mask, unit_type: str, is_
denoising_strength, batch_count, batch_size,
video_skip_frames, video_type, video_duration, video_loop, video_pad, video_interpolate,
ip_adapter, ip_scale, ip_image,
*input_script_args
*input_script_args # pylint: disable=unused-argument
):
global pipe, original_pipeline # pylint: disable=global-statement
debug(f'Control: type={unit_type} input={inputs} init={inits} type={input_type}')
@@ -403,10 +403,10 @@ def control_run(units: List[unit.Unit], inputs, inits, mask, unit_type: str, is_
return msg
elif unit_type == 'controlnet' and input_type == 1: # Init image same as control
p.init_images = input_image
p.task_args['control_image'] = p.override or input_image
p.task_args['control_image'] = p.override or input_image # TODO multiple units
p.task_args['strength'] = p.denoising_strength
elif unit_type == 'controlnet' and input_type == 2: # Separate init image
p.task_args['control_image'] = init_image
p.task_args['control_image'] = init_image # TODO multiple units
p.task_args['strength'] = init_image
if init_image is None:
shared.log.warning('Control: separate init image not provided')
+94 -117
View File
@@ -153,127 +153,105 @@ class EmbeddingDatabase:
vec = shared.sd_model.cond_stage_model.encode_embedding_init_text(",", 1)
return vec.shape[1]
def load_diffusers_embedding(
self,
filename: Union[str, List[str]],
path: Optional[Union[str, List[str]]] = None,
):
def load_diffusers_embedding(self, filename: Union[str, List[str]], path: Optional[Union[str, List[str]]] = None):
_loaded_pre = len(self.word_embeddings)
embeddings_to_load = []
loaded_embeddings = {}
skipped_embeddings = []
if shared.sd_model is None:
return 0
pipe = shared.sd_model
tokenizer = getattr(pipe, 'tokenizer', None)
tokenizer_2 = getattr(pipe, 'tokenizer_2', None)
clip_l = getattr(pipe, 'text_encoder', None) # clip_l
clip_g = getattr(pipe, 'text_encoder_2', None) # clip_g
if clip_l is None or tokenizer is None:
return 0
filenames = (
[filename]
if not isinstance(filename, list)
else filename
)
exts = [".SAFETENSORS", '.BIN', '.PT', '.PNG', '.WEBP', '.JXL', '.AVIF'] # SDXL only uses safetensors
filename_paths = zip(filenames, len(filenames) * [path] if (isinstance(path, str) or path is None) else path)
model_type = None
if clip_g is None and tokenizer_2 is None:
model_type = 'SD'
elif clip_g and tokenizer_2:
tokenizer = getattr(shared.sd_model, 'tokenizer', None)
tokenizer_2 = getattr(shared.sd_model, 'tokenizer_2', None)
clip_l = getattr(shared.sd_model, 'text_encoder', None)
clip_g = getattr(shared.sd_model, 'text_encoder_2', None)
if clip_g and tokenizer_2:
model_type = 'SDXL'
elif clip_l and tokenizer:
model_type = 'SD'
else:
model_type = 'UNDEFINED'
try:
unk_token_id = tokenizer.convert_tokens_to_ids(tokenizer.unk_token)
for _filename, _path in filename_paths:
if _path is None:
_path = _filename
_filename = os.path.basename(_path)
fn, ext = os.path.splitext(_filename)
name = os.path.basename(fn)
embedding = Embedding(vec=None, name=name, filename=_path)
try:
ext = ext.upper()
_, _ext = os.path.splitext(_path)
_ext = _ext.upper()
if ext != _ext:
raise ValueError(f'filename and path extensions do not match: `{ext}` != `{_ext}`')
if ext not in exts:
raise ValueError(f'extension `{ext}` is invalid, expected one of: {exts}')
if name in tokenizer.get_vocab() or f"{name}_1" in tokenizer.get_vocab():
raise ValueError(f'token already exists in the tokenizer vocabulary: `{name}`')
embeddings_to_load.append(embedding)
except Exception:
skipped_embeddings.append(embedding)
continue
embeddings_to_load = sorted(embeddings_to_load, key=lambda e: exts.index(os.path.splitext(e.filename)[1].upper()))
tokens_to_add = {}
tokenizer_vocab = tokenizer.get_vocab()
for embedding in embeddings_to_load:
try:
name = embedding.name
if name in tokenizer_vocab:
raise UserWarning(f'token `{name}` already in Model Vocabulary')
if name in tokens_to_add or name in loaded_embeddings:
raise UserWarning('duplicate Embedding Token')
embeddings_dict = {}
_, ext = os.path.splitext(embedding.filename)
ext = ext.upper()
if ext in ['.SAFETENSORS']:
with safetensors.torch.safe_open(embedding.filename, framework="pt") as f: # type: ignore
for k in f.keys():
embeddings_dict[k] = f.get_tensor(k)
else: # fallback for sd1.5 pt embeddings
embeddings_dict["clip_l"] = self.load_from_file(embedding.filename, embedding.filename)
if 'clip_l' not in embeddings_dict:
raise ValueError('Invalid Embedding, dict missing required key `clip_l`')
if 'clip_g' not in embeddings_dict and model_type == "SDXL" and shared.opts.diffusers_convert_embed:
embeddings_dict["clip_g"] = convert_embedding(embeddings_dict["clip_l"], clip_l, clip_g)
if 'clip_g' in embeddings_dict:
embedding_type = 'SDXL'
else:
embedding_type = 'SD'
if embedding_type != model_type:
raise ValueError(f'Unable to load {embedding_type} Embedding "{embedding.name}" into {model_type} Model')
_tokens_to_add = {}
for i in range(len(embeddings_dict["clip_l"])):
if len(clip_l.get_input_embeddings().weight.data[0]) == len(embeddings_dict["clip_l"][i]):
token = name if i == 0 else f"{name}_{i}"
if token in tokenizer_vocab:
raise RuntimeError(f'Multi-Vector Embedding would add pre-existing Token in Vocabulary: {token}')
if token in tokens_to_add:
raise RuntimeError(f'Multi-Vector Embedding would add duplicate Token to Add: {token}')
_tokens_to_add[token] = TokenToAdd(
embeddings_dict["clip_l"][i],
embeddings_dict["clip_g"][i] if 'clip_g' in embeddings_dict else None
)
if not _tokens_to_add:
raise ValueError('no valid tokens to add')
tokens_to_add.update(_tokens_to_add)
return 0
filenames = [filename] if not isinstance(filename, list) else filename
exts = [".SAFETENSORS", '.BIN', '.PT', '.PNG', '.WEBP', '.JXL', '.AVIF'] # SDXL only uses safetensors
filename_paths = zip(filenames, len(filenames) * [path] if (isinstance(path, str) or path is None) else path)
unk_token_id = tokenizer.convert_tokens_to_ids(tokenizer.unk_token)
for filename, fullname in filename_paths:
debug(f'Embedding check: {filename}')
if fullname is None:
fullname = filename
filename = os.path.basename(fullname)
fn, ext = os.path.splitext(filename)
name = os.path.basename(fn)
embedding = Embedding(vec=None, name=name, filename=fullname)
try:
if ext.upper() not in exts:
raise ValueError(f'extension `{ext}` is invalid, expected one of: {exts}')
if name in tokenizer.get_vocab() or f"{name}_1" in tokenizer.get_vocab():
loaded_embeddings[name] = embedding
except Exception as e:
debug(f"TI Loading: {e}")
continue
if len(tokens_to_add) > 0:
tokenizer.add_tokens(list(tokens_to_add.keys()))
clip_l.resize_token_embeddings(len(tokenizer))
if model_type == 'SDXL':
tokenizer_2.add_tokens(list(tokens_to_add.keys())) # type: ignore
clip_g.resize_token_embeddings(len(tokenizer_2)) # type: ignore
for token, data in tokens_to_add.items():
token_id = tokenizer.convert_tokens_to_ids(token)
if token_id > unk_token_id:
clip_l.get_input_embeddings().weight.data[token_id] = data.clip_l
if model_type == 'SDXL':
clip_g.get_input_embeddings().weight.data[token_id] = data.clip_g # type: ignore
except Exception as e:
errors.display(e, 'Embedding Load Failure')
debug(f'Embedding already loaded: {name}')
embeddings_to_load.append(embedding)
except Exception as e:
skipped_embeddings.append(embedding)
debug(f'Embedding check: {name} {e}')
continue
embeddings_to_load = sorted(embeddings_to_load, key=lambda e: exts.index(os.path.splitext(e.filename)[1].upper()))
tokens_to_add = {}
tokenizer_vocab = tokenizer.get_vocab()
for embedding in embeddings_to_load:
try:
debug(f'Embedding load: {embedding.name} file={embedding.filename}')
if embedding.name in tokens_to_add or embedding.name in loaded_embeddings:
raise ValueError('duplicate token')
embeddings_dict = {}
_, ext = os.path.splitext(embedding.filename)
if ext.upper() in ['.SAFETENSORS']:
with safetensors.torch.safe_open(embedding.filename, framework="pt") as f: # type: ignore
for k in f.keys():
embeddings_dict[k] = f.get_tensor(k)
else: # fallback for sd1.5 pt embeddings
embeddings_dict["clip_l"] = self.load_from_file(embedding.filename, embedding.filename)
if 'clip_l' not in embeddings_dict:
raise ValueError('Invalid Embedding, dict missing required key `clip_l`')
if 'clip_g' not in embeddings_dict and model_type == "SDXL" and shared.opts.diffusers_convert_embed:
embeddings_dict["clip_g"] = convert_embedding(embeddings_dict["clip_l"], clip_l, clip_g)
if 'clip_g' in embeddings_dict:
embedding_type = 'SDXL'
else:
embedding_type = 'SD'
if embedding_type != model_type:
raise ValueError(f'Unable to load {embedding_type} Embedding "{embedding.name}" into {model_type} Model')
_tokens_to_add = {}
for i in range(len(embeddings_dict["clip_l"])):
if len(clip_l.get_input_embeddings().weight.data[0]) == len(embeddings_dict["clip_l"][i]):
token = embedding.name if i == 0 else f"{embedding.name}_{i}"
if token in tokenizer_vocab:
raise RuntimeError(f'Multi-Vector Embedding would add pre-existing Token in Vocabulary: {token}')
if token in tokens_to_add:
raise RuntimeError(f'Multi-Vector Embedding would add duplicate Token to Add: {token}')
_tokens_to_add[token] = TokenToAdd(
embeddings_dict["clip_l"][i],
embeddings_dict["clip_g"][i] if 'clip_g' in embeddings_dict else None
)
if not _tokens_to_add:
raise ValueError('no valid tokens to add')
tokens_to_add.update(_tokens_to_add)
loaded_embeddings[name] = embedding
except Exception as e:
debug(f"Embedding loading: {embedding.filename} {e}")
continue
if len(tokens_to_add) > 0:
tokenizer.add_tokens(list(tokens_to_add.keys()))
clip_l.resize_token_embeddings(len(tokenizer))
if model_type == 'SDXL':
tokenizer_2.add_tokens(list(tokens_to_add.keys())) # type: ignore
clip_g.resize_token_embeddings(len(tokenizer_2)) # type: ignore
for token, data in tokens_to_add.items():
token_id = tokenizer.convert_tokens_to_ids(token)
if token_id > unk_token_id:
clip_l.get_input_embeddings().weight.data[token_id] = data.clip_l
if model_type == 'SDXL':
clip_g.get_input_embeddings().weight.data[token_id] = data.clip_g # type: ignore
for embedding in loaded_embeddings.values():
if not embedding:
@@ -287,11 +265,12 @@ class EmbeddingDatabase:
continue
self.skipped_embeddings[embedding.name] = embedding
try:
debug(f"TI Loading: Text Encoder total embeddings={shared.sd_model.text_encoder.get_input_embeddings().weight.data.shape[0]}")
if model_type == 'SD':
debug(f"Embeddings loaded: text-encoder={shared.sd_model.text_encoder.get_input_embeddings().weight.data.shape[0]}")
if model_type == 'SDXL':
debug(f"Embeddings loaded: text-encoder-1={shared.sd_model.text_encoder.get_input_embeddings().weight.data.shape[0]} text-encoder-2={shared.sd_model.text_encoder_2.get_input_embeddings().weight.data.shape[0]}")
except Exception:
pass
if model_type == 'SDXL':
debug(f"TI Loading: Text Encoder 2 total embeddings={shared.sd_model.text_encoder_2.get_input_embeddings().weight.data.shape[0]}")
return len(self.word_embeddings) - _loaded_pre
def load_from_file(self, path, filename):
@@ -361,12 +340,10 @@ class EmbeddingDatabase:
else:
for file_path in file_paths:
try:
if os.stat(file_path).st_size == 0:
continue
fn = os.path.basename(file_path)
self.load_from_file(file_path, fn)
except Exception as e:
errors.display(e, f'embedding load {fn}')
errors.display(e, f'Load embeding={fn}')
continue
def load_textual_inversion_embeddings(self, force_reload=False):