diff --git a/sgl-kernel/csrc/expert_specialization/es_fp8_blockwise_functor.cuh b/sgl-kernel/csrc/expert_specialization/es_fp8_blockwise_functor.cuh index c4af33447..db7f430f2 100644 --- a/sgl-kernel/csrc/expert_specialization/es_fp8_blockwise_functor.cuh +++ b/sgl-kernel/csrc/expert_specialization/es_fp8_blockwise_functor.cuh @@ -126,7 +126,7 @@ struct Fp8BlockwiseGroupedGemmProblemSizeFilterFunctor { Fp8BlockwiseGroupedGemmProblemSizeFilterFunctor(int* _problem_sizes) : problem_sizes(_problem_sizes) {} void CUTE_DEVICE operator()(int64_t expert_id, int m, int n, int k) { - if (m <= 48) { + if (m < 64) { // Swap A/B problem_sizes[expert_id * 3 + 0] = n; problem_sizes[expert_id * 3 + 1] = m; @@ -168,7 +168,7 @@ struct Fp8BlockwiseGroupedGemmProblemSizeFilterFunctor { Fp8BlockwiseGroupedGemmProblemSizeFilterFunctor(int* _problem_sizes) : problem_sizes(_problem_sizes) {} void CUTE_DEVICE operator()(int64_t expert_id, int m, int n, int k) { - if (m > 48 && m <= 96) { + if (m >= 64 && m < 128) { problem_sizes[expert_id * 3 + 0] = m; problem_sizes[expert_id * 3 + 1] = n; problem_sizes[expert_id * 3 + 2] = k; @@ -208,7 +208,7 @@ struct Fp8BlockwiseGroupedGemmProblemSizeFilterFunctor { Fp8BlockwiseGroupedGemmProblemSizeFilterFunctor(int* _problem_sizes) : problem_sizes(_problem_sizes) {} void CUTE_DEVICE operator()(int64_t expert_id, int m, int n, int k) { - if (m > 96) { + if (m >= 128) { problem_sizes[expert_id * 3 + 0] = m; problem_sizes[expert_id * 3 + 1] = n; problem_sizes[expert_id * 3 + 2] = k; diff --git a/sgl-kernel/csrc/expert_specialization/es_fp8_blockwise_launcher.cuh b/sgl-kernel/csrc/expert_specialization/es_fp8_blockwise_launcher.cuh index f6ab33eb0..3ed98821d 100644 --- a/sgl-kernel/csrc/expert_specialization/es_fp8_blockwise_launcher.cuh +++ b/sgl-kernel/csrc/expert_specialization/es_fp8_blockwise_launcher.cuh @@ -232,7 +232,7 @@ void es_sm90_fp8_blockwise_scaled_group_mm_distpatch_out_dtype( workspace, stream); } else { - launch_sm90_fp8_blockwise_scaled_group_mm( + launch_sm90_fp8_blockwise_scaled_group_mm( out_ptrs, a_ptrs, b_ptrs,