Files
facefusion/facefusion/inference_manager.py
T

140 lines
5.6 KiB
Python

import importlib
import random
from functools import lru_cache
from time import sleep, time
from typing import List, Optional
from onnxruntime import InferenceSession
from facefusion import logger, process_manager, state_manager, store_creator, translator
from facefusion.common_helper import is_windows
from facefusion.execution import create_inference_providers, get_onnxruntime_version, has_execution_provider
from facefusion.exit_helper import fatal_exit
from facefusion.filesystem import get_file_name, is_file
from facefusion.session_manager import resolve_owner_id
from facefusion.time_helper import calculate_end_time
from facefusion.types import DownloadSet, ExecutionProvider, InferencePool, InferenceProvider, Store
INFERENCE_POOL_STORE : Store = store_creator.create_store({})
def init() -> None:
owner_id = resolve_owner_id()
store_creator.init_content(INFERENCE_POOL_STORE, owner_id)
def get_inference_pool(module_name : str, model_names : List[str], model_source_set : DownloadSet) -> InferencePool:
while process_manager.is_checking():
sleep(0.5)
owner_id = resolve_owner_id()
inference_pool_set = store_creator.get_content(INFERENCE_POOL_STORE, owner_id)
execution_device_ids = state_manager.get_item('execution_device_ids')
execution_providers = state_manager.get_item('execution_providers')
has_arena_leak = has_execution_provider('cuda') and get_onnxruntime_version() > (1, 24, 4)
for execution_device_id in execution_device_ids:
inference_context = get_inference_context(module_name, model_names, execution_device_id, execution_providers)
if not has_arena_leak:
inference_pool = find_inference_pool(inference_context)
if inference_pool:
inference_pool_set[inference_context] = inference_pool
if not inference_pool_set.get(inference_context):
inference_providers = resolve_static_inference_providers(module_name, execution_device_id)
inference_pool_set[inference_context] = create_inference_pool(model_source_set, inference_providers)
current_inference_context = get_inference_context(module_name, model_names, random.choice(execution_device_ids), execution_providers)
return inference_pool_set.get(current_inference_context)
def find_inference_pool(inference_context : str) -> Optional[InferencePool]:
inference_pool_values = list(INFERENCE_POOL_STORE.get('content_set').values())
for inference_pool_set in inference_pool_values:
if inference_pool_set.get(inference_context):
return inference_pool_set.get(inference_context)
return None
def create_inference_pool(model_source_set : DownloadSet, inference_providers : List[InferenceProvider]) -> InferencePool:
inference_pool : InferencePool = {}
for model_name in model_source_set.keys():
model_path = model_source_set.get(model_name).get('path')
if is_file(model_path):
inference_pool[model_name] = create_inference_session(model_path, inference_providers)
return inference_pool
def clear_inference_pool(module_name : str, model_names : List[str]) -> None:
owner_id = resolve_owner_id()
inference_pool_set = store_creator.get_content(INFERENCE_POOL_STORE, owner_id)
execution_device_ids = state_manager.get_item('execution_device_ids')
execution_providers = state_manager.get_item('execution_providers')
if is_windows() and has_execution_provider('directml'):
inference_pool_set.clear()
for execution_device_id in execution_device_ids:
inference_context = get_inference_context(module_name, model_names, execution_device_id, execution_providers)
if inference_pool_set.get(inference_context):
del inference_pool_set[inference_context]
def clear() -> None:
owner_id = resolve_owner_id()
store_creator.init_content(INFERENCE_POOL_STORE, owner_id)
def create_inference_session(model_path : str, inference_providers : List[InferenceProvider]) -> InferenceSession:
model_file_name = get_file_name(model_path)
start_time = time()
try:
inference_session = InferenceSession(model_path, providers = inference_providers)
logger.debug(translator.get('loading_model_succeeded').format(model_name = model_file_name, seconds = calculate_end_time(start_time)), __name__)
return inference_session
except Exception:
logger.error(translator.get('loading_model_failed').format(model_name = model_file_name), __name__)
fatal_exit(1)
def get_inference_context(module_name : str, model_names : List[str], execution_device_id : int, execution_providers : List[ExecutionProvider]) -> str:
inference_context = '.'.join([ module_name ] + model_names + [ str(execution_device_id) ] + list(execution_providers))
return inference_context
@lru_cache()
def resolve_static_inference_providers(module_name : str, execution_device_id : int) -> List[InferenceProvider]:
module = importlib.import_module(module_name)
execution_providers = state_manager.get_item('execution_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])
return inference_providers
return create_inference_providers(execution_device_id, execution_providers)