From b5d396f4b1876f5233a57f1db31304c993706921 Mon Sep 17 00:00:00 2001 From: Henry Ruhs Date: Sat, 12 Sep 2026 12:48:42 +0200 Subject: [PATCH] refactor inference manager, kill app context (#1238) --- facefusion/apis/endpoints/session.py | 4 +- facefusion/app_context.py | 18 ------- facefusion/core.py | 3 +- facefusion/inference_manager.py | 57 ++++++++++++-------- facefusion/types.py | 2 - tests/test_face_creator.py | 3 +- tests/test_face_detector.py | 3 +- tests/test_face_tracker.py | 3 +- tests/test_inference_manager.py | 79 ++++++++++++++++++++++------ 9 files changed, 111 insertions(+), 61 deletions(-) delete mode 100644 facefusion/app_context.py diff --git a/facefusion/apis/endpoints/session.py b/facefusion/apis/endpoints/session.py index 7761aeec..59f09770 100644 --- a/facefusion/apis/endpoints/session.py +++ b/facefusion/apis/endpoints/session.py @@ -4,7 +4,7 @@ from starlette.requests import Request from starlette.responses import JSONResponse from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_401_UNAUTHORIZED, HTTP_404_NOT_FOUND -from facefusion import content_store, process_manager, session_context, session_manager, state_manager, translator +from facefusion import content_store, inference_manager, process_manager, session_context, session_manager, state_manager, translator from facefusion.apis import asset_store from facefusion.apis.session_helper import validate_api_key from facefusion.filesystem import is_directory, remove_directory @@ -21,6 +21,7 @@ async def create_session(request : Request) -> JSONResponse: state_manager.init() content_store.init() + inference_manager.init() process_manager.init() return JSONResponse( @@ -83,6 +84,7 @@ async def destroy_session(request : Request) -> JSONResponse: state_manager.clear() content_store.clear() + inference_manager.clear() process_manager.clear() return JSONResponse( diff --git a/facefusion/app_context.py b/facefusion/app_context.py deleted file mode 100644 index 13bcdfa1..00000000 --- a/facefusion/app_context.py +++ /dev/null @@ -1,18 +0,0 @@ -import os -import sys - -from facefusion.types import AppContext - - -def detect_app_context() -> AppContext: - jobs_path = os.path.join('facefusion', 'jobs') - apis_path = os.path.join('facefusion', 'apis') - frame = sys._getframe(1) - - while frame: - if jobs_path in frame.f_code.co_filename: - return 'cli' - if apis_path in frame.f_code.co_filename: - return 'api' - frame = frame.f_back - return 'cli' diff --git a/facefusion/core.py b/facefusion/core.py index 222b1076..36317590 100755 --- a/facefusion/core.py +++ b/facefusion/core.py @@ -8,7 +8,7 @@ from time import time import uvicorn import facefusion.apis.core -from facefusion import args_helper, benchmarker, cli_helper, content_analyser, content_store, hash_helper, logger, process_manager, state_manager, translator +from facefusion import args_helper, benchmarker, cli_helper, content_analyser, content_store, hash_helper, inference_manager, logger, process_manager, state_manager, translator from facefusion.args_helper import apply_args from facefusion.download import conditional_download_hashes, conditional_download_sources from facefusion.exit_helper import hard_exit, signal_exit @@ -39,6 +39,7 @@ def cli() -> None: logger.init(state_manager.get_item('log_level')) content_store.init() + inference_manager.init() process_manager.init() route(args) diff --git a/facefusion/inference_manager.py b/facefusion/inference_manager.py index 57210cf5..7bd9ca59 100644 --- a/facefusion/inference_manager.py +++ b/facefusion/inference_manager.py @@ -2,50 +2,59 @@ import importlib import random from functools import lru_cache from time import sleep, time -from typing import List +from typing import List, Optional from onnxruntime import InferenceSession -from facefusion import logger, process_manager, state_manager, translator -from facefusion.app_context import detect_app_context +from facefusion import logger, process_manager, state_manager, store_creator, translator from facefusion.common_helper import is_windows from facefusion.execution import create_inference_providers, get_onnxruntime_version, has_execution_provider from facefusion.exit_helper import fatal_exit from facefusion.filesystem import get_file_name, is_file +from facefusion.session_context import get_session_id from facefusion.time_helper import calculate_end_time -from facefusion.types import DownloadSet, ExecutionProvider, InferencePool, InferencePoolSet, InferenceProvider +from facefusion.types import DownloadSet, ExecutionProvider, InferencePool, InferenceProvider, Store -INFERENCE_POOL_SET : InferencePoolSet =\ -{ - 'cli': {}, - 'api': {} -} +INFERENCE_POOL_STORE : Store = store_creator.create_store({}) + + +def init() -> None: + session_id = get_session_id() + store_creator.init_content(INFERENCE_POOL_STORE, session_id) def get_inference_pool(module_name : str, model_names : List[str], model_source_set : DownloadSet) -> InferencePool: while process_manager.is_checking(): sleep(0.5) + session_id = get_session_id() + inference_pool_set = store_creator.get_content(INFERENCE_POOL_STORE, session_id) execution_device_ids = state_manager.get_item('execution_device_ids') execution_providers = state_manager.get_item('execution_providers') has_arena_leak = has_execution_provider('cuda') and get_onnxruntime_version() > (1, 24, 4) - app_context = detect_app_context() for execution_device_id in execution_device_ids: inference_context = get_inference_context(module_name, model_names, execution_device_id, execution_providers) if not has_arena_leak: - if app_context == 'cli' and INFERENCE_POOL_SET.get('api').get(inference_context): - INFERENCE_POOL_SET['cli'][inference_context] = INFERENCE_POOL_SET.get('api').get(inference_context) - if app_context == 'api' and INFERENCE_POOL_SET.get('cli').get(inference_context): - INFERENCE_POOL_SET['api'][inference_context] = INFERENCE_POOL_SET.get('cli').get(inference_context) + inference_pool = find_inference_pool(inference_context) - if not INFERENCE_POOL_SET.get(app_context).get(inference_context): + if inference_pool: + inference_pool_set[inference_context] = inference_pool + + if not inference_pool_set.get(inference_context): inference_providers = resolve_static_inference_providers(module_name, execution_device_id) - INFERENCE_POOL_SET[app_context][inference_context] = create_inference_pool(model_source_set, inference_providers) + inference_pool_set[inference_context] = create_inference_pool(model_source_set, inference_providers) current_inference_context = get_inference_context(module_name, model_names, random.choice(execution_device_ids), execution_providers) - return INFERENCE_POOL_SET.get(app_context).get(current_inference_context) + return inference_pool_set.get(current_inference_context) + + +def find_inference_pool(inference_context : str) -> Optional[InferencePool]: + for inference_pool_set in INFERENCE_POOL_STORE.get('content_set').values(): + if inference_pool_set.get(inference_context): + return inference_pool_set.get(inference_context) + return None def create_inference_pool(model_source_set : DownloadSet, inference_providers : List[InferenceProvider]) -> InferencePool: @@ -61,18 +70,24 @@ def create_inference_pool(model_source_set : DownloadSet, inference_providers : def clear_inference_pool(module_name : str, model_names : List[str]) -> None: + session_id = get_session_id() + inference_pool_set = store_creator.get_content(INFERENCE_POOL_STORE, session_id) execution_device_ids = state_manager.get_item('execution_device_ids') execution_providers = state_manager.get_item('execution_providers') - app_context = detect_app_context() if is_windows() and has_execution_provider('directml'): - INFERENCE_POOL_SET[app_context].clear() + inference_pool_set.clear() for execution_device_id in execution_device_ids: inference_context = get_inference_context(module_name, model_names, execution_device_id, execution_providers) - if INFERENCE_POOL_SET.get(app_context).get(inference_context): - del INFERENCE_POOL_SET[app_context][inference_context] + if inference_pool_set.get(inference_context): + del inference_pool_set[inference_context] + + +def clear() -> None: + session_id = get_session_id() + store_creator.init_content(INFERENCE_POOL_STORE, session_id) def create_inference_session(model_path : str, inference_providers : List[InferenceProvider]) -> InferenceSession: diff --git a/facefusion/types.py b/facefusion/types.py index 303844ae..667f5f65 100755 --- a/facefusion/types.py +++ b/facefusion/types.py @@ -501,10 +501,8 @@ DownloadSet : TypeAlias = Dict[str, Download] VideoMemoryStrategy = Literal['strict', 'moderate', 'tolerant'] ApiSecurityStrategy = Literal['strict', 'moderate'] -AppContext = Literal['cli', 'api'] InferencePool : TypeAlias = Dict[str, InferenceSession] -InferencePoolSet : TypeAlias = Dict[AppContext, Dict[str, InferencePool]] JobOutputSet : TypeAlias = Dict[str, List[str]] JobStatus = Literal['drafted', 'queued', 'completed', 'failed'] diff --git a/tests/test_face_creator.py b/tests/test_face_creator.py index 8d76bab3..b24e6cc2 100644 --- a/tests/test_face_creator.py +++ b/tests/test_face_creator.py @@ -2,7 +2,7 @@ import numpy import pytest -from facefusion import face_aligner, face_classifier, face_detector, face_recognizer, ffmpeg, ffmpeg_builder, process_manager, state_manager +from facefusion import face_aligner, face_classifier, face_detector, face_recognizer, ffmpeg, ffmpeg_builder, inference_manager, process_manager, state_manager from facefusion.download import conditional_download from facefusion.face_creator import average_face_geometry, get_many_faces, get_one_face, refill_faces from facefusion.face_store import clear_faces @@ -13,6 +13,7 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: state_manager.init() + inference_manager.init() process_manager.start() conditional_download(get_test_examples_directory(), diff --git a/tests/test_face_detector.py b/tests/test_face_detector.py index c73afad4..398bfdb0 100644 --- a/tests/test_face_detector.py +++ b/tests/test_face_detector.py @@ -1,7 +1,7 @@ import pytest -from facefusion import face_detector, ffmpeg, ffmpeg_builder, process_manager, state_manager +from facefusion import face_detector, ffmpeg, ffmpeg_builder, inference_manager, process_manager, state_manager from facefusion.download import conditional_download from facefusion.face_detector import detect_with_retinaface, detect_with_scrfd, detect_with_yolo_face, detect_with_yunet from facefusion.face_helper import apply_nms, get_nms_threshold @@ -12,6 +12,7 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: state_manager.init() + inference_manager.init() process_manager.start() conditional_download(get_test_examples_directory(), diff --git a/tests/test_face_tracker.py b/tests/test_face_tracker.py index ce4e2ee6..55524ac4 100644 --- a/tests/test_face_tracker.py +++ b/tests/test_face_tracker.py @@ -1,7 +1,7 @@ import numpy import pytest -from facefusion import face_aligner, face_classifier, face_detector, face_recognizer, state_manager +from facefusion import face_aligner, face_classifier, face_detector, face_recognizer, inference_manager, state_manager from facefusion.common_helper import get_first, get_last from facefusion.download import conditional_download from facefusion.face_creator import get_many_faces, get_one_face @@ -14,6 +14,7 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: state_manager.init() + inference_manager.init() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_inference_manager.py b/tests/test_inference_manager.py index 13de7066..431f45db 100644 --- a/tests/test_inference_manager.py +++ b/tests/test_inference_manager.py @@ -1,12 +1,13 @@ from types import SimpleNamespace +from typing import Iterator from unittest.mock import Mock, patch import pytest from onnxruntime import InferenceSession -from facefusion import content_analyser, state_manager +from facefusion import content_analyser, session_context, state_manager, store_creator from facefusion.execution import resolve_cache_path -from facefusion.inference_manager import get_inference_pool, resolve_static_inference_providers +from facefusion.inference_manager import INFERENCE_POOL_STORE, clear, get_inference_pool, init, resolve_static_inference_providers @pytest.fixture(scope = 'module', autouse = True) @@ -17,33 +18,81 @@ def before_all() -> None: state_manager.init_item('download_providers', [ 'github' ]) +@pytest.fixture(scope = 'function', autouse = True) +def before_each() -> Iterator[None]: + local_id = session_context.resolve_local_id() + + for session_id in list(INFERENCE_POOL_STORE.get('content_set').keys()): + store_creator.delete_content(INFERENCE_POOL_STORE, session_id) + + session_context.set_session_id(local_id) + init() + + yield + + session_context.set_session_id(local_id) + + +def test_init() -> None: + local_id = session_context.resolve_local_id() + model_names = [ 'nsfw_1', 'nsfw_2', 'nsfw_3' ] + _, model_source_set = content_analyser.collect_model_downloads() + + local_inference_pool = get_inference_pool('facefusion.content_analyser', model_names, model_source_set) + session_context.set_session_id('session-a') + state_manager.init() + init() + session_inference_pool = get_inference_pool('facefusion.content_analyser', model_names, model_source_set) + + assert session_inference_pool.get('nsfw_1') is local_inference_pool.get('nsfw_1') + + session_context.set_session_id(local_id) + + assert get_inference_pool('facefusion.content_analyser', model_names, model_source_set) is local_inference_pool + + def test_get_inference_pool() -> None: model_names = [ 'nsfw_1', 'nsfw_2', 'nsfw_3' ] _, model_source_set = content_analyser.collect_model_downloads() with patch('facefusion.inference_manager.has_execution_provider', return_value = True): with patch('facefusion.inference_manager.get_onnxruntime_version', return_value = (1, 26, 0)): + session_context.set_session_id('session-a') + state_manager.init() + init() + session_a_inference_pool = get_inference_pool('facefusion.content_analyser', model_names, model_source_set) - with patch('facefusion.inference_manager.detect_app_context', return_value = 'cli'): - cli_inference_pool = get_inference_pool('facefusion.content_analyser', model_names, model_source_set) + assert isinstance(session_a_inference_pool.get('nsfw_1'), InferenceSession) - assert isinstance(cli_inference_pool.get('nsfw_1'), InferenceSession) + session_context.set_session_id('session-b') + state_manager.init() + init() + session_b_inference_pool = get_inference_pool('facefusion.content_analyser', model_names, model_source_set) - with patch('facefusion.inference_manager.detect_app_context', return_value = 'api'): - api_inference_pool = get_inference_pool('facefusion.content_analyser', model_names, model_source_set) - - assert isinstance(api_inference_pool.get('nsfw_1'), InferenceSession) - - assert not (cli_inference_pool.get('nsfw_1') is api_inference_pool.get('nsfw_1')) + assert isinstance(session_b_inference_pool.get('nsfw_1'), InferenceSession) + assert not session_a_inference_pool.get('nsfw_1') is session_b_inference_pool.get('nsfw_1') with patch('facefusion.inference_manager.get_onnxruntime_version', return_value = (1, 24, 4)): + session_context.set_session_id('session-c') + state_manager.init() + init() + session_c_inference_pool = get_inference_pool('facefusion.content_analyser', model_names, model_source_set) - with patch('facefusion.inference_manager.detect_app_context', return_value = 'api'): - api_inference_pool = get_inference_pool('facefusion.content_analyser', model_names, model_source_set) + assert isinstance(session_c_inference_pool.get('nsfw_1'), InferenceSession) + assert session_c_inference_pool.get('nsfw_1') is session_a_inference_pool.get('nsfw_1') - assert isinstance(api_inference_pool.get('nsfw_1'), InferenceSession) - assert cli_inference_pool.get('nsfw_1') is api_inference_pool.get('nsfw_1') +def test_clear() -> None: + model_names = [ 'nsfw_1', 'nsfw_2', 'nsfw_3' ] + _, model_source_set = content_analyser.collect_model_downloads() + + session_context.set_session_id('session-a') + state_manager.init() + init() + get_inference_pool('facefusion.content_analyser', model_names, model_source_set) + clear() + + assert store_creator.get_content(INFERENCE_POOL_STORE, 'session-a') == {} @pytest.fixture