# -*- coding: utf-8 -*- """Shared accelerator adapter primitives.""" from __future__ import annotations import os import re import shutil import subprocess from dataclasses import dataclass, field from typing import Optional, Protocol @dataclass(frozen=True) class AcceleratorInfo: """Normalized hardware/runtime information used by the application.""" vendor: str runtime: str device: str device_count: int = 0 visible_devices: tuple[str, ...] = () total_memory_gb: float = 0.0 smi_command: Optional[str] = None available: bool = False reason: str = "" metadata: dict[str, object] = field(default_factory=dict) @property def is_gpu(self) -> bool: return self.vendor not in {"cpu", "unknown"} and self.available def as_dict(self) -> dict[str, object]: return { "vendor": self.vendor, "runtime": self.runtime, "device": self.device, "device_count": self.device_count, "visible_devices": list(self.visible_devices), "total_memory_gb": self.total_memory_gb, "smi_command": self.smi_command, "available": self.available, "reason": self.reason, "metadata": self.metadata, } @property def supports_sharded(self) -> bool: value = self.metadata.get("supports_sharded") return bool(value) class AcceleratorAdapter(Protocol): vendor: str runtime: str def detect(self) -> AcceleratorInfo: """Return normalized accelerator info.""" def command_exists(command: str) -> bool: return shutil.which(command) is not None def run_command(command: list[str], timeout: float = 3.0) -> str: try: completed = subprocess.run( command, check=False, capture_output=True, text=True, timeout=timeout, ) except (OSError, subprocess.TimeoutExpired): return "" if completed.returncode != 0: return "" return completed.stdout.strip() def parse_visible_devices(raw: str | None) -> tuple[str, ...]: value = (raw or "").strip() if not value or value.lower() in {"all", "none", "void"}: return () return tuple(part.strip() for part in value.split(",") if part.strip()) def first_env_devices(names: tuple[str, ...]) -> tuple[str, ...]: shared = parse_visible_devices(os.getenv("ASR_VISIBLE_DEVICES")) if shared: return shared for name in names: devices = parse_visible_devices(os.getenv(name)) if devices: return devices return () def parse_memory_gb(text: str) -> float: """Parse the smallest memory value from common smi outputs.""" values: list[float] = [] for number, unit in re.findall(r"([0-9]+(?:\.[0-9]+)?)\s*(GiB|GB|MiB|MB)", text, re.I): value = float(number) normalized_unit = unit.lower() if normalized_unit in {"mib", "mb"}: value = value / 1024 values.append(value) return min(values) if values else 0.0