mirror of
https://github.com/facefusion/facefusion.git
synced 2026-07-28 04:50:58 +02:00
Revamp execution provider overrides/adjustments (#1206)
* Split provider hooks into override/adjust with cached CoreML base Replace the single resolve_inference_providers processor hook with two: override_inference_providers (full replacement) and adjust_inference_providers (merge options onto the base providers built by create_inference_providers). This lets CoreML processors inherit ModelCacheDirectory + SpecializationStrategy from the base while layering ModelFormat/MLComputeUnits on top. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01HTQCZiYjJyUX11bDpbRSiB * fix caching for execution provider by having override and adjust ways * fix caching for execution provider by having override and adjust ways * fix lint * use proper pytest fixtures --------- Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
aba04180ed
commit
7990faa38c
@@ -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)
|
||||
|
||||
@@ -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':
|
||||
|
||||
@@ -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'
|
||||
})
|
||||
]
|
||||
|
||||
|
||||
@@ -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') ]
|
||||
|
||||
|
||||
@@ -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'
|
||||
})
|
||||
]
|
||||
|
||||
|
||||
@@ -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() }) ]
|
||||
|
||||
Reference in New Issue
Block a user