From aa209ddd2171f5565a02395e00f3404017d0625c Mon Sep 17 00:00:00 2001 From: biondizzle Date: Fri, 15 May 2026 22:59:44 +0000 Subject: [PATCH] debug: add SF remap roundtrip verifier Checks that forward remap wrote the correct bytes by comparing src[mn*stride_mn + k_sf*stride_ksf] against dst[layout_sf(make_coord(mn, k_sf*16, 0))]. Prints error count for SFA and SFB on first GEMM call. --- .../cutlass_nvfp4_gemm/cutlass_nvfp4_gemm.cu | 58 +++++++++++++++++++ 1 file changed, 58 insertions(+) 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 66024447..39077d4c 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 @@ -141,6 +141,35 @@ __global__ void remap_sf_to_cutlass_kernel( dst[dst_idx] = src[mn * src_stride_mn + k_sf * src_stride_ksf]; } +// Roundtrip verifier: check that forward remap wrote the correct bytes +template +__global__ void check_sf_forward_kernel( + const cutlass::float_ue4m3_t* src, + const cutlass::float_ue4m3_t* dst, + LayoutSF layout_sf, + int MN, + int K_sf, + int src_stride_mn, + int src_stride_ksf, + int* errors +) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= MN * K_sf) return; + + int mn = tid / K_sf; + int k_sf = tid % K_sf; + + int src_idx = mn * src_stride_mn + k_sf * src_stride_ksf; + int dst_idx = layout_sf(cute::make_coord(mn, k_sf * 16, 0)); + + auto* src_u8 = reinterpret_cast(src); + auto* dst_u8 = reinterpret_cast(dst); + + if (src_u8[src_idx] != dst_u8[dst_idx]) { + atomicAdd(errors, 1); + } +} + ///////////////////////////////////////////////////////////////////////////////////////////////// // C API ///////////////////////////////////////////////////////////////////////////////////////////////// @@ -187,6 +216,35 @@ int cutlass_nvfp4_gemm_run( static_cast(SFB_ptr), sfb_cutlass.get(), layout_SFB, N, K_sf, 1, N); // SFB source: row-major (K_sf, N) after transpose + // One-time roundtrip verification of SF remap + static bool verified = false; + if (!verified) { + verified = true; + cudaStreamSynchronize(stream); + int* d_errors; + cudaMalloc(&d_errors, sizeof(int)); + cudaMemset(d_errors, 0, sizeof(int)); + + check_sf_forward_kernel<<<(sfa_src_total + block - 1) / block, block, 0, stream>>>( + static_cast(SFA_ptr), sfa_cutlass.get(), layout_SFA, + M, K_sf, K_sf, 1, d_errors); + int sfa_errors = 0; + cudaMemcpyAsync(&sfa_errors, d_errors, sizeof(int), cudaMemcpyDeviceToHost, stream); + + cudaMemset(d_errors, 0, sizeof(int)); + check_sf_forward_kernel<<<(sfb_src_total + block - 1) / block, block, 0, stream>>>( + static_cast(SFB_ptr), sfb_cutlass.get(), layout_SFB, + N, K_sf, 1, N, d_errors); + int sfb_errors = 0; + cudaMemcpyAsync(&sfb_errors, d_errors, sizeof(int), cudaMemcpyDeviceToHost, stream); + + cudaStreamSynchronize(stream); + printf("[SF-VERIFY] M=%d N=%d K=%d K_sf=%d sfa_errors=%d sfb_errors=%d " + "sfa_size=%d sfb_size=%d\n", + M, N, K, K_sf, sfa_errors, sfb_errors, sfa_size, sfb_size); + cudaFree(d_errors); + } + typename Gemm::Arguments arguments { cutlass::gemm::GemmUniversalMode::kGemm, {M, N, K, 1},