[bugfix] guard NVIDIA SM-capability checks with is_cuda() for AMD/ROCm (#28486)

This commit is contained in:
Oxana Korzh
2026-06-16 22:51:42 -07:00
committed by GitHub
parent 0f5e14e1d9
commit c01f62e341
3 changed files with 6 additions and 5 deletions
@@ -17,6 +17,7 @@ import torch
from sglang.srt.layers.attention.linear.kernels.kernel_backend import ( from sglang.srt.layers.attention.linear.kernels.kernel_backend import (
LinearAttnKernelBase, LinearAttnKernelBase,
) )
from sglang.srt.utils import is_cuda
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -54,7 +55,7 @@ def _get_flashinfer_gdn_kernels():
_flashinfer_gated_delta_rule_mtp_bf16 = gated_delta_rule_mtp_bf16 _flashinfer_gated_delta_rule_mtp_bf16 = gated_delta_rule_mtp_bf16
_flashinfer_gated_delta_rule_decode = gated_delta_rule_decode_pretranspose _flashinfer_gated_delta_rule_decode = gated_delta_rule_decode_pretranspose
_flashinfer_gdn_available = ( _flashinfer_gdn_available = (
torch.cuda.is_available() and torch.cuda.get_device_capability()[0] >= 9 is_cuda() and torch.cuda.get_device_capability()[0] >= 9
) )
if _flashinfer_gdn_available: if _flashinfer_gdn_available:
logger.info("FlashInfer GDN kernels loaded successfully") logger.info("FlashInfer GDN kernels loaded successfully")
+2 -2
View File
@@ -14,7 +14,7 @@ from sglang.srt.layers.dp_attention import (
get_attention_dp_size, get_attention_dp_size,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.utils import is_npu from sglang.srt.utils import is_cuda, is_npu
_is_npu = is_npu() _is_npu = is_npu()
@@ -288,7 +288,7 @@ def initialize_moe_config(server_args: ServerArgs):
DEEPEP_CONFIG = server_args.deepep_config or "" DEEPEP_CONFIG = server_args.deepep_config or ""
IS_TBO_ENABLED = server_args.enable_two_batch_overlap IS_TBO_ENABLED = server_args.enable_two_batch_overlap
IS_SBO_ENABLED = server_args.enable_single_batch_overlap IS_SBO_ENABLED = server_args.enable_single_batch_overlap
if IS_SBO_ENABLED and torch.cuda.is_available(): if IS_SBO_ENABLED and is_cuda():
if torch.cuda.get_device_capability()[0] == 9: if torch.cuda.get_device_capability()[0] == 9:
raise ValueError( raise ValueError(
"SBO (single batch overlap) is not supported on SM90 GPUs with latest sgl-deep-gemm wheel. Please try removing --enable-single-batch-overlap argument." "SBO (single batch overlap) is not supported on SM90 GPUs with latest sgl-deep-gemm wheel. Please try removing --enable-single-batch-overlap argument."
+2 -2
View File
@@ -3399,7 +3399,7 @@ class ServerArgs:
if ( if (
decode == "flashinfer" decode == "flashinfer"
and self.mamba_ssm_dtype != "bfloat16" and self.mamba_ssm_dtype != "bfloat16"
and torch.cuda.is_available() and is_cuda()
and torch.cuda.get_device_capability()[0] >= 10 and torch.cuda.get_device_capability()[0] >= 10
): ):
raise ValueError( raise ValueError(
@@ -3415,7 +3415,7 @@ class ServerArgs:
cuda_major = int(cuda_version.split(".")[0]) if cuda_version is not None else 0 cuda_major = int(cuda_version.split(".")[0]) if cuda_version is not None else 0
if ( if (
prefill == "flashinfer" prefill == "flashinfer"
and torch.cuda.is_available() and is_cuda()
and torch.cuda.get_device_capability()[0] >= 10 and torch.cuda.get_device_capability()[0] >= 10
and cuda_major < 13 and cuda_major < 13
): ):