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:
@@ -500,7 +500,7 @@ class PredicatedTileAccessIterator<Shape_, Element_, layout::PitchLinear,
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization of PredicatedTileAccessIterator for pitch-linear data.
|
||||
/// Specialization of PredicatedTileAccessIterator for column-major data.
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
@@ -676,7 +676,7 @@ class PredicatedTileAccessIterator<Shape_, Element_, layout::ColumnMajor,
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization of PredicatedTileAccessIterator for pitch-linear data.
|
||||
/// Specialization of PredicatedTileAccessIterator for row-major data.
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
@@ -852,8 +852,8 @@ class PredicatedTileAccessIterator<Shape_, Element_, layout::RowMajor,
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization of PredicatedTileAccessIterator for interleaved data. It
|
||||
/// is mapped to the congruous layout.
|
||||
/// Specialization of PredicatedTileAccessIterator for column-major interleaved data.
|
||||
/// It is mapped to the congruous layout.
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
@@ -1032,8 +1032,8 @@ class PredicatedTileAccessIterator<Shape_, Element_,
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization of PredicatedTileAccessIterator for interleaved data. It
|
||||
/// is mapped to the congruous layout.
|
||||
/// Specialization of PredicatedTileAccessIterator for row-major interleaved data.
|
||||
// It is mapped to the congruous layout.
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
|
||||
@@ -805,268 +805,6 @@ class RegularTileAccessIterator<Shape_, Element_,
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for k interleaved arrangements for TensorOps
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank, typename ThreadMap_, int InterleavedK, int Alignment>
|
||||
class RegularTileAccessIterator<
|
||||
Shape_, Element_,
|
||||
layout::TensorOpMultiplicandRowMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for pitch-linear iterator may along advance along the "
|
||||
"contiguous(rank=0) or strided(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout =
|
||||
layout::TensorOpMultiplicandRowMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
static int const kAlignment = Alignment;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using ThreadMap = ThreadMap_;
|
||||
|
||||
/// Internal details made public to facilitate introspection
|
||||
struct Detail {
|
||||
/// This iterator is specialized for an access size that is 128 bits in
|
||||
/// length.
|
||||
static int const kAccessSizeInBits = 128;
|
||||
|
||||
static_assert(sizeof_bits<Element_>::value * ThreadMap::kElementsPerAccess ==
|
||||
kAccessSizeInBits,
|
||||
"This iterator requires a policy whose access size is 128bs");
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
/// Element type per access
|
||||
using AccessType = Array<Element, Layout::kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Internal pointer to first access of tile
|
||||
AccessType *pointer_;
|
||||
|
||||
/// Internal byte offset
|
||||
Index byte_offset_;
|
||||
|
||||
/// Iteration in the contiguous dimension
|
||||
int iteration_contiguous_;
|
||||
|
||||
/// Iteration in the strided dimension
|
||||
int iteration_strided_;
|
||||
|
||||
public:
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator(TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: byte_offset_(0) {
|
||||
layout::PitchLinearCoord thread_offset_base =
|
||||
ThreadMap::initial_offset(thread_id);
|
||||
|
||||
// initialize pointer
|
||||
pointer_ = reinterpret_cast<AccessType *>(
|
||||
ref.data() + ref.offset(thread_offset_base));
|
||||
|
||||
set_iteration_index(0);
|
||||
}
|
||||
|
||||
/// Overrides the internal iteration index
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_iteration_index(int index) {
|
||||
iteration_contiguous_ = index % ThreadMap::Iterations::kContiguous;
|
||||
iteration_strided_ = index / ThreadMap::Iterations::kContiguous;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
byte_offset_ += pointer_offset * sizeof(Element);
|
||||
}
|
||||
|
||||
/// Returns a pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
AccessType *get() const {
|
||||
AccessType *access_ptr = pointer_;
|
||||
|
||||
int access_offset =
|
||||
(iteration_strided_ * ThreadMap::Delta::kStrided * Layout::kInterleavedK +
|
||||
iteration_contiguous_ * ThreadMap::Delta::kContiguous) / ThreadMap::kElementsPerAccess;
|
||||
|
||||
char *access_byte_ptr =
|
||||
reinterpret_cast<char *>(access_ptr + access_offset);
|
||||
|
||||
return reinterpret_cast<AccessType *>(access_byte_ptr + byte_offset_);
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator &operator++() {
|
||||
++iteration_contiguous_;
|
||||
|
||||
if (iteration_contiguous_ < ThreadMap::Iterations::kContiguous)
|
||||
return *this;
|
||||
|
||||
// Enter here only if (iteration_contiguous_ ==
|
||||
// ThreadMap::Iteration::kContiguous)
|
||||
iteration_contiguous_ = 0;
|
||||
++iteration_strided_;
|
||||
|
||||
if (iteration_strided_ < ThreadMap::Iterations::kStrided) {
|
||||
return *this;
|
||||
}
|
||||
|
||||
// Enter here only if (iteration_strided_ == ThreadMap::Iteration::kStrided)
|
||||
// which means we enter the next tile.
|
||||
iteration_strided_ = 0;
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator operator++(int) {
|
||||
RegularTileAccessIterator prev(*this);
|
||||
this->operator++();
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
add_pointer_offset(coord.contiguous() * Shape::kCount);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for k interleaved arrangements for TensorOps
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
|
||||
template <typename Shape_, typename Element_, int AdvanceRank, typename ThreadMap_, int InterleavedK, int Alignment>
|
||||
class RegularTileAccessIterator<
|
||||
Shape_, Element_,
|
||||
layout::TensorOpMultiplicandColumnMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for pitch-linear iterator may along advance along the "
|
||||
"contiguous(rank=0) or strided(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout =
|
||||
layout::TensorOpMultiplicandColumnMajorInterleaved<sizeof_bits<Element_>::value,
|
||||
InterleavedK>;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
static int const kAlignment = Alignment;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using ThreadMap = ThreadMap_;
|
||||
|
||||
/// Underlying iterator type
|
||||
using UnderlyingIterator = RegularTileAccessIterator<
|
||||
cutlass::MatrixShape<Shape::kColumn, Shape::kRow>,
|
||||
Element,
|
||||
layout::TensorOpMultiplicandRowMajorInterleaved<sizeof_bits<Element_>::value, InterleavedK>,
|
||||
(kAdvanceRank == 1 ? 0 : 1),
|
||||
ThreadMap
|
||||
>;
|
||||
|
||||
private:
|
||||
|
||||
/// Element type per access
|
||||
using AccessType = Array<Element, Layout::kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
|
||||
/// Underlying iterator
|
||||
UnderlyingIterator iterator_;
|
||||
|
||||
public:
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator(TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: iterator_({ref.data(), ref.stride()}, thread_id) {}
|
||||
|
||||
/// Overrides the internal iteration index
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_iteration_index(int index) {
|
||||
iterator_.set_iteration_index(index);
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Returns a pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
AccessType *get() const {
|
||||
return iterator_.get();
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileAccessIterator operator++(int) {
|
||||
RegularTileAccessIterator prev(*this);
|
||||
++iterator_;
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.strided(), coord.contiguous()});
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
|
||||
Reference in New Issue
Block a user