[Fix] Fall back to triton MOE for GPT-OSS on Blackwell with driver >= 595 (#21780)

Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Baizhou Zhang
2026-03-31 15:52:10 -07:00
committed by GitHub
co-authored by Claude Opus 4.6
parent 9191b02eda
commit f60f2ccc10
4 changed files with 61 additions and 32 deletions
+5 -15
View File
@@ -194,22 +194,12 @@ class GPUEnv(BaseEnv):
""" """
Get CUDA driver version. Get CUDA driver version.
""" """
versions = set() from sglang.srt.utils.common import get_nvidia_driver_version_str
try:
output = subprocess.check_output( ver = get_nvidia_driver_version_str()
[ if ver is None:
"nvidia-smi",
"--query-gpu=driver_version",
"--format=csv,noheader,nounits",
]
)
versions = set(output.decode().strip().split("\n"))
if len(versions) == 1:
return {"CUDA Driver Version": versions.pop()}
else:
return {"CUDA Driver Versions": ", ".join(sorted(versions))}
except subprocess.SubprocessError:
return {"CUDA Driver Version": "Not Available"} return {"CUDA Driver Version": "Not Available"}
return {"CUDA Driver Version": ver}
def get_topology(self): def get_topology(self):
""" """
+5 -13
View File
@@ -77,19 +77,11 @@ def _run_smi(query, query_type="gpu"):
def _get_smi_version(): def _get_smi_version():
"""Return nvidia-smi driver version and CUDA version, or None on failure.""" """Return nvidia-smi driver version and GPU name, or None on failure."""
try: from sglang.srt.utils.common import get_nvidia_driver_version_str
out = subprocess.check_output(
[ driver = get_nvidia_driver_version_str()
"nvidia-smi", if driver is None:
"--query-gpu=driver_version",
"--format=csv,noheader,nounits",
],
text=True,
timeout=10,
)
driver = out.strip().splitlines()[0].strip() if out.strip() else "unknown"
except (subprocess.SubprocessError, FileNotFoundError, IndexError):
return None return None
try: try:
out = subprocess.check_output( out = subprocess.check_output(
+12
View File
@@ -42,6 +42,7 @@ from sglang.srt.utils.common import (
get_device_name, get_device_name,
get_device_sm, get_device_sm,
get_int_env_var, get_int_env_var,
get_nvidia_driver_version,
get_quantization_config, get_quantization_config,
human_readable_int, human_readable_int,
is_blackwell_supported, is_blackwell_supported,
@@ -1757,6 +1758,17 @@ class ServerArgs:
and is_triton_kernels_available() and is_triton_kernels_available()
and self.quantization is None and self.quantization is None
): ):
# The triton_kernels package segfaults on Blackwell (B200)
# with NVIDIA driver >= 595. Fall back to triton backend.
if is_blackwell_supported() and get_nvidia_driver_version() >= (
595,
):
self.moe_runner_backend = "triton"
logger.warning(
"Detected GPT-OSS model on Blackwell with driver >= 595, "
"using triton MOE kernel to avoid triton_kernels SIGSEGV."
)
else:
self.moe_runner_backend = "triton_kernel" self.moe_runner_backend = "triton_kernel"
logger.warning( logger.warning(
"Detected GPT-OSS model, enabling triton_kernels MOE kernel." "Detected GPT-OSS model, enabling triton_kernels MOE kernel."
+35
View File
@@ -3399,6 +3399,41 @@ def is_triton_kernels_available() -> bool:
return importlib.util.find_spec("triton_kernels") is not None return importlib.util.find_spec("triton_kernels") is not None
@lru_cache(maxsize=1)
def get_nvidia_driver_version() -> tuple:
"""Return the NVIDIA driver version as a tuple of ints, e.g. (595, 58, 3).
Returns (0,) on failure."""
version_str = get_nvidia_driver_version_str()
if version_str is None:
return (0,)
try:
return tuple(int(x) for x in version_str.split("."))
except ValueError:
return (0,)
@lru_cache(maxsize=1)
def get_nvidia_driver_version_str() -> str:
"""Return the NVIDIA driver version string, e.g. '595.58.03'.
Returns None on failure."""
try:
result = subprocess.run(
[
"nvidia-smi",
"--query-gpu=driver_version",
"--format=csv,noheader,nounits",
],
capture_output=True,
text=True,
check=True,
timeout=10,
)
version_str = result.stdout.strip().split("\n")[0].strip()
return version_str if version_str else None
except (subprocess.CalledProcessError, FileNotFoundError, ValueError):
return None
def check_cuda_result(raw_output): def check_cuda_result(raw_output):
import cuda.bindings.runtime as cuda_rt import cuda.bindings.runtime as cuda_rt