From d63bc05b6ccc2b7ca2d22a6846599f4c121317a4 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Sat, 27 Sep 2025 11:12:13 +0900 Subject: [PATCH] update rocm.py to detect rocm-sdk packages --- installer.py | 8 ++- modules/rocm.py | 112 ++++++++++++++++++++++++++++--------- modules/zluda_installer.py | 11 ++-- 3 files changed, 98 insertions(+), 33 deletions(-) diff --git a/installer.py b/installer.py index f05a5bb32..bc0612393 100644 --- a/installer.py +++ b/installer.py @@ -677,6 +677,10 @@ def install_rocm_zluda(): if args.skip_all or args.skip_requirements: return torch_command from modules import rocm + if rocm.err is not None: + log.warning(f'ROCm: error checking ROCm toolkit: {rocm.err}') + log.info('Using CPU-only torch') + return os.environ.get('TORCH_COMMAND', 'torch torchvision') if not rocm.is_installed: log.warning('ROCm: could not find ROCm toolkit installed') log.info('Using CPU-only torch') @@ -720,8 +724,8 @@ def install_rocm_zluda(): if sys.platform == "win32": #check_python(supported_minors=[10, 11, 12, 13], reason='ZLUDA backend requires a Python version between 3.10 and 3.13') - if args.use_rocm and args.experimental and (sys.version_info.major, sys.version_info.minor) == (3, 12): # TODO install: switch to pytorch source when it becomes available - torch_command = os.environ.get('TORCH_COMMAND', '--no-cache-dir https://repo.radeon.com/rocm/windows/rocm-rel-6.4.4/torch-2.8.0a0%2Bgitfc14c65-cp312-cp312-win_amd64.whl https://repo.radeon.com/rocm/windows/rocm-rel-6.4.4/torchvision-0.24.0a0%2Bc85f008-cp312-cp312-win_amd64.whl') + if args.use_rocm: # TODO install: switch to pytorch source when it becomes available + torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://rocm.nightlies.amd.com/v2-staging/{rocm.get_distribution(device)}') else: if args.device_id is not None: if os.environ.get('HIP_VISIBLE_DEVICES', None) is not None: diff --git a/modules/rocm.py b/modules/rocm.py index 554c3af55..bf06ba3b9 100644 --- a/modules/rocm.py +++ b/modules/rocm.py @@ -6,6 +6,7 @@ import subprocess import importlib.metadata from typing import Union, List from enum import Enum +from installer import installed def resolve_link(path_: str) -> str: @@ -40,6 +41,33 @@ def conceal(): os.environ["PATH"] = ";".join(paths_no_rocm) +class Environment: + hip: ctypes.CDLL + + def __init__(self, path: str): + self.hip = ctypes.CDLL(path) + + +# rocm is installed system-wide +class ROCmEnvironment(Environment): + path: str + + def __init__(self, path: str): + for v in (6, 7): + lib = os.path.join(path, "bin", f"amdhip64_{v}.dll") + if os.path.exists(lib): + super().__init__(lib) + break + self.path = path + + +# rocm-sdk package is installed +class PythonPackageEnvironment(Environment): + def __init__(self): + import _rocm_sdk_core + super().__init__(os.path.join(_rocm_sdk_core.__path__[0], "bin", "amdhip64_7.dll")) + + class MicroArchitecture(Enum): GCN = "gcn" RDNA = "rdna" @@ -81,7 +109,7 @@ class Agent: self.blaslt_supported = os.path.exists(os.path.join(blaslt_tensile_libpath, f"Kernels.so-000-{name}.hsaco" if sys.platform == "win32" else f"extop_{name}.co")) def get_gfx_version(self) -> Union[str, None]: - if self.gfx_version >= 0x1101 and self.gfx_version < 0x1200: + if self.gfx_version >= 0x1100 and self.gfx_version < 0x1200: return "11.0.0" elif self.gfx_version != 0x1030 and self.gfx_version >= 0x1000 and self.gfx_version < 0x1100: # gfx1010 users had to override gfx version to 10.3.0 in Linux @@ -101,15 +129,23 @@ def get_version_torch() -> Union[str, None]: return version_.split("+rocm")[1] -if sys.platform == "win32": - def find() -> Union[str, None]: - hip_path = shutil.which("hipconfig") - if hip_path is not None: - return dirname(resolve_link(hip_path), 2) +def find() -> Union[Environment, None]: + hip_path = shutil.which("hipconfig") + if hip_path is not None: + py_path = os.path.dirname(sys.executable) + if hip_path.startswith(py_path): + try: + import _rocm_sdk_core # pylint: disable=unused-import + return PythonPackageEnvironment() + except ImportError: + pass + else: + return ROCmEnvironment(dirname(resolve_link(hip_path), 2)) + if sys.platform == "win32": hip_path = os.environ.get("HIP_PATH", None) if hip_path is not None: - return hip_path + return ROCmEnvironment(hip_path) program_files = os.environ.get('ProgramFiles', r'C:\Program Files') hip_path = rf'{program_files}\AMD\ROCm' @@ -146,29 +182,47 @@ if sys.platform == "win32": if latest is None: return None - return os.path.join(hip_path, str(latest)) + return ROCmEnvironment(os.path.join(hip_path, str(latest))) + else: + if not os.path.exists("/opt/rocm"): + return None + return ROCmEnvironment(resolve_link("/opt/rocm")) - def get_version() -> str: # cannot just run hipconfig as it requires Perl installed on Windows. - return os.path.basename(path) or os.path.basename(os.path.dirname(path)) +def get_version() -> str: + version = ctypes.c_int() + environment.hip.hipRuntimeGetVersion(ctypes.byref(version)) + major = version.value // 10000000 + minor = (version.value // 100000) % 100 + patch = version.value % 100000 + return f"{major}.{minor}.{patch}" + + +if sys.platform == "win32": def get_agents() -> List[Agent]: - return [Agent(x.split(' ')[-1].strip()) for x in spawn("hipinfo", cwd=os.path.join(path, 'bin')).split("\n") if x.startswith('gcnArchName:')] + if isinstance(environment, ROCmEnvironment): + out = spawn("amdgpu-arch", cwd=os.path.join(environment.path, 'bin')) + else: + out = spawn("amdgpu-arch") + out = out.strip() + return [Agent(x.split(' ')[-1].strip()) for x in out.split("\n")] + + def get_distribution(agent: Agent) -> str: + if agent.gfx_version >= 0x1100 and agent.gfx_version < 0x1110: + return "gfx110X-dgpu" + if agent.gfx_version == 0x1151: + return "gfx1151" + if agent.gfx_version >= 0x1200 and agent.gfx_version < 0x1210: + return "gfx120X-all" + if agent.gfx_version >= 0x940 and agent.gfx_version < 0x950: + return "gfx94X-dcgpu" + if agent.gfx_version == 0x950: + return "gfx950-dcgpu" + raise Exception(f"Unsupported GPU architecture: {agent.name}") is_wsl: bool = False version_torch = None else: - def find() -> Union[str, None]: - rocm_path = shutil.which("hipconfig") - if rocm_path is not None: - return dirname(resolve_link(rocm_path), 2) - if not os.path.exists("/opt/rocm"): - return None - return resolve_link("/opt/rocm") - - def get_version() -> str: - arr = spawn("hipconfig --version", cwd=os.path.join(path, 'bin')).split(".") - return f'{arr[0]}.{arr[1]}' if len(arr) >= 2 else None - def get_agents() -> List[Agent]: try: agents = spawn("rocm_agent_enumerator").split("\n") @@ -208,11 +262,17 @@ else: is_wsl: bool = os.environ.get('WSL_DISTRO_NAME', 'unknown' if spawn('wslpath -w /') else None) is not None version_torch = get_version_torch() -path = find() +environment = None +err = None +try: + environment = find() +except Exception as e: + err = e blaslt_tensile_libpath = "" is_installed = False version = None -if path is not None: - blaslt_tensile_libpath = os.environ.get("HIPBLASLT_TENSILE_LIBPATH", os.path.join(path, "bin" if sys.platform == "win32" else "lib", "hipblaslt", "library")) +if environment is not None: + if isinstance(environment, ROCmEnvironment): + blaslt_tensile_libpath = os.environ.get("HIPBLASLT_TENSILE_LIBPATH", os.path.join(environment.path, "bin" if sys.platform == "win32" else "lib", "hipblaslt", "library")) is_installed = True version = get_version() diff --git a/modules/zluda_installer.py b/modules/zluda_installer.py index da0ee290f..2922f196f 100644 --- a/modules/zluda_installer.py +++ b/modules/zluda_installer.py @@ -119,12 +119,13 @@ def link_or_copy(src: os.PathLike, dst: os.PathLike): def load(): + assert isinstance(rocm.environment, rocm.ROCmEnvironment) global core, ml, hipBLASLt_enabled, MIOpen_enabled # pylint: disable=global-statement core = Core(ctypes.windll.LoadLibrary(os.path.join(path, 'nvcuda.dll'))) ml = ZLUDALibrary(ctypes.windll.LoadLibrary(os.path.join(path, 'nvml.dll'))) is_nightly = core.get_nightly_flag() == 1 - hipBLASLt_enabled = is_nightly and os.path.exists(rocm.blaslt_tensile_libpath) and os.path.exists(os.path.join(rocm.path, "bin", "hipblaslt.dll")) and default_agent is not None - MIOpen_enabled = is_nightly and os.path.exists(os.path.join(rocm.path, "bin", "MIOpen.dll")) + hipBLASLt_enabled = is_nightly and os.path.exists(rocm.blaslt_tensile_libpath) and os.path.exists(os.path.join(rocm.environment.path, "bin", "hipblaslt.dll")) and default_agent is not None + MIOpen_enabled = is_nightly and os.path.exists(os.path.join(rocm.environment.path, "bin", "MIOpen.dll")) if hipBLASLt_enabled: if not default_agent.blaslt_supported: @@ -147,19 +148,19 @@ def load(): os.environ["ZLUDA_NVRTC_LIB"] = os.path.join([v for v in site.getsitepackages() if v.endswith("site-packages")][0], "torch", "lib", "nvrtc64_112_0.dll") for v in HIPSDK_TARGETS: - ctypes.windll.LoadLibrary(os.path.join(rocm.path, 'bin', v)) + ctypes.windll.LoadLibrary(os.path.join(rocm.environment.path, 'bin', v)) for v in DLL_MAPPING.values(): ctypes.windll.LoadLibrary(os.path.join(path, v)) if hipBLASLt_enabled: os.environ.setdefault("DISABLE_ADDMM_CUDA_LT", "0") - ctypes.windll.LoadLibrary(os.path.join(rocm.path, 'bin', 'hipblaslt.dll')) + ctypes.windll.LoadLibrary(os.path.join(rocm.environment.path, 'bin', 'hipblaslt.dll')) ctypes.windll.LoadLibrary(os.path.join(path, 'cublasLt64_11.dll')) else: os.environ["DISABLE_ADDMM_CUDA_LT"] = "1" if MIOpen_enabled: - ctypes.windll.LoadLibrary(os.path.join(rocm.path, 'bin', 'MIOpen.dll')) + ctypes.windll.LoadLibrary(os.path.join(rocm.environment.path, 'bin', 'MIOpen.dll')) ctypes.windll.LoadLibrary(os.path.join(path, 'cudnn64_9.dll')) def conceal():