From 36001151bbbcbe817ff2103579f5df930571461e Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sat, 9 Sep 2023 19:39:21 -0400 Subject: [PATCH] ti fixes --- javascript/promptChecker.js | 2 +- modules/prompt_parser_diffusers.py | 3 ++- modules/sd_hijack_clip.py | 2 +- .../textual_inversion/textual_inversion.py | 20 ++++++------------- .../ui_extra_networks_textual_inversion.py | 4 ++++ 5 files changed, 14 insertions(+), 17 deletions(-) diff --git a/javascript/promptChecker.js b/javascript/promptChecker.js index 66fa4442d..d18a030ec 100644 --- a/javascript/promptChecker.js +++ b/javascript/promptChecker.js @@ -25,7 +25,7 @@ function setupBracketChecking(idPrompt, idCounter) { const textarea = gradioApp().querySelector(`#${idPrompt} > label > textarea`); const counter = gradioApp().getElementById(idCounter); if (!textarea || !counter) return; - if (!promptCheckerInitialized) log('promptChecker'); + if (!promptCheckerInitialized) log('initPromptChecker'); promptCheckerInitialized = true; textarea.addEventListener('input', () => checkBrackets(textarea, counter)); } diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index 6a8b1bef0..57cbf7f17 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -62,7 +62,8 @@ class DiffusersTextualInversionManager(BaseTextualInversionManager): 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)) + if hasattr(self.pipe, 'embedding_db'): + 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]: diff --git a/modules/sd_hijack_clip.py b/modules/sd_hijack_clip.py index 73fc00c9c..045833de4 100644 --- a/modules/sd_hijack_clip.py +++ b/modules/sd_hijack_clip.py @@ -179,7 +179,7 @@ class FrozenCLIPEmbedderWithCustomWordsBase(torch.nn.Module): used_embeddings[embedding.name] = embedding z = self.process_tokens(tokens, multipliers) zs.append(z) - self.hijack.embedding_db.embeddings_used = [name for name, embedding in used_embeddings.items()] + self.hijack.embedding_db.embeddings_used = [name for name in used_embeddings.keys()] return torch.hstack(zs) def process_tokens(self, remade_batch_tokens, batch_multipliers): diff --git a/modules/textual_inversion/textual_inversion.py b/modules/textual_inversion/textual_inversion.py index a0090cdd3..c8cee595d 100644 --- a/modules/textual_inversion/textual_inversion.py +++ b/modules/textual_inversion/textual_inversion.py @@ -21,13 +21,10 @@ textual_inversion_templates = {} def list_textual_inversion_templates(): textual_inversion_templates.clear() - for root, _dirs, fns in os.walk(shared.opts.embeddings_templates_dir): for fn in fns: path = os.path.join(root, fn) - textual_inversion_templates[fn] = TextualInversionTemplate(fn, path) - return textual_inversion_templates @@ -35,6 +32,7 @@ class Embedding: def __init__(self, vec, name, step=None): self.vec = vec self.name = name + self.tag = name self.step = step self.shape = None self.vectors = 0 @@ -81,13 +79,11 @@ class DirWithTextualInversionEmbeddings: def has_changed(self): if not os.path.isdir(self.path): return False - return directory_mtime(self.path) != self.mtime def update(self): if not os.path.isdir(self.path): return - self.mtime = directory_mtime(self.path) @@ -177,19 +173,14 @@ class EmbeddingDatabase: return if ext in ['.PNG', '.WEBP', '.JXL', '.AVIF']: - _, second_ext = os.path.splitext(name) - if second_ext.upper() == '.PREVIEW': + if '.preview' in filename.lower(): return embed_image = Image.open(path) if hasattr(embed_image, 'text') and 'sd-ti-embedding' in embed_image.text: data = embedding_from_b64(embed_image.text['sd-ti-embedding']) - name = data.get('name', name) else: data = extract_image_data_embed(embed_image) - if data: - name = data.get('name', name) - else: - # if data is None, means this is not an embeding, just a preview image + if not data: # if data is None, means this is not an embeding, just a preview image return elif ext in ['.BIN', '.PT']: data = torch.load(path, map_location="cpu") @@ -207,17 +198,18 @@ 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") self.skipped_embeddings[name] = Embedding(None, name) return emb = next(iter(data.values())) if len(emb.shape) == 1: emb = emb.unsqueeze(0) else: - raise RuntimeError(f"Couldn't identify {filename} as neither textual inversion embedding nor diffuser concept.") + raise RuntimeError(f"Couldn't identify {filename} as textual inversion embedding") vec = emb.detach().to(devices.device, dtype=torch.float32) + # name = data.get('name', name) embedding = Embedding(vec, name) + 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) diff --git a/modules/ui_extra_networks_textual_inversion.py b/modules/ui_extra_networks_textual_inversion.py index 4ad9ac10a..51d077855 100644 --- a/modules/ui_extra_networks_textual_inversion.py +++ b/modules/ui_extra_networks_textual_inversion.py @@ -35,6 +35,9 @@ class ExtraNetworksPageTextualInversion(ui_extra_networks.ExtraNetworksPage): embeddings = [] for embedding in embeddings: path, _ext = os.path.splitext(embedding.filename) + tags = {} + if embedding.tag is not None: + tags[embedding.tag]=1 yield { "name": os.path.splitext(embedding.name)[0], "filename": embedding.filename, @@ -44,6 +47,7 @@ class ExtraNetworksPageTextualInversion(ui_extra_networks.ExtraNetworksPage): "search_term": self.search_terms_from_path(embedding.filename), "prompt": json.dumps(os.path.splitext(embedding.name)[0]), "local_preview": f"{path}.preview.{shared.opts.samples_format}", + "tags": tags, } def allowed_directories_for_previews(self):