[MUSA][1/N] sglang.check_env (#16959)
Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
This commit is contained in:
@@ -66,7 +66,6 @@ runtime_common = [
|
|||||||
"grpcio==1.75.1", # keep it align with compile_proto.py
|
"grpcio==1.75.1", # keep it align with compile_proto.py
|
||||||
"grpcio-tools==1.75.1", # keep it align with compile_proto.py
|
"grpcio-tools==1.75.1", # keep it align with compile_proto.py
|
||||||
"grpcio-reflection==1.75.1", # required by srt/entrypoints/grpc_server.py
|
"grpcio-reflection==1.75.1", # required by srt/entrypoints/grpc_server.py
|
||||||
"bidict",
|
|
||||||
]
|
]
|
||||||
|
|
||||||
tracing = [
|
tracing = [
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ from collections import OrderedDict, defaultdict
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.utils import is_hip, is_npu
|
from sglang.srt.utils import is_hip, is_musa, is_npu
|
||||||
|
|
||||||
|
|
||||||
def is_cuda_v2():
|
def is_cuda_v2():
|
||||||
@@ -423,6 +423,96 @@ class NPUEnv(BaseEnv):
|
|||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
|
||||||
|
class MUSAEnv(BaseEnv):
|
||||||
|
"""Environment checker for MThreads GPU"""
|
||||||
|
|
||||||
|
def get_info(self):
|
||||||
|
musa_info = {"MUSA available": torch.musa.is_available()}
|
||||||
|
|
||||||
|
if musa_info["MUSA available"]:
|
||||||
|
musa_info.update(self.get_device_info())
|
||||||
|
musa_info.update(self._get_musa_version_info())
|
||||||
|
|
||||||
|
return musa_info
|
||||||
|
|
||||||
|
def _get_musa_version_info(self):
|
||||||
|
"""
|
||||||
|
Get MUSA version information.
|
||||||
|
"""
|
||||||
|
from torch_musa.utils.musa_extension import MUSA_HOME
|
||||||
|
|
||||||
|
musa_info = {"MUSA_HOME": MUSA_HOME}
|
||||||
|
|
||||||
|
if MUSA_HOME and os.path.isdir(MUSA_HOME):
|
||||||
|
musa_info.update(self._get_mcc_info())
|
||||||
|
musa_info.update(self._get_musa_driver_version())
|
||||||
|
|
||||||
|
return musa_info
|
||||||
|
|
||||||
|
def _get_mcc_info(self):
|
||||||
|
"""
|
||||||
|
Get MCC version information.
|
||||||
|
"""
|
||||||
|
from torch_musa.utils.musa_extension import MUSA_HOME
|
||||||
|
|
||||||
|
try:
|
||||||
|
mcc = os.path.join(MUSA_HOME, "bin/mcc")
|
||||||
|
mcc_output = (
|
||||||
|
subprocess.check_output(f'"{mcc}" --version', shell=True)
|
||||||
|
.decode("utf-8")
|
||||||
|
.strip()
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"MCC": mcc_output[
|
||||||
|
mcc_output.rfind("mcc version") : mcc_output.rfind("Target")
|
||||||
|
].strip()
|
||||||
|
}
|
||||||
|
except subprocess.SubprocessError:
|
||||||
|
return {"MCC": "Not Available"}
|
||||||
|
|
||||||
|
def _get_musa_driver_version(self):
|
||||||
|
"""
|
||||||
|
Get MUSA driver version.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
output = subprocess.check_output(
|
||||||
|
[
|
||||||
|
"mthreads-gmi",
|
||||||
|
"-q",
|
||||||
|
],
|
||||||
|
text=True,
|
||||||
|
)
|
||||||
|
driver_version = None
|
||||||
|
for line in output.splitlines():
|
||||||
|
if "Driver Version" in line:
|
||||||
|
driver_version = line.split(":", 1)[1].strip()
|
||||||
|
break
|
||||||
|
|
||||||
|
return {"MUSA Driver Version": driver_version}
|
||||||
|
except subprocess.SubprocessError:
|
||||||
|
return {"MUSA Driver Version": "Not Available"}
|
||||||
|
|
||||||
|
def get_topology(self):
|
||||||
|
"""
|
||||||
|
Get GPU topology information.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
result = subprocess.run(
|
||||||
|
["mthreads-gmi", "topo", "-m"],
|
||||||
|
stdout=subprocess.PIPE,
|
||||||
|
stderr=subprocess.PIPE,
|
||||||
|
text=True,
|
||||||
|
check=True,
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"MTHREADS Topology": (
|
||||||
|
"\n" + result.stdout if result.returncode == 0 else None
|
||||||
|
)
|
||||||
|
}
|
||||||
|
except subprocess.SubprocessError:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
if is_cuda_v2():
|
if is_cuda_v2():
|
||||||
env = GPUEnv()
|
env = GPUEnv()
|
||||||
@@ -430,4 +520,6 @@ if __name__ == "__main__":
|
|||||||
env = HIPEnv()
|
env = HIPEnv()
|
||||||
elif is_npu():
|
elif is_npu():
|
||||||
env = NPUEnv()
|
env = NPUEnv()
|
||||||
|
elif is_musa():
|
||||||
|
env = MUSAEnv()
|
||||||
env.check_env()
|
env.check_env()
|
||||||
|
|||||||
@@ -185,6 +185,15 @@ def is_cpu() -> bool:
|
|||||||
return os.getenv("SGLANG_USE_CPU_ENGINE", "0") == "1" and is_host_cpu_supported
|
return os.getenv("SGLANG_USE_CPU_ENGINE", "0") == "1" and is_host_cpu_supported
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
|
def is_musa() -> bool:
|
||||||
|
try:
|
||||||
|
import torchada # noqa: F401
|
||||||
|
except ImportError:
|
||||||
|
return False
|
||||||
|
return hasattr(torch.version, "musa") and torch.version.musa is not None
|
||||||
|
|
||||||
|
|
||||||
def is_float4_e2m1fn_x2(dtype) -> bool:
|
def is_float4_e2m1fn_x2(dtype) -> bool:
|
||||||
"""Check if dtype is float4_e2m1fn_x2 and CUDA is available."""
|
"""Check if dtype is float4_e2m1fn_x2 and CUDA is available."""
|
||||||
target_dtype = getattr(torch, "float4_e2m1fn_x2", None)
|
target_dtype = getattr(torch, "float4_e2m1fn_x2", None)
|
||||||
|
|||||||
Reference in New Issue
Block a user