diffusers better ti reload

This commit is contained in:
Vladimir Mandic
2023-09-26 17:40:29 -04:00
parent 2f17cef6f6
commit 6bb1e59de6
6 changed files with 22 additions and 19 deletions
+2 -1
View File
@@ -82,7 +82,8 @@ def wrap_gradio_call(func, extra_outputs=None, add_stats=False, name=None):
vram_html = ''
if not shared.mem_mon.disabled:
vram = {k: -(v//-(1024*1024)) for k, v in shared.mem_mon.read().items()}
vram_html += f" | <p class='vram'>GPU active {max(vram['active_peak'], vram['reserved_peak'])} MB reserved {vram['reserved']} | used {vram['used']} MB free {vram['free']} MB total {vram['total']} MB | retries {vram['retries']} oom {vram['oom']}</p>"
if vram.get('active_peak', 0) > 0:
vram_html = f" | <p class='vram'>GPU active {max(vram['active_peak'], vram['reserved_peak'])} MB reserved {vram['reserved']} | used {vram['used']} MB free {vram['free']} MB total {vram['total']} MB | retries {vram['retries']} oom {vram['oom']}</p>"
res[-1] += f"<div class='performance'><p class='time'>Time: {elapsed_text}</p>{vram_html}</div>"
return tuple(res)
return f
+3 -3
View File
@@ -7,11 +7,11 @@ import torch
from modules.shared import state
from modules import sd_samplers_common, prompt_parser, shared
import modules.uni_pc
import modules.unipc
samplers_data_compvis = [
sd_samplers_common.SamplerData('UniPC', lambda model: VanillaStableDiffusionSampler(modules.uni_pc.UniPCSampler, model), [], {}),
sd_samplers_common.SamplerData('UniPC', lambda model: VanillaStableDiffusionSampler(modules.unipc.UniPCSampler, model), [], {}),
sd_samplers_common.SamplerData('DDIM', lambda model: VanillaStableDiffusionSampler(ldm.models.diffusion.ddim.DDIMSampler, model), [], {"default_eta_is_0": True}),
sd_samplers_common.SamplerData('PLMS', lambda model: VanillaStableDiffusionSampler(ldm.models.diffusion.plms.PLMSSampler, model), [], {}),
]
@@ -22,7 +22,7 @@ class VanillaStableDiffusionSampler:
self.sampler = constructor(sd_model)
self.is_ddim = hasattr(self.sampler, 'p_sample_ddim')
self.is_plms = hasattr(self.sampler, 'p_sample_plms')
self.is_unipc = isinstance(self.sampler, modules.uni_pc.UniPCSampler)
self.is_unipc = isinstance(self.sampler, modules.unipc.UniPCSampler)
self.orig_p_sample_ddim = None
if self.is_plms:
self.orig_p_sample_ddim = self.sampler.p_sample_plms
+17 -15
View File
@@ -138,8 +138,13 @@ class EmbeddingDatabase:
done = False
if hasattr(pipe,"load_textual_inversion"):
try:
pipe.load_textual_inversion(path, cache_dir=shared.opts.diffusers_dir, local_files_only=True)
done = True
token_ids = pipe.tokenizer.convert_tokens_to_ids(name)
if token_ids > 49407: # already loaded
done = True
else:
pipe.load_textual_inversion(path, token=name, cache_dir=shared.opts.diffusers_dir, local_files_only=True)
done = True
self.register_embedding(embedding, shared.sd_model)
except Exception:
pass
if not done and "safetensors" in path:
@@ -148,30 +153,27 @@ class EmbeddingDatabase:
with safe_open(path, framework="pt") as f:
for k in f.keys():
embeddings_dict[k] = f.get_tensor(k)
clip_l = pipe.text_encoder.get_input_embeddings().weight if hasattr(pipe, 'text_encoder') and hasattr(pipe.text_encoder, "resize_token_embeddings") else None
clip_g = pipe.text_encoder_2.get_input_embeddings().weight if hasattr(pipe, 'text_encoder_2') and hasattr(pipe.text_encoder_2, "resize_token_embeddings") else None
tokens = []
for i in range(len(embeddings_dict["clip_l"])):
tokens.append(name if i == 0 else f"{name}_{i}")
if clip_l is not None and len(clip_l.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)
clip_l = None
clip_g = None
if hasattr(pipe.text_encoder, "resize_token_embeddings"):
if clip_l is not None:
pipe.text_encoder.resize_token_embeddings(len(pipe.tokenizer))
clip_l = pipe.text_encoder.get_input_embeddings().weight
if hasattr(pipe.text_encoder_2, "resize_token_embeddings"):
pipe.text_encoder_2.resize_token_embeddings(len(pipe.tokenizer))
clip_g = pipe.text_encoder_2.get_input_embeddings().weight
for i in range(len(token_ids)):
if clip_l is not None:
for i in range(len(token_ids)):
clip_l.data[token_ids[i]] = embeddings_dict["clip_l"][i]
if clip_g is not None:
if clip_g is not None:
pipe.text_encoder_2.resize_token_embeddings(len(pipe.tokenizer))
for i in range(len(token_ids)):
clip_g.data[token_ids[i]] = embeddings_dict["clip_g"][i]
self.register_embedding(embedding, shared.sd_model)
else:
raise NotImplementedError
# self.word_embeddings[name] = embedding
self.register_embedding(embedding, shared.sd_model)
except Exception:
self.skipped_embeddings[name] = embedding