refactor inference manager, kill app context (#1238)

This commit is contained in:
Henry Ruhs
2026-09-12 12:48:42 +02:00
committed by GitHub
parent fb21d1d1b5
commit b5d396f4b1
9 changed files with 111 additions and 61 deletions
+3 -1
View File
@@ -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(
-18
View File
@@ -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'
+2 -1
View File
@@ -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)
+36 -21
View File
@@ -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:
-2
View File
@@ -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']
+2 -1
View File
@@ -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(),
+2 -1
View File
@@ -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(),
+2 -1
View File
@@ -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(),
[
+64 -15
View File
@@ -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