mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
sort out embeddings loading in dffusers
This commit is contained in:
@@ -14,7 +14,7 @@ from einops import rearrange, repeat
|
||||
from ldm.util import default
|
||||
from modules import devices, processing, sd_models, shared, sd_samplers, hashes, sd_hijack_checkpoint, errors
|
||||
import modules.textual_inversion.dataset
|
||||
from modules.textual_inversion import textual_inversion, logging
|
||||
from modules.textual_inversion import textual_inversion, ti_logging
|
||||
from modules.textual_inversion.learn_schedule import LearnRateScheduler
|
||||
|
||||
|
||||
@@ -438,13 +438,13 @@ def statistics(data):
|
||||
std = 0
|
||||
else:
|
||||
std = stdev(data)
|
||||
total_information = f"loss:{mean(data):.3f}" + u"\u00B1" + f"({std/ (len(data) ** 0.5):.3f})"
|
||||
total_information = f"loss:{mean(data):.3f}" + "\u00B1" + f"({std/ (len(data) ** 0.5):.3f})"
|
||||
recent_data = data[-32:]
|
||||
if len(recent_data) < 2:
|
||||
std = 0
|
||||
else:
|
||||
std = stdev(recent_data)
|
||||
recent_information = f"recent 32 loss:{mean(recent_data):.3f}" + u"\u00B1" + f"({std / (len(recent_data) ** 0.5):.3f})"
|
||||
recent_information = f"recent 32 loss:{mean(recent_data):.3f}" + "\u00B1" + f"({std / (len(recent_data) ** 0.5):.3f})"
|
||||
return total_information, recent_information
|
||||
|
||||
|
||||
@@ -557,7 +557,7 @@ def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradi
|
||||
model_name=checkpoint.model_name, model_hash=checkpoint.shorthash, num_of_dataset_images=len(ds),
|
||||
**{field: getattr(hypernetwork, field) for field in ['layer_structure', 'activation_func', 'weight_init', 'add_layer_norm', 'use_dropout', ]}
|
||||
)
|
||||
logging.save_settings_to_file(log_directory, {**saved_params, **locals()})
|
||||
ti_logging.save_settings_to_file(log_directory, {**saved_params, **locals()})
|
||||
|
||||
latent_sampling_method = ds.latent_sampling_method
|
||||
|
||||
|
||||
@@ -42,9 +42,7 @@ def get_crop_region(mask, pad=0):
|
||||
def expand_crop_region(crop_region, processing_width, processing_height, image_width, image_height):
|
||||
"""expands crop region get_crop_region() to match the ratio of the image the region will processed in; returns expanded region
|
||||
for example, if user drew mask in a 128x32 region, and the dimensions for processing are 512x512, the region will be expanded to 128x128."""
|
||||
|
||||
x1, y1, x2, y2 = crop_region
|
||||
|
||||
ratio_crop_region = (x2 - x1) / (y2 - y1)
|
||||
ratio_processing = processing_width / processing_height
|
||||
|
||||
@@ -82,17 +80,12 @@ def expand_crop_region(crop_region, processing_width, processing_height, image_w
|
||||
|
||||
def fill(image, mask):
|
||||
"""fills masked regions with colors from image using blur. Not extremely effective."""
|
||||
|
||||
image_mod = Image.new('RGBA', (image.width, image.height))
|
||||
|
||||
image_masked = Image.new('RGBa', (image.width, image.height))
|
||||
image_masked.paste(image.convert("RGBA").convert("RGBa"), mask=ImageOps.invert(mask.convert('L')))
|
||||
|
||||
image_masked = image_masked.convert('RGBa')
|
||||
|
||||
for radius, repeats in [(256, 1), (64, 1), (16, 2), (4, 4), (2, 2), (0, 1)]:
|
||||
blurred = image_masked.filter(ImageFilter.GaussianBlur(radius)).convert('RGBA')
|
||||
for _ in range(repeats):
|
||||
image_mod.alpha_composite(blurred)
|
||||
|
||||
return image_mod.convert("RGB")
|
||||
|
||||
+13
-8
@@ -17,7 +17,6 @@ from blendmodes.blend import blendLayers, BlendType
|
||||
from installer import git_commit
|
||||
import modules.sd_hijack
|
||||
from modules import devices, prompt_parser, masking, sd_samplers, lowvram, generation_parameters_copypaste, script_callbacks, extra_networks, sd_vae_approx, scripts, sd_samplers_common # pylint: disable=unused-import
|
||||
from modules.sd_hijack import model_hijack
|
||||
import modules.shared as shared
|
||||
import modules.paths as paths
|
||||
import modules.face_restoration
|
||||
@@ -515,6 +514,8 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts, all_seeds, all_su
|
||||
if 'color' in p.ops:
|
||||
args["Color correction"] = True
|
||||
|
||||
if hasattr(modules.sd_hijack.model_hijack, 'embedding_db') and len(modules.sd_hijack.model_hijack.embedding_db.embeddings_used) > 0: # this is for original hijaacked models only, diffusers are handled separately
|
||||
args["Embeddings"] = ', '.join(modules.sd_hijack.model_hijack.embedding_db.embeddings_used)
|
||||
|
||||
# tome
|
||||
token_merging_ratio = p.get_token_merging_ratio()
|
||||
@@ -661,9 +662,8 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed:
|
||||
p.all_subseeds = subseed
|
||||
else:
|
||||
p.all_subseeds = [int(subseed) + x for x in range(len(p.all_prompts))]
|
||||
|
||||
if os.path.exists(shared.opts.embeddings_dir) and not p.do_not_reload_embeddings:
|
||||
model_hijack.embedding_db.load_textual_inversion_embeddings()
|
||||
modules.sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings()
|
||||
if p.scripts is not None:
|
||||
p.scripts.process(p)
|
||||
infotexts = []
|
||||
@@ -736,8 +736,8 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed:
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
uc = get_conds_with_caching(prompt_parser.get_learned_conditioning, p.negative_prompts, p.steps * step_multiplier, cached_uc)
|
||||
c = get_conds_with_caching(prompt_parser.get_multicond_learned_conditioning, p.prompts, p.steps * step_multiplier, cached_c)
|
||||
if len(model_hijack.comments) > 0:
|
||||
for comment in model_hijack.comments:
|
||||
if len(modules.sd_hijack.model_hijack.comments) > 0:
|
||||
for comment in modules.sd_hijack.model_hijack.comments:
|
||||
comments[comment] = 1
|
||||
with devices.without_autocast() if devices.unet_needs_upcast else devices.autocast():
|
||||
samples_ddim = p.sample(conditioning=c, unconditional_conditioning=uc, seeds=p.seeds, subseeds=p.subseeds, subseed_strength=p.subseed_strength, prompts=p.prompts)
|
||||
@@ -1119,10 +1119,15 @@ class StableDiffusionProcessingImg2Img(StableDiffusionProcessing):
|
||||
self.width = image.width
|
||||
self.height = image.height
|
||||
if image_mask is not None:
|
||||
image_masked = Image.new('RGBa', (image.width, image.height))
|
||||
image_masked.paste(image.convert("RGBA").convert("RGBa"), mask=ImageOps.invert(self.mask_for_overlay.convert('L')) if self.mask_for_overlay is not None else None)
|
||||
try:
|
||||
image_masked = Image.new('RGBa', (image.width, image.height))
|
||||
image_to_paste = image.convert("RGBA").convert("RGBa")
|
||||
image_to_mask = ImageOps.invert(self.mask_for_overlay.convert('L')) if self.mask_for_overlay is not None else None
|
||||
image_masked.paste(image_to_paste, mask=image_to_mask)
|
||||
self.overlay_images.append(image_masked.convert('RGBA'))
|
||||
except Exception as e:
|
||||
shared.log.error(f"Failed to apply mask to image: {e}")
|
||||
self.mask = image_mask # assign early for diffusers
|
||||
self.overlay_images.append(image_masked.convert('RGBA'))
|
||||
# crop_region is not None if we are doing inpaint full res
|
||||
if crop_region is not None:
|
||||
image = image.crop(crop_region)
|
||||
|
||||
@@ -304,6 +304,9 @@ def process_diffusers(p: StableDiffusionProcessing, seeds, prompts, negative_pro
|
||||
except AssertionError as e:
|
||||
shared.log.info(e)
|
||||
|
||||
if hasattr(shared.sd_model, 'embedding_db') and len(shared.sd_model.embedding_db.embeddings_used) > 0:
|
||||
p.extra_generation_params['Embeddings'] = ', '.join(shared.sd_model.embedding_db.embeddings_used)
|
||||
|
||||
if lora_state['active']:
|
||||
p.extra_generation_params['LoRA method'] = shared.opts.diffusers_lora_loader
|
||||
unload_diffusers_lora()
|
||||
|
||||
@@ -38,6 +38,8 @@ CLIP_SKIP_MAPPING = {
|
||||
class DiffusersTextualInversionManager(BaseTextualInversionManager):
|
||||
def __init__(self, pipe):
|
||||
self.pipe = pipe
|
||||
if hasattr(self.pipe, 'embedding_db'):
|
||||
self.pipe.embedding_db.embeddings_used.clear()
|
||||
|
||||
#from https://github.com/huggingface/diffusers/blob/705c592ea98ba4e288d837b9cba2767623c78603/src/diffusers/loaders.py#L599
|
||||
def maybe_convert_prompt(self, prompt: typing.Union[str, typing.List[str]], tokenizer = "PreTrainedTokenizer"):
|
||||
@@ -52,12 +54,15 @@ class DiffusersTextualInversionManager(BaseTextualInversionManager):
|
||||
unique_tokens = set(tokens)
|
||||
for token in unique_tokens:
|
||||
if token in tokenizer.added_tokens_encoder:
|
||||
if hasattr(self.pipe, 'embedding_db'):
|
||||
self.pipe.embedding_db.embeddings_used.append(token)
|
||||
replacement = token
|
||||
i = 1
|
||||
while f"{token}_{i}" in tokenizer.added_tokens_encoder:
|
||||
replacement += f" {token}_{i}"
|
||||
i += 1
|
||||
prompt = prompt.replace(token, replacement)
|
||||
self.pipe.embedding_db.embeddings_used = list(set(self.pipe.embedding_db.embeddings_used))
|
||||
return prompt
|
||||
|
||||
def expand_textual_inversion_token_ids_if_necessary(self, token_ids: typing.List[int]) -> typing.List[int]:
|
||||
|
||||
@@ -200,7 +200,7 @@ class StableDiffusionModelHijack:
|
||||
torch._dynamo.config.suppress_errors = opts.cuda_compile_errors # pylint: disable=protected-access
|
||||
torch.backends.cudnn.benchmark = True
|
||||
if opts.cuda_compile_backend == 'hidet':
|
||||
import hidet
|
||||
import hidet # pylint: disable=import-error
|
||||
hidet.torch.dynamo_config.use_tensor_core(True)
|
||||
hidet.torch.dynamo_config.search_space(2)
|
||||
m.model = torch.compile(m.model, mode=opts.cuda_compile_mode, backend=opts.cuda_compile_backend, fullgraph=opts.cuda_compile_fullgraph, dynamic=False)
|
||||
@@ -209,7 +209,6 @@ class StableDiffusionModelHijack:
|
||||
shared.log.warning(f"Model compile not supported: {err}")
|
||||
|
||||
self.optimization_method = apply_optimizations()
|
||||
|
||||
self.clip = m.cond_stage_model
|
||||
|
||||
def flatten(el):
|
||||
@@ -226,20 +225,16 @@ class StableDiffusionModelHijack:
|
||||
return # not ldm model
|
||||
if type(m.cond_stage_model) == xlmr.BertSeriesModelWithTransformation:
|
||||
m.cond_stage_model = m.cond_stage_model.wrapped
|
||||
|
||||
elif type(m.cond_stage_model) == sd_hijack_clip.FrozenCLIPEmbedderWithCustomWords:
|
||||
m.cond_stage_model = m.cond_stage_model.wrapped
|
||||
|
||||
model_embeddings = m.cond_stage_model.transformer.text_model.embeddings
|
||||
if type(model_embeddings.token_embedding) == EmbeddingsWithFixes:
|
||||
model_embeddings.token_embedding = model_embeddings.token_embedding.wrapped
|
||||
elif type(m.cond_stage_model) == sd_hijack_open_clip.FrozenOpenCLIPEmbedderWithCustomWords:
|
||||
m.cond_stage_model.wrapped.model.token_embedding = m.cond_stage_model.wrapped.model.token_embedding.wrapped
|
||||
m.cond_stage_model = m.cond_stage_model.wrapped
|
||||
|
||||
undo_optimizations()
|
||||
undo_weighted_forward(m)
|
||||
|
||||
self.apply_circular(False)
|
||||
self.layers = None
|
||||
self.clip = None
|
||||
@@ -247,9 +242,7 @@ class StableDiffusionModelHijack:
|
||||
def apply_circular(self, enable):
|
||||
if self.circular_enabled == enable:
|
||||
return
|
||||
|
||||
self.circular_enabled = enable
|
||||
|
||||
for layer in [layer for layer in self.layers if type(layer) == torch.nn.Conv2d]:
|
||||
layer.padding_mode = 'circular' if enable else 'zeros'
|
||||
|
||||
|
||||
@@ -179,9 +179,7 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module):
|
||||
used_embeddings[embedding.name] = embedding
|
||||
z = self.process_tokens(tokens, multipliers)
|
||||
zs.append(z)
|
||||
if len(used_embeddings) > 0:
|
||||
embeddings_list = ", ".join([f'{name} [{embedding.checksum()}]' for name, embedding in used_embeddings.items()])
|
||||
self.hijack.comments.append(f"Used embeddings: {embeddings_list}")
|
||||
self.hijack.embedding_db.embeddings_used = [name for name, embedding in used_embeddings.items()]
|
||||
return torch.hstack(zs)
|
||||
|
||||
def process_tokens(self, remade_batch_tokens, batch_multipliers):
|
||||
|
||||
@@ -9,7 +9,6 @@ tokenizer = open_clip.tokenizer._tokenizer # pylint: disable=protected-access
|
||||
class FrozenOpenCLIPEmbedderWithCustomWords(sd_hijack_clip.FrozenCLIPEmbedderWithCustomWordsBase):
|
||||
def __init__(self, wrapped, hijack):
|
||||
super().__init__(wrapped, hijack)
|
||||
|
||||
self.comma_token = [v for k, v in tokenizer.encoder.items() if k == ',</w>'][0]
|
||||
self.id_start = tokenizer.encoder["<start_of_text>"]
|
||||
self.id_end = tokenizer.encoder["<end_of_text>"]
|
||||
@@ -17,17 +16,14 @@ class FrozenOpenCLIPEmbedderWithCustomWords(sd_hijack_clip.FrozenCLIPEmbedderWit
|
||||
|
||||
def tokenize(self, texts):
|
||||
tokenized = [tokenizer.encode(text) for text in texts]
|
||||
|
||||
return tokenized
|
||||
|
||||
def encode_with_transformers(self, tokens):
|
||||
z = self.wrapped.encode_with_transformer(tokens)
|
||||
|
||||
return z
|
||||
|
||||
def encode_embedding_init_text(self, init_text, nvpt):
|
||||
ids = tokenizer.encode(init_text)
|
||||
ids = torch.asarray([ids], device=devices.device, dtype=torch.int)
|
||||
embedded = self.wrapped.model.token_embedding.wrapped(ids).squeeze(0)
|
||||
|
||||
return embedded
|
||||
|
||||
@@ -878,15 +878,14 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
shared.log.error("Failed to load diffusers model")
|
||||
errors.display(e, "loading Diffusers model")
|
||||
|
||||
from modules.textual_inversion import textual_inversion
|
||||
sd_model.embedding_db = textual_inversion.EmbeddingDatabase()
|
||||
if op == 'refiner':
|
||||
model_data.sd_refiner = sd_model
|
||||
else:
|
||||
model_data.sd_model = sd_model
|
||||
|
||||
from modules.textual_inversion import textual_inversion
|
||||
embedding_db = textual_inversion.EmbeddingDatabase()
|
||||
embedding_db.add_embedding_dir(shared.opts.embeddings_dir)
|
||||
embedding_db.load_textual_inversion_embeddings(force_reload=True)
|
||||
sd_model.embedding_db.add_embedding_dir(shared.opts.embeddings_dir)
|
||||
sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True)
|
||||
|
||||
timer.record("load")
|
||||
shared.log.info(f"Model loaded in {timer.summary()} native={get_native(sd_model)}")
|
||||
|
||||
@@ -12,7 +12,7 @@ from modules import shared, devices, sd_hijack, processing, sd_models, images, s
|
||||
import modules.textual_inversion.dataset
|
||||
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.logging import save_settings_to_file
|
||||
from modules.textual_inversion.ti_logging import save_settings_to_file
|
||||
from modules.modelloader import directory_files, extension_filter, directory_mtime
|
||||
|
||||
TextualInversionTemplate = namedtuple("TextualInversionTemplate", ["name", "path"])
|
||||
@@ -53,9 +53,7 @@ class Embedding:
|
||||
"sd_checkpoint": self.sd_checkpoint,
|
||||
"sd_checkpoint_name": self.sd_checkpoint_name,
|
||||
}
|
||||
|
||||
torch.save(embedding_data, filename)
|
||||
|
||||
if shared.opts.save_optimizer_state and self.optimizer_state_dict is not None:
|
||||
optimizer_saved_dict = {
|
||||
'hash': self.checksum(),
|
||||
@@ -66,13 +64,11 @@ class Embedding:
|
||||
def checksum(self):
|
||||
if self.cached_checksum is not None:
|
||||
return self.cached_checksum
|
||||
|
||||
def const_hash(a):
|
||||
r = 0
|
||||
for v in a:
|
||||
r = (r * 281 ^ int(v) * 997) & 0xFFFFFFFF
|
||||
return r
|
||||
|
||||
self.cached_checksum = f'{const_hash(self.vec.reshape(-1) * 100) & 0xffff:04x}'
|
||||
return self.cached_checksum
|
||||
|
||||
@@ -103,6 +99,7 @@ class EmbeddingDatabase:
|
||||
self.expected_shape = -1
|
||||
self.embedding_dirs = {}
|
||||
self.previously_displayed_embeddings = ()
|
||||
self.embeddings_used = []
|
||||
|
||||
def add_embedding_dir(self, path):
|
||||
self.embedding_dirs[path] = DirWithTextualInversionEmbeddings(path)
|
||||
@@ -139,13 +136,13 @@ class EmbeddingDatabase:
|
||||
return
|
||||
name = os.path.basename(fn)
|
||||
embedding = Embedding(vec=None, name=name)
|
||||
embedding.filename = path
|
||||
try:
|
||||
if hasattr(pipe,"load_textual_inversion"):
|
||||
pipe.load_textual_inversion(path, cache_dir=shared.opts.diffusers_dir, local_files_only=True)
|
||||
elif "safetensors" in path:
|
||||
embeddings_dict = {}
|
||||
from safetensors.torch import safe_open
|
||||
|
||||
with safe_open(path, framework="pt") as f:
|
||||
for k in f.keys():
|
||||
embeddings_dict[k] = f.get_tensor(k)
|
||||
@@ -161,8 +158,9 @@ class EmbeddingDatabase:
|
||||
pipe.text_encoder.get_input_embeddings().weight.data[token_id] = embeddings_dict["clip_l"][i]
|
||||
pipe.text_encoder_2.get_input_embeddings().weight.data[token_id] = embeddings_dict["clip_g"][i]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
self.word_embeddings[name] = embedding
|
||||
raise NotImplementedError
|
||||
# self.word_embeddings[name] = embedding
|
||||
self.register_embedding(embedding, shared.sd_model)
|
||||
except Exception:
|
||||
self.skipped_embeddings[name] = embedding
|
||||
try:
|
||||
@@ -209,7 +207,7 @@ class EmbeddingDatabase:
|
||||
# diffuser concepts
|
||||
elif type(data) == dict and type(next(iter(data.values()))) == torch.Tensor:
|
||||
if len(data.keys()) != 1:
|
||||
shared.log.warning(f"Skipping embedding: {filename} multiple keys found")
|
||||
# shared.log.warning(f"Skipping embedding: {filename} multiple keys found")
|
||||
self.skipped_embeddings[name] = Embedding(None, name)
|
||||
return
|
||||
emb = next(iter(data.values()))
|
||||
@@ -226,7 +224,6 @@ class EmbeddingDatabase:
|
||||
embedding.vectors = vec.shape[0]
|
||||
embedding.shape = vec.shape[-1]
|
||||
embedding.filename = path
|
||||
|
||||
if self.expected_shape == -1 or self.expected_shape == embedding.shape:
|
||||
self.register_embedding(embedding, shared.sd_model)
|
||||
else:
|
||||
@@ -238,10 +235,8 @@ 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:
|
||||
@@ -261,12 +256,11 @@ class EmbeddingDatabase:
|
||||
break
|
||||
if not need_reload:
|
||||
return
|
||||
|
||||
self.ids_lookup.clear()
|
||||
self.word_embeddings.clear()
|
||||
self.skipped_embeddings.clear()
|
||||
self.embeddings_used.clear()
|
||||
self.expected_shape = self.get_expected_shape()
|
||||
|
||||
for embdir in self.embedding_dirs.values():
|
||||
self.load_from_dir(embdir)
|
||||
embdir.update()
|
||||
|
||||
@@ -188,14 +188,9 @@ class ExtraNetworksPage:
|
||||
subdir = subdir[1:]
|
||||
if not self.is_empty(tgt):
|
||||
subdirs[subdir] = 1
|
||||
if subdirs:
|
||||
subdirs = OrderedDict(sorted(subdirs.items()))
|
||||
subdirs = {"": 1, **subdirs}
|
||||
subdirs_html = "".join([f"""
|
||||
<button class='lg secondary gradio-button custom-button{" search-all" if subdir=="" else ""}' onclick='extraNetworksSearchButton(event)'>
|
||||
{html.escape(subdir) if subdir!="" else "all"}
|
||||
</button><br>""" for subdir in subdirs])
|
||||
# try:
|
||||
subdirs = OrderedDict(sorted(subdirs.items()))
|
||||
subdirs_html = "<button class='lg secondary gradio-button custom-button search-all' onclick='extraNetworksSearchButton(event)'>all</button><br>"
|
||||
subdirs_html += "".join([f"<button class='lg secondary gradio-button custom-button' onclick='extraNetworksSearchButton(event)'>{html.escape(subdir)}</button><br>" for subdir in subdirs if subdir != ''])
|
||||
if len(self.html) > 0:
|
||||
res = f"<div id='{tabname}_{self_name_id}_subdirs' class='extra-network-subdirs'>{subdirs_html}</div><div id='{tabname}_{self_name_id}_cards' class='extra-network-cards'>{self.html}</div>"
|
||||
return res
|
||||
@@ -221,9 +216,6 @@ class ExtraNetworksPage:
|
||||
shared.log.debug(f'Extra networks: {self.name} items={len(self.items)} subdirs={len(subdirs)} time={round(t1-t0, 2)}')
|
||||
threading.Thread(target=self.create_thumb).start()
|
||||
return res
|
||||
# except Exception as e:
|
||||
# shared.log.error(f'Extra networks page error: title={self.title} tab={tabname} class={e.__class__.__name__} {e}')
|
||||
# return f"<div id='{tabname}_{self_name_id}_subdirs' class='extra-network-subdirs'></div><div id='{tabname}_{self_name_id}_cards' class='extra-network-cards'>Extra network error<br>{e}</div>"
|
||||
|
||||
def list_items(self):
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import json
|
||||
import os
|
||||
|
||||
from modules import ui_extra_networks, sd_hijack, shared
|
||||
from modules import shared, sd_hijack, sd_models, ui_extra_networks
|
||||
from modules.textual_inversion.textual_inversion import Embedding
|
||||
|
||||
|
||||
@@ -14,14 +14,20 @@ class ExtraNetworksPageTextualInversion(ui_extra_networks.ExtraNetworksPage):
|
||||
sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings(force_reload=True)
|
||||
|
||||
def list_items(self):
|
||||
embeddings = list(sd_hijack.model_hijack.embedding_db.word_embeddings.values())
|
||||
if len(embeddings) == 0: # maybe not loaded yet, so lets just look them up
|
||||
if sd_models.model_data.sd_model is None:
|
||||
embeddings = []
|
||||
for root, _dirs, fns in os.walk(shared.opts.embeddings_dir, followlinks=True):
|
||||
for fn in fns:
|
||||
if fn.lower().endswith(".pt"):
|
||||
if fn.lower().endswith(".pt") or fn.lower().endswith(".safetensors"):
|
||||
embedding = Embedding(0, fn)
|
||||
embedding.filename = os.path.join(root, fn)
|
||||
embeddings.append(embedding)
|
||||
elif shared.backend == shared.Backend.ORIGINAL:
|
||||
embeddings = list(sd_hijack.model_hijack.embedding_db.word_embeddings.values())
|
||||
elif hasattr(sd_models.model_data.sd_model, 'embedding_db'):
|
||||
embeddings = list(sd_models.model_data.sd_model.embedding_db.word_embeddings.values())
|
||||
else:
|
||||
embeddings = []
|
||||
for embedding in embeddings:
|
||||
path, _ext = os.path.splitext(embedding.filename)
|
||||
yield {
|
||||
|
||||
Reference in New Issue
Block a user