Fix TI Loading

This commit is contained in:
AI-Casanova
2024-01-17 20:27:05 -06:00
committed by Vladimir Mandic
parent 65a976a6e7
commit 997644776c
2 changed files with 114 additions and 261 deletions
-132
View File
@@ -1,132 +0,0 @@
from typing import TYPE_CHECKING, Dict, List, Optional, Union
import torch
from torch import nn
from diffusers.loaders.textual_inversion import TextualInversionLoaderMixin, load_textual_inversion_state_dicts
from modules import shared
from modules.patches import patch_method
if TYPE_CHECKING:
from transformers import PreTrainedModel, PreTrainedTokenizer
try:
from accelerate.hooks import AlignDevicesHook, CpuOffload, remove_hook_from_module
except Exception:
pass
@patch_method(TextualInversionLoaderMixin)
def load_textual_inversion(
self: TextualInversionLoaderMixin,
pretrained_model_name_or_path: Union[str, List[str], Dict[str, torch.Tensor], List[Dict[str, torch.Tensor]]],
token: Optional[Union[str, List[str]]] = None,
tokenizer: Optional["PreTrainedTokenizer"] = None,
text_encoder: Optional["PreTrainedModel"] = None,
**kwargs, # pylint: disable=W0613
):
# 1. Set correct tokenizer and text encoder
tokenizer: PreTrainedTokenizer = tokenizer or getattr(self, "tokenizer", None)
text_encoder: PreTrainedModel = text_encoder or getattr(self, "text_encoder", None)
loaded_model_names_or_paths = {}
assert tokenizer and text_encoder, 'Can not resolve `tokenizer` or `text_encoder`'
# 2. Normalize inputs
pretrained_model_name_or_paths = (
[pretrained_model_name_or_path]
if not isinstance(pretrained_model_name_or_path, list)
else pretrained_model_name_or_path
)
tokens = len(pretrained_model_name_or_paths) * [token] if (isinstance(token, str) or token is None) else token
assert len(tokens) == len(pretrained_model_name_or_paths), f'Number of Models ({len(pretrained_model_name_or_paths)}) and Tokens ({len(tokens)}) must be equal.'
number_of_models = len(pretrained_model_name_or_paths)
token_data = {}
# Build a unique list of Shape-Token/Embedding-Shape pairs, mapped with an associated `name_or_path` for reporting
expected_emb_dim = text_encoder.get_input_embeddings().weight.shape[-1]
for idx in range(number_of_models):
try:
name_or_path = pretrained_model_name_or_paths[idx]
token = tokens[idx]
_embedding_data = {
shape_token: (embedding_shape, name_or_path )
for shape_token, embedding_shape in zip(
# 5. Extend tokens and embeddings for multi vector
*self._extend_tokens_and_embeddings( # pylint: disable=W0212
# 4. Retrieve tokens and embeddings
*TextualInversionLoaderMixin._retrieve_tokens_and_embeddings( # pylint: disable=W0212
[token],
# 3. Load state dicts of textual embeddings
load_textual_inversion_state_dicts(
[name_or_path],
cache_dir=shared.opts.diffusers_dir,
local_files_only=True
),
tokenizer
),
tokenizer
)
)
}
# 6. Make sure all embeddings have the correct size
for embedding_shape, _ in _embedding_data.values():
if expected_emb_dim != embedding_shape.shape[-1]:
#debug.error(f'Incorrect Shape: {embedding_shape.shape[-1]} vs {expected_emb_dim}')
raise ValueError(
"Loaded embeddings are of incorrect shape. Expected each textual inversion embedding "
"to be of shape {embedding_shape.shape[-1]}, but are {embeddings.shape[-1]} "
)
token_data.update(_embedding_data)
except Exception:
if number_of_models == 1:
raise
# 7. Now we can be sure that loading the embedding matrix works
# < Unsafe code:
# 7.1 Offload all hooks in case the pipeline was cpu offloaded before make sure, we offload and onload again
is_model_cpu_offload = False
is_sequential_cpu_offload = False
for _, component in self.components.items():
if isinstance(component, nn.Module):
if hasattr(component, "_hf_hook"):
is_model_cpu_offload = isinstance(getattr(component, "_hf_hook"), CpuOffload) # noqa: B009
is_sequential_cpu_offload = isinstance(getattr(component, "_hf_hook"), AlignDevicesHook) # noqa: B009
shared.log.debug("Accelerate hooks detected. Since you have called `load_textual_inversion()`, the previous hooks will be first removed. Then the textual inversion parameters will be loaded and the hooks will be applied again.")
remove_hook_from_module(component, recurse=is_sequential_cpu_offload)
# 7.2 save expected device and dtype
device = text_encoder.device
dtype = text_encoder.dtype
# 7.2 Add Tokens to the Tokenizer
tokens_to_add = list(token_data)
tokenizer.add_tokens(tokens_to_add)
tokenizer_size = len(tokenizer)
# 7.3 Increase token embedding matrix
text_encoder.resize_token_embeddings(tokenizer_size)
input_embeddings = text_encoder.get_input_embeddings().weight
unk_token_id = tokenizer.convert_tokens_to_ids(tokenizer.unk_token)
# 7.4 Load token and embedding
for token_id, load_token in zip(tokenizer.convert_tokens_to_ids(tokens_to_add), tokens_to_add):
if token_id <= unk_token_id:
raise RuntimeError(f'Processed Shape-Token `{load_token}` does not resolve to a new Token ID: {token_id} <= {unk_token_id} ({tokenizer.unk_token})')
embedding = token_data[load_token][0]
path = token_data[load_token][1]
input_embeddings.data[token_id] = embedding
loaded_model_names_or_paths[path] = True
input_embeddings.to(dtype=dtype, device=device)
# 7.5 Offload the model again
if is_model_cpu_offload:
self.enable_model_cpu_offload()
elif is_sequential_cpu_offload:
self.enable_sequential_cpu_offload()
return loaded_model_names_or_paths.keys() if number_of_models != 1 else None
# / Unsafe Code >
+114 -129
View File
@@ -11,12 +11,13 @@ import numpy as np
from PIL import Image, PngImagePlugin
from modules import shared, devices, processing, sd_models, images, errors
import modules.textual_inversion.dataset
import modules.textual_inversion.loaders
from modules.textual_inversion.learn_schedule import LearnRateScheduler
from modules.textual_inversion.image_embedding import embedding_to_b64, embedding_from_b64, insert_image_data_embed, extract_image_data_embed, caption_image_overlay
from modules.textual_inversion.ti_logging import save_settings_to_file
from modules.files_cache import directory_files, directory_mtime, extension_filter
debug = shared.log.trace if os.environ.get('SD_TI_DEBUG', None) is not None else lambda *args, **kwargs: None
debug('Trace: TEXTUAL INVERSION')
TokenToAdd = namedtuple("TokenToAdd", ["clip_l", "clip_g"])
TextualInversionTemplate = namedtuple("TextualInversionTemplate", ["name", "path"])
@@ -146,136 +147,117 @@ class EmbeddingDatabase:
embeddings_to_load = []
loaded_embeddings = {}
skipped_embeddings = []
if shared.sd_model is not None:
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
filenames = (
[filename]
if not isinstance(filename, list)
else filename
)
exts = [".SAFETENSORS", ".PT"]
filename_paths = zip(filenames, len(filenames) * [path] if (isinstance(path, str) or path is None) else path)
model_type = None
if clip_l and tokenizer:
if clip_g is None and tokenizer_2 is None:
model_type = 'SD'
elif clip_g and tokenizer_2:
model_type = 'SD-XL'
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()))
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 and tokenizer is None:
return 0
if model_type == 'SD':
loaded_filenames = pipe.load_textual_inversion(
[embedding.filename for embedding in embeddings_to_load],
token=[embedding.name for embedding in embeddings_to_load],
tokenizer=tokenizer,
text_encoder=clip_l
)
_len = len(embeddings_to_load)
for embedding in embeddings_to_load.copy():
if embedding.filename in loaded_filenames:
loaded_embeddings[embedding.name] = embedding
embeddings_to_load.remove(embedding)
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)
"""
# The following note has been here a while (as of 11/05/23), go or no-go?
# alternatively could disable load_textual_inversion and load everything here
elif ext.lower() in ['.PT', '.BIN']:
data = torch.load(path, map_location="cpu")
embedding.tag = data.get('name', None)
embedding.step = data.get('step', None)
embedding.sd_checkpoint = data.get('sd_checkpoint', None)
embedding.sd_checkpoint_name = data.get('sd_checkpoint_name', None)
param_dict = data.get('string_to_param', None)
embeddings_dict['clip_l'] = []
for tokens in param_dict.values():
for vec in tokens:
embeddings_dict['clip_l'].append(vec)
"""
else:
raise NotImplementedError(f'extension {ext} not supported')
if 'clip_l' not in embeddings_dict:
raise ValueError('Invalid Embedding, dict missing required key `clip_l`')
if 'clip_g' in embeddings_dict:
embedding_type = 'SD-XL'
else:
embedding_type = 'SD'
if embedding_type != model_type:
raise ValueError(f'Unable to load `{embedding_type}` Embedding 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)
loaded_embeddings[name] = embedding
except Exception:
continue
if len(tokens_to_add) > 0:
clip_l.resize_token_embeddings(len(tokenizer))
if model_type == 'SD-XL':
tokenizer_2.add_tokens(list(tokens_to_add.keys())) # type: ignore
clip_g.resize_token_embeddings(len(tokenizer)) # 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 == 'SD-XL':
clip_g.get_input_embeddings().weight.data[token_id] = data.clip_g # type: ignore
filenames = (
[filename]
if not isinstance(filename, list)
else filename
)
exts = [".SAFETENSORS", ".PT"]
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:
model_type = 'SD-XL'
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:
raise NotImplementedError(f'extension {ext} not supported')
if 'clip_l' not in embeddings_dict:
raise ValueError('Invalid Embedding, dict missing required key `clip_l`')
if 'clip_g' in embeddings_dict:
embedding_type = 'SD-XL'
else:
embedding_type = 'SD'
if embedding_type != model_type:
raise ValueError(f'Unable to load `{embedding_type}` Embedding 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)
loaded_embeddings[name] = embedding
except Exception as e:
errors.display(e, 'Embedding Load Failure')
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 == 'SD-XL':
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 == 'SD-XL':
clip_g.get_input_embeddings().weight.data[token_id] = data.clip_g # type: ignore
except Exception as e:
errors.display(e, 'Embedding Load Failure')
for embedding in loaded_embeddings.values():
if not embedding:
continue
@@ -287,6 +269,9 @@ class EmbeddingDatabase:
if loaded_embeddings.get(embedding.name, None) == embedding:
continue
self.skipped_embeddings[embedding.name] = embedding
debug(f"TI Loading: Text Encoder total embeddings={shared.sd_model.text_encoder.get_input_embeddings().weight.data.shape[0]}")
if model_type == 'SD-XL':
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):