CUTLASS 2.4 (Implicit GEMM convolution) (#147)
CUTLASS 2.4 (Implicit GEMM Convolution) Co-authored-by: Manish Gupta <manigupta@nvidia.com>, Haicheng Wu <haichengw@nvidia.com>, Dustyn Blasig <dblasig@nvidia.com>, Andrew Kerr <akerr@nvidia.com>
This commit is contained in:
co-authored by
Manish Gupta <manigupta@nvidia.com>, Haicheng Wu <haichengw@nvidia.com>, Dustyn Blasig <dblasig@nvidia.com>, Andrew Kerr <akerr@nvidia.com>
parent
c2b80ad4e4
commit
6615010cd0
@@ -43,6 +43,7 @@
|
||||
#include "cutlass/epilogue/threadblock/output_tile_thread_map.h"
|
||||
#include "cutlass/arch/arch.h"
|
||||
#include "cutlass/arch/memory.h"
|
||||
#include "cutlass/epilogue/threadblock/predicated_tile_iterator_params.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -102,68 +103,20 @@ public:
|
||||
// Parameters struct
|
||||
//
|
||||
|
||||
struct Params {
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
LongIndex stride; ///< stride in bytes between rows
|
||||
|
||||
LongIndex increment_row; ///< increment quantity (in bytes) to advance when moving between rows
|
||||
LongIndex increment_group; ///< increment quantity (in bytes) to advance when moving to the next group
|
||||
LongIndex increment_cluster; ///< increment quantity (in bytes) to advance when moving to the next cluster
|
||||
|
||||
LongIndex advance_row; ///< amount to add to move to the next 'row' position
|
||||
LongIndex advance_group; ///< amount to add to move to the next 'group' position
|
||||
LongIndex advance_cluster; ///< amount to add to move to the next 'cluster' position
|
||||
LongIndex advance_tile; ///< amount to add to move to the next 'tile'
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
/// Uses a non-template class
|
||||
struct Params : PredicatedTileIteratorParams {
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Status initialize(Index stride_) {
|
||||
|
||||
stride = LongIndex(stride_);
|
||||
|
||||
increment_row = stride * ThreadMap::Delta::kRow;
|
||||
|
||||
increment_group = stride * ThreadMap::Delta::kGroup
|
||||
- stride * ThreadMap::Delta::kRow * (ThreadMap::Iterations::kRow - 1);
|
||||
|
||||
increment_cluster = stride * ThreadMap::Delta::kCluster
|
||||
- stride * ThreadMap::Delta::kGroup * (ThreadMap::Iterations::kGroup - 1)
|
||||
- stride * ThreadMap::Delta::kRow * (ThreadMap::Iterations::kRow - 1);
|
||||
|
||||
advance_row = stride * ThreadMap::Shape::kRow;
|
||||
|
||||
advance_group = stride * (ThreadMap::Shape::kGroup - 1) * ThreadMap::Shape::kRow * ThreadMap::Count::kRow;
|
||||
|
||||
advance_cluster =
|
||||
stride *
|
||||
ThreadMap::Count::kGroup * ThreadMap::Shape::kGroup * ThreadMap::Count::kRow * ThreadMap::Shape::kRow;;
|
||||
|
||||
advance_tile =
|
||||
stride *
|
||||
ThreadMap::Shape::kGroup *
|
||||
ThreadMap::Shape::kRow *
|
||||
ThreadMap::Shape::kCluster *
|
||||
ThreadMap::Shape::kTile;
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
Params() { }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params() {
|
||||
initialize(0);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Layout const &layout) {
|
||||
|
||||
initialize(layout.stride(0) * int(sizeof(AccessType)) / kElementsPerAccess);
|
||||
Params(Layout const &layout):
|
||||
PredicatedTileIteratorParams(
|
||||
layout.stride(0) * int(sizeof(AccessType)) / kElementsPerAccess,
|
||||
make_OutputTileThreadMapDesc<ThreadMap>()
|
||||
)
|
||||
{
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
@@ -207,7 +160,7 @@ private:
|
||||
//
|
||||
|
||||
/// Parameters structure containing reference and precomputed state.
|
||||
Params params_;
|
||||
PredicatedTileIteratorParams params_;
|
||||
|
||||
/// Byte-level pointer
|
||||
uint8_t *byte_pointer_;
|
||||
@@ -239,12 +192,13 @@ public:
|
||||
/// Constructor
|
||||
CUTLASS_DEVICE
|
||||
PredicatedTileIterator(
|
||||
Params const & params,
|
||||
PredicatedTileIteratorParams const & params,
|
||||
Element *pointer,
|
||||
TensorCoord extent,
|
||||
int thread_idx,
|
||||
TensorCoord threadblock_offset = TensorCoord()
|
||||
): params_(params)
|
||||
):
|
||||
params_(params)
|
||||
{
|
||||
|
||||
TensorCoord thread_offset = ThreadMap::initial_offset(thread_idx) + threadblock_offset;
|
||||
@@ -745,6 +699,309 @@ public:
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator used to load output tile from shared memory in epilogue.
|
||||
///
|
||||
/// Satisfies: ReadableTileIterator | InterleavedMaskedTileIterator | ForwardTileIterator
|
||||
///
|
||||
template <
|
||||
typename ThreadMap_, ///< Thread map (conept: OutputTileThreadMap)
|
||||
typename Element_, ///< Element data type
|
||||
int InterleavedN ///< Number of Interleaved N
|
||||
>
|
||||
class InterleavedConvPredicatedTileIterator {
|
||||
public:
|
||||
using ThreadMap = ThreadMap_;
|
||||
|
||||
using Element = Element_;
|
||||
|
||||
using Layout = layout::TensorNCxHWx<InterleavedN>;
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using ConstTensorRef = typename TensorRef::ConstTensorRef;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
using TensorCoord = Tensor4DCoord;
|
||||
|
||||
static int const kElementsPerAccess = ThreadMap::kElementsPerAccess;
|
||||
static int const kThreads = ThreadMap::kThreads;
|
||||
static int const kIterations = ThreadMap::Iterations::kCount;
|
||||
|
||||
/// Fragment object
|
||||
using Fragment = Array<Element, ThreadMap::kElementsPerAccess>;
|
||||
|
||||
/// Memory access size
|
||||
using AccessType = AlignedArray<Element, ThreadMap::kElementsPerAccess>;
|
||||
|
||||
//
|
||||
// Parameters struct
|
||||
//
|
||||
|
||||
struct Params {
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
LongIndex stride_col; ///< stride in bytes between columns
|
||||
LongIndex stride_row; ///< stride in bytes between rows
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Status initialize(typename Layout::Stride stride_) {
|
||||
stride_col = stride_[1];
|
||||
stride_row = stride_[2];
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params() {
|
||||
initialize(cutlass::make_Coord(0, 0, 0));
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Layout const &layout) {
|
||||
|
||||
initialize(layout.stride());
|
||||
}
|
||||
};
|
||||
|
||||
/// Mask object
|
||||
struct Mask {
|
||||
static int const kCount =
|
||||
(ThreadMap::Iterations::kRow < 8) ? 8 : ThreadMap::Iterations::kRow;
|
||||
|
||||
/// Predicate state
|
||||
bool predicates[kCount];
|
||||
|
||||
//
|
||||
// Mask
|
||||
//
|
||||
CUTLASS_HOST_DEVICE
|
||||
Mask() {
|
||||
enable();
|
||||
}
|
||||
|
||||
///< Efficiently disables all accesses guarded by mask
|
||||
CUTLASS_HOST_DEVICE void clear() {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kCount; ++i) {
|
||||
predicates[i] = false;
|
||||
}
|
||||
}
|
||||
|
||||
///< CUTLASS_HOST_DEVICE enables all accesses guarded by mask
|
||||
CUTLASS_DEVICE void enable() {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kCount; ++i) {
|
||||
predicates[i] = true;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Parameters structure containing reference and precomputed state.
|
||||
Params params_;
|
||||
|
||||
/// Byte-level pointer
|
||||
uint8_t *byte_pointer_;
|
||||
|
||||
/// Array of boolean values to contain steady-state predicates
|
||||
Mask mask_;
|
||||
|
||||
/// Extent of the matrix tile in columns
|
||||
Index extent_col_;
|
||||
|
||||
/// Extent of the matrix tile in rows
|
||||
Index extent_row_;
|
||||
|
||||
/// Extent of the matrix tile in pq
|
||||
Index extent_pq_;
|
||||
|
||||
/// A thread's starting row position (assuming steady-state predicates have
|
||||
/// been computed)
|
||||
Index thread_start_row_;
|
||||
|
||||
/// A thread's starting column position (assuming steady-state predicates have
|
||||
/// been computed)
|
||||
Index thread_start_col_;
|
||||
|
||||
/// Internal iteration counter
|
||||
LongIndex iteration_row_;
|
||||
LongIndex iteration_col_;
|
||||
|
||||
uint32_t pq_mul_;
|
||||
|
||||
uint32_t pq_shr_;
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Constructor
|
||||
CUTLASS_DEVICE
|
||||
InterleavedConvPredicatedTileIterator(
|
||||
Params const & params,
|
||||
Element *pointer,
|
||||
TensorCoord extent,
|
||||
int thread_idx,
|
||||
MatrixCoord threadblock_offset
|
||||
):
|
||||
params_(params) {
|
||||
MatrixCoord thread_offset = ThreadMap::initial_offset(thread_idx) + threadblock_offset;
|
||||
|
||||
extent_col_ = extent.c();
|
||||
extent_pq_ = extent.h() * extent.w();
|
||||
extent_row_ = extent.n() * extent_pq_;
|
||||
|
||||
find_divisor(pq_mul_, pq_shr_, extent_pq_);
|
||||
|
||||
thread_start_row_ = thread_offset.row();
|
||||
thread_start_col_ = thread_offset.column();
|
||||
|
||||
// Initialize predicates
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int r = 0; r < ThreadMap::Iterations::kRow; ++r) {
|
||||
mask_.predicates[r] =
|
||||
((thread_offset.row() + ThreadMap::Delta::kRow * r) < extent_row_);
|
||||
}
|
||||
|
||||
// Initialize pointer
|
||||
byte_pointer_ = reinterpret_cast<uint8_t *>(pointer) +
|
||||
((thread_start_col_ / InterleavedN) * params_.stride_col +
|
||||
(thread_start_col_ % InterleavedN)) *
|
||||
sizeof_bits<Element>::value / 8;
|
||||
|
||||
// Initialize internal state counter
|
||||
iteration_row_ = iteration_col_ = 0;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
byte_pointer_ += pointer_offset * sizeof_bits<Element>::value / 8;
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
|
||||
int col_offset = iteration_col_ * ThreadMap::Delta::kColumn;
|
||||
bool col_guard = ((thread_start_col_ + col_offset) < extent_col_);
|
||||
bool guard = col_guard && mask_.predicates[iteration_row_];
|
||||
|
||||
int n, pq_rem;
|
||||
|
||||
fast_divmod(n, pq_rem,
|
||||
thread_start_row_ + iteration_row_ * ThreadMap::Delta::kRow,
|
||||
extent_pq_, pq_mul_, pq_shr_);
|
||||
|
||||
uint8_t *byte_pointer =
|
||||
byte_pointer_ + (n * params_.stride_row + pq_rem * InterleavedN) *
|
||||
sizeof_bits<Element>::value / 8;
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
AccessType const *memory_pointer =
|
||||
reinterpret_cast<AccessType const *>(byte_pointer);
|
||||
|
||||
cutlass::arch::global_load<
|
||||
AccessType,
|
||||
sizeof(AccessType)
|
||||
>(
|
||||
*frag_ptr,
|
||||
(void *)memory_pointer,
|
||||
guard);
|
||||
}
|
||||
|
||||
/// Stores a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
|
||||
int col_offset = iteration_col_ * ThreadMap::Delta::kColumn;
|
||||
bool col_guard = ((thread_start_col_ + col_offset) < extent_col_);
|
||||
bool guard = col_guard && mask_.predicates[iteration_row_];
|
||||
|
||||
int n, pq_rem;
|
||||
|
||||
fast_divmod(n, pq_rem,
|
||||
thread_start_row_ + iteration_row_ * ThreadMap::Delta::kRow,
|
||||
extent_pq_, pq_mul_, pq_shr_);
|
||||
|
||||
uint8_t *byte_pointer =
|
||||
byte_pointer_ + (n * params_.stride_row + pq_rem * InterleavedN) *
|
||||
sizeof_bits<Element>::value / 8;
|
||||
AccessType const *frag_ptr = reinterpret_cast<AccessType const *>(&frag);
|
||||
AccessType *memory_pointer = reinterpret_cast<AccessType *>(byte_pointer);
|
||||
|
||||
if (guard) {
|
||||
*memory_pointer = *frag_ptr;
|
||||
}
|
||||
}
|
||||
|
||||
/// Overrides the internal iteration index
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_iteration_index(int iteration) {
|
||||
iteration_row_ = iteration % ThreadMap::Iterations::kRow;
|
||||
iteration_col_ = iteration / ThreadMap::Iterations::kRow;
|
||||
}
|
||||
|
||||
/// Advances to the next position to load or store
|
||||
CUTLASS_HOST_DEVICE
|
||||
InterleavedConvPredicatedTileIterator &operator++() {
|
||||
|
||||
++iteration_row_;
|
||||
|
||||
if (iteration_row_ == ThreadMap::Iterations::kRow) {
|
||||
|
||||
iteration_row_ = 0;
|
||||
++iteration_col_;
|
||||
byte_pointer_ += params_.stride_col;
|
||||
|
||||
if (iteration_col_ == ThreadMap::Iterations::kColumn) {
|
||||
iteration_col_ = 0;
|
||||
}
|
||||
}
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
///< Efficiently disables all accesses guarded by mask
|
||||
CUTLASS_DEVICE void clear_mask() {
|
||||
mask_.clear();
|
||||
}
|
||||
|
||||
///< Efficiently enables all accesses guarded by mask
|
||||
CUTLASS_DEVICE void enable_mask() {
|
||||
mask_.enable();
|
||||
}
|
||||
|
||||
///< Sets the mask
|
||||
CUTLASS_DEVICE void get_mask(Mask &mask) {
|
||||
return mask_;
|
||||
}
|
||||
|
||||
///< Sets the mask
|
||||
CUTLASS_DEVICE void set_mask(Mask const &mask) {
|
||||
mask_ = mask;
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
|
||||
Reference in New Issue
Block a user