dougvk's picture
Publish Unlimited-OCR RDNA 4 runtime v0.1.0
f340984 verified
Raw
History Blame Contribute Delete
5.39 kB
from __future__ import annotations
import os
import sys
from dataclasses import asdict, dataclass
from importlib.metadata import PackageNotFoundError, version
from typing import Any
from .constants import (
HARDWARE_VERIFIED_ARCHITECTURES,
SUPPORTED_ARCHITECTURES,
VERIFIED_HIP_VERSION,
VERIFIED_PACKAGE_VERSIONS,
)
from .errors import RuntimeEnvironmentError
@dataclass(frozen=True)
class DeviceInfo:
index: int
name: str
architecture: str | None
uuid: str | None
@dataclass(frozen=True)
class RuntimeInfo:
torch_version: str
hip_version: str | None
cuda_available: bool
bf16_supported: bool
visible_devices: int
devices: tuple[DeviceInfo, ...]
package_versions: dict[str, str | None]
accepted_architecture: bool
hardware_verified: bool
software_verified: bool
warnings: tuple[str, ...]
def to_dict(self) -> dict[str, Any]:
return asdict(self)
def select_device(device: str | None) -> None:
if device is None:
return
if "torch" in sys.modules:
raise RuntimeEnvironmentError("GPU selection must happen before importing PyTorch")
existing_rocr = os.environ.get("ROCR_VISIBLE_DEVICES")
if existing_rocr is not None and existing_rocr != device:
raise RuntimeEnvironmentError(
f"--device {device!r} conflicts with existing ROCR_VISIBLE_DEVICES={existing_rocr!r}; unset one explicitly"
)
conflicts = [name for name in ("HIP_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES") if name in os.environ]
if conflicts:
raise RuntimeEnvironmentError(
f"--device cannot be combined with existing {', '.join(conflicts)}; "
"unset the conflicting visibility variable"
)
os.environ["ROCR_VISIBLE_DEVICES"] = device
def _installed_versions() -> dict[str, str | None]:
installed: dict[str, str | None] = {}
for package in VERIFIED_PACKAGE_VERSIONS:
try:
installed[package] = version(package)
except PackageNotFoundError:
installed[package] = None
return installed
def runtime_issues(info: RuntimeInfo, *, require_single: bool) -> tuple[str, ...]:
issues: list[str] = []
if info.hip_version is None or not info.cuda_available:
issues.append("PyTorch does not see a ROCm GPU; a CUDA-only wheel is not sufficient")
if require_single and info.visible_devices != 1:
issues.append(f"expected exactly one visible GPU, found {info.visible_devices}; pass --device INDEX_OR_UUID")
if not info.accepted_architecture:
architectures = ", ".join(device.architecture or "unknown" for device in info.devices) or "none"
issues.append(f"expected RDNA 4 (gfx1200/gfx1201), found: {architectures}")
if info.cuda_available and not info.bf16_supported:
issues.append("the selected GPU/runtime does not report BF16 support")
return tuple(issues)
def inspect_runtime(*, device: str | None = None, require_single: bool = False, validate: bool = True) -> RuntimeInfo:
select_device(device)
try:
import torch
except ImportError as exc:
raise RuntimeEnvironmentError(
"ROCm PyTorch is not installed. Run scripts/bootstrap-rocm.sh from the repository checkout."
) from exc
hip_version = getattr(torch.version, "hip", None)
cuda_available = bool(torch.cuda.is_available())
count = torch.cuda.device_count() if cuda_available else 0
devices: list[DeviceInfo] = []
for index in range(count):
properties = torch.cuda.get_device_properties(index)
devices.append(
DeviceInfo(
index=index,
name=properties.name,
architecture=getattr(properties, "gcnArchName", None),
uuid=str(getattr(properties, "uuid", "")) or None,
)
)
accepted_architecture = bool(devices) and all(device.architecture in SUPPORTED_ARCHITECTURES for device in devices)
hardware_verified = bool(devices) and all(
device.architecture in HARDWARE_VERIFIED_ARCHITECTURES for device in devices
)
package_versions = _installed_versions()
software_verified = hip_version == VERIFIED_HIP_VERSION and all(
package_versions[name] == expected for name, expected in VERIFIED_PACKAGE_VERSIONS.items()
)
warnings: list[str] = []
if accepted_architecture and not hardware_verified:
warnings.append(
"gfx1200 is accepted by the RDNA 4 guard but has not completed this project's GPU acceptance gate"
)
if not software_verified:
warnings.append("the installed software stack differs from the exact validated versions")
info = RuntimeInfo(
torch_version=torch.__version__,
hip_version=hip_version,
cuda_available=cuda_available,
bf16_supported=bool(torch.cuda.is_bf16_supported()) if cuda_available else False,
visible_devices=count,
devices=tuple(devices),
package_versions=package_versions,
accepted_architecture=accepted_architecture,
hardware_verified=hardware_verified,
software_verified=software_verified,
warnings=tuple(warnings),
)
issues = runtime_issues(info, require_single=require_single)
if validate and issues:
raise RuntimeEnvironmentError("; ".join(issues))
return info