Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion python/pyproject_other.toml
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,6 @@ runtime_common = [
"grpcio==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
"bidict",
]

tracing = [
Expand Down
94 changes: 93 additions & 1 deletion python/sglang/check_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@

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():
Expand Down Expand Up @@ -423,11 +423,103 @@ def get_topology(self):
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()
)
Comment thread
yeahdongcn marked this conversation as resolved.
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
)
Comment thread
yeahdongcn marked this conversation as resolved.
}
except subprocess.SubprocessError:
return {}


if __name__ == "__main__":
if is_cuda_v2():
env = GPUEnv()
elif is_hip():
env = HIPEnv()
elif is_npu():
env = NPUEnv()
elif is_musa():
env = MUSAEnv()
env.check_env()
Comment thread
yeahdongcn marked this conversation as resolved.
9 changes: 9 additions & 0 deletions python/sglang/srt/utils/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -185,6 +185,15 @@ def is_cpu() -> bool:
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:
"""Check if dtype is float4_e2m1fn_x2 and CUDA is available."""
target_dtype = getattr(torch, "float4_e2m1fn_x2", None)
Expand Down
Loading