feat: auto detect gpu providers for pytorch

This commit is contained in:
Jeffrey Wu
2024-07-29 14:26:31 +08:00
parent d5c9d3c426
commit 46f045e570
5 changed files with 26 additions and 21 deletions
+12 -1
View File
@@ -1,3 +1,5 @@
import torch
from functools import lru_cache
import subprocess
import xml.etree.ElementTree as ElementTree
@@ -5,9 +7,11 @@ from typing import List, Any
from .typing import ValueAndUnit, ExecutionDevice
def apply_execution_provider_options(execution_providers: List[str]) -> List[Any]:
def apply_execution_provider_options(execution_providers: List[str] | None = None) -> List[Any]:
execution_providers_with_options : List[Any] = []
if execution_providers is None:
execution_providers = get_default_providers()
for execution_provider in execution_providers:
if execution_provider == 'CUDAExecutionProvider':
execution_providers_with_options.append((execution_provider,
@@ -18,6 +22,13 @@ def apply_execution_provider_options(execution_providers: List[str]) -> List[Any
execution_providers_with_options.append(execution_provider)
return execution_providers_with_options
def get_default_providers() -> List[str]:
if torch.cuda.is_available():
return ['CUDAExecutionProvider', 'CPUExecutionProvider']
elif torch.backends.mps.is_available():
return ['CoreMLExecutionProvider', 'CPUExecutionProvider']
return ['CPUExecutionProvider']
def use_exhaustive() -> bool:
execution_devices = detect_static_execution_devices()
+2 -5
View File
@@ -10,9 +10,6 @@ from .typing import FaceLandmark68, VisionFrame, Mask, Padding, FaceMaskRegion,
from .execution import apply_execution_provider_options
from .filesystem import resolve_relative_path
# TODO load from options
execution_providers = ['CoreMLExecutionProvider', 'CPUExecutionProvider']
FACE_OCCLUDER = None
FACE_PARSER = None
THREAD_LOCK : threading.Lock = threading.Lock()
@@ -50,7 +47,7 @@ def get_face_occluder() -> Any:
with THREAD_LOCK:
if FACE_OCCLUDER is None:
model_path = MODELS['face_occluder']['path']
FACE_OCCLUDER = onnxruntime.InferenceSession(model_path, providers = apply_execution_provider_options(execution_providers))
FACE_OCCLUDER = onnxruntime.InferenceSession(model_path, providers = apply_execution_provider_options())
return FACE_OCCLUDER
@@ -60,7 +57,7 @@ def get_face_parser() -> Any:
with THREAD_LOCK:
if FACE_PARSER is None:
model_path = MODELS['face_parser']['path']
FACE_PARSER = onnxruntime.InferenceSession(model_path, providers = apply_execution_provider_options(execution_providers))
FACE_PARSER = onnxruntime.InferenceSession(model_path, providers = apply_execution_provider_options())
return FACE_PARSER
+10 -11
View File
@@ -24,7 +24,6 @@ face_recognizer_model: Optional[FaceRecognizerModel] = 'arcface_inswapper'
face_detector_score = 0.5
face_landmarker_score = 0.5
face_detector_size = '640x640'
execution_providers = ['CoreMLExecutionProvider', 'CPUExecutionProvider']
face_analyser_order: FaceAnalyserOrder = 'left-right'
face_analyser_age: Optional[FaceAnalyserAge] = None
face_analyser_gender: Optional[FaceAnalyserGender] = None
@@ -98,24 +97,24 @@ def get_face_analyser() -> Any:
with THREAD_LOCK:
if FACE_ANALYSER is None:
if face_detector_model in [ 'many', 'retinaface' ]:
face_detectors['retinaface'] = onnxruntime.InferenceSession(MODELS['face_detector_retinaface']['path'], providers = apply_execution_provider_options(execution_providers))
face_detectors['retinaface'] = onnxruntime.InferenceSession(MODELS['face_detector_retinaface']['path'], providers = apply_execution_provider_options())
if face_detector_model in [ 'many', 'scrfd' ]:
face_detectors['scrfd'] = onnxruntime.InferenceSession(MODELS['face_detector_scrfd']['path'], providers = apply_execution_provider_options(execution_providers))
face_detectors['scrfd'] = onnxruntime.InferenceSession(MODELS['face_detector_scrfd']['path'], providers = apply_execution_provider_options())
if face_detector_model in [ 'many', 'yoloface' ]:
face_detectors['yoloface'] = onnxruntime.InferenceSession(MODELS['face_detector_yoloface']['path'], providers = apply_execution_provider_options(execution_providers))
face_detectors['yoloface'] = onnxruntime.InferenceSession(MODELS['face_detector_yoloface']['path'], providers = apply_execution_provider_options())
if face_detector_model in [ 'yunet' ]:
face_detectors['yunet'] = cv2.FaceDetectorYN.create(MODELS['face_detector_yunet']['path'], '', (0, 0))
if face_recognizer_model == 'arcface_blendswap':
face_recognizer = onnxruntime.InferenceSession(MODELS['face_recognizer_arcface_blendswap']['path'], providers = apply_execution_provider_options(execution_providers))
face_recognizer = onnxruntime.InferenceSession(MODELS['face_recognizer_arcface_blendswap']['path'], providers = apply_execution_provider_options())
if face_recognizer_model == 'arcface_inswapper':
face_recognizer = onnxruntime.InferenceSession(MODELS['face_recognizer_arcface_inswapper']['path'], providers = apply_execution_provider_options(execution_providers))
face_recognizer = onnxruntime.InferenceSession(MODELS['face_recognizer_arcface_inswapper']['path'], providers = apply_execution_provider_options())
if face_recognizer_model == 'arcface_simswap':
face_recognizer = onnxruntime.InferenceSession(MODELS['face_recognizer_arcface_simswap']['path'], providers = apply_execution_provider_options(execution_providers))
face_recognizer = onnxruntime.InferenceSession(MODELS['face_recognizer_arcface_simswap']['path'], providers = apply_execution_provider_options())
if face_recognizer_model == 'arcface_uniface':
face_recognizer = onnxruntime.InferenceSession(MODELS['face_recognizer_arcface_uniface']['path'], providers = apply_execution_provider_options(execution_providers))
face_landmarkers['68'] = onnxruntime.InferenceSession(MODELS['face_landmarker_68']['path'], providers = apply_execution_provider_options(execution_providers))
face_landmarkers['68_5'] = onnxruntime.InferenceSession(MODELS['face_landmarker_68_5']['path'], providers = apply_execution_provider_options(execution_providers))
gender_age = onnxruntime.InferenceSession(MODELS['gender_age']['path'], providers = apply_execution_provider_options(execution_providers))
face_recognizer = onnxruntime.InferenceSession(MODELS['face_recognizer_arcface_uniface']['path'], providers = apply_execution_provider_options())
face_landmarkers['68'] = onnxruntime.InferenceSession(MODELS['face_landmarker_68']['path'], providers = apply_execution_provider_options())
face_landmarkers['68_5'] = onnxruntime.InferenceSession(MODELS['face_landmarker_68_5']['path'], providers = apply_execution_provider_options())
gender_age = onnxruntime.InferenceSession(MODELS['gender_age']['path'], providers = apply_execution_provider_options())
FACE_ANALYSER =\
{
'face_detectors': face_detectors,
+1 -2
View File
@@ -84,7 +84,6 @@ class FaceRestoration:
self._execution_thread_count = 4
self._execution_queue_count = 1
self._execution_providers = ['CoreMLExecutionProvider', 'CPUExecutionProvider']
self._frame_processor = None
@@ -212,7 +211,7 @@ class FaceRestoration:
with THREAD_LOCK:
if self._frame_processor is None:
model_path = get_faceless_model_path('face_restoration', self._model_name)
self._frame_processor = onnxruntime.InferenceSession(model_path, providers = apply_execution_provider_options(self._execution_providers))
self._frame_processor = onnxruntime.InferenceSession(model_path, providers = apply_execution_provider_options())
return self._frame_processor
def _blend_frame(self, temp_vision_frame : VisionFrame, paste_vision_frame : VisionFrame) -> VisionFrame:
+1 -2
View File
@@ -86,7 +86,6 @@ class FaceSwapper:
self._execution_queue_count = 1
self._execution_thread_count = 4
self._execution_providers = ['CoreMLExecutionProvider', 'CPUExecutionProvider']
self._frame_processor = None
self._model_initializer = None
@@ -192,7 +191,7 @@ class FaceSwapper:
model_path = get_faceless_model_path('face_swapper', self._model_name)
if model_path is None:
raise Exception("can not get model path")
self._frame_processor = onnxruntime.InferenceSession(model_path, providers = apply_execution_provider_options(self._execution_providers))
self._frame_processor = onnxruntime.InferenceSession(model_path, providers = apply_execution_provider_options())
return self._frame_processor
def _swap_face(self, source_face: Face, target_face: Face, source_vision_frame, target_vision_frame: VisionFrame) -> VisionFrame: