# -*- 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, }, )