diff --git a/tests/test_inference_manager.py b/tests/test_inference_manager.py index 1dfbaa4c..6e7e1e60 100644 --- a/tests/test_inference_manager.py +++ b/tests/test_inference_manager.py @@ -1,10 +1,12 @@ -from unittest.mock import patch +from types import SimpleNamespace +from unittest.mock import Mock, patch import pytest from onnxruntime import InferenceSession from facefusion import content_analyser, state_manager -from facefusion.inference_manager import INFERENCE_POOL_SET, get_inference_pool +from facefusion.execution import resolve_cache_path +from facefusion.inference_manager import get_inference_pool, resolve_static_inference_providers @pytest.fixture(scope = 'module', autouse = True) @@ -13,21 +15,60 @@ def before_all() -> None: state_manager.init_item('execution_providers', [ 'cpu' ]) state_manager.init_item('download_providers', [ 'github' ]) - content_analyser.pre_check() - 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.detect_app_context', return_value = 'cli'): - get_inference_pool('facefusion.content_analyser', model_names, model_source_set) + with patch('facefusion.inference_manager.has_execution_provider', return_value = True): + with patch('facefusion.inference_manager.get_onnxruntime_version', return_value = (1, 26, 0)): - assert isinstance(INFERENCE_POOL_SET.get('cli').get('facefusion.content_analyser.nsfw_1.nsfw_2.nsfw_3.0.cpu').get('nsfw_1'), InferenceSession) + 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) - with patch('facefusion.inference_manager.detect_app_context', return_value = 'api'): - get_inference_pool('facefusion.content_analyser', model_names, model_source_set) + assert isinstance(cli_inference_pool.get('nsfw_1'), InferenceSession) - assert isinstance(INFERENCE_POOL_SET.get('cli').get('facefusion.content_analyser.nsfw_1.nsfw_2.nsfw_3.0.cpu').get('nsfw_1'), InferenceSession) + 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 INFERENCE_POOL_SET.get('cli').get('facefusion.content_analyser.nsfw_1.nsfw_2.nsfw_3.0.cpu').get('nsfw_1') == INFERENCE_POOL_SET.get('api').get('facefusion.content_analyser.nsfw_1.nsfw_2.nsfw_3.0.cpu').get('nsfw_1') + assert isinstance(api_inference_pool.get('nsfw_1'), InferenceSession) + + assert not (cli_inference_pool.get('nsfw_1') is api_inference_pool.get('nsfw_1')) + + with patch('facefusion.inference_manager.get_onnxruntime_version', return_value = (1, 24, 4)): + + 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 cli_inference_pool.get('nsfw_1') is api_inference_pool.get('nsfw_1') + + +@pytest.fixture +def override_module() -> SimpleNamespace: + return SimpleNamespace(override_inference_providers = Mock(return_value = [ ('CoreMLExecutionProvider', { 'ModelFormat': 'MLProgram' }) ])) + + +@pytest.fixture +def adjust_module() -> SimpleNamespace: + return SimpleNamespace(adjust_inference_providers = Mock(return_value = [ ('CoreMLExecutionProvider', { 'ModelFormat': 'MLProgram' }) ])) + + +def test_resolve_static_inference_providers(override_module : SimpleNamespace, adjust_module : SimpleNamespace) -> None: + state_manager.init_item('execution_providers', ['coreml']) + resolve_static_inference_providers.cache_clear() + + with patch('facefusion.inference_manager.importlib', Mock(import_module = Mock(return_value = override_module))): + inference_providers = resolve_static_inference_providers('override_module', 0) + + assert inference_providers == [ ('CoreMLExecutionProvider', { 'ModelFormat': 'MLProgram' }) ] + + with patch('facefusion.inference_manager.importlib', Mock(import_module = Mock(return_value = adjust_module))): + inference_providers = resolve_static_inference_providers('adjust_module', 0) + + assert inference_providers == [ ('CoreMLExecutionProvider', { 'SpecializationStrategy': 'FastPrediction', 'ModelCacheDirectory': resolve_cache_path(), 'ModelFormat': 'MLProgram' }) ] + + inference_providers = resolve_static_inference_providers('test', 0) + + assert inference_providers == [ ('CoreMLExecutionProvider', { 'SpecializationStrategy': 'FastPrediction', 'ModelCacheDirectory': resolve_cache_path() }) ]