mirror of
https://github.com/facefusion/facefusion.git
synced 2026-09-15 20:15:28 +02:00
refactor inference manager, kill app context (#1238)
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,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(),
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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(),
|
||||
[
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user