simplify async code with ThreadPoolExecutor

This commit is contained in:
Alex "mcmonkey" Goodwin 2024-09-19 17:12:50 +09:00
parent b8022cf02b
commit 5d8a0b7afd

View File

@ -8,6 +8,7 @@ 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'}
@ -197,34 +198,24 @@ def recursive_search(directory: str, excluded_dir_names: list[str] | None=None)
logging.debug("recursive file list on directory {}".format(directory)) logging.debug("recursive file list on directory {}".format(directory))
async def proc_subdir(path: str): with ThreadPoolExecutor() as executor:
dirs[path] = await AsyncFiles.getmtime(path)
def proc_thread(): def proc_subdir(path: str):
asyncio.set_event_loop(asyncio.new_event_loop()) dirs[path] = os.path.getmtime(path)
calls = []
async def handle(file): def handle(file):
if not await AsyncFiles.isdir(file): if not os.path.isdir(file):
relative_path = await AsyncFiles.relpath(file, directory) relative_path = os.path.relpath(file, directory)
result.append(relative_path) result.append(relative_path)
return return
calls.append(proc_subdir(file)) executor.submit(lambda: proc_subdir(file))
for subdir in await AsyncFiles.listdir(file): for subdir in os.listdir(file):
path = os.path.join(file, subdir) path = os.path.join(file, subdir)
if subdir not in excluded_dir_names: if subdir not in excluded_dir_names:
calls.append(handle(path)) executor.submit(lambda: handle(path))
calls.append(handle(directory))
while len(calls) > 0: executor.submit(lambda: handle(directory))
future = asyncio.gather(*calls) executor.shutdown(wait=True)
calls = []
asyncio.get_event_loop().run_until_complete(future)
asyncio.get_event_loop().close()
thread = threading.Thread(target=proc_thread)
thread.start()
thread.join()
logging.debug("found {} files".format(len(result))) logging.debug("found {} files".format(len(result)))
return result, dirs return result, dirs
@ -281,35 +272,26 @@ def cached_filename_list_(folder_name: str) -> tuple[list[str], dict[str, float]
must_invalidate = threading.Event() must_invalidate = threading.Event()
folders = folder_names_and_paths[folder_name] folders = folder_names_and_paths[folder_name]
async def check_folder_mtime(folder: str, time_modified: float): with ThreadPoolExecutor() as executor:
if await AsyncFiles.getmtime(folder) != time_modified:
must_invalidate.set()
async def check_new_dirs(x: str): def check_folder_mtime(folder: str, time_modified: float):
if await AsyncFiles.isdir(x): if os.path.getmtime(folder) != time_modified:
if x not in out[1]:
must_invalidate.set() must_invalidate.set()
def proc_thread(): def check_new_dirs(x: str):
asyncio.set_event_loop(asyncio.new_event_loop()) if os.path.isdir(x):
calls = [] if x not in out[1]:
must_invalidate.set()
for x in out[1]: for x in out[1]:
time_modified = out[1][x] time_modified = out[1][x]
call = check_folder_mtime(x, time_modified) executor.submit(lambda: check_folder_mtime(x, time_modified))
calls.append(call)
for x in folders[0]: for x in folders[0]:
call = check_new_dirs(x) executor.submit(lambda: check_new_dirs(x))
calls.append(call)
future = asyncio.gather(*calls) executor.shutdown(wait=True)
asyncio.get_event_loop().run_until_complete(future)
asyncio.get_event_loop().close()
thread = threading.Thread(target=proc_thread)
thread.start()
thread.join()
if must_invalidate.is_set(): if must_invalidate.is_set():
return None return None
@ -370,35 +352,3 @@ def get_save_image_path(filename_prefix: str, output_dir: str, image_width=0, im
os.makedirs(full_output_folder, exist_ok=True) os.makedirs(full_output_folder, exist_ok=True)
counter = 1 counter = 1
return full_output_folder, filename, counter, subfolder, filename_prefix return full_output_folder, filename, counter, subfolder, filename_prefix
def aio_wrap(func):
@wraps(func)
async def run(*args, loop=None, executor=None, **kwargs):
if loop is None:
loop = asyncio.get_running_loop()
pfunc = partial(func, *args, **kwargs)
return await loop.run_in_executor(executor, pfunc)
return run
class AsyncFiles:
@staticmethod
@aio_wrap
def listdir(path: str) -> list[str]:
return os.listdir(path)
@staticmethod
@aio_wrap
def isdir(path: str) -> bool:
return os.path.isdir(path)
@staticmethod
@aio_wrap
def getmtime(path: str) -> float:
return os.path.getmtime(path)
@staticmethod
@aio_wrap
def relpath(file: str, directory: str) -> str:
return os.path.relpath(file, directory)