107 lines
3.5 KiB
Python
107 lines
3.5 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""Iluvatar/Tianshu accelerator adapter."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import re
|
|
|
|
from .base import (
|
|
AcceleratorInfo,
|
|
command_exists,
|
|
first_env_devices,
|
|
parse_memory_gb,
|
|
run_command,
|
|
)
|
|
|
|
|
|
class IluvatarAcceleratorAdapter:
|
|
vendor = "iluvatar"
|
|
runtime = "ix"
|
|
smi_command = "ixsmi"
|
|
visible_env_names = (
|
|
"ILUVATAR_VISIBLE_DEVICES",
|
|
"IX_VISIBLE_DEVICES",
|
|
"CUDA_VISIBLE_DEVICES",
|
|
)
|
|
|
|
def _query_device_count(self) -> tuple[int, str]:
|
|
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))
|
|
if indexes:
|
|
return len(indexes), table_output
|
|
if re.search(r"\bIluvatar\b|\b天数\b|\bIX\b", table_output, re.I):
|
|
return 1, table_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="ixsmi 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 "ixsmi did not report any device",
|
|
metadata={
|
|
"torch_cuda_available": torch_cuda_available,
|
|
"torch_version": torch_version,
|
|
"supports_sharded": available,
|
|
},
|
|
)
|