mirror of
https://github.com/facefusion/facefusion.git
synced 2026-09-20 22:40:42 +02:00
refactor inference manager, kill app context (#1238)
This commit is contained in:
@@ -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