mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 08:44:33 +02:00
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.
This commit is contained in:
+20
-2
@@ -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)
|
||||
|
||||
@@ -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 >
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user