From 0d57fa3016404c53855fdfb9f64c1c70ed80ae5b Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Thu, 8 Aug 2024 14:44:53 +0900 Subject: [PATCH] fix zluda torch cpp_extension --- installer.py | 34 ++++++++++------------------------ modules/zluda_hijacks.py | 15 --------------- modules/zluda_installer.py | 15 +++++++++++++++ 3 files changed, 25 insertions(+), 39 deletions(-) diff --git a/installer.py b/installer.py index 506d71a1a..d19a4b079 100644 --- a/installer.py +++ b/installer.py @@ -517,9 +517,6 @@ def install_rocm_zluda(torch_command): log.info("For ZLUDA support specify '--use-zluda'") log.info('Using CPU-only torch') torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision') - - # conceal ROCm installed - rocm.conceal() else: if rocm.version is None or float(rocm.version) > 6.1: # assume the latest if version check fails torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision --index-url https://download.pytorch.org/whl/rocm6.1') @@ -596,16 +593,6 @@ def install_openvino(torch_command): return torch_command -def is_rocm_available(allow_rocm): - if not allow_rocm: - return False - if installed('torch-directml', quiet=True): - log.debug('DirectML installation is detected. Skipping HIP SDK check.') - return False - from modules.rocm import is_installed - return is_installed - - def install_torch_addons(): xformers_package = os.environ.get('XFORMERS_PACKAGE', '--pre xformers') if opts.get('cross_attention_optimization', '') == 'xFormers' or args.use_xformers else 'none' triton_command = os.environ.get('TRITON_COMMAND', 'triton') if sys.platform == 'linux' else None @@ -648,6 +635,7 @@ def check_torch(): if args.profile: pr = cProfile.Profile() pr.enable() + from modules import rocm allow_cuda = not (args.use_rocm or args.use_directml or args.use_ipex or args.use_openvino) allow_rocm = not (args.use_cuda or args.use_directml or args.use_ipex or args.use_openvino) allow_ipex = not (args.use_cuda or args.use_rocm or args.use_directml or args.use_openvino) @@ -663,15 +651,8 @@ def check_torch(): log.info('nVidia CUDA toolkit detected: nvidia-smi present') torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision --index-url https://download.pytorch.org/whl/cu121') install('onnxruntime-gpu', 'onnxruntime-gpu', ignore=True, quiet=True) - elif is_rocm_available(allow_rocm): + elif allow_rocm and rocm.is_installed: torch_command = install_rocm_zluda(torch_command) - - from modules import rocm - if rocm.is_wsl: # WSL ROCm - try: - rocm.load_hsa_runtime() - except OSError: - log.error("Failed to preload HSA Runtime library.") elif is_ipex_available(allow_ipex): torch_command = install_ipex(torch_command) elif allow_openvino and args.use_openvino: @@ -686,9 +667,6 @@ def check_torch(): if 'torch' in torch_command and not args.version: install(torch_command, 'torch torchvision') install('onnxruntime-directml', 'onnxruntime-directml', ignore=True) - from modules import rocm - if rocm.is_installed: - rocm.conceal() else: if args.use_zluda: log.warning("ZLUDA failed to initialize: no HIP SDK found") @@ -734,6 +712,14 @@ def check_torch(): log.error(f'Could not load torch: {e}') if not args.ignore: sys.exit(1) + if rocm.is_installed: + if sys.platform == "win32": # CPU, DirectML, ZLUDA + rocm.conceal() + elif rocm.is_wsl: # WSL ROCm + try: + rocm.load_hsa_runtime() + except OSError: + log.error("Failed to preload HSA Runtime library.") if args.version: return if not args.skip_all: diff --git a/modules/zluda_hijacks.py b/modules/zluda_hijacks.py index eef4aab14..f7af35724 100644 --- a/modules/zluda_hijacks.py +++ b/modules/zluda_hijacks.py @@ -1,5 +1,3 @@ -import os -import sys import torch @@ -10,19 +8,6 @@ def topk(tensor: torch.Tensor, *args, **kwargs): return torch.return_types.topk((values.to(device), indices.to(device),)) -def _join_rocm_home(*paths) -> str: - from torch.utils.cpp_extension import ROCM_HOME - return os.path.join(ROCM_HOME, *paths) - - def do_hijack(): torch.version.hip = "5.7" torch.topk = topk - platform = sys.platform - sys.platform = "" - from torch.utils import cpp_extension - sys.platform = platform - cpp_extension.IS_WINDOWS = platform == "win32" - cpp_extension.IS_MACOS = False - cpp_extension.IS_LINUX = platform.startswith('linux') - cpp_extension._join_rocm_home = _join_rocm_home # pylint: disable=protected-access diff --git a/modules/zluda_installer.py b/modules/zluda_installer.py index 11d543eb5..7c20d3d79 100644 --- a/modules/zluda_installer.py +++ b/modules/zluda_installer.py @@ -1,4 +1,5 @@ import os +import sys import ctypes import shutil import zipfile @@ -61,3 +62,17 @@ def load(zluda_path: os.PathLike) -> None: ctypes.windll.LoadLibrary(os.path.join(zluda_path, v)) for v in DLL_MAPPING.values(): ctypes.windll.LoadLibrary(os.path.join(zluda_path, v)) + + def conceal(): + import torch # pylint: disable=unused-import + platform = sys.platform + sys.platform = "" + from torch.utils import cpp_extension + sys.platform = platform + cpp_extension.IS_WINDOWS = platform == "win32" + cpp_extension.IS_MACOS = False + cpp_extension.IS_LINUX = platform.startswith('linux') + def _join_rocm_home(*paths) -> str: + return os.path.join(cpp_extension.ROCM_HOME, *paths) + cpp_extension._join_rocm_home = _join_rocm_home # pylint: disable=protected-access + rocm.conceal = conceal