test/app/core/accelerators/base.py

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