feat: auto detect gpu providers for pytorch
This commit is contained in:
+12
-1
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user