Changed logic for using onnxruntime
This commit is contained in:
@@ -3,6 +3,9 @@ import json
|
||||
import shutil
|
||||
import inspect
|
||||
from server import PromptServer
|
||||
import subprocess
|
||||
import sys
|
||||
from importlib import metadata
|
||||
|
||||
config = None
|
||||
|
||||
@@ -143,6 +146,47 @@ def install_js():
|
||||
log("Copying JS files")
|
||||
shutil.copytree(src_dir, dst_dir, dirs_exist_ok=True)
|
||||
|
||||
def has_nvidia_cuda():
|
||||
try:
|
||||
output = subprocess.check_output(
|
||||
['nvidia-smi', '--query-gpu=name', '--format=csv,noheader'],
|
||||
stderr=subprocess.DEVNULL,
|
||||
encoding='utf-8'
|
||||
).strip()
|
||||
return bool(output)
|
||||
except (subprocess.CalledProcessError, FileNotFoundError):
|
||||
return False
|
||||
|
||||
def is_installed(pkg_name: str, min_version: str = '') -> bool:
|
||||
try:
|
||||
ver = metadata.version(pkg_name)
|
||||
if min_version:
|
||||
return tuple(map(int, ver.split('.'))) >= tuple(map(int, min_version.split('.')))
|
||||
return True
|
||||
except metadata.PackageNotFoundError:
|
||||
return False
|
||||
|
||||
def install(pkg_spec: str):
|
||||
subprocess.check_call([sys.executable, '-m', 'pip', 'install', pkg_spec])
|
||||
|
||||
def uninstall(pkg: str):
|
||||
subprocess.check_call([sys.executable, '-m', 'pip', 'uninstall', '-y', pkg])
|
||||
|
||||
def install_onnxruntime():
|
||||
if has_nvidia_cuda():
|
||||
if is_installed('onnxruntime'):
|
||||
#uninstall("onnxruntime")
|
||||
print(f"LP >>> Your python has the 'onnxruntime' library installed, although your computer supports 'onnxruntime-gpu'.")
|
||||
print(f"LP >>> Solution: If other node packages do not use the 'onnxruntime' library, then remove the 'onnxruntime' library for your python.")
|
||||
print(f"LP >>> Otherwise the Image Remove Background node will not work effectively (the node will use the CPU instead of the GPU).")
|
||||
if not is_installed('onnxruntime-gpu'):
|
||||
install("onnxruntime-gpu>=1.22")
|
||||
else:
|
||||
if is_installed('onnxruntime-gpu'):
|
||||
uninstall("onnxruntime-gpu")
|
||||
if not is_installed('onnxruntime'):
|
||||
install("onnxruntime>=1.22")
|
||||
|
||||
def init(check_imports=None):
|
||||
log("Init")
|
||||
|
||||
@@ -156,4 +200,5 @@ def init(check_imports=None):
|
||||
return False
|
||||
|
||||
install_js()
|
||||
install_onnxruntime()
|
||||
return True
|
||||
|
||||
Reference in New Issue
Block a user