From 30b6c894247969f1d09e39a2f696d60d82f4c6be Mon Sep 17 00:00:00 2001 From: biondizzle Date: Fri, 15 May 2026 20:44:46 +0000 Subject: [PATCH] fix: correct SF remap coordinate extraction - First flattened group IS M/N (not K as previously assumed) - mn = f0 + 32*f1 + 128*f2 - k_sf = f4 + 4*f5 (f3 is stride-0 inner K, ignored) - The atom stride-0 dimension (f3) maps to offset 0, not a meaningful K sub-index. The actual k_sf comes from f4 (sub_k) + f5*4 (tile_k) - Original code had group assignment right but k_sf extraction wrong --- .../cutlass_nvfp4_gemm/cutlass_nvfp4_gemm.cu | 42 ++++++++++--------- 1 file changed, 22 insertions(+), 20 deletions(-) diff --git a/src/nvfp4_megamoe_kernel/cutlass_nvfp4_gemm/cutlass_nvfp4_gemm.cu b/src/nvfp4_megamoe_kernel/cutlass_nvfp4_gemm/cutlass_nvfp4_gemm.cu index 4e38a3ec..fbac8d4a 100644 --- a/src/nvfp4_megamoe_kernel/cutlass_nvfp4_gemm/cutlass_nvfp4_gemm.cu +++ b/src/nvfp4_megamoe_kernel/cutlass_nvfp4_gemm/cutlass_nvfp4_gemm.cu @@ -118,7 +118,7 @@ using LayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFB; template __global__ void remap_sf_to_cutlass_kernel( - const cutlass::float_ue4m3_t* __restrict__ src, // (MN, K_sf) or (K_sf, MN) depending on col_major_src + const cutlass::float_ue4m3_t* __restrict__ src, // (MN, K_sf) or (K_sf, MN) row-major cutlass::float_ue4m3_t* __restrict__ dst, // CUTLASS interleaved layout (zero-initialized) LayoutSF layout_sf, // CuTe layout for dst int MN, int K_sf, // Source dimensions (in SF groups) @@ -133,28 +133,30 @@ __global__ void remap_sf_to_cutlass_kernel( constexpr int R = cute::rank_v; - int m = 0, k_sf = 0; + int mn = 0; + int k_sf = 0; - if constexpr (R == 6) { - // K-major SfAtom tiled with Step<(2,1,3)>: - // Shape: ((32, 4, K_tiles), (SFVecSize, 4, M_tiles)) - // First group covers K dimension, second covers M dimension - // get<0..2> = K group, get<3..5> = M group - // K element index = get<0> + get<1>*32 + get<2>*128 - // k_sf = K_element_index / SFVecSize - // M element index = get<3> + get<4>*SFVecSize + get<5>*(SFVecSize*4) - k_sf = (cute::get<0>(flat) + cute::get<1>(flat) * 32 + cute::get<2>(flat) * 128) / InputSFVectorSize; - m = cute::get<3>(flat) + cute::get<4>(flat) * InputSFVectorSize + cute::get<5>(flat) * (InputSFVectorSize * 4); - } else if constexpr (R == 8) { - // With batch/L dimension: ((32, 4, K_tiles), (SFVecSize, 4, M_tiles), (1, 1, L)) - k_sf = (cute::get<0>(flat) + cute::get<1>(flat) * 32 + cute::get<2>(flat) * 128) / InputSFVectorSize; - m = cute::get<3>(flat) + cute::get<4>(flat) * InputSFVectorSize + cute::get<5>(flat) * (InputSFVectorSize * 4); - } else { - m = 0; k_sf = 0; + if constexpr (R >= 6) { + // K-major SF atom: Shape, Shape> + // Stride, Stride<_0,_1>> + // SFA tiled as make_shape(M,K,...), SFB as make_shape(N,K,...) + // First flattened group is M/N: ((32,4), tile_mn) + mn = + int(cute::get<0>(flat)) + + 32 * int(cute::get<1>(flat)) + + 128 * int(cute::get<2>(flat)); + + // Second flattened group is K: (SFVecSize, 4, tile_k) + // get<3> is the stride-0 k-within-SF-vector coordinate — ignore it. + k_sf = + int(cute::get<4>(flat)) + + 4 * int(cute::get<5>(flat)); } - if (m < MN && k_sf < K_sf) { - dst[dst_idx] = col_major_src ? src[k_sf * MN + m] : src[m * K_sf + k_sf]; + if (mn < MN && k_sf < K_sf) { + dst[dst_idx] = col_major_src + ? src[k_sf * MN + mn] // SFB source: (K_sf, N) + : src[mn * K_sf + k_sf]; // SFA source: (M, K_sf) } }