Merge 67a69da55f6748b4eea9fdced6aa53d8761e0a2d into 614377abd6c018dec4aeb5700c0e203e2f91f9b0

This commit is contained in:
Alex "mcmonkey" Goodwin 2024-10-08 10:15:54 +13:00 committed by GitHub
commit 87718c7414
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 69 additions and 24 deletions

View File

@ -1,11 +1,13 @@
from __future__ import annotations from __future__ import annotations
import threading
import os import os
import time import time
import mimetypes import mimetypes
import logging import logging
from typing import Set, List, Dict, Tuple, Literal from typing import Set, List, Dict, Tuple, Literal
from collections.abc import Collection from collections.abc import Collection
from concurrent.futures import ThreadPoolExecutor
supported_pt_extensions: set[str] = {'.ckpt', '.pt', '.bin', '.pth', '.safetensors', '.pkl', '.sft'} supported_pt_extensions: set[str] = {'.ckpt', '.pt', '.bin', '.pth', '.safetensors', '.pkl', '.sft'}
@ -46,6 +48,8 @@ user_directory = os.path.join(os.path.dirname(os.path.realpath(__file__)), "user
filename_list_cache: dict[str, tuple[list[str], dict[str, float], float]] = {} filename_list_cache: dict[str, tuple[list[str], dict[str, float], float]] = {}
async_executor = ThreadPoolExecutor(32)
class CacheHelper: class CacheHelper:
""" """
Helper class for managing file list cache data. Helper class for managing file list cache data.
@ -210,6 +214,22 @@ def get_folder_paths(folder_name: str) -> list[str]:
folder_name = map_legacy(folder_name) folder_name = map_legacy(folder_name)
return folder_names_and_paths[folder_name][0][:] return folder_names_and_paths[folder_name][0][:]
def prebuild_lists():
start_time = time.perf_counter()
with ThreadPoolExecutor(32) as executor:
calls = []
for folder_name in folder_names_and_paths:
calls.append(executor.submit(lambda: get_filename_list(folder_name)))
for call in calls:
call.result()
end_time = time.perf_counter()
logging.info(f"Scanned model lists in {end_time - start_time:.2f} seconds")
def recursive_search(directory: str, excluded_dir_names: list[str] | None=None) -> tuple[list[str], dict[str, float]]: def recursive_search(directory: str, excluded_dir_names: list[str] | None=None) -> tuple[list[str], dict[str, float]]:
if not os.path.isdir(directory): if not os.path.isdir(directory):
return [], {} return [], {}
@ -225,33 +245,43 @@ def recursive_search(directory: str, excluded_dir_names: list[str] | None=None)
dirs[directory] = os.path.getmtime(directory) dirs[directory] = os.path.getmtime(directory)
except FileNotFoundError: except FileNotFoundError:
logging.warning(f"Warning: Unable to access {directory}. Skipping this path.") logging.warning(f"Warning: Unable to access {directory}. Skipping this path.")
return [], {}
logging.debug("recursive file list on directory {}".format(directory)) logging.debug("recursive file list on directory {}".format(directory))
dirpath: str
subdirs: list[str]
filenames: list[str]
for dirpath, subdirs, filenames in os.walk(directory, followlinks=True, topdown=True): calls = []
subdirs[:] = [d for d in subdirs if d not in excluded_dir_names]
for file_name in filenames: def proc_subdir(path: str):
relative_path = os.path.relpath(os.path.join(dirpath, file_name), directory) dirs[path] = os.path.getmtime(path)
result.append(relative_path)
def handle(file):
try:
if not os.path.isdir(file):
relative_path = os.path.relpath(file, directory)
result.append(relative_path)
return
calls.append(async_executor.submit(lambda f=file: proc_subdir(f)))
for subdir in os.listdir(file):
if subdir not in excluded_dir_names:
path = os.path.join(file, subdir)
calls.append(async_executor.submit(lambda p=path: handle(p)))
except Exception as e:
logging.error(f"recursive_search encountered error while handling '{file}': {e}")
calls.append(async_executor.submit(lambda: handle(directory)))
while len(calls) > 0:
calls.pop().result()
for d in subdirs:
path: str = os.path.join(dirpath, d)
try:
dirs[path] = os.path.getmtime(path)
except FileNotFoundError:
logging.warning(f"Warning: Unable to access {path}. Skipping this path.")
continue
logging.debug("found {} files".format(len(result))) logging.debug("found {} files".format(len(result)))
return result, dirs return result, dirs
def filter_files_extensions(files: Collection[str], extensions: Collection[str]) -> list[str]: def filter_files_extensions(files: Collection[str], extensions: Collection[str]) -> list[str]:
return sorted(list(filter(lambda a: os.path.splitext(a)[-1].lower() in extensions or len(extensions) == 0, files))) return sorted(list(filter(lambda a: os.path.splitext(a)[-1].lower() in extensions or len(extensions) == 0, files)))
def get_full_path(folder_name: str, filename: str) -> str | None: def get_full_path(folder_name: str, filename: str) -> str | None:
global folder_names_and_paths global folder_names_and_paths
folder_name = map_legacy(folder_name) folder_name = map_legacy(folder_name)
@ -300,19 +330,32 @@ def cached_filename_list_(folder_name: str) -> tuple[list[str], dict[str, float]
if folder_name not in filename_list_cache: if folder_name not in filename_list_cache:
return None return None
out = filename_list_cache[folder_name] out = filename_list_cache[folder_name]
must_invalidate = threading.Event()
folders = folder_names_and_paths[folder_name]
def check_folder_mtime(folder: str, time_modified: float):
if os.path.getmtime(folder) != time_modified:
must_invalidate.set()
def check_new_dirs(x: str):
if os.path.isdir(x):
if x not in out[1]:
must_invalidate.set()
calls = []
for x in out[1]: for x in out[1]:
time_modified = out[1][x] time_modified = out[1][x]
folder = x calls.append(async_executor.submit(lambda f=x, t=time_modified: check_folder_mtime(f, t)))
if os.path.getmtime(folder) != time_modified:
return None
folders = folder_names_and_paths[folder_name]
for x in folders[0]: for x in folders[0]:
if os.path.isdir(x): calls.append(async_executor.submit(lambda f=x: check_new_dirs(f)))
if x not in out[1]:
return None
for call in calls:
call.result()
if must_invalidate.is_set():
return None
return out return out
def get_filename_list(folder_name: str) -> list[str]: def get_filename_list(folder_name: str) -> list[str]:

View File

@ -212,6 +212,8 @@ if __name__ == "__main__":
nodes.init_extra_nodes(init_custom_nodes=not args.disable_all_custom_nodes) nodes.init_extra_nodes(init_custom_nodes=not args.disable_all_custom_nodes)
folder_paths.prebuild_lists()
cuda_malloc_warning() cuda_malloc_warning()
server.add_routes() server.add_routes()