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 c5852e5d..c4ab503f 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 @@ -159,47 +159,20 @@ __global__ void remap_sf_to_cutlass_kernel( int m = 0, k_sf = 0; if constexpr (R == 8) { - // 8 flattened coordinates: 4 for M, 4 for K - // M: (inner_32, inner_4, tile_interleave, tile_m) - // K: (inner_16, inner_4, tile_interleave_k, tile_k) + // 8 flattened coordinates from idx2crd: + // f0 = inner_m (0..31) — varies fastest within M atom + // f1 = sub_m (0..3) — second M sub-coordinate + // f2 = tile_m (0..) — M tile index + // f3 = step_m stride (degenerate — always equals total) — NOT a coordinate + // f4 = sub_k (0..3) — K sub-coordinate within atom + // f5 = tile_k (0..) — K tile index + // f6, f7 = 0 (unused higher tiling) // - // SfAtom shape (32,4) with stride (16,4): - // inner_32 * 16 + inner_4 * 4 = local offset within atom - // m = inner_32 * 4 + inner_4 + (tile_interleave + tile_m * 2) * 128 - // (step=2 means M tiles interleave, so tile_interleave is 0 or 1) - // k_sf = inner_16 is within one SF group (0..15), ignored - // inner_4 is the sub-K (0..3) - // k_sf = inner_4 + tile_interleave_k * 4 + tile_k * 4 - // Actually inner_16 covers 16 elements = 1 SF group, so it doesn't - // contribute to k_sf. k_sf = sub_k + (tile_k * 4) - // But wait, with the tile_interleave_k, it might be: - // k_sf = inner_4 + tile_interleave_k * 4 + tile_k * 8 - // - // Let me just compute m and k_sf empirically from the flat values: - int f0 = cute::get<0>(flat); - int f1 = cute::get<1>(flat); - int f2 = cute::get<2>(flat); - int f3 = cute::get<3>(flat); - int f4 = cute::get<4>(flat); - int f5 = cute::get<5>(flat); - int f6 = cute::get<6>(flat); - int f7 = cute::get<7>(flat); - - // M = f0..f3, K = f4..f7 - // f0 = inner_32 (0..31), f1 = inner_4 (0..3) - // These two give local_m = f0 * 4 + f1 (0..127) - // f2, f3 are M tiling: with Step<2>, f2 is interleave (0..1), f3 is tile (0..) - // m = (f3 * 2 + f2) * 128 + f0 * 4 + f1 - m = (f3 * 2 + f2) * 128 + f0 * 4 + f1; - - // f4 = inner_16 (0..15, within one SF group, doesn't contribute to k_sf) - // f5 = inner_4_k (0..3, K sub-index within atom) - // f6, f7 are K tiling - // k_sf = f5 + f6 * 4 + f7 * 8 (guessing the tiling pattern) - // Actually with Step<1> on K, the tiling is simpler: - // k_sf = f5 + (f6 + f7 * n_k_interleave) * 4 - // Let me try: k_sf = f5 + f6 * 4 (assuming f7 is outer tile) - k_sf = f5 + f6 * 4; // This may need adjustment based on printf output + // CuTe "first sub varies fastest" for Shape<32, 4>: + // m = f0 + f1 * 32 + f2 * 128 + // k_sf = f4 + f5 * 4 + m = cute::get<0>(flat) + cute::get<1>(flat) * 32 + cute::get<2>(flat) * 128; + k_sf = cute::get<4>(flat) + cute::get<5>(flat) * 4; } else { // Fallback: index 0 and 1 m = 0; k_sf = 0;