test/app/core/accelerators/metax.py

175 lines
6.1 KiB
Python

# -*- coding: utf-8 -*-
"""MetaX/MuXi MACA accelerator adapter."""
from __future__ import annotations
import json
import re
from typing import Any
from .base import (
AcceleratorInfo,
command_exists,
first_env_devices,
parse_memory_gb,
run_command,
)
class MetaxAcceleratorAdapter:
vendor = "metax"
runtime = "maca"
smi_command = "mx-smi"
visible_env_names = (
"METAX_VISIBLE_DEVICES",
"MACA_VISIBLE_DEVICES",
"MX_VISIBLE_DEVICES",
)
def _json_output(self, command: list[str]) -> Any:
output = run_command(command)
if not output:
return None
try:
return json.loads(output)
except json.JSONDecodeError:
return None
def _walk_json(self, value: Any):
if isinstance(value, dict):
yield value
for child in value.values():
yield from self._walk_json(child)
elif isinstance(value, list):
for item in value:
yield from self._walk_json(item)
def _count_from_json(self, value: Any) -> int:
max_index = -1
device_like = 0
for item in self._walk_json(value):
keys = {str(key).lower(): key for key in item.keys()}
if any(key in keys for key in ("gpu id", "gpu_id", "gpu", "device id", "device_id", "index")):
device_like += 1
for key_name in ("gpu id", "gpu_id", "device id", "device_id", "index", "id"):
original = keys.get(key_name)
if original is None:
continue
try:
max_index = max(max_index, int(item[original]))
except (TypeError, ValueError):
continue
if max_index >= 0:
return max_index + 1
return device_like
def _memory_from_json(self, value: Any) -> float:
memory_values: list[float] = []
for item in self._walk_json(value):
for raw_key, raw_value in item.items():
key = str(raw_key).lower()
if not any(token in key for token in ("memory", "mem", "hbm", "vram", "容量")):
continue
if isinstance(raw_value, (int, float)):
numeric = float(raw_value)
# mx-smi reports are commonly MiB for memory counters.
if numeric > 1024:
numeric = numeric / 1024
memory_values.append(numeric)
continue
if isinstance(raw_value, str):
parsed = parse_memory_gb(raw_value)
if parsed > 0:
memory_values.append(parsed)
return min(memory_values) if memory_values else 0.0
def _query_device_count(self) -> tuple[int, str]:
json_info = self._json_output([self.smi_command, "-j"])
json_count = self._count_from_json(json_info)
if json_count > 0:
return json_count, json.dumps(json_info, ensure_ascii=False)
list_output = run_command([self.smi_command, "-L"])
if list_output:
lines = [line for line in list_output.splitlines() if line.strip()]
gpu_lines = [
line
for line in lines
if re.search(r"\b(gpu|device|card)\b", line, re.I)
]
return len(gpu_lines or lines), list_output
table_output = run_command([self.smi_command])
if table_output:
indexes = set(re.findall(r"(?:GPU|Device|Card)\s*[:#]?\s*([0-9]+)", table_output, re.I))
if not indexes:
indexes = set(re.findall(r"^\s*\|\s*([0-9]+)\s+", table_output, re.M))
return len(indexes), table_output
return 0, ""
def _query_memory_gb(self, fallback_output: str) -> float:
for command in (
[self.smi_command, "--show-memory", "-j"],
[self.smi_command, "--show-hwinfo", "-j"],
[self.smi_command, "-j"],
):
parsed = self._memory_from_json(self._json_output(command))
if parsed > 0:
return parsed
memory_output = run_command([self.smi_command, "--show-memory"])
hwinfo_output = run_command([self.smi_command, "--show-hwinfo"])
return (
parse_memory_gb(memory_output)
or parse_memory_gb(hwinfo_output)
or parse_memory_gb(fallback_output)
)
def detect(self) -> AcceleratorInfo:
visible_devices = first_env_devices(self.visible_env_names)
smi_available = command_exists(self.smi_command)
if not smi_available:
return AcceleratorInfo(
vendor=self.vendor,
runtime=self.runtime,
device="cpu",
visible_devices=visible_devices,
smi_command=None,
available=False,
reason="mx-smi not found",
)
device_count, raw_output = self._query_device_count()
if visible_devices:
device_count = min(device_count or len(visible_devices), len(visible_devices))
total_memory_gb = self._query_memory_gb(raw_output)
torch_cuda_available = False
torch_version = ""
try:
import torch
torch_version = getattr(torch, "__version__", "")
torch_cuda_available = bool(torch.cuda.is_available())
except Exception:
pass
available = device_count > 0
return AcceleratorInfo(
vendor=self.vendor,
runtime=self.runtime,
device="cuda:0" if available else "cpu",
device_count=device_count,
visible_devices=visible_devices,
total_memory_gb=total_memory_gb,
smi_command=self.smi_command,
available=available,
reason="" if available else "mx-smi did not report any device",
metadata={
"torch_cuda_available": torch_cuda_available,
"torch_version": torch_version,
"supports_sharded": available,
},
)