CUTLASS 2.10 (#615)

Co-authored-by: Aniket Shivam <ashivam@nvidia.com>
This commit is contained in:
ANIKET SHIVAM
2022-09-03 18:48:46 -04:00
committed by GitHub
co-authored by Aniket Shivam
parent ca23ff7924
commit b72cbf957d
289 changed files with 43708 additions and 2513 deletions
@@ -43,6 +43,7 @@
#include "cutlass/array.h"
#include "cutlass/layout/matrix.h"
#include "cutlass/layout/tensor.h"
#include "cutlass/layout/permute.h"
#include "cutlass/matrix_shape.h"
#include "cutlass/tensor_ref.h"
#include "cutlass/transform/pitch_linear_thread_map.h"
@@ -70,6 +71,7 @@ template <
typename ThreadMap_, ///< Thread map (conept: OutputTileThreadMap)
typename Element_, ///< Element data type
bool ScatterD = false, ///< Scatter D operand or not
typename PermuteDLayout = layout::NoPermute, ///< Permute D operand or not
bool UseCUDAStore = false
>
class PredicatedTileIterator {
@@ -173,9 +175,12 @@ private:
/// Parameters structure containing reference and precomputed state.
PredicatedTileIteratorParams params_;
/// Byte-level pointer
/// Byte-level pointer. This pointer is usually for both load() and store(), unless PermuteD is performed. When having PermuteD, byte_pointer_ is only for load().
uint8_t *byte_pointer_;
/// Byte-level pointer for store(). Due to PermuteD Op, store_byte_pointer_ may be with different address computation compared to byte_pointer_.
uint8_t *store_byte_pointer_;
/// Array of boolean values to contain steady-state predicates
Mask mask_;
@@ -196,6 +201,11 @@ private:
/// Scatter indices
int const *indices_;
/// Whether to perform Permute Op
bool PermuteD;
/// PermuteDLayout
mutable PermuteDLayout permute_layout_;
//
// Static asserts about internal strides
@@ -255,7 +265,7 @@ public:
mask_.clear();
}
// Initialize pointer
// Initialize byte_pointer_
byte_pointer_ = reinterpret_cast<uint8_t *>(pointer) +
LongIndex(thread_offset.row()) * LongIndex(params_.stride) +
LongIndex(thread_offset.column()) * sizeof(AccessType) / kElementsPerAccess;
@@ -265,6 +275,19 @@ public:
LongIndex(thread_offset.column()) * sizeof(AccessType) / kElementsPerAccess;
}
// store_byte_pointer_ is set to be the same with byte_pointer_ unless PermuteD is used.
store_byte_pointer_ = byte_pointer_;
// Initialize PermuteD. If PermuteD is true, store_byte_pointer_ is initialized accordingly.
if (platform::is_same<PermuteDLayout, layout::NoPermute>::value) {
PermuteD = false;
}else{
PermuteD = true;
store_byte_pointer_ = reinterpret_cast<uint8_t *>(pointer);
permute_layout_ = PermuteDLayout(extent,
params_.stride * kElementsPerAccess / sizeof(AccessType));
}
// Initialize internal state counter
state_[0] = state_[1] = state_[2] = 0;
}
@@ -272,6 +295,7 @@ public:
/// Adds a pointer offset in units of Element
CUTLASS_HOST_DEVICE
void add_pointer_offset(LongIndex pointer_offset) {
store_byte_pointer_ += pointer_offset * sizeof_bits<Element>::value / 8;
byte_pointer_ += pointer_offset * sizeof_bits<Element>::value / 8;
}
@@ -353,7 +377,7 @@ public:
/// Stores a fragment to memory
CUTLASS_DEVICE
void store_with_byte_offset(Fragment const &frag, int64_t byte_offset) const {
uint8_t *byte_pointer = byte_pointer_;
uint8_t *byte_pointer = store_byte_pointer_;
AccessType const *frag_ptr = reinterpret_cast<AccessType const *>(&frag);
CUTLASS_PRAGMA_UNROLL
@@ -388,21 +412,38 @@ public:
bool guard = row_guard && mask_.predicates[column];
int col_offset = column * ThreadMap::Delta::kColumn;
if (PermuteD) {
int col = col_offset + thread_start_column_;
int row = row_offset + thread_start_row_;
TensorCoord init_coord(row, col);
// Locate memory_pointer
memory_pointer = reinterpret_cast<AccessType *>(byte_pointer + byte_offset
+ permute_layout_(init_coord) * sizeof(AccessType) / kElementsPerAccess);
}
if (UseCUDAStore) {
if (guard) {
memory_pointer[column * ThreadMap::Delta::kColumn / kElementsPerAccess] =
memory_pointer[0] =
frag_ptr[frag_row_idx * ThreadMap::Iterations::kColumn + column];
}
} else {
cutlass::arch::global_store<AccessType, sizeof(AccessType)>(
frag_ptr[frag_row_idx * ThreadMap::Iterations::kColumn + column],
(void *)&memory_pointer[column * ThreadMap::Delta::kColumn / kElementsPerAccess],
(void *)&memory_pointer[0],
guard);
}
if (!PermuteD) {
memory_pointer += (ThreadMap::Delta::kColumn / kElementsPerAccess);
}
}
if (row + 1 < ThreadMap::Iterations::kRow) {
if (!ScatterD) {
if (!ScatterD && !PermuteD) {
byte_pointer += params_.increment_row;
}
}
@@ -605,6 +646,10 @@ public:
++state_[0];
if (!ScatterD && !PermuteD) {
store_byte_pointer_ += params_.advance_row;
}
if (!ScatterD) {
byte_pointer_ += params_.advance_row;
}
@@ -616,6 +661,7 @@ public:
state_[0] = 0;
++state_[1];
byte_pointer_ += params_.advance_group;
store_byte_pointer_ += params_.advance_group;
thread_start_row_ += (ThreadMap::Shape::kGroup - 1) *
ThreadMap::Shape::kRow * ThreadMap::Count::kRow;
@@ -625,6 +671,7 @@ public:
state_[1] = 0;
++state_[2];
byte_pointer_ += params_.advance_cluster;
store_byte_pointer_ += params_.advance_group;
thread_start_row_ += ThreadMap::Count::kGroup *
ThreadMap::Shape::kGroup * ThreadMap::Count::kRow * ThreadMap::Shape::kRow;
@@ -632,6 +679,7 @@ public:
if (state_[2] == ThreadMap::Count::kCluster) {
state_[2] = 0;
byte_pointer_ += params_.advance_tile;
store_byte_pointer_ += params_.advance_group;
}
}
}