CUTLASS 3.2.1 (#1113)

* Updates for 3.2.1 release.

* Minor fix in gemm op profiler for raster order.

* Add scheduler mapping for raster order in the kernels.
This commit is contained in:
ANIKET SHIVAM
2023-09-26 17:24:26 -04:00
committed by GitHub
parent e0aaa3c3b3
commit 90d3b0fb18
428 changed files with 22252 additions and 21761 deletions
+47 -8
View File
@@ -219,7 +219,6 @@ public :
ThreadCategory role = ThreadCategory::NonParticipant;
uint32_t is_leader = 0;
uint32_t num_consumers = 0;
cute::tuple<int, int> active_warps = {0, 0};
};
// Constructor
@@ -232,7 +231,7 @@ public :
int warp_idx = canonical_warp_idx();
int lane_predicate = cute::elect_one_sync();
auto cluster_shape = ClusterShape{};
if (warp_idx == cute::get<0>(params.active_warps) && lane_predicate == 1) {
if (warp_idx == 0 && lane_predicate == 1) {
// Barrier FULL init
for (int i = 0; i < Stages; ++i) {
full_barrier_ptr_[i].init(1);
@@ -350,6 +349,11 @@ public :
return consumer_try_wait(state.index(), state.phase(), skip_wait);
}
CUTLASS_DEVICE
ConsumerToken consumer_test_wait(PipelineState state, uint32_t skip_wait = false) {
return consumer_test_wait(state.index(), state.phase(), skip_wait);
}
CUTLASS_DEVICE
void consumer_wait(PipelineState state) {
consumer_wait(state.index(), state.phase());
@@ -440,6 +444,15 @@ private :
return {static_cast<BarrierStatus>(barrier_status)};
}
CUTLASS_DEVICE
ConsumerToken consumer_test_wait(uint32_t stage, uint32_t phase, uint32_t skip_wait) {
if (skip_wait) {
return {BarrierStatus::WaitDone};
}
uint32_t barrier_status = full_barrier_ptr_[stage].test_wait(phase);
return {static_cast<BarrierStatus>(barrier_status)};
}
// Wait for producer to commit transactions (done by TMA)
CUTLASS_DEVICE
void consumer_wait(uint32_t stage, uint32_t phase) {
@@ -621,7 +634,6 @@ public :
uint32_t producer_arv_count = 1;
uint32_t consumer_arv_count = 1;
uint32_t dst_blockid = cute::block_rank_in_cluster();
cute::tuple<int, int> active_warps = {0, 0};
};
// Constructor
@@ -636,7 +648,7 @@ public :
// Barrier FULL, EMPTY init
// Init is done only by thread 0 of the block
if (warp_idx == cute::get<0>(params.active_warps) && lane_predicate == 1) {
if (warp_idx == 0 && lane_predicate == 1) {
for (int i = 0; i < Stages; ++i) {
full_barrier_ptr_[i].init(params.producer_arv_count);
empty_barrier_ptr_[i].init(params.consumer_arv_count);
@@ -708,6 +720,11 @@ public :
return consumer_try_wait(state.index(), state.phase(), skip_wait);
}
CUTLASS_DEVICE
ConsumerToken consumer_test_wait(PipelineState state, uint32_t skip_wait = false) {
return consumer_test_wait(state.index(), state.phase(), skip_wait);
}
CUTLASS_DEVICE
void consumer_wait(PipelineState state, ConsumerToken barrier_token = {BarrierStatus::WaitAgain}) {
consumer_wait(state.index(), state.phase(), barrier_token);
@@ -764,6 +781,15 @@ private:
return {static_cast<BarrierStatus>(barrier_status)};
}
CUTLASS_DEVICE
ConsumerToken consumer_test_wait(uint32_t stage, uint32_t phase, uint32_t skip_wait) {
if (skip_wait) {
return {BarrierStatus::WaitDone};
}
uint32_t barrier_status = full_barrier_ptr_[stage].test_wait(phase);
return {static_cast<BarrierStatus>(barrier_status)};
}
CUTLASS_DEVICE
void consumer_wait(uint32_t stage, uint32_t phase, ConsumerToken barrier_token) {
if (barrier_token == BarrierStatus::WaitAgain) {
@@ -809,7 +835,6 @@ public :
uint32_t producer_arv_count = 1;
uint32_t consumer_arv_count = 1;
uint32_t dst_blockid = cute::block_rank_in_cluster();
cute::tuple<int, int> active_warps = {0, 0};
};
// Default assumption when only storage is passed is :
@@ -831,7 +856,7 @@ public :
// Barrier FULL, EMPTY init
// Init is done only by thread 0 of the block
if (warp_idx == cute::get<0>(params.active_warps) && lane_predicate == 1) {
if (warp_idx == 0 && lane_predicate == 1) {
for (int i = 0; i < Stages; ++i) {
full_barrier_ptr_[i].init(params.producer_arv_count);
empty_barrier_ptr_[i].init(params.consumer_arv_count);
@@ -897,6 +922,11 @@ public :
return consumer_try_wait(state.index(), state.phase(), skip_wait);
}
CUTLASS_DEVICE
ConsumerToken consumer_test_wait(PipelineState state, uint32_t skip_wait = false) {
return consumer_test_wait(state.index(), state.phase(), skip_wait);
}
CUTLASS_DEVICE
void consumer_wait(PipelineState state, ConsumerToken barrier_token = {BarrierStatus::WaitAgain}) {
consumer_wait(state.index(), state.phase(), barrier_token);
@@ -947,6 +977,15 @@ private:
return {static_cast<BarrierStatus>(barrier_status)};
}
CUTLASS_DEVICE
ConsumerToken consumer_test_wait(uint32_t stage, uint32_t phase, uint32_t skip_wait) {
if (skip_wait) {
return {BarrierStatus::WaitDone};
}
uint32_t barrier_status = full_barrier_ptr_[stage].test_wait(phase);
return {static_cast<BarrierStatus>(barrier_status)};
}
CUTLASS_DEVICE
void consumer_wait(uint32_t stage, uint32_t phase) {
uint32_t done = full_barrier_ptr_[stage].test_wait(phase);
@@ -968,6 +1007,7 @@ private:
}
};
///////////////////////////////////////////////////////////////////////////////////////////////////
//
// Barrier to ensure an Ordered Sequence between
@@ -989,7 +1029,6 @@ public :
struct Params {
uint32_t group_id;
uint32_t group_size;
cute::tuple<int, int> active_warps = {0, 0};
};
private :
@@ -1020,7 +1059,7 @@ public:
// Barrier FULL, EMPTY init
// Init is done only by the one elected thread of the block
if (warp_idx == cute::get<0>(params.active_warps) && lane_predicate == 1) {
if (warp_idx == 0 && lane_predicate == 1) {
for (int d = 0; d < Depth; ++d) {
for (int l = 0; l < Length; ++l) {
barrier_ptr_[d * Length + l].init(params.group_size);