Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
84 commits
Select commit Hold shift + click to select a range
96225e8
Start lowering serial grid reduction
jacobhinkle Dec 6, 2023
f9d2d01
Add grid_serialization.{cpp,h}
jacobhinkle Dec 8, 2023
6454b06
Add requestSerialGridReduction
jacobhinkle Dec 8, 2023
1142e44
Disable previous changes to indexing pass.
jacobhinkle Dec 8, 2023
1368ba8
Call insertGridSerializationSyncs pass
jacobhinkle Dec 8, 2023
b97db45
Remove file added by mistake
jacobhinkle Dec 8, 2023
ac2da9d
Fix formatting lintrunner messed up
jacobhinkle Dec 8, 2023
ebef797
Bump num_reduction_op_attr
jacobhinkle Dec 8, 2023
dcd8606
Add test
jacobhinkle Dec 8, 2023
e196fa7
Fix sync insertion in lowering pass.
jacobhinkle Dec 11, 2023
8be320a
Fix missing allocation of sync flag buffer
jacobhinkle Dec 11, 2023
41b125f
Allocate global work buffer. Index is zero for now
jacobhinkle Dec 12, 2023
e465d96
Merge remote-tracking branch 'origin/main' into lower_serial_reduction
jacobhinkle Dec 12, 2023
34d623c
Taking a stab at replay/indexing of intermediate
jacobhinkle Dec 13, 2023
910ff09
Use fullSelfReplay and getGlobalConsumerStridedIndices
jacobhinkle Dec 13, 2023
507cf47
Infer shape using allocation domain instead of root
jacobhinkle Dec 13, 2023
e44ef7e
Update comments
jacobhinkle Dec 13, 2023
8a3134e
Hoist index scalar.
jacobhinkle Dec 13, 2023
327573c
Merge remote-tracking branch 'origin/main' into lower_serial_reduction
jacobhinkle Dec 13, 2023
d27a675
Clean up comments.
jacobhinkle Dec 13, 2023
46c70b6
Update NVFuserTest.Pipeline_CUDA
jacobhinkle Dec 13, 2023
6d2089d
Merge branch 'main' into lower_serial_reduction
jacobhinkle Dec 13, 2023
751f326
Merge remote-tracking branch 'origin/main' into lower_serial_reduction
jacobhinkle Dec 20, 2023
f4ad5ff
Merge branch 'main' into lower_serial_reduction
jacobhinkle Jan 4, 2024
1149b43
Clean up sum val computation
jacobhinkle Jan 4, 2024
94e55b8
Merge remote-tracking branch 'origin/main' into lower_serial_reduction
jacobhinkle Jan 8, 2024
f2d7461
Clean up comments and reset sync pattern properly
jacobhinkle Jan 10, 2024
a864184
Fix compile error
jacobhinkle Jan 11, 2024
c236810
Merge branch 'main' into lower_serial_reduction
jacobhinkle Jan 12, 2024
819a6e0
Re-use TensorDomain instead of replaying
jacobhinkle Jan 16, 2024
cf42527
Merge remote-tracking branch 'origin/main' into lower_serial_reduction
jacobhinkle Jan 18, 2024
449a296
Copy domains to create new TensorDomain instead of reusing
jacobhinkle Jan 18, 2024
27e449a
Allocate work buffer like leaf of output
jacobhinkle Jan 19, 2024
c6ddebf
Merge remote-tracking branch 'origin/main' into lower_serial_reduction
jacobhinkle Jan 19, 2024
2c3278b
Use serial grid reduction in split-K
jacobhinkle Dec 13, 2023
4e40d68
Set proper dtype for init in MmaOp
jacobhinkle Dec 13, 2023
7bfa709
Restore split-k benchmarks
jacobhinkle Dec 20, 2023
fc07a9a
Fix after rebase
jacobhinkle Jan 19, 2024
49b7feb
Set proper dtype for init in MmaOp
jacobhinkle Dec 13, 2023
fe8fbf5
Naive first step. Doesn't work yet
jacobhinkle Dec 15, 2023
ef4b3f3
Fix typo
jacobhinkle Jan 2, 2024
d5febec
Add debug output
jacobhinkle Jan 2, 2024
d358071
TMP add more debug output, comments
jacobhinkle Jan 4, 2024
362bde9
Fix up to get kernel running, but still have bank conflicts
jacobhinkle Jan 5, 2024
4c71cee
Delint test, add printout of bank conflict
jacobhinkle Jan 5, 2024
96fa361
Undo special handling for bothWays
jacobhinkle Jan 5, 2024
5ee9421
Vectorize smem store
jacobhinkle Jan 8, 2024
21ece7c
Remove debug print
jacobhinkle Jan 8, 2024
8bc4d0e
Use smem in all splitk tests (with/wo batch/bias)
jacobhinkle Jan 8, 2024
49edbfe
[WIP] Add benchmarks for nanogpt bwd split-K sizes
jacobhinkle Jan 9, 2024
bc50c8d
Vectorized serial grid reduction
jacobhinkle Dec 14, 2023
f740f12
Remove debug prints
jacobhinkle Dec 14, 2023
71bae3a
Use loadGlobalToLocal instead of loadGenericVolatile
jacobhinkle Dec 18, 2023
747b466
Use smem_epilogue when possible in benchmarks
jacobhinkle Jan 9, 2024
0b13a59
Vectorize innermost non-trivial dimension.
jacobhinkle Jan 19, 2024
21203ea
Clean up check for use_smem_epilogue in benchmark
jacobhinkle Jan 19, 2024
aef008d
Set manual seed before generating inputs in SingleMatmulBase
jacobhinkle Jan 19, 2024
e36740d
Disable double buffering smem writes
jacobhinkle Jan 19, 2024
90b0eae
Merge remote-tracking branch 'origin/main' into splitk_smem_epilogue
jacobhinkle Feb 7, 2024
907e877
Add smem_epilogue:{0,1} cases in benchmarks
jacobhinkle Feb 7, 2024
5976c43
Merge remote-tracking branch 'origin/main' into splitk_smem_epilogue
jacobhinkle Feb 12, 2024
b04a23f
Split inner axis if vec width is too large
jacobhinkle Feb 12, 2024
3fcdfed
Delint
jacobhinkle Feb 12, 2024
687c13a
Delint benchmark
jacobhinkle Feb 12, 2024
7f23082
Guard against limited smem in test
jacobhinkle Feb 12, 2024
b6b5287
Fix missing stage number in test smem check
jacobhinkle Feb 12, 2024
c0dcd76
Merge remote-tracking branch 'origin/main' into splitk_smem_epilogue
jacobhinkle Feb 26, 2024
e848abe
Merge remote-tracking branch 'origin/main' into splitk_smem_epilogue
jacobhinkle Mar 10, 2024
af710d6
Fix batch split-K tests.
jacobhinkle Mar 11, 2024
1c11d1f
Merge remote-tracking branch 'origin/main' into splitk_smem_epilogue
jacobhinkle Mar 11, 2024
629bcb1
Merge remote-tracking branch 'origin/main' into splitk_smem_epilogue
jacobhinkle Mar 12, 2024
27db63e
Revert inadvertent change to benchmark
jacobhinkle Mar 18, 2024
3ebc733
Comment on bool options
jacobhinkle Mar 18, 2024
c7cfcaf
Move definition of smem_epilogue further down
jacobhinkle Mar 19, 2024
5672f33
Merge remote-tracking branch 'origin/main' into splitk_smem_epilogue
jacobhinkle Mar 19, 2024
e6377f4
Refactor splitk_sum scheduling into separate function
jacobhinkle Mar 19, 2024
c7f8bac
Clean up comments
jacobhinkle Mar 19, 2024
b20cc73
Simplify scheduleSplitKSum and fix comment formatting
jacobhinkle Mar 20, 2024
26dc559
Fix comments
jacobhinkle Mar 20, 2024
83ec210
Fix comment
jacobhinkle Mar 20, 2024
65da699
Merge remote-tracking branch 'origin/main' into splitk_smem_epilogue
jacobhinkle Mar 20, 2024
9929255
Fix comment formatting
jacobhinkle Mar 21, 2024
9fea8d6
Merge branch 'main' into splitk_smem_epilogue
jacobhinkle Mar 21, 2024
590b9a8
Fix another clang-format wrapped comment
jacobhinkle Mar 21, 2024
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
50 changes: 41 additions & 9 deletions benchmarks/cpp/matmul.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
#include <scheduler/all_schedulers.h>
#include <scheduler/matmul.h>
#include <scheduler/matmul_heuristic.h>
#include <scheduler/mma_utils.h>
#include <utils.h>

#include <benchmark/benchmark.h>
Expand Down Expand Up @@ -145,6 +146,9 @@ static void SingleMatmulBase(
int64_t n = benchmark_state.range(1);
int64_t k = benchmark_state.range(2);

// inputs
at::manual_seed(0);

// Tensor inputs
auto inputs = matmulAtInput2D(m, n, k, layout);
auto expected_output = atMatmul(
Expand All @@ -159,9 +163,6 @@ static void SingleMatmulBase(

preseg_passes::OptimizationPass<preseg_passes::PreSegmenter>::runPass(fusion);

// inputs
at::manual_seed(0);

KernelArgumentHolder args = KernelArgumentHolder::createKernelArgumentHolder(
{inputs.first, inputs.second});

Expand Down Expand Up @@ -250,6 +251,14 @@ MatmulParams getMatmulParams(
params.double_buffer_options.double_buffer_smem_read = true;
params.double_buffer_options.smem_double_buffer_stage = stage_number;
params.splitk_factor = splitk_factor;
std::tie(params.use_smem_epilogue, params.promote_prologue_smem_reuse) =
mma_utils::generateSharedMemoryEpilogueHeuristics(
gemm_tile,
stage_number,
{DataType::Half, DataType::Half, DataType::Float},
/*smem_a_reuse_guaranteed=*/true,
/*smem_b_reuse_guaranteed=*/true,
/*ignore_occupancy_drop=*/true);

return params;
}
Expand Down Expand Up @@ -362,7 +371,8 @@ static void NvFuserScheduler_Matmul(
benchmark::State& benchmark_state,
MmaLayout layout,
int splitk_factor = 1,
bool partitionedk = false) {
bool partitionedk = false,
bool use_smem_epilogue = false) {
int num_warps = benchmark_state.range(3);
int number_of_stage = benchmark_state.range(4);

Expand All @@ -376,6 +386,15 @@ static void NvFuserScheduler_Matmul(

auto params = getMatmulParams(
cta_tile, number_of_stage, layout, partitionedk ? 1 : splitk_factor);
if (use_smem_epilogue) {
if (!params.use_smem_epilogue) {
benchmark_state.SkipWithError(
"Insufficient shared mem for smem epilogue");
}
} else {
params.use_smem_epilogue = false;
params.promote_prologue_smem_reuse = false;
}

NVFUSER_BENCHMARK_ARCH_SMEM_GUARD(
8, 0, getSmemSize(cta_tile, number_of_stage), benchmark_state);
Expand Down Expand Up @@ -627,13 +646,23 @@ static void MatmulShapeWarpStageAutoSplitK(benchmark::internal::Benchmark* b) {
// Use this for manual splitk.
static void MatmulShapeWarpStageSpecificSplitK(
benchmark::internal::Benchmark* b) {
b->ArgNames({"M", "N", "K", "warps", "stages", "splitk_factor"});
b->ArgNames(
{"M", "N", "K", "warps", "stages", "splitk_factor", "smem_epilogue"});
for (long int num_warps : NumWarps) {
for (long int num_stages : NumStages) {
for (auto [m, n, k] :
std::vector<std::tuple<int, int, int>>(SplitKSpecificShapes)) {
for (auto splitk_factor : {2, 3, 4, 5, 6}) {
b->Args({m, n, k, num_warps, num_stages, splitk_factor});
for (bool use_smem_epilogue : {false, true}) {
b->Args(
{m,
n,
k,
num_warps,
num_stages,
splitk_factor,
use_smem_epilogue});
}
}
}
}
Expand Down Expand Up @@ -724,8 +753,13 @@ static void NvFuserScheduler_Matmul_Manual(
benchmark::State& benchmark_state,
MmaLayout layout) {
int splitk_factor = benchmark_state.range(5);
bool use_smem_epilogue = benchmark_state.range(6);
NvFuserScheduler_Matmul(
benchmark_state, layout, splitk_factor, /*partitionedk=*/false);
benchmark_state,
layout,
splitk_factor,
/*partitionedk=*/false,
use_smem_epilogue);
}

#define SpecificSplitKBenchmark(layout) \
Expand All @@ -739,9 +773,7 @@ static void NvFuserScheduler_Matmul_Manual(

ForAllLayouts(EagerModeBenchmark);
ForAllLayouts(NvfuserMatmulBenchmark);
ForAllLayouts(AutoSplitKBenchmark);
ForAllLayouts(SpecificSplitKBenchmark);
ForAllLayouts(AutoPartitionedKBenchmark);

// Note: SplitK Reduction benchmarks are parametrized only by M, N. The splitk
// factor is deduced automatically from N
Expand Down
14 changes: 12 additions & 2 deletions csrc/device_lower/validation.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -388,6 +388,16 @@ class VectorizeValidator : public OptInDispatch {
if (r_id->isReduction() || r_id->isBroadcast()) {
continue;
}
if ((tv->getMemoryType() == MemoryType::Shared ||
tv->getMemoryType() == MemoryType::Local) &&
r_id->isBlockDim()) {
// Inner-most parallelized dimensions don't count in allocation of
// shared and local tensors.
continue;
}
Comment on lines +391 to +397

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I am vectorizing the store to smem, but without this I hit

Vectorized dim for consumer has to be from a contiguous inner most position. tv: T13_s[ iblockIdx.x56{( ceilDiv(( (( (( getMetaData(T0) )).logical_size ))[0] ), 128) )}, iblockIdx.y58{( ceilDiv(( (( (( getMetaData(T1) )).logical_size ))[1] ), 128) )}, iblockIdx.z55{2}, ithreadIdx.z353{( ceilDiv(( ( ceilDiv(128, 4) ) * 4 ), 64) )}, ithreadIdx.y355{( ceilDiv(( ( ( ceilDiv(( ceilDiv(128, 8) ), 4) ) * 4 ) * 8 ), 64) )}, iS357{( ceilDiv(64, 16) )}, iS359{( ceilDiv(64, 8) )}, ithreadIdx.x370{( ( ( ceilDiv(( ceilDiv(16, 8) ), 2) ) * 8 ) * ( ceilDiv(8, 2) ) )}, iS366{( ceilDiv(8, 8) )}, iS364{2}, iV369{2} ] ca_pos( 2 ) produce_pos( 3 ), allocation domain: iS412{( (( (( getMetaData(T0) )).logical_size ))[0] )}, iS413{( (( (( getMetaData(T1) )).logical_size ))[1] )}, iblockIdx.z55{2}, vectorized id: iS413{( (( (( getMetaData(T1) )).logical_size ))[1] )}, innermost id: iblockIdx.z55{2}, contiguity: t

For a local or shared tensor, block-parallelized dimensions don't affect the allocation, so we should just skip them here I think.

if (tv->getMemoryType() == MemoryType::Local && r_id->isThreadDim()) {
continue;
}
last_alloc_dim = r_id;
last_alloc_dim_pos = i - 1;
break;
Expand Down Expand Up @@ -416,9 +426,9 @@ class VectorizeValidator : public OptInDispatch {
", allocation domain: ",
ir_utils::toString(tv->getMaybeAllocationDomain()),
", vectorized id: ",
validator.vectorized_id_,
validator.vectorized_id_->toString(),
", innermost id: ",
last_alloc_dim,
last_alloc_dim->toString(),
", contiguity: ",
contiguity.has_value() ? (*contiguity ? "t" : "f") : "n");
}
Expand Down
Loading