From 0d3089601547dd99412dd51427733c304ab2084a Mon Sep 17 00:00:00 2001 From: blake-snc Date: Sun, 15 Feb 2026 08:44:47 -0800 Subject: [PATCH] fix(sgl-kernel): use >= 120 for SM12x CUDA kernel dispatch (#18750) Co-authored-by: Claude Opus 4.6 --- sgl-kernel/csrc/gemm/fp8_blockwise_gemm_kernel.cu | 2 +- sgl-kernel/csrc/gemm/nvfp4_scaled_mm_kernels.cu | 2 +- sgl-kernel/csrc/moe/nvfp4_blockwise_moe.cu | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/sgl-kernel/csrc/gemm/fp8_blockwise_gemm_kernel.cu b/sgl-kernel/csrc/gemm/fp8_blockwise_gemm_kernel.cu index b8b23c427..1274b1ac9 100644 --- a/sgl-kernel/csrc/gemm/fp8_blockwise_gemm_kernel.cu +++ b/sgl-kernel/csrc/gemm/fp8_blockwise_gemm_kernel.cu @@ -448,7 +448,7 @@ torch::Tensor fp8_blockwise_scaled_mm( #if defined(CUTLASS_ARCH_MMA_SM120A_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) #if defined(CUDA_VERSION) && CUDA_VERSION >= 12080 - if (sm_version == 120) { + if (sm_version >= 120) { if (out_dtype == torch::kBFloat16) { sm120_fp8_blockwise_dispatch_shape( out_padded, mat_a_padded, mat_b, scales_a_padded, scales_b); diff --git a/sgl-kernel/csrc/gemm/nvfp4_scaled_mm_kernels.cu b/sgl-kernel/csrc/gemm/nvfp4_scaled_mm_kernels.cu index 40d320ac1..28e86f7c8 100644 --- a/sgl-kernel/csrc/gemm/nvfp4_scaled_mm_kernels.cu +++ b/sgl-kernel/csrc/gemm/nvfp4_scaled_mm_kernels.cu @@ -663,7 +663,7 @@ void cutlass_scaled_fp4_mm_sm100a_sm120a( // Check SM version and dispatch accordingly auto sm_version = getSMVersion(); - if (sm_version == 120) { + if (sm_version >= 120) { // Use SM120 specific dispatch if (out_dtype == at::ScalarType::Half) { cutlass_fp4_f16_gemm_dispatch_sm120(D, A, B, A_sf, B_sf, alpha, m, n, k, stream); diff --git a/sgl-kernel/csrc/moe/nvfp4_blockwise_moe.cu b/sgl-kernel/csrc/moe/nvfp4_blockwise_moe.cu index 6dbfb7bf2..ba9da1d77 100644 --- a/sgl-kernel/csrc/moe/nvfp4_blockwise_moe.cu +++ b/sgl-kernel/csrc/moe/nvfp4_blockwise_moe.cu @@ -665,7 +665,7 @@ void cutlass_fp4_group_mm( N, K); } - } else if (sm_version == 120) { + } else if (sm_version >= 120) { if (output.scalar_type() == torch::kBFloat16) { run_fp4_blockwise_scaled_group_mm_sm120( output,