| import dataclasses | |
| from functools import lru_cache | |
| import logging | |
| import platform | |
| import re | |
| import subprocess | |
| from typing import Optional | |
| import torch | |
| class CUDASpecs: | |
| highest_compute_capability: tuple[int, int] | |
| cuda_version_string: str | |
| cuda_version_tuple: tuple[int, int] | |
| def has_imma(self) -> bool: | |
| return torch.version.hip or self.highest_compute_capability >= (7, 5) | |
| def get_compute_capabilities() -> list[tuple[int, int]]: | |
| return sorted(torch.cuda.get_device_capability(torch.cuda.device(i)) for i in range(torch.cuda.device_count())) | |
| def get_cuda_version_tuple() -> Optional[tuple[int, int]]: | |
| """Get CUDA/HIP version as a tuple of (major, minor).""" | |
| try: | |
| if torch.version.cuda: | |
| version_str = torch.version.cuda | |
| elif torch.version.hip: | |
| version_str = torch.version.hip | |
| else: | |
| return None | |
| parts = version_str.split(".") | |
| if len(parts) >= 2: | |
| return tuple(map(int, parts[:2])) | |
| return None | |
| except (AttributeError, ValueError, IndexError): | |
| return None | |
| def get_cuda_version_string() -> Optional[str]: | |
| """Get CUDA/HIP version as a string.""" | |
| version_tuple = get_cuda_version_tuple() | |
| if version_tuple is None: | |
| return None | |
| major, minor = version_tuple | |
| return f"{major}{minor}" | |
| def get_cuda_specs() -> Optional[CUDASpecs]: | |
| """Get CUDA/HIP specifications.""" | |
| if not torch.cuda.is_available(): | |
| return None | |
| try: | |
| compute_capabilities = get_compute_capabilities() | |
| if not compute_capabilities: | |
| return None | |
| version_tuple = get_cuda_version_tuple() | |
| if version_tuple is None: | |
| return None | |
| version_string = get_cuda_version_string() | |
| if version_string is None: | |
| return None | |
| return CUDASpecs( | |
| highest_compute_capability=compute_capabilities[-1], | |
| cuda_version_string=version_string, | |
| cuda_version_tuple=version_tuple, | |
| ) | |
| except Exception: | |
| return None | |
| def get_rocm_gpu_arch() -> str: | |
| """Get ROCm GPU architecture.""" | |
| logger = logging.getLogger(__name__) | |
| try: | |
| if torch.version.hip: | |
| # On Windows, use hipinfo.exe; on Linux, use rocminfo | |
| if platform.system() == "Windows": | |
| cmd = ["hipinfo.exe"] | |
| arch_pattern = r"gcnArchName:\s+gfx([a-zA-Z\d]+)" | |
| else: | |
| cmd = ["rocminfo"] | |
| arch_pattern = r"Name:\s+gfx([a-zA-Z\d]+)" | |
| result = subprocess.run(cmd, capture_output=True, text=True) | |
| match = re.search(arch_pattern, result.stdout) | |
| if match: | |
| return "gfx" + match.group(1) | |
| else: | |
| return "unknown" | |
| else: | |
| return "unknown" | |
| except Exception as e: | |
| logger.error(f"Could not detect ROCm GPU architecture: {e}") | |
| if torch.cuda.is_available(): | |
| logger.warning( | |
| """ | |
| ROCm GPU architecture detection failed despite ROCm being available. | |
| """, | |
| ) | |
| return "unknown" | |