From 08ecb7dbbdf3f51643256b2795c6fc946eb3155d Mon Sep 17 00:00:00 2001 From: Midcoastal Date: Wed, 3 Jan 2024 18:19:34 -0500 Subject: [PATCH] Embeddings Load Refactor Allow loading Embeddings in batches, as opposed to 1 at a time. Loading 1 at a time (as before) causes repeated Trie rebuilds, which is expensive. --- modules/patches.py | 22 +- modules/textual_inversion/loaders.py | 143 ++++++++++ .../textual_inversion/textual_inversion.py | 262 ++++++++++++------ 3 files changed, 340 insertions(+), 87 deletions(-) create mode 100644 modules/textual_inversion/loaders.py diff --git a/modules/patches.py b/modules/patches.py index cff6bfd64..b773a8f90 100644 --- a/modules/patches.py +++ b/modules/patches.py @@ -1,7 +1,8 @@ from collections import defaultdict +from typing import Optional -def patch(key, obj, field, replacement): +def patch(key, obj, field, replacement, add_if_not_exists:bool = False): """Replaces a function in a module or a class. Also stores the original function in this module, possible to be retrieved via original(key, obj, field). @@ -21,7 +22,10 @@ def patch(key, obj, field, replacement): if patch_key in originals[key]: raise RuntimeError(f"patch for {field} is already applied") - original_func = getattr(obj, field) + if not hasattr(obj, field) and not add_if_not_exists: + raise AttributeError(f"type {type(obj)} '{type.__name__}' has no attribute '{field}'") + + original_func = getattr(obj, field, None) originals[key][patch_key] = original_func setattr(obj, field, replacement) @@ -49,6 +53,8 @@ def undo(key, obj, field): raise RuntimeError(f"there is no patch for {field} to undo") original_func = originals[key].pop(patch_key) + if original_func is None: + delattr(obj, field) setattr(obj, field, original_func) return None @@ -60,4 +66,16 @@ def original(key, obj, field): return originals[key].get(patch_key, None) +def patch_method(cls, key:Optional[str]=None): + def decorator(func): + patch(func.__module__ if key is None else key, cls, func.__name__, func) + return decorator + + +def add_method(cls, key:Optional[str]=None): + def decorator(func): + patch(func.__module__ if key is None else key, cls, func.__name__, func, True) + return decorator + + originals = defaultdict(dict) diff --git a/modules/textual_inversion/loaders.py b/modules/textual_inversion/loaders.py new file mode 100644 index 000000000..81ed233bf --- /dev/null +++ b/modules/textual_inversion/loaders.py @@ -0,0 +1,143 @@ + +import torch +from modules import shared +from modules.patches import patch_method +from typing import Dict, List, Optional, Union +from diffusers.loaders.textual_inversion import TextualInversionLoaderMixin, logger, nn, load_textual_inversion_state_dicts +from transformers import PreTrainedTokenizer, PreTrainedModel + +try: + from accelerate.hooks import AlignDevicesHook, CpuOffload, remove_hook_from_module +except: + pass + +#debug = shared.log.env('SD_LOAD_TI_DEBUG').prefix(f'[{__name__}]') + +@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, # noqa: F821 + text_encoder: Optional["PreTrainedModel"] = None, # noqa: F821 + **kwargs, +): + #_debug = debug.prefix('pipe.load_textual_inversion ') + #_debug(f'processing {len(pretrained_model_name_or_path)} Embeddings') + #_debug.debug(f'pipe.load_textual_inversion: {pretrained_model_name_or_path}') + + # 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( + # 4. Retrieve tokens and embeddings + *TextualInversionLoaderMixin._retrieve_tokens_and_embeddings( + [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 + ) + ) + } + #_debug(f'Token `{token}` has {len(_embedding_data)} Shapes: \n{[ st for st in _embedding_data.keys()]}') + # 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: + 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) + is_sequential_cpu_offload = isinstance(getattr(component, "_hf_hook"), AlignDevicesHook) + logger.info( + "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 + _initial_tokenizer_size = len(tokenizer) + tokens_to_add = [token for token in token_data] + tokens_added = tokenizer.add_tokens(tokens_to_add) + tokenizer_size = len(tokenizer) + + #_debug(f'Added {tokens_added} Tokens: tokenizer_size={tokenizer_size}, _initial_tokenizer_size={_initial_tokenizer_size}') + #_debug.debug(f'Added Tokens: tokens={tokens_to_add}') + + # 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, 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 `{token}` does not resolve to a new Token ID') + embedding = token_data[token][0] + path = token_data[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 [ name_or_path for name_or_path in loaded_model_names_or_paths ] if number_of_models != 1 else None + # / Unsafe Code > \ No newline at end of file diff --git a/modules/textual_inversion/textual_inversion.py b/modules/textual_inversion/textual_inversion.py index 0688e37be..0af6e596f 100644 --- a/modules/textual_inversion/textual_inversion.py +++ b/modules/textual_inversion/textual_inversion.py @@ -15,6 +15,10 @@ 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.modelloader import directory_files, extension_filter, directory_mtime +from typing import List, Optional, Union +import modules.textual_inversion.loaders + +TokenToAdd = namedtuple("TokenToAdd", ["clip_l", "clip_g"]) TextualInversionTemplate = namedtuple("TextualInversionTemplate", ["name", "path"]) textual_inversion_templates = {} @@ -29,6 +33,12 @@ def list_textual_inversion_templates(): return textual_inversion_templates +def list_embeddings(*dirs): + is_ext = extension_filter(['.SAFETENSORS', '.PT' ] + ( ['.PNG', '.WEBP', '.JXL', '.AVIF', '.BIN' ] if shared.backend != shared.Backend.DIFFUSERS else [] )) + is_not_preview = lambda fp: not next(iter(os.path.splitext(fp))).upper().endswith('.PREVIEW') # pylint: disable=unnecessary-lambda-assignment + return list(filter(lambda fp: is_ext(fp) and is_not_preview(fp) and os.stat(fp).st_size > 0, directory_files(*dirs))) + + class Embedding: def __init__(self, vec, name, filename=None, step=None): self.vec = vec @@ -127,83 +137,163 @@ class EmbeddingDatabase: return 0 vec = shared.sd_model.cond_stage_model.encode_embedding_init_text(",", 1) return vec.shape[1] - - def load_diffusers_embedding(self, filename: str, path: str): - if shared.sd_model is None: - return - fn, ext = os.path.splitext(filename) - if ext.lower() != ".pt" and ext.lower() != ".safetensors": - return - pipe = shared.sd_model - name = os.path.basename(fn) - embedding = Embedding(vec=None, name=name, filename=path) - if not hasattr(pipe, "tokenizer") or not hasattr(pipe, 'text_encoder'): - self.skipped_embeddings[name] = embedding - return - try: - is_xl = hasattr(pipe, 'text_encoder_2') - try: - if not is_xl: # only use for sd15/sd21 - pipe.load_textual_inversion(path, token=name, cache_dir=shared.opts.diffusers_dir, local_files_only=True) - self.register_embedding(embedding, shared.sd_model) - except Exception: - pass - 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: - embeddings_dict = {} - if ext.lower() in ['.safetensors']: - with safetensors.torch.safe_open(path, framework="pt") as f: - for k in f.keys(): - embeddings_dict[k] = f.get_tensor(k) + + 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 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: - raise NotImplementedError - """ - # 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) - """ - 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.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 - clip_l.resize_token_embeddings(len(pipe.tokenizer)) - for i in range(len(token_ids)): - clip_l.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_l"][i] - elif is_xl: - pipe.tokenizer_2.add_tokens(tokens) - clip_l.resize_token_embeddings(len(pipe.tokenizer)) - clip_g.resize_token_embeddings(len(pipe.tokenizer)) - for i in range(len(token_ids)): - 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 - except Exception: - self.skipped_embeddings[name] = embedding + 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 as e: + 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 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 Exception(f'token `{name}` already in Model Vocabulary') + if name in tokens_to_add or name in loaded_embeddings: + raise Exception('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 Exception(f'extension {ext} not supported') + continue + if 'clip_l' not in embeddings_dict: + raise ValueError(f'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') + did_add = False + _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: + continue + if len(tokens_to_add) > 0: + _tokenizer_len = len(tokenizer) + num_added = tokenizer.add_tokens([k for k in tokens_to_add.keys()]) + clip_l.resize_token_embeddings(len(tokenizer)) + if model_type == 'SD-XL': + tokenizer_2.add_tokens([k for k in 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 + except Exception as e: + errors.display(e, f'Embedding Load Failure') + for name, embedding in loaded_embeddings.items(): + if not embedding: + continue + self.register_embedding(embedding, shared.sd_model) + if embedding in embeddings_to_load: + embeddings_to_load.remove(embedding) + skipped_embeddings.extend(embeddings_to_load) + for embedding in skipped_embeddings: + if loaded_embeddings.get(embedding.name, None) == embedding: + continue + self.skipped_embeddings[embedding.name] = embedding + return len(self.word_embeddings) - _loaded_pre + def load_from_file(self, path, filename): name, ext = os.path.splitext(filename) ext = ext.upper() @@ -265,17 +355,19 @@ class EmbeddingDatabase: return if not os.path.isdir(embdir.path): return - is_ext = extension_filter(['.PNG', '.WEBP', '.JXL', '.AVIF', '.BIN', '.PT', '.SAFETENSORS']) - is_not_preview = lambda fp: not next(iter(os.path.splitext(fp))).upper().endswith('.PREVIEW') # pylint: disable=unnecessary-lambda-assignment - for file_path in [*filter(lambda fp: is_ext(fp) and is_not_preview(fp), directory_files(embdir.path))]: - try: - if os.stat(file_path).st_size == 0: + file_paths = list_embeddings(embdir.path) + if shared.backend == shared.Backend.DIFFUSERS: + self.load_diffusers_embedding(file_paths) + 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}') 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}') - continue def load_textual_inversion_embeddings(self, force_reload=False): if shared.sd_model is None: