This commit is contained in:
Vladimir Mandic
2023-09-09 19:39:21 -04:00
parent 7bda411738
commit 36001151bb
5 changed files with 14 additions and 17 deletions
+1 -1
View File
@@ -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));
}
+2 -1
View File
@@ -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]:
+1 -1
View File
@@ -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):
+6 -14
View File
@@ -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)
@@ -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):