175 lines
6.1 KiB
Python
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,
|
|
},
|
|
)
|