Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
81 changes: 80 additions & 1 deletion examples/91_fp4_gemv/91_fp4_gemv.cu
Original file line number Diff line number Diff line change
Expand Up @@ -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<float> 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,
Expand Down Expand Up @@ -728,6 +801,7 @@ struct TestbedGemvFp4SFD : public TestbedGemvFp4SFDBase<

struct Options {
bool help = false;
bool test_k_boundaries = false;

int m = 4096;
int k = 2048;
Expand All @@ -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);
Expand All @@ -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=<int> Sets the M extent of the GEMM\n"
<< " --k=<int> Sets the K extent of the GEMM\n"
Expand Down Expand Up @@ -839,6 +915,9 @@ run_fp4_gemv_device(Options const& options)
GemvBlockScaled<ElementA, LayoutA, ElementB, ElementD, ElementAccumulatorMainloop, EpilogueOp, kElementsPerAccess>>;

TestbedGemvFp4SFD<Gemv> testbed;
if (options.test_k_boundaries) {
return testbed.test_k_boundaries({options.m, options.k}, options.batch);
}

bool pass = true;

Expand Down Expand Up @@ -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);
Expand Down
35 changes: 35 additions & 0 deletions examples/91_fp4_gemv/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -28,9 +28,44 @@

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)

# 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
TEST_K_1152
TEST_K_1280
TEST_K_1312
TEST_K_2048
TEST_K_2176
)

endif()
41 changes: 25 additions & 16 deletions include/cutlass/gemm/kernel/gemv_blockscaled.h
Original file line number Diff line number Diff line change
Expand Up @@ -318,8 +318,9 @@ struct GemvBlockScaled<ElementA_,
// Local aliases
const int tileA_k_local = kThreadsPerRow * kElementsPerAccess;
const int total_tiles = gemm_k / tileA_k_local;
const int mainloop_k = total_tiles * tileA_k_local;

int unroll_col_k = 0; // total K elements consumed so far by this thread
int unroll_col_k = 0; // K position of the next global-memory prefetch
const int thread_id = threadIdx.y * kThreadsPerRow + threadIdx.x;
const bool is_even_thread = (threadIdx.x % 2 == 0);
const bool load_b = (threadIdx.y == 0);
Expand All @@ -343,7 +344,7 @@ struct GemvBlockScaled<ElementA_,
// Only one row of threads (threadIdx.y == 0) loads B
const int smem_offset_B = threadIdx.x * (kElementsPerAccess / kPackedElementsB);

// PROLOGUE – prime first kStageCount-1 stages into buffer 0
// PROLOGUE - prime the first buffer with complete K tiles only
CUTLASS_PRAGMA_UNROLL
for (int b = 0; b < kBufferCount - 1; ++b) {
// Load all stages using the helper function
Expand All @@ -358,7 +359,7 @@ struct GemvBlockScaled<ElementA_,
smem_sf_write_offset,
is_even_thread,
load_b,
true, // valid_tile = true for prologue
mainloop_k, // exclusive end of complete K tiles
ptr_A,
ptr_B,
ptr_SF_A,
Expand Down Expand Up @@ -421,7 +422,7 @@ struct GemvBlockScaled<ElementA_,
auto k_block_next = (k_block + Int<1>{}) % 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],
Expand All @@ -436,9 +437,6 @@ struct GemvBlockScaled<ElementA_,
// Copy gmem to smem before computing gemm on each k-pipe
if (k_block == 0)
{
// Use predicate instead of branch for cp_async
bool valid_tile = (global_k < gemm_k);

// Load all stages using the helper function
load_stages_gmem_to_smem(
smem_pipe_write, // buffer_idx
Expand All @@ -451,7 +449,7 @@ struct GemvBlockScaled<ElementA_,
smem_sf_write_offset,
is_even_thread,
load_b,
valid_tile,
mainloop_k,
ptr_A,
ptr_B,
ptr_SF_A,
Expand All @@ -466,6 +464,7 @@ struct GemvBlockScaled<ElementA_,
smem_pipe_read = (smem_pipe_read == kBufferCount) ? 0 : smem_pipe_read;
}

// Invalid stages contain zero A/B and scale factors.
{
int frag_idx = k_block & 1;

Expand All @@ -484,9 +483,10 @@ struct GemvBlockScaled<ElementA_,
cutlass::arch::cp_async_wait<0>();
__syncthreads();

// Tail elements that don't fill a full tile
if (unroll_col_k + idx_col_k * kPackedElementsA < gemm_k) {
accum += process_tail_elements(unroll_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,
Expand Down Expand Up @@ -522,7 +522,7 @@ struct GemvBlockScaled<ElementA_,
int smem_sf_write_offset,
bool is_even_thread,
bool load_b,
bool valid_tile,
int mainloop_k,
ElementA const* ptr_A,
ElementB const* ptr_B,
ElementSFA const* ptr_SF_A,
Expand All @@ -531,6 +531,9 @@ struct GemvBlockScaled<ElementA_,

CUTLASS_PRAGMA_UNROLL
for (int s = 0; s < num_stages; ++s) {
// Recompute for each stage, both in the prologue and in lookahead
// buffers. Partial tiles are handled separately by the scalar tail.
bool valid_tile = unroll_col_k < mainloop_k;
// Load scaling factors using cp.async - only even threads participate
// Calculate SF indices for this thread
int SF_idx = global_k / kSFVecSize;
Expand All @@ -539,20 +542,26 @@ struct GemvBlockScaled<ElementA_,
void *smem_ptr_SFA = &shared_storage.smem_SFA[buffer_idx][s][smem_sf_write_offset];
const void *gmem_ptr_SFA = ptr_SF_A + SF_offset_by_k;
// Load 4 FP8 values (32 bits) - for this thread and next thread
cutlass::arch::cp_async<sizeof(uint32_t)>(smem_ptr_SFA, gmem_ptr_SFA, valid_tile && is_even_thread);
if (is_even_thread) {
cutlass::arch::cp_async_zfill<sizeof(uint32_t)>(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<sizeof(uint32_t)>(smem_ptr_SFB, gmem_ptr_SFB, valid_tile && load_b && is_even_thread);
if (load_b && is_even_thread) {
cutlass::arch::cp_async_zfill<sizeof(uint32_t)>(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<sizeof(FragmentA)>(smem_ptr_A, gmem_ptr_A, valid_tile);
cutlass::arch::cp_async_zfill<sizeof(FragmentA)>(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<sizeof(FragmentB)>(smem_ptr_B, gmem_ptr_B, valid_tile && load_b);
if (load_b) {
cutlass::arch::cp_async_zfill<sizeof(FragmentB)>(smem_ptr_B, gmem_ptr_B, valid_tile);
}

unroll_col_k += tileA_k_local;
global_k += tileA_k_local;
Expand Down