mirror of
https://github.com/facefusion/facefusion.git
synced 2026-08-06 00:58:37 +02:00
3.8.1 (#1214)
* fix image to image for coreml * fix image to image for coreml * Update metadata.py * modernize onnxruntime * introduce cuda@12 and cuda@13 for installer * introduce cuda@12 and cuda@13 for installer * temporary fix for inference pool disaster * temporary fix for inference pool disaster * use onnxruntime 1.24.4 for cuda@12 * fix test * fix test * fix test * add test for job list * switch to has_arena_leak guard * remove directml hack for inference pool * uninstall every onnxruntime first * uninstall every onnxruntime first * restore the directml fix * restore the directml fix * fix directml again * simplify condition
This commit is contained in:
@@ -6,7 +6,7 @@ from onnxruntime import InferenceSession
|
||||
|
||||
from facefusion import content_analyser, state_manager
|
||||
from facefusion.execution import resolve_cache_path
|
||||
from facefusion.inference_manager import INFERENCE_POOL_SET, get_inference_pool, resolve_static_inference_providers
|
||||
from facefusion.inference_manager import get_inference_pool, resolve_static_inference_providers
|
||||
|
||||
|
||||
@pytest.fixture(scope = 'module', autouse = True)
|
||||
@@ -20,17 +20,29 @@ 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 = 'ui'):
|
||||
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 = 'ui'):
|
||||
ui_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('ui').get('facefusion.content_analyser.nsfw_1.nsfw_2.nsfw_3.0.cpu').get('nsfw_1')
|
||||
assert isinstance(ui_inference_pool.get('nsfw_1'), InferenceSession)
|
||||
|
||||
assert not (cli_inference_pool.get('nsfw_1') is ui_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 = 'ui'):
|
||||
ui_inference_pool = get_inference_pool('facefusion.content_analyser', model_names, model_source_set)
|
||||
|
||||
assert isinstance(ui_inference_pool.get('nsfw_1'), InferenceSession)
|
||||
|
||||
assert cli_inference_pool.get('nsfw_1') is ui_inference_pool.get('nsfw_1')
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -48,9 +60,15 @@ def test_resolve_static_inference_providers(override_module : SimpleNamespace, a
|
||||
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' }) ]
|
||||
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))):
|
||||
assert resolve_static_inference_providers('adjust_module', 0) == [ ('CoreMLExecutionProvider', { 'SpecializationStrategy': 'FastPrediction', 'ModelCacheDirectory': resolve_cache_path(), 'ModelFormat': 'MLProgram' }) ]
|
||||
inference_providers = resolve_static_inference_providers('adjust_module', 0)
|
||||
|
||||
assert resolve_static_inference_providers('test', 0) == [ ('CoreMLExecutionProvider', { 'SpecializationStrategy': 'FastPrediction', 'ModelCacheDirectory': resolve_cache_path() }) ]
|
||||
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() }) ]
|
||||
|
||||
Reference in New Issue
Block a user