diff --git a/extensions-builtin/Lora/networks.py b/extensions-builtin/Lora/networks.py index 0453fb597..ae33e030d 100644 --- a/extensions-builtin/Lora/networks.py +++ b/extensions-builtin/Lora/networks.py @@ -16,7 +16,7 @@ import network_glora import lora_convert import torch import diffusers.models.lora -from modules import shared, devices, sd_models, sd_models_compile, errors, scripts +from modules import shared, devices, sd_models, sd_models_compile, errors, scripts, files_cache debug = os.environ.get('SD_LORA_DEBUG', None) is not None @@ -435,20 +435,20 @@ def network_MultiheadAttention_load_state_dict(self, *args, **kwargs): def list_available_networks(): + global available_networks, available_network_aliases, forbidden_network_aliases, available_network_hash_lookup available_networks.clear() available_network_aliases.clear() forbidden_network_aliases.clear() available_network_hash_lookup.clear() forbidden_network_aliases.update({"none": 1, "Addams": 1}) os.makedirs(shared.cmd_opts.lora_dir, exist_ok=True) - candidates = [] + directories = [] if os.path.exists(shared.cmd_opts.lora_dir): - candidates += list(shared.walk_files(shared.cmd_opts.lora_dir, allowed_extensions=[".pt", ".ckpt", ".safetensors"])) + directories.append(shared.cmd_opts.lora_dir) else: shared.log.warning('LoRA directory not found: path="{shared.cmd_opts.lora_dir}"') if os.path.exists(shared.cmd_opts.lyco_dir): - candidates += list(shared.walk_files(shared.cmd_opts.lyco_dir, allowed_extensions=[".pt", ".ckpt", ".safetensors"])) - + directories.append(shared.cmd_opts.lyco_dir) def add_network(filename): if os.path.isdir(filename): return @@ -466,8 +466,9 @@ def list_available_networks(): shared.log.error(f"Failed to load network {name} from {filename} {e}") with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor: - for fn in candidates: + for fn in files_cache.list_files(*directories, ext_filter=[".pt", ".ckpt", ".safetensors"]): executor.submit(add_network, fn) + print(f'Lora/LyCORIS Networks: networks={len(available_networks)} directories={directories}') def infotext_pasted(infotext, params): # pylint: disable=W0613 diff --git a/modules/extensions.py b/modules/extensions.py index fb20ce563..1d6f287e0 100644 --- a/modules/extensions.py +++ b/modules/extensions.py @@ -1,7 +1,7 @@ import os from datetime import datetime import git -from modules import shared, errors +from modules import shared, errors, files_cache from modules.paths import extensions_dir, extensions_builtin_dir diff --git a/modules/files_cache.py b/modules/files_cache.py new file mode 100644 index 000000000..00721a6b5 --- /dev/null +++ b/modules/files_cache.py @@ -0,0 +1,389 @@ +import itertools +import os.path as path +from collections import UserDict +from dataclasses import dataclass, field +from os import scandir +from typing import Callable, Dict, Iterator, List, Optional, Union + +from installer import print_dict + +WasDirty = bool +DidDelete = bool +IsDirectory = bool +DirectoryExists = bool +IsDirectory = bool +IsDirty = bool +CachedDirectoryIsStale = bool +MTime = float +IsHidden = bool + +FilePath = str +FilePathList = List[FilePath] +FilePathIterator = Iterator[FilePath] + +DirectoryPath = str +DirectoryPathList = List[DirectoryPath] +DirectoryPathIterator = Iterator[DirectoryPath] + +class Directory: + ... + +DirectoryList = List[Directory] +DirectoryIterator = Iterator[Directory] +DirectoryCollection = Dict[DirectoryPath, Directory] + +ExtensionFilter = Callable +ExtensionList = list[str] + +RecursiveType = Union[bool,Callable] + + +def real_path(directory_path:DirectoryPath) -> DirectoryPath | None: + try: + return path.abspath(path.expanduser(directory_path)) + except Exception: + pass + return None + + +@dataclass(slots=True,frozen=True) +class Directory(Directory): # pylint: disable=E0102 + + + path: DirectoryPath = field(default_factory=str) + mtime: float = field(default_factory=float, init=False) + files: FilePathList = field(default_factory=list) + directories: DirectoryPathList = field(default_factory=list) + + + def __post_init__(self): + object.__setattr__(self, 'mtime', self.live_mtime) + + + @classmethod + def from_dict(cls, dict_object: dict) -> Directory: + directory = cls.__new__(cls) + object.__setattr__(directory, 'path', dict_object.get('path')) + object.__setattr__(directory, 'mtime', dict_object.get('mtime')) + object.__setattr__(directory, 'files', dict_object.get('files')) + object.__setattr__(directory, 'directories', dict_object.get('directories')) + return directory + + + def clear(self) -> None: + self._update(Directory.from_dict({ + 'path': None, + 'mtime': float(), + 'files': [], + 'directories': [] + })) + + + def update(self, source_directory: Directory) -> Directory: + if source_directory is not self: + self._update(source_directory) + return self + + + def _update(self, source:Directory) -> None: + assert not source.path or source.path == self.path, f'When updating a directory, the paths must match. Attemped to update Directory `{self.path}` with `{source.path}`' + for dead_path in self.directories: + if dead_path not in source.directories: + delete_cached_directory(dead_path) + self.directories[:] = source.directories + self.files[:] = source.files + object.__setattr__(self, 'mtime', source.mtime) + + + def __str__(self) -> str: + return str(print_dict(self, path=self.path, mtime=self.mtime, files=len(self.files), directories=len(self.directories))) + + + @property + def exists(self) -> DirectoryExists: + return self.path and path.exists(self.path) + + + @property + def is_directory(self) -> IsDirectory: + return self.exists and path.isdir(self.path) + + + @property + def live_mtime(self) -> MTime: + return path.getmtime(self.path) if self.is_directory else 0 + + + @property + def is_stale(self) -> CachedDirectoryIsStale: + return not self.is_directory or self.mtime != self.live_mtime + + +class DirectoryCache(UserDict, DirectoryCollection): + def __delattr__(self, directory_path: str) -> None: + directory: Directory = get_directory(directory_path, fetch=False) + if directory: + map(delete_cached_directory, directory.directories) + directory.clear() + del self.data[directory_path] + + +def clean_directory(directory: Directory, /, recursive: RecursiveType=False) -> bool: + if not directory.is_directory: + is_clean = False + delete_cached_directory(directory.path) + else: + is_clean = not directory.is_stale + if not is_clean: + directory.update(fetch_directory(directory.path)) + else: + for directory_path in directory.directories[:]: + try: + recurse = recursive and (not callable(recursive) or recursive(directory.path)) + directory = get_directory(directory_path, fetch=recurse) + if directory: + if directory.is_directory: + if recurse: + is_clean = clean_directory(directory, recursive=recurse) and is_clean + continue + delete_cached_directory(directory_path) + # If we had intended to fetch this directory, but didn't, that means it doesn't exist. Purge. + if recurse: + directory.directories.remove(directory_path) + is_clean = False + except Exception: + pass + return is_clean + + +def get_directory(directory_or_path: DirectoryPath, /, fetch:bool=True) -> Directory | None: + if isinstance(directory_or_path, Directory): + if directory_or_path.is_directory: + return directory_or_path + else: + directory_or_path = directory_or_path.path + global cache_folders # pylint: disable=W0602 + directory_or_path = real_path(directory_or_path) + if not cache_folders.get(directory_or_path, None): + if fetch: + directory = fetch_directory(directory_path=directory_or_path) + if directory: + cache_folders[directory_or_path] = directory + else: + clean_directory(cache_folders[directory_or_path]) + return cache_folders[directory_or_path] if directory_or_path in cache_folders else None + + +def fetch_directory(directory_path: DirectoryPath) -> Directory | None: + directory: Directory + for directory in _walk(directory_path, lambda e, path: delete_cached_directory(path), recurse=False): + return directory # The return is intentional, we get a generator, we only need the one + return None + + +def _walk(top, onerror:Callable=None, /, recurse:RecursiveType=True) -> Directory: + # A near-exact copy of `path.walk()`, trimmed slightly. Probably not nessesary for most people's collections, but makes a difference on really large datasets. + nondirs = [] + walk_dirs = [] + try: + scandir_it = scandir(top) + except OSError as error: + if callable(onerror): + onerror(error, top) + return + with scandir_it: + while True: + try: + try: + entry = next(scandir_it) + except StopIteration: + break + except OSError as error: + if callable(onerror): + onerror(error, top) + return + try: + is_dir = entry.is_dir() + except OSError: + is_dir = False + if not is_dir: + nondirs.append(entry.path) + else: + try: + if entry.is_symlink() and not path.exists(entry.path): + raise NotADirectoryError('Broken Symlink') + walk_dirs.append(entry.path) + except OSError as error: + if callable(onerror): + onerror(error, entry.path) + yield Directory(top, nondirs, walk_dirs) + if recurse: + # Recurse into sub-directories + for new_path in walk_dirs: + if path.basename(new_path).startswith('models--'): + continue + if callable(recurse) and not recurse(new_path): + continue + yield from _walk(new_path, onerror, recurse=recurse) + + +def _cached_walk(top, onerror:Callable=None, /, recurse:RecursiveType=True) -> Directory: + top = get_directory(top) + if not top: + return + yield top + if recurse: + for child_directory in top.directories: + if path.basename(child_directory).startswith('models--'): + continue + if callable(recurse) and not recurse(child_directory): + continue + yield from _cached_walk(child_directory, onerror, recurse=recurse) + +def walk(top, onerror:Callable=None, /, recurse:RecursiveType=True, cached=True) -> Directory: + if cached: + yield from _cached_walk(top, onerror, recurse=recurse) + else: + yield from _walk(top, onerror, recurse=recurse) + + +def delete_cached_directory(directory_path:DirectoryPath) -> DidDelete: + global cache_folders # pylint: disable=W0602 + if directory_path in cache_folders: + del cache_folders[directory_path] + + +def is_directory(dir_path:DirectoryPath) -> IsDirectory: + return dir_path and path.exists(dir_path) and path.isdir(dir_path) + + +def directory_mtime(directory_path:DirectoryPath, /, recursive:RecursiveType=True) -> MTime: + return float(max(0, *[directory.mtime for directory in get_directories(directory_path, recursive=recursive)])) + + +def unique_directories(directories:DirectoryPathList, /, recursive:RecursiveType=True) -> DirectoryPathIterator: + '''Ensure no empty, or duplicates''' + '''If we are going recursive, then directories that are children of other directories are redundant''' + ''' @todo this is incredibly inneficient. the hit is small, but it is ugly, no? ''' + directories = sorted(unique_paths(directories), reverse=True) + while directories: + directory = directories.pop() + yield directory + if not recursive: + continue + _directory = path.join(directory, '') + child_directory = None + while directories and directories[-1].startswith(_directory): + if not callable(recursive) or not child_directory: + directories.pop() + continue + child_directory = directories[-1][len(directory):] + if child_directory: + next_directory = _directory + if not callable(recursive): + _remove_directory = next_directory + else: + for sub_directory in child_directory.split(path.sep): + next_directory = path.join(next_directory, sub_directory) + if recursive(next_directory): + _remove_directory = path.join(next_directory, '') + break + while _remove_directory and directories: + _d = directories.pop() + if not directories[-1].startswith(_remove_directory): + del _remove_directory + + +def unique_paths(directory_paths:DirectoryPathList) -> DirectoryPathIterator: + realpaths = ( + real_path(directory_path) + for directory_path + in filter(bool, directory_paths) + ) + return { + real_directory_path: True + for real_directory_path + in filter(bool, realpaths) + }.keys() + + +def get_directories(*directory_paths: DirectoryPathList, fetch:bool=True, recursive:RecursiveType=True) -> DirectoryCollection: + directory_paths = unique_directories( + directory_paths, recursive=recursive + ) + directories = ( + get_directory(directory_path, fetch=fetch) + for directory_path + in directory_paths + ) + return filter( + bool, + directories + ) + + +def directory_files(*directories_or_paths: DirectoryPathList|DirectoryList, recursive: RecursiveType=True) -> FilePathIterator: + return itertools.chain.from_iterable( + itertools.chain( + directory_object.files, + [] + if not recursive + else itertools.chain.from_iterable( + directory_files(directory, recursive=recursive) + for directory + in filter( + bool, + map( + get_directory, + filter( + ( + ( bool if recursive else False ) + if not callable(recursive) + else recursive + ), + directory_object.directories + ) + ) + ) + ) + ) + for directory_object + in filter( + bool, + map( + get_directory, + directories_or_paths + ) + ) + ) + + +def extension_filter(ext_filter: Optional[ExtensionList]=None, ext_blacklist: Optional[ExtensionList]=None) -> ExtensionFilter: + if ext_filter: + ext_filter = [*map(str.upper, ext_filter)] + if ext_blacklist: + ext_blacklist = [*map(str.upper, ext_blacklist)] + def filter_functon(fp:str): + return (not ext_filter or any(fp.upper().endswith(ew) for ew in ext_filter)) and (not ext_blacklist or not any(fp.upper().endswith(ew) for ew in ext_blacklist)) + return filter_functon + + +def not_hidden(filepath: FilePath) -> IsHidden: + return not path.basename(filepath).startswith('.') + + +def filter_files(file_paths: FilePathList, ext_filter: Optional[ExtensionList]=None, ext_blacklist: Optional[ExtensionList]=None) -> FilePathIterator: + return filter(extension_filter(ext_filter, ext_blacklist), file_paths) + + +def list_files(*directory_paths:DirectoryPathList, ext_filter: Optional[ExtensionList]=None, ext_blacklist: Optional[ExtensionList]=None, recursive:RecursiveType=True) -> FilePathIterator: + return filter_files(itertools.chain.from_iterable( + directory_files(directory, recursive=recursive) + for directory + in get_directories( + *directory_paths, recursive=recursive + ) + ), ext_filter, ext_blacklist) + + +cache_folders = DirectoryCache({}) diff --git a/modules/hypernetworks/hypernetwork.py b/modules/hypernetworks/hypernetwork.py index 5afa572be..7352f3bc4 100644 --- a/modules/hypernetworks/hypernetwork.py +++ b/modules/hypernetworks/hypernetwork.py @@ -11,7 +11,7 @@ from torch import einsum from torch.nn.init import normal_, xavier_normal_, xavier_uniform_, kaiming_normal_, kaiming_uniform_, zeros_ from einops import rearrange, repeat from ldm.util import default -from modules import devices, processing, sd_models, shared, hashes, errors +from modules import devices, processing, sd_models, shared, hashes, errors, files_cache import modules.textual_inversion.dataset from modules.textual_inversion import textual_inversion, ti_logging from modules.textual_inversion.learn_schedule import LearnRateScheduler @@ -282,18 +282,16 @@ class Hypernetwork: def list_hypernetworks(path): - res = {} - def list_folder(folder): - for filename in os.listdir(folder): - fn = os.path.join(folder, filename) - if os.path.isfile(fn) and fn.lower().endswith(".pt"): - name = os.path.splitext(os.path.basename(fn))[0] - res[name] = fn - elif os.path.isdir(fn) and not fn.startswith('.'): - list_folder(fn) - - list_folder(path) - return res + hypernetworks = { + os.path.splitext(os.path.basename(hypernetwork_path))[0]: hypernetwork_path + for hypernetwork_path + in files_cache.list_files( + path, + ext_filter=['.pt'], + recursive=files_cache.not_hidden + ) + } + return hypernetworks def load_hypernetwork(name): diff --git a/modules/interrogate.py b/modules/interrogate.py index ba216a146..cd43653f1 100644 --- a/modules/interrogate.py +++ b/modules/interrogate.py @@ -8,7 +8,7 @@ import torch.hub # pylint: disable=ungrouped-imports from PIL import Image from torchvision import transforms from torchvision.transforms.functional import InterpolationMode -from modules import devices, paths, shared, lowvram, modelloader, errors +from modules import devices, paths, shared, lowvram, errors blip_image_eval_size = 384 @@ -80,6 +80,7 @@ class InterrogateModels: def load_blip_model(self): self.create_fake_fairscale() import models.blip # pylint: disable=no-name-in-module + import modules.modelloader as modelloader model_path = os.path.join(paths.models_path, "BLIP") download_name='model_base_caption_capfilt_large.pth', shared.log.debug(f'Model interrogate load: type=BLiP model={download_name} path={model_path}') diff --git a/modules/modelloader.py b/modules/modelloader.py index bab1ff306..b8f4662a9 100644 --- a/modules/modelloader.py +++ b/modules/modelloader.py @@ -9,54 +9,12 @@ import rich.progress as p from modules import shared, errors from modules.upscaler import Upscaler, UpscalerLanczos, UpscalerNearest, UpscalerNone from modules.paths import script_path, models_path +from modules.files_cache import list_files, unique_directories +from installer import print_dict diffuser_repos = [] -def walk(top, onerror:callable=None): - # A near-exact copy of `os.path.walk()`, trimmed slightly. Probably not nessesary for most people's collections, but makes a difference on really large datasets. - nondirs = [] - walk_dirs = [] - try: - scandir_it = os.scandir(top) - except OSError as error: - if onerror is not None: - onerror(error, top) - return - with scandir_it: - while True: - try: - try: - entry = next(scandir_it) - except StopIteration: - break - except OSError as error: - if onerror is not None: - onerror(error, top) - return - try: - is_dir = entry.is_dir() - except OSError: - is_dir = False - if not is_dir: - nondirs.append(entry.name) - else: - try: - if entry.is_symlink() and not os.path.exists(entry.path): - raise NotADirectoryError('Broken Symlink') - walk_dirs.append(entry.path) - except OSError as error: - if onerror is not None: - onerror(error, entry.path) - # Recurse into sub-directories - for new_path in walk_dirs: - if os.path.basename(new_path).startswith('models--'): - continue - yield from walk(new_path, onerror) - # Yield after recursion if going bottom up - yield top, nondirs - - def download_civit_meta(model_path: str, model_id): fn = os.path.splitext(model_path)[0] + '.json' url = f'https://civitai.com/api/v1/models/{model_id}' @@ -358,93 +316,6 @@ def load_civitai(model: str, url: str): return None -cache_folders = {} -cache_last = 0 -cache_time = 1 - - -def directory_updated(path:str, *, recursive:bool=True) -> bool: # pylint: disable=redefined-builtin - try: - path = os.path.abspath(path) - if path not in cache_folders: - return True - if cache_last > (time.time() - cache_time): - return False - if not (os.path.exists(path) and os.path.isdir(path) and os.path.getmtime(path) == cache_folders[path][0]): - return True - if recursive: - for folder in cache_folders: - if folder.startswith(path) and folder != path and not (os.path.exists(folder) and os.path.isdir(folder) and os.path.getmtime(folder) == cache_folders[folder][0]): - return True - except Exception as e: - shared.log.error(f"Filesystem Error: {e.__class__.__name__}({e})") - return True - return False - - -def directory_list(path:str, *, recursive:bool=True) -> dict[str,tuple[float,list[str]]]: # pylint: disable=redefined-builtin - path = os.path.abspath(path) - res = {} - if not os.path.exists(path): - return res - if directory_updated(path, recursive=recursive): - for folder in list(cache_folders): - del cache_folders[folder] - if os.path.exists(folder) or os.path.isdir(folder): - continue - for folder, files in walk(path, lambda e, path: shared.log.debug(f"FS walk error: {e} {path}")): - if not os.path.exists(folder): - continue - try: - mtime = os.path.getmtime(folder) - if folder not in cache_folders or mtime != cache_folders[folder][0]: - cache_folders[folder] = (mtime, [os.path.join(folder, fn) for fn in files]) - except Exception as e: - shared.log.error(f"Filesystem Error: {e.__class__.__name__}({e})") - del cache_folders[folder] - for folder in cache_folders: - if folder == path or (recursive and folder.startswith(path)): - res[folder] = cache_folders[folder] - if not recursive: - break - return res - - -def directory_mtime(path:str, *, recursive:bool=True) -> float: # pylint: disable=redefined-builtin - return float(max(0, *[mtime for mtime, _ in directory_list(path, recursive=recursive).values()])) - - -def directories_file_paths(directories:dict) -> list[str]: - return sum([dat[1] for dat in directories.values()],[]) - - -def directories_unique(directories:list[str], *, recursive:bool=True) -> list[str]: - '''Ensure no empty, or duplicates''' - directories = { os.path.abspath(path): True for path in directories if path }.keys() - if recursive: - '''If we are going recursive, then directories that are children of other directories are redundant''' - directories = [path for path in directories if not any(d != path and path.startswith(os.path.join(d,'')) for d in directories)] - return directories - - -def unique_paths(paths:list[str]) -> list[str]: - return { fp: True for fp in paths }.keys() - - -def directory_files(*directories:list[str], recursive:bool=True) -> list[str]: - return unique_paths(sum([[*directories_file_paths(directory_list(d, recursive=recursive))] for d in directories_unique(directories, recursive=recursive)],[])) - - -def extension_filter(ext_filter=None, ext_blacklist=None): - if ext_filter: - ext_filter = [*map(str.upper, ext_filter)] - if ext_blacklist: - ext_blacklist = [*map(str.upper, ext_blacklist)] - def filter(fp:str): # pylint: disable=redefined-builtin - return (not ext_filter or any(fp.upper().endswith(ew) for ew in ext_filter)) and (not ext_blacklist or not any(fp.upper().endswith(ew) for ew in ext_blacklist)) - return filter - - def download_url_to_file(url: str, dst: str): # based on torch.hub.download_url_to_file import uuid @@ -513,10 +384,10 @@ def load_models(model_path: str, model_url: str = None, command_path: str = None @param ext_filter: An optional list of filename extensions to filter by @return: A list of paths containing the desired model(s) """ - places = directories_unique([model_path, command_path]) + places = [model_path, command_path] output = [] try: - output:list = [*filter(extension_filter(ext_filter, ext_blacklist), directory_files(*places))] + output:list = [*list_files(*places, ext_filter=ext_filter, ext_blacklist=ext_blacklist)] if model_url is not None and len(output) == 0: if download_name is not None: dl = load_file_from_url(model_url, model_dir=places[0], progress=True, file_name=download_name) @@ -524,7 +395,7 @@ def load_models(model_path: str, model_url: str = None, command_path: str = None else: output.append(model_url) except Exception as e: - shared.log.error(f"Error listing models: {places} {e}") + errors.display(e,f"Error listing models: {unique_directories(places)}") return output diff --git a/modules/sd_models.py b/modules/sd_models.py index 0a98a5cc4..da83d11d2 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -143,7 +143,7 @@ def list_models(): ext_filter = [".safetensors"] else: ext_filter = [".ckpt", ".safetensors"] - model_list = modelloader.load_models(model_path=model_path, model_url=None, command_path=shared.opts.ckpt_dir, ext_filter=ext_filter, download_name=None, ext_blacklist=[".vae.ckpt", ".vae.safetensors"]) + model_list = list(modelloader.load_models(model_path=model_path, model_url=None, command_path=shared.opts.ckpt_dir, ext_filter=ext_filter, download_name=None, ext_blacklist=[".vae.ckpt", ".vae.safetensors"])) if shared.backend == shared.Backend.DIFFUSERS: model_list += modelloader.load_diffusers_models(model_path=os.path.join(models_path, 'Diffusers'), command_path=shared.opts.diffusers_dir, clear=True) for filename in sorted(model_list, key=str.lower): diff --git a/modules/textual_inversion/textual_inversion.py b/modules/textual_inversion/textual_inversion.py index 97467df2b..81fbb26d5 100644 --- a/modules/textual_inversion/textual_inversion.py +++ b/modules/textual_inversion/textual_inversion.py @@ -15,7 +15,8 @@ import modules.textual_inversion.loaders from modules.textual_inversion.learn_schedule import LearnRateScheduler from modules.textual_inversion.image_embedding import embedding_to_b64, embedding_from_b64, insert_image_data_embed, extract_image_data_embed, caption_image_overlay from modules.textual_inversion.ti_logging import save_settings_to_file -from modules.modelloader import directory_files, directory_mtime, extension_filter +from typing import List, Optional, Union +from modules.files_cache import directory_files, directory_mtime, extension_filter TokenToAdd = namedtuple("TokenToAdd", ["clip_l", "clip_g"]) diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index 84cc5afca..bc6e16217 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -15,7 +15,7 @@ from collections import OrderedDict import gradio as gr from PIL import Image from starlette.responses import FileResponse, JSONResponse -from modules import paths, shared, scripts, modelloader, errors +from modules import paths, shared, scripts, files_cache, errors from modules.ui_components import ToolButton import modules.ui_symbols as symbols @@ -149,11 +149,7 @@ class ExtraNetworksPage: return preview def is_empty(self, folder): - for f in shared.listdir(folder): - _fn, ext = os.path.splitext(f) - if ext.lower() in ['.ckpt', '.safetensors', '.pt', '.json'] or os.path.isdir(os.path.join(folder, f)): - return False - return True + return any(files_cache.list_files(folder, ext_filter=['.ckpt', '.safetensors', '.pt', '.json'])) def create_thumb(self): debug(f'EN create-thumb: {self.name}') @@ -216,8 +212,9 @@ class ExtraNetworksPage: return f"
Extra network page not ready
Click refresh to try again
" subdirs = {} allowed_folders = [os.path.abspath(x) for x in self.allowed_directories_for_previews()] - for parentdir, dirs in {d: modelloader.directory_list(d) for d in allowed_folders}.items(): - for tgt in dirs.keys(): + for parentdir, dirs in {d: files_cache.walk(d, cached=True, recurse=files_cache.not_hidden) for d in allowed_folders}.items(): + for tgt in dirs: + tgt = tgt.path if os.path.join(paths.models_path, 'Reference') in tgt: subdirs['Reference'] = 1 if shared.backend == shared.Backend.DIFFUSERS and shared.opts.diffusers_dir in tgt: @@ -227,9 +224,10 @@ class ExtraNetworksPage: subdir = tgt[len(parentdir):].replace("\\", "/") while subdir.startswith("/"): subdir = subdir[1:] + if not subdir: + continue # if not self.is_empty(tgt): - if not subdir.startswith("."): - subdirs[subdir] = 1 + subdirs[subdir] = 1 debug(f"Extra networks: page='{self.name}' subfolders={list(subdirs)}") subdirs = OrderedDict(sorted(subdirs.items())) if self.name == 'model': @@ -295,17 +293,17 @@ class ExtraNetworksPage: return "" def find_preview_file(self, path): + exts = ["jpg", "jpeg", "png", "webp", "tiff", "jp2"] if path is None: return 'html/card-no-preview.png' if shared.opts.diffusers_dir in path: path = os.path.relpath(path, shared.opts.diffusers_dir) ref = os.path.join('models', 'Reference') fn = os.path.join(ref, path.replace('models--', '').replace('\\', '/').split('/')[0]) - files = shared.listdir(ref) + files = list(files_cache.list_files(ref, ext_filter=exts, recursive=False)) else: - files = shared.listdir(os.path.dirname(path)) + files = list(files_cache.list_files(os.path.dirname(path), ext_filter=exts, recursive=False)) fn = os.path.splitext(path)[0] - exts = ["jpg", "jpeg", "png", "webp", "tiff", "jp2"] for file in [f'{fn}{mid}{ext}' for ext in exts for mid in ['.thumb.', '.', '.preview.']]: if file in files: if 'Reference' not in file and '.thumb.' not in file: @@ -328,7 +326,7 @@ class ExtraNetworksPage: self.text += '\n' fn = os.path.splitext(path)[0] + '.txt' - if fn in shared.listdir(os.path.dirname(path)): + if os.path.exists(fn): try: with open(fn, "r", encoding="utf-8", errors="replace") as f: txt = f.read() @@ -348,7 +346,7 @@ class ExtraNetworksPage: def find_info(self, path): fn = os.path.splitext(path)[0] + '.json' data = {} - if fn in shared.listdir(os.path.dirname(path)): + if os.path.exists(fn): t0 = time.time() data = shared.readfile(fn, silent=True) if type(data) is list: diff --git a/modules/ui_extra_networks_textual_inversion.py b/modules/ui_extra_networks_textual_inversion.py index 5db980db0..4eadfd34f 100644 --- a/modules/ui_extra_networks_textual_inversion.py +++ b/modules/ui_extra_networks_textual_inversion.py @@ -1,7 +1,7 @@ import json import os import concurrent -from modules import shared, sd_hijack, sd_models, ui_extra_networks +from modules import shared, sd_hijack, sd_models, ui_extra_networks, files_cache from modules.textual_inversion.textual_inversion import Embedding @@ -45,20 +45,16 @@ class ExtraNetworksPageTextualInversion(ui_extra_networks.ExtraNetworksPage): return record def list_items(self): - - def list_folder(folder): - for filename in os.listdir(folder): - fn = os.path.join(folder, filename) - if os.path.isfile(fn) and (fn.lower().endswith(".pt") or fn.lower().endswith(".safetensors")): - embedding = Embedding(vec=0, name=os.path.basename(fn), filename=fn) - embedding.filename = fn - self.embeddings.append(embedding) - elif os.path.isdir(fn) and not fn.startswith('.'): - list_folder(fn) - if sd_models.model_data.sd_model is None: - self.embeddings = [] - list_folder(shared.opts.embeddings_dir) + self.embeddings = [ + Embedding(vec=0, name=os.path.basename(embedding_path), filename=embedding_path) + for embedding_path + in files_cache.list_files( + shared.opts.embeddings_dir, + ext_filter=['.pt', '.safetensors'], + recursive=files_cache.not_hidden + ) + ] elif shared.backend == shared.Backend.ORIGINAL: self.embeddings = list(sd_hijack.model_hijack.embedding_db.word_embeddings.values()) elif hasattr(sd_models.model_data.sd_model, 'embedding_db'):