mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 01:04:32 +02:00
6b7c43cd29
add_network keys available_networks by the basename with dots mangled to underscores, and create_item looked the name up in that table using the name as loaded, so any file with a dot in its stem missed and logged "not registered". The tag pass then dropped the network's own tags and used the filename instead. The lookup now falls back to the alias table, which already carries the natural basename and the subfolder path.
140 lines
5.8 KiB
Python
140 lines
5.8 KiB
Python
import os
|
|
import json
|
|
import concurrent.futures
|
|
from modules import shared, ui_extra_networks, modelstats
|
|
from modules.logger import log
|
|
from modules.lora import lora_load
|
|
|
|
|
|
debug = os.environ.get('SD_LORA_DEBUG', None) is not None
|
|
|
|
|
|
class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage):
|
|
def __init__(self):
|
|
super().__init__('Lora')
|
|
self.list_time = 0
|
|
|
|
def refresh(self):
|
|
lora_load.list_available_networks()
|
|
|
|
@staticmethod
|
|
def get_tags(l, info, version):
|
|
tags = {}
|
|
try:
|
|
if l.metadata is not None:
|
|
modelspec_tags = l.metadata.get('modelspec.tags', {})
|
|
possible_tags = l.metadata.get('ss_tag_frequency', {}) # tags from model metedata
|
|
if isinstance(possible_tags, str):
|
|
possible_tags = {}
|
|
if isinstance(modelspec_tags, str):
|
|
modelspec_tags = {}
|
|
if len(list(modelspec_tags)) > 0:
|
|
possible_tags.update(modelspec_tags)
|
|
for k, v in possible_tags.items():
|
|
words = k.split('_', 1) if '_' in k else [v, k]
|
|
words = [str(w).replace('.json', '') for w in words]
|
|
if words[0] == '{}':
|
|
words[0] = 0
|
|
tag = ' '.join(words[1:]).lower()
|
|
tags[tag] = words[0]
|
|
|
|
possible_tags = version.get('trainedWords', [])
|
|
if isinstance(possible_tags, list):
|
|
for tag_str in possible_tags:
|
|
for tag in tag_str.split(','):
|
|
tag = tag.strip().lower()
|
|
if tag not in tags:
|
|
tags[tag] = 0
|
|
|
|
possible_tags = info.get('tags', []) # tags from info json
|
|
if not isinstance(possible_tags, list):
|
|
possible_tags = list(possible_tags.values())
|
|
for tag in possible_tags:
|
|
tag = tag.strip().lower()
|
|
if tag not in tags:
|
|
tags[tag] = 0
|
|
except Exception:
|
|
pass
|
|
|
|
# cleanup tags: remove model names, bad words, and special characters
|
|
model_words = ['ltx', 'minimax', 'h3', 'sdxl', 'klein', 'wan', 'flux', 'qwen', 'vace', 'lcm', 'slider']
|
|
bad_words = ['concept', 'style', 'styles', 'base model', 'video', 'audio', 'turbo', 'distill', 'assets', 'action', 'enhancer', 'detail', 'tool', 'dir', 'all']
|
|
bad_parts = ['lora', 'comfyui', 't2i', 'i2i', 't2v', 'i2v', 'steps']
|
|
bad_chars = [';', ':', '<', ">", "*", '?', '\'', '\"', '(', ')', '[', ']', '{', '}', '\\', '/']
|
|
clean_tags = {}
|
|
for k, v in tags.items():
|
|
if k in bad_words:
|
|
continue
|
|
if any(k.startswith(s) for s in model_words):
|
|
continue
|
|
if any(s in k for s in bad_parts):
|
|
continue
|
|
tag = ''.join(i for i in k if i not in bad_chars).strip()
|
|
clean_tags[tag] = v
|
|
|
|
clean_tags.pop('img', None)
|
|
clean_tags.pop('dataset', None)
|
|
return clean_tags
|
|
|
|
_VERSION_DISPLAY = {
|
|
'f1': 'Flux', 'sd1': 'SD 1.5', 'sd2': 'SD 2', 'xl': 'SDXL',
|
|
'sd3': 'SD3', 'sc': 'Cascade', 'hv': 'HunyuanVideo',
|
|
'chroma': 'Chroma', 'zimage': 'zImage', 'qwen': 'Qwen',
|
|
'krea2': 'Krea 2',
|
|
}
|
|
|
|
def cleanup_version(self, dct, lora):
|
|
ver = dct.get("baseModel", lora.sd_version)
|
|
ver = self._VERSION_DISPLAY.get(ver, ver)
|
|
for suffix in (' 0.9', ' 1.0'): # strip uninformative minor versions
|
|
ver = ver.replace(suffix, '')
|
|
return ver
|
|
|
|
def create_item(self, name):
|
|
l = lora_load.available_networks.get(name) or lora_load.available_network_aliases.get(name) # that table mangles dots in the stem to underscores; the aliases carry the natural basename and the subfolder path
|
|
if l is None:
|
|
log.warning(f'Networks: type=lora registered={len(list(lora_load.available_networks))} file="{name}" not registered')
|
|
return None
|
|
try:
|
|
# path, _ext = os.path.splitext(l.filename)
|
|
name = os.path.splitext(os.path.relpath(l.filename, shared.cmd_opts.lora_dir))[0]
|
|
size, mtime = modelstats.stat(l.filename)
|
|
info = self.find_info(l.filename)
|
|
ver_dct = self.find_version(l, info)
|
|
item = {
|
|
"type": 'Lora',
|
|
"name": name,
|
|
"alias": os.path.splitext(os.path.basename(l.filename))[0],
|
|
"filename": l.filename,
|
|
"hash": l.shorthash,
|
|
"prompt": json.dumps(f" <lora:{l.get_alias()}:{shared.opts.extra_networks_default_multiplier}>"),
|
|
"metadata": json.dumps(l.metadata, indent=4) if l.metadata else None,
|
|
"mtime": mtime,
|
|
"size": size,
|
|
"version": self.cleanup_version(ver_dct, l),
|
|
"info": info,
|
|
"description": self.find_description(l.filename, info),
|
|
"tags": self.get_tags(l, info, ver_dct),
|
|
}
|
|
return item
|
|
except Exception as e:
|
|
log.error(f'Networks: type=lora file="{name}" {e}')
|
|
if debug:
|
|
from modules import errors
|
|
errors.display(e, 'Lora')
|
|
return None
|
|
|
|
def list_items(self):
|
|
items = []
|
|
with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor:
|
|
future_items = {executor.submit(self.create_item, net): net for net in lora_load.available_networks}
|
|
for future in concurrent.futures.as_completed(future_items):
|
|
item = future.result()
|
|
if item is not None:
|
|
items.append(item)
|
|
self.update_all_previews(items)
|
|
return items
|
|
|
|
def allowed_directories_for_previews(self):
|
|
return [shared.cmd_opts.lora_dir]
|