From 92d769b3d9ab9c611d8f3b7e7300e488a58ba5ac Mon Sep 17 00:00:00 2001 From: TANGBUDU <79953480+TANGBUDU@users.noreply.github.com> Date: Sat, 29 Aug 2026 01:59:14 +0800 Subject: [PATCH 1/2] Fix dropped K tail in blockscaled GEMV The tail guard used unroll_col_k, which is the gmem prefetch issue cursor. The prologue primes one buffer and each mainloop iteration issues one more ahead of consumption, so at loop exit unroll_col_k leads the accumulated K range by kStageCount * tileA_k_local. When floor(gemm_k / tileA_k_local) is a multiple of kStageCount and gemm_k is not tile aligned, the guard is false for every thread, process_tail_elements() is skipped, and up to a tile of K is dropped from the dot product. tile_idx advances by kStageCount per iteration and the FMA sequence within an iteration is unconditional, so tile_idx * tileA_k_local at loop exit is the K range that was accumulated. Use it as the tail origin. ptr_A and ptr_B are not advanced by the mainloop, so the absolute-K addressing inside process_tail_elements() is unaffected. ctest cases on example 91: k=1056, 1152 and 2176 fail before this change and pass after. k=288, 1280 and 1312 are controls against using total_tiles * tileA_k_local as the origin, which would double count. Addresses #3536, reported by @VaggelisGian. Not covered here: the same issue notes loads issued past the current K range. Since the FMA sequence is unconditional, the mainloop always runs a whole multiple of kStageCount tiles and accumulates past gemm_k when floor(gemm_k / tileA_k_local) % kStageCount != 0. Visible as wrong results for batch > 1, e.g. --m=256 --k=1280 --batch=2, and needs per-stage load/FMA predication. --- examples/91_fp4_gemv/CMakeLists.txt | 19 +++++++++++++++++++ .../cutlass/gemm/kernel/gemv_blockscaled.h | 7 ++++--- 2 files changed, 23 insertions(+), 3 deletions(-) diff --git a/examples/91_fp4_gemv/CMakeLists.txt b/examples/91_fp4_gemv/CMakeLists.txt index 04e1c23124..5223973f8e 100644 --- a/examples/91_fp4_gemv/CMakeLists.txt +++ b/examples/91_fp4_gemv/CMakeLists.txt @@ -28,9 +28,28 @@ if (NOT MSVC) +# Cover exact, partial-stage, and stage-boundary K tails. +set(TEST_K_288 --m=256 --k=288 --batch=1 --epilogue_st=1.0 --profiling=false) +set(TEST_K_1024 --m=256 --k=1024 --batch=1 --epilogue_st=1.0 --profiling=false) +set(TEST_K_1056 --m=256 --k=1056 --batch=1 --epilogue_st=1.0 --profiling=false) +set(TEST_K_1152 --m=256 --k=1152 --batch=1 --epilogue_st=1.0 --profiling=false) +set(TEST_K_1280 --m=256 --k=1280 --batch=1 --epilogue_st=1.0 --profiling=false) +set(TEST_K_1312 --m=256 --k=1312 --batch=1 --epilogue_st=1.0 --profiling=false) +set(TEST_K_2048 --m=256 --k=2048 --batch=1 --epilogue_st=1.0 --profiling=false) +set(TEST_K_2176 --m=256 --k=2176 --batch=1 --epilogue_st=1.0 --profiling=false) + cutlass_example_add_executable( 91_fp4_gemv 91_fp4_gemv.cu + TEST_COMMAND_OPTIONS + TEST_K_288 + TEST_K_1024 + TEST_K_1056 + TEST_K_1152 + TEST_K_1280 + TEST_K_1312 + TEST_K_2048 + TEST_K_2176 ) endif() diff --git a/include/cutlass/gemm/kernel/gemv_blockscaled.h b/include/cutlass/gemm/kernel/gemv_blockscaled.h index 9a09de9f28..461809a412 100644 --- a/include/cutlass/gemm/kernel/gemv_blockscaled.h +++ b/include/cutlass/gemm/kernel/gemv_blockscaled.h @@ -319,7 +319,7 @@ struct GemvBlockScaled Date: Sun, 27 Sep 2026 00:50:34 +0800 Subject: [PATCH 2/2] Bound blockscaled GEMV stages and cover batch isolation --- examples/91_fp4_gemv/91_fp4_gemv.cu | 81 ++++++++++++++++++- examples/91_fp4_gemv/CMakeLists.txt | 16 ++++ .../cutlass/gemm/kernel/gemv_blockscaled.h | 42 ++++++---- 3 files changed, 121 insertions(+), 18 deletions(-) diff --git a/examples/91_fp4_gemv/91_fp4_gemv.cu b/examples/91_fp4_gemv/91_fp4_gemv.cu index 6b9e6a9550..d75109898f 100644 --- a/examples/91_fp4_gemv/91_fp4_gemv.cu +++ b/examples/91_fp4_gemv/91_fp4_gemv.cu @@ -582,6 +582,79 @@ struct TestbedGemvFp4SFDBase return true; } + + // Exercise the K boundary with nonzero isolated contributions, then check + // that changing batch 1 cannot affect the logical D/SFD of batch 0. + bool test_k_boundaries(cutlass::MatrixCoord problem_size, int32_t batch_count) { + int m = problem_size.row(), k = problem_size.column(); + int tile_k = Gemv::kThreadsPerRow * Gemv::GemvKernel::kElementsPerAccess; + int full_end = k / tile_k * tile_k; + if (batch_count < 2 || full_end == 0 || full_end == k) { + std::cerr << "Boundary tests require at least two batches, a full K tile and a scalar tail.\n"; + return false; + } + if (!initialize(problem_size, batch_count)) return false; + cutlass::reference::host::TensorFill(tensor_B.host_view(), ElementB(1)); + cutlass::reference::host::TensorFill(tensor_SFA.host_view(), ElementSFA(1)); + cutlass::reference::host::TensorFill(tensor_SFB.host_view(), ElementSFB(1)); + tensor_B.sync_device(); + tensor_SFA.sync_device(); + tensor_SFB.sync_device(); + + auto set_active_columns = [&](int begin, int end) { + for (int row = 0; row < batch_count * m; ++row) { + for (int col = 0; col < k; ++col) { + tensor_A.host_view().at({row, col}) = ElementA(col >= begin && col < end ? 1 : 0); + } + } + tensor_A.sync_device(); + }; + auto output_batch0 = [&]() { + auto shape = cute::make_shape(m, 1, k, batch_count); + auto sf = cute::make_tensor(tensor_SFD.host_data(), + Sm1xxBlockScaledOutputConfig::tile_atom_to_shape_SFD(shape)); + std::vector result; + for (int row = 0; row < m; ++row) { + result.push_back(float(tensor_D.host_view().at({row, 0}))); + result.push_back(float(sf(row, 0, 0))); + } + return result; + }; + auto verify = [&]() { + if (!run_gemv(problem_size, batch_count, 1.0f, 0.0f, 1.0f, false, 1) || + !run_reference(problem_size, batch_count, 1.0f, 0.0f, 1.0f) || + !compare_reference()) return false; + auto output = output_batch0(); + for (size_t i = 0; i < output.size(); i += 2) { + if (output[i] != 0 && output[i + 1] != 0) return true; + } + std::cerr << "Boundary test produced only zero outputs.\n"; + return false; + }; + + set_active_columns(full_end - tile_k, full_end); + if (!verify()) return false; + set_active_columns(full_end, k); + if (!verify()) return false; + set_active_columns(0, k); + if (!verify()) return false; + auto before = output_batch0(); + for (int col = 0; col < k; ++col) { + tensor_B.host_view().at({k + col, 0}) = ElementB(2); + for (int row = m; row < 2 * m; ++row) { + tensor_A.host_view().at({row, col}) = ElementA(2); + } + } + tensor_A.sync_device(); + tensor_B.sync_device(); + if (!verify()) return false; + if (before != output_batch0()) { + std::cerr << "Batch 0 changed when only batch 1 inputs were modified.\n"; + return false; + } + return true; + } + bool profile(cutlass::MatrixCoord problem_size, int32_t batch_count, ElementCompute alpha, @@ -728,6 +801,7 @@ struct TestbedGemvFp4SFD : public TestbedGemvFp4SFDBase< struct Options { bool help = false; + bool test_k_boundaries = false; int m = 4096; int k = 2048; @@ -750,6 +824,7 @@ struct Options { return; } + test_k_boundaries = cmd.check_cmd_line_flag("test-k-boundaries"); cmd.get_cmd_line_argument("m", m); cmd.get_cmd_line_argument("k", k); cmd.get_cmd_line_argument("batch", batch); @@ -766,6 +841,7 @@ struct Options { out << "91_fp4_gemv\n\n" << " FP4 GEMV with block-scaled inputs and outputs.\n\n" << "Options:\n\n" + << " --test-k-boundaries Run deterministic K-boundary and batch-isolation checks (alpha=1, beta=0, ST=1)\n" << " --help If specified, displays this usage statement\n\n" << " --m= Sets the M extent of the GEMM\n" << " --k= Sets the K extent of the GEMM\n" @@ -839,6 +915,9 @@ run_fp4_gemv_device(Options const& options) GemvBlockScaled>; TestbedGemvFp4SFD testbed; + if (options.test_k_boundaries) { + return testbed.test_k_boundaries({options.m, options.k}, options.batch); + } bool pass = true; @@ -877,7 +956,7 @@ main(int argc, char const** argv) } - if (options.profiling) { + if (options.profiling && !options.test_k_boundaries) { // Start profiling printf("\nProfiling...\n"); passed = run_fp4_gemv_device(options); diff --git a/examples/91_fp4_gemv/CMakeLists.txt b/examples/91_fp4_gemv/CMakeLists.txt index 5223973f8e..8bc79cafa6 100644 --- a/examples/91_fp4_gemv/CMakeLists.txt +++ b/examples/91_fp4_gemv/CMakeLists.txt @@ -38,10 +38,26 @@ set(TEST_K_1312 --m=256 --k=1312 --batch=1 --epilogue_st=1.0 --profiling=false) set(TEST_K_2048 --m=256 --k=2048 --batch=1 --epilogue_st=1.0 --profiling=false) set(TEST_K_2176 --m=256 --k=2176 --batch=1 --epilogue_st=1.0 --profiling=false) +# Cover incomplete prologue/lookahead buffers and the remaining scalar tail. +set(TEST_PROLOGUE_1 --m=256 --k=256 --batch=2 --epilogue_st=1.0 --profiling=false) +set(TEST_PROLOGUE_2 --m=256 --k=512 --batch=2 --epilogue_st=1.0 --profiling=false) +set(TEST_PROLOGUE_3 --m=256 --k=768 --batch=3 --epilogue_st=1.0 --profiling=false) +set(TEST_PREFETCH_1 --m=256 --k=1280 --batch=2 --epilogue_st=1.0 --profiling=false) +set(TEST_PREFETCH_3 --m=256 --k=1792 --batch=3 --epilogue_st=1.0 --profiling=false) +set(TEST_PARTIAL_TILE --m=256 --k=1312 --batch=3 --epilogue_st=1.0 --profiling=false) +set(TEST_K_BOUNDARIES --m=256 --k=1312 --batch=2 --test-k-boundaries) + cutlass_example_add_executable( 91_fp4_gemv 91_fp4_gemv.cu TEST_COMMAND_OPTIONS + TEST_PROLOGUE_1 + TEST_PROLOGUE_2 + TEST_PROLOGUE_3 + TEST_PREFETCH_1 + TEST_PREFETCH_3 + TEST_PARTIAL_TILE + TEST_K_BOUNDARIES TEST_K_288 TEST_K_1024 TEST_K_1056 diff --git a/include/cutlass/gemm/kernel/gemv_blockscaled.h b/include/cutlass/gemm/kernel/gemv_blockscaled.h index 461809a412..cd24a16f9f 100644 --- a/include/cutlass/gemm/kernel/gemv_blockscaled.h +++ b/include/cutlass/gemm/kernel/gemv_blockscaled.h @@ -318,8 +318,9 @@ struct GemvBlockScaled{}) % kStageCount; int frag_idx_next = (k_block + 1) & 1; - // Prefetch next kblock data using saved pipe index + // Every stage is initialized, including zero-filled lookahead. load_smem_fragments( fragA_reg[frag_idx_next], fragB_reg[frag_idx_next], @@ -436,9 +437,6 @@ struct GemvBlockScaled(); __syncthreads(); - // Tail elements that don't fill a full tile - const int tail_col_k = tile_idx * tileA_k_local; - if (tail_col_k + idx_col_k * kPackedElementsA < gemm_k) { - accum += process_tail_elements(tail_col_k, idx_col_k, gemm_k, + // Only complete tiles were accumulated; the prefetch cursor may + // have advanced beyond them. Start the scalar tail at mainloop_k. + if (mainloop_k + idx_col_k * kPackedElementsA < gemm_k) { + accum += process_tail_elements(mainloop_k, idx_col_k, gemm_k, ptr_A, ptr_B, ptr_SF_A, ptr_SF_B, A_converter, B_converter, @@ -523,7 +522,7 @@ struct GemvBlockScaled(smem_ptr_SFA, gmem_ptr_SFA, valid_tile && is_even_thread); + if (is_even_thread) { + cutlass::arch::cp_async_zfill(smem_ptr_SFA, gmem_ptr_SFA, valid_tile); + } void *smem_ptr_SFB = &shared_storage.smem_SFB[buffer_idx][s][(threadIdx.x / 2) * 4]; const void *gmem_ptr_SFB = ptr_SF_B + SF_offset_by_k; // Load 4 FP8 values (32 bits) - for this thread and next thread, only if threadIdx.y == 0 - cutlass::arch::cp_async(smem_ptr_SFB, gmem_ptr_SFB, valid_tile && load_b && is_even_thread); + if (load_b && is_even_thread) { + cutlass::arch::cp_async_zfill(smem_ptr_SFB, gmem_ptr_SFB, valid_tile); + } void *smem_ptr_A = &shared_storage.smem_A[buffer_idx][s][smem_offset_A]; const void *gmem_ptr_A = ptr_A + unroll_col_k / kPackedElementsA; - cutlass::arch::cp_async(smem_ptr_A, gmem_ptr_A, valid_tile); + cutlass::arch::cp_async_zfill(smem_ptr_A, gmem_ptr_A, valid_tile); void *smem_ptr_B = &shared_storage.smem_B[buffer_idx][s][smem_offset_B]; const void *gmem_ptr_B = ptr_B + unroll_col_k / kPackedElementsB; - cutlass::arch::cp_async(smem_ptr_B, gmem_ptr_B, valid_tile && load_b); + if (load_b) { + cutlass::arch::cp_async_zfill(smem_ptr_B, gmem_ptr_B, valid_tile); + } unroll_col_k += tileA_k_local; global_k += tileA_k_local;