diff --git a/facefusion/inference_manager.py b/facefusion/inference_manager.py index 3b10eed3..aed662f3 100644 --- a/facefusion/inference_manager.py +++ b/facefusion/inference_manager.py @@ -93,10 +93,23 @@ def resolve_static_inference_providers(module_name : str, execution_device_id : module = importlib.import_module(module_name) execution_providers = state_manager.get_item('execution_providers') - if hasattr(module, 'resolve_inference_providers'): - inference_providers = getattr(module, 'resolve_inference_providers')() + if hasattr(module, 'override_inference_providers'): + override_inference_providers = getattr(module, 'override_inference_providers')() + + if override_inference_providers: + return override_inference_providers + + if hasattr(module, 'adjust_inference_providers'): + adjust_inference_providers = getattr(module, 'adjust_inference_providers')() + + if adjust_inference_providers: + inference_providers = create_inference_providers(execution_device_id, execution_providers) + + for adjust_inference_provider in adjust_inference_providers: + for inference_provider in inference_providers: + if inference_provider[0] == adjust_inference_provider[0] and inference_provider[1]: + inference_provider[1].update(adjust_inference_provider[1]) - if inference_providers: return inference_providers return create_inference_providers(execution_device_id, execution_providers) diff --git a/facefusion/processors/modules/background_remover/core.py b/facefusion/processors/modules/background_remover/core.py index 2ff10aa9..f6af0472 100644 --- a/facefusion/processors/modules/background_remover/core.py +++ b/facefusion/processors/modules/background_remover/core.py @@ -479,7 +479,7 @@ def clear_inference_pool() -> None: inference_manager.clear_inference_pool(__name__, model_names) -def resolve_inference_providers() -> List[InferenceProvider]: +def override_inference_providers() -> List[InferenceProvider]: model_type = get_model_options().get('type') if is_macos() and has_execution_provider('coreml') or is_windows() and has_execution_provider('directml') and model_type == 'corridor_key': diff --git a/facefusion/processors/modules/face_swapper/core.py b/facefusion/processors/modules/face_swapper/core.py index 31dec34f..8f4954fd 100755 --- a/facefusion/processors/modules/face_swapper/core.py +++ b/facefusion/processors/modules/face_swapper/core.py @@ -502,7 +502,7 @@ def clear_inference_pool() -> None: inference_manager.clear_inference_pool(__name__, model_names) -def resolve_inference_providers() -> List[InferenceProvider]: +def adjust_inference_providers() -> List[InferenceProvider]: model_precision = get_model_options().get('precision') model_type = get_model_options().get('type') @@ -512,8 +512,7 @@ def resolve_inference_providers() -> List[InferenceProvider]: [ (facefusion.choices.execution_provider_set.get('coreml'), { - 'ModelFormat': 'MLProgram', - 'SpecializationStrategy': 'FastPrediction' + 'ModelFormat': 'MLProgram' }) ] diff --git a/facefusion/processors/modules/frame_colorizer/core.py b/facefusion/processors/modules/frame_colorizer/core.py index 87a9a357..c3a9b718 100644 --- a/facefusion/processors/modules/frame_colorizer/core.py +++ b/facefusion/processors/modules/frame_colorizer/core.py @@ -172,7 +172,7 @@ def clear_inference_pool() -> None: inference_manager.clear_inference_pool(__name__, model_names) -def resolve_inference_providers() -> List[InferenceProvider]: +def override_inference_providers() -> List[InferenceProvider]: if is_macos() and has_execution_provider('coreml'): return [ facefusion.choices.execution_provider_set.get('cpu') ] diff --git a/facefusion/processors/modules/frame_enhancer/core.py b/facefusion/processors/modules/frame_enhancer/core.py index 8b1506a5..e849f57c 100644 --- a/facefusion/processors/modules/frame_enhancer/core.py +++ b/facefusion/processors/modules/frame_enhancer/core.py @@ -558,7 +558,7 @@ def clear_inference_pool() -> None: inference_manager.clear_inference_pool(__name__, model_names) -def resolve_inference_providers() -> List[InferenceProvider]: +def adjust_inference_providers() -> List[InferenceProvider]: model_precision = get_model_options().get('precision') if is_macos() and has_execution_provider('coreml') and model_precision == 'fp16': @@ -566,8 +566,7 @@ def resolve_inference_providers() -> List[InferenceProvider]: [ (facefusion.choices.execution_provider_set.get('coreml'), { - 'ModelFormat': 'MLProgram', - 'SpecializationStrategy': 'FastPrediction' + 'ModelFormat': 'MLProgram' }) ] diff --git a/tests/test_inference_manager.py b/tests/test_inference_manager.py index 8f90e9f9..047093cf 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 INFERENCE_POOL_SET, get_inference_pool, resolve_static_inference_providers @pytest.fixture(scope = 'module', autouse = True) @@ -12,7 +14,6 @@ def before_all() -> None: state_manager.init_item('execution_device_ids', [ 0 ]) state_manager.init_item('execution_providers', [ 'cpu' ]) state_manager.init_item('download_providers', [ 'github' ]) - content_analyser.pre_check() def test_get_inference_pool() -> None: @@ -30,3 +31,26 @@ def test_get_inference_pool() -> None: assert isinstance(INFERENCE_POOL_SET.get('cli').get('facefusion.content_analyser.nsfw_1.nsfw_2.nsfw_3.0.cpu').get('nsfw_1'), InferenceSession) 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('ui').get('facefusion.content_analyser.nsfw_1.nsfw_2.nsfw_3.0.cpu').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))): + assert resolve_static_inference_providers('override_module', 0) == [ ('CoreMLExecutionProvider', { 'ModelFormat': 'MLProgram' }) ] + + with patch('facefusion.inference_manager.importlib', Mock(import_module = Mock(return_value = adjust_module))): + assert resolve_static_inference_providers('adjust_module', 0) == [ ('CoreMLExecutionProvider', { 'SpecializationStrategy': 'FastPrediction', 'ModelCacheDirectory': resolve_cache_path(), 'ModelFormat': 'MLProgram' }) ] + + assert resolve_static_inference_providers('test', 0) == [ ('CoreMLExecutionProvider', { 'SpecializationStrategy': 'FastPrediction', 'ModelCacheDirectory': resolve_cache_path() }) ]