110 lines
3.0 KiB
Python
110 lines
3.0 KiB
Python
# -*- 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
|