mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
ti fixes
This commit is contained in:
@@ -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));
|
||||
}
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user