115 lines
3.7 KiB
Python
115 lines
3.7 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""Moore Threads / MUSA accelerator adapter."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import re
|
|
|
|
from .base import (
|
|
AcceleratorInfo,
|
|
command_exists,
|
|
first_env_devices,
|
|
parse_memory_gb,
|
|
run_command,
|
|
)
|
|
|
|
|
|
class MThreadsAcceleratorAdapter:
|
|
vendor = "mthreads"
|
|
runtime = "musa"
|
|
smi_command = "mthreads-gmi"
|
|
visible_env_names = (
|
|
"MTHREADS_VISIBLE_DEVICES",
|
|
"MUSA_VISIBLE_DEVICES",
|
|
"CUDA_VISIBLE_DEVICES",
|
|
)
|
|
|
|
def _query_device_count(self) -> tuple[int, str]:
|
|
for command in (
|
|
[self.smi_command, "-L"],
|
|
[self.smi_command, "list"],
|
|
[self.smi_command],
|
|
):
|
|
output = run_command(command)
|
|
if not output:
|
|
continue
|
|
lines = [line for line in output.splitlines() if line.strip()]
|
|
gpu_lines = [
|
|
line
|
|
for line in lines
|
|
if re.search(r"\b(gpu|device|card|musa|mthreads|moore)\b", line, re.I)
|
|
]
|
|
if gpu_lines:
|
|
return len(gpu_lines), output
|
|
|
|
indexes = set(
|
|
re.findall(r"(?:GPU|Device|Card)\s*[:#]?\s*([0-9]+)", output, re.I)
|
|
)
|
|
if not indexes:
|
|
indexes = set(re.findall(r"^\s*\|\s*([0-9]+)\s+", output, re.M))
|
|
if indexes:
|
|
return len(indexes), output
|
|
|
|
if re.search(r"\bMUSA\b|\bMoore\s+Threads\b|\bMThreads\b", output, re.I):
|
|
return max(len(lines), 1), output
|
|
|
|
return 0, ""
|
|
|
|
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="mthreads-gmi 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))
|
|
|
|
torch_cuda_available = False
|
|
torch_count = 0
|
|
torch_memory_gb = 0.0
|
|
torch_version = ""
|
|
try:
|
|
import torch
|
|
|
|
torch_version = getattr(torch, "__version__", "")
|
|
torch_cuda_available = bool(torch.cuda.is_available())
|
|
if torch_cuda_available:
|
|
torch_count = int(torch.cuda.device_count())
|
|
torch_memory_gb = min(
|
|
torch.cuda.get_device_properties(i).total_memory / (1024**3)
|
|
for i in range(torch_count)
|
|
)
|
|
except Exception:
|
|
pass
|
|
|
|
if torch_count:
|
|
device_count = min(torch_count, len(visible_devices)) if visible_devices else torch_count
|
|
total_memory_gb = torch_memory_gb or parse_memory_gb(raw_output)
|
|
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 "mthreads-gmi did not report any device",
|
|
metadata={
|
|
"torch_cuda_available": torch_cuda_available,
|
|
"torch_version": torch_version,
|
|
"supports_sharded": available,
|
|
},
|
|
)
|