From 2fafefb7b9c7e233c1cfb258f38139a6406d1511 Mon Sep 17 00:00:00 2001 From: Qi Yuhang <45795032+HydraQYH@users.noreply.github.com> Date: Fri, 23 Jan 2026 15:56:52 +0800 Subject: [PATCH] [Bug Fix]Set NumSplitsM to 1 when TileShapeM < 128 in sm90 fp8 blockwise scaling CollectiveMma (#2965) * Fix NumSplitsM when TileShapeM < 128. * Use cute::conditional_t to replace std::conditional_t. --- ...array_tma_gmma_ss_warpspecialized_fp8_blockwise_scaling.hpp | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/include/cutlass/gemm/collective/sm90_mma_array_tma_gmma_ss_warpspecialized_fp8_blockwise_scaling.hpp b/include/cutlass/gemm/collective/sm90_mma_array_tma_gmma_ss_warpspecialized_fp8_blockwise_scaling.hpp index f27731bc..20c40956 100644 --- a/include/cutlass/gemm/collective/sm90_mma_array_tma_gmma_ss_warpspecialized_fp8_blockwise_scaling.hpp +++ b/include/cutlass/gemm/collective/sm90_mma_array_tma_gmma_ss_warpspecialized_fp8_blockwise_scaling.hpp @@ -212,7 +212,8 @@ struct CollectiveMma< static_assert(cute::is_same_v, "ElementAccumulator and ElementBlockScale should be same datatype"); - using NumSplitsM = cute::C(TileShape_{}) / 128>; + // For TileShapeM < 128, NumSplitsM should be 1 + using NumSplitsM = cute::conditional_t(TileShape_{}) < _128{}, _1, cute::C(TileShape_{}) / 128>>; static_assert(NumSplitsM{} == 1 || NumSplitsM{} == 2); struct SharedStorage {