sort out embeddings loading in dffusers

This commit is contained in:
Vladimir Mandic
2023-09-05 11:06:25 -04:00
parent 9058bfa250
commit a3033dc65f
13 changed files with 52 additions and 68 deletions
+4 -4
View File
@@ -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
-7
View File
@@ -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
View File
@@ -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)
+3
View File
@@ -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()
+5
View File
@@ -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]:
+1 -8
View File
@@ -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'
+1 -3
View File
@@ -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):
-4
View File
@@ -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
+4 -5
View File
@@ -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)}")
+8 -14
View File
@@ -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()
+3 -11
View File
@@ -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
+10 -4
View File
@@ -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 {