CUTLASS 2.0 (#62)
CUTLASS 2.0 Substantially refactored for - Better performance, particularly for native Turing Tensor Cores - Robust and durable templates spanning the design space - Encapsulated functionality embodying modern C++11 programming techniques - Optimized containers and data types for efficient, generic, portable device code Updates to: - Quick start guide - Documentation - Utilities - CUTLASS Profiler Native Turing Tensor Cores - Efficient GEMM kernels targeting Turing Tensor Cores - Mixed-precision floating point, 8-bit integer, 4-bit integer, and binarized operands Coverage of existing CUTLASS functionality: - GEMM kernels targeting CUDA and Tensor Cores in NVIDIA GPUs - Volta Tensor Cores through native mma.sync and through WMMA API - Optimizations such as parallel reductions, threadblock rasterization, and intra-threadblock reductions - Batched GEMM operations - Complex-valued GEMMs Note: this commit and all that follow require a host compiler supporting C++11 or greater.
This commit is contained in:
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,838 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
*modification, are permitted provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice,
|
||||
*this list of conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright
|
||||
*notice, this list of conditions and the following disclaimer in the
|
||||
*documentation and/or other materials provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its
|
||||
*contributors may be used to endorse or promote products derived from this
|
||||
*software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
*AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
*IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
*DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE FOR ANY DIRECT,
|
||||
*INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE,
|
||||
*DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY
|
||||
*OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TOR (INCLUDING
|
||||
*NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE,
|
||||
*EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates calculating the address and predicates to the load of tiles
|
||||
from pitch-linear rank=2 tensors.
|
||||
|
||||
This iterator uses masks to guard out-of-bounds accesses and visits the last
|
||||
"residue" tile first, with the objective of minimizing predicate mask updates
|
||||
during steady-state operation.
|
||||
|
||||
A precomputed "Params" object minimizes the amount of state that must be
|
||||
stored in registers, and integer addition is used to advance the pointer
|
||||
through memory.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/coord.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/layout/pitch_linear.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/predicate_vector.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/tensor_view.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace transform {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// PredicatedTileAccessIterator2dThreadTile
|
||||
///
|
||||
template <typename Shape, typename Element, typename Layout, int AdvanceRank,
|
||||
typename ThreadMap, typename AccessType>
|
||||
class PredicatedTileAccessIterator2dThreadTile;
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization of PredicatedTileAccessIterator2dThreadTile for pitch-linear data.
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, typename AccessType_>
|
||||
class PredicatedTileAccessIterator2dThreadTile<Shape_, Element_, layout::PitchLinear,
|
||||
AdvanceRank, ThreadMap_, AccessType_> {
|
||||
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::PitchLinear;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
using AccessType = AccessType_;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorView = TensorView<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using Pointer = Element *;
|
||||
using NonConstPointer = typename platform::remove_const<Element>::type *;
|
||||
|
||||
static int const kPredicatesPerByte = 4;
|
||||
static int const kPredicatesPerWord = 4 * kPredicatesPerByte;
|
||||
|
||||
/// Number of 32b words containing predicates
|
||||
static int const kPredicateByteCount = (ThreadMap::Iterations::kCount * ThreadMap::ThreadAccessShape::kStrided + kPredicatesPerByte - 1) / kPredicatesPerByte;
|
||||
static int const kPredicateWordCount = (kPredicateByteCount + 3) / 4;
|
||||
|
||||
static unsigned const kPredicateMask = (1u << kPredicatesPerByte) - 1u;
|
||||
|
||||
static_assert(kPredicateWordCount <= 4, "Too many predicates.");
|
||||
|
||||
/// Predicate vector stores mask to guard accesses
|
||||
using Mask = Array<uint32_t, kPredicateWordCount>;
|
||||
|
||||
/// Parameters object is precomputed state and is host-constructible
|
||||
class Params {
|
||||
public:
|
||||
friend PredicatedTileAccessIterator2dThreadTile;
|
||||
|
||||
private:
|
||||
/// stride of pitch-linear layout (units of Element)
|
||||
int stride_;
|
||||
/// amount (in byte) to increment pointer to move to next access along
|
||||
/// strided dimension
|
||||
int inc_strided_;
|
||||
/// amount (in byte) to increment pointer from last access to first access
|
||||
/// of next tile
|
||||
int inc_next_;
|
||||
/// amount (in byte) to increment pointer from first access of current tile
|
||||
/// to first access of next tile
|
||||
int inc_advance_;
|
||||
|
||||
public:
|
||||
|
||||
// Default ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(): stride_(0), inc_strided_(0), inc_next_(0), inc_advance_(0) { }
|
||||
|
||||
/// Construct the Params object given a pitch-linear tensor's layout
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Layout const &layout) : stride_(layout.stride(0)) {
|
||||
|
||||
inc_strided_ =
|
||||
(stride_ * ThreadMap::Delta::kStrided) * int(sizeof(Element));
|
||||
|
||||
if (kAdvanceRank) {
|
||||
// advance along strided dimension
|
||||
inc_advance_ = Shape::kStrided * stride_ * int(sizeof(Element));
|
||||
} else {
|
||||
// advance along contiguous dimension
|
||||
inc_advance_ = Shape::kContiguous * int(sizeof(Element));
|
||||
}
|
||||
|
||||
inc_next_ = inc_advance_ - (ThreadMap::Iterations::kStrided - 1) *
|
||||
ThreadMap::Delta::kStrided * stride_ *
|
||||
int(sizeof(Element));
|
||||
};
|
||||
};
|
||||
|
||||
private:
|
||||
/// Internal pointer type permits fast address arithmetic
|
||||
using BytePointer = char *;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Parameters object with precomputed internal state
|
||||
Params const ¶ms_;
|
||||
|
||||
/// Internal pointer to first access of tile
|
||||
BytePointer pointer_;
|
||||
|
||||
/// Guard predicates
|
||||
uint32_t predicates_[kPredicateWordCount];
|
||||
|
||||
/// Size of tensor
|
||||
TensorCoord extent_;
|
||||
|
||||
/// Initial offset for each thread
|
||||
TensorCoord thread_offset_;
|
||||
|
||||
/// Index of residue tile
|
||||
int residue_tile_idx_;
|
||||
|
||||
/// Used for out-of-order visitation
|
||||
bool is_residue_tile_;
|
||||
|
||||
/// Iteration in the contiguous dimension
|
||||
int iteration_contiguous_;
|
||||
|
||||
/// Iteration in the strided dimension
|
||||
int iteration_strided_;
|
||||
|
||||
/// Tracks iterations within the thread loop
|
||||
int iteration_thread_;
|
||||
|
||||
private:
|
||||
/// Computes predicates based on internally tracked per-thread offset.
|
||||
CUTLASS_HOST_DEVICE
|
||||
void compute_predicates_(
|
||||
/// optionally, simplify predicate calculation during 'steady state' phase
|
||||
bool is_steady_state = false) {
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kPredicateWordCount; ++i) {
|
||||
predicates_[i] = 0u;
|
||||
}
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ts = 0; ts < ThreadMap::ThreadAccessShape::kStrided; ts++) {
|
||||
|
||||
TensorCoord iteration_coord(c * ThreadMap::Delta::kContiguous,
|
||||
ts + s * ThreadMap::Delta::kStrided);
|
||||
|
||||
TensorCoord coord = thread_offset_ + iteration_coord;
|
||||
|
||||
bool guard;
|
||||
|
||||
if (is_steady_state) {
|
||||
if (kAdvanceRank == 0) {
|
||||
guard = (coord.strided() < extent_.strided());
|
||||
} else {
|
||||
guard = (coord.contiguous() < extent_.contiguous());
|
||||
}
|
||||
} else {
|
||||
guard = (coord.strided() < extent_.strided() &&
|
||||
coord.contiguous() < extent_.contiguous());
|
||||
}
|
||||
|
||||
int pred_idx = ts + c * ThreadMap::ThreadAccessShape::kStrided + s * ThreadMap::Iterations::kContiguous * ThreadMap::ThreadAccessShape::kStrided;
|
||||
int word_idx = pred_idx / kPredicatesPerWord;
|
||||
int residual = pred_idx % kPredicatesPerWord;
|
||||
int byte_idx = residual / kPredicatesPerByte;
|
||||
int bit_idx = residual % kPredicatesPerByte;
|
||||
|
||||
predicates_[word_idx] |= (unsigned(guard) << (byte_idx * 8 + bit_idx));
|
||||
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
public:
|
||||
/// Constructs a TileIterator from its precomputed state, threadblock offset,
|
||||
/// and thread ID
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile(
|
||||
/// Precomputed parameters object
|
||||
Params const ¶ms,
|
||||
/// Pointer to start of tensor
|
||||
Pointer pointer,
|
||||
/// Extent of tensor
|
||||
TensorCoord extent,
|
||||
/// ID of each participating thread
|
||||
int thread_id,
|
||||
/// Initial offset of threadblock
|
||||
TensorCoord const &threadblock_offset)
|
||||
: params_(params),
|
||||
pointer_(reinterpret_cast<BytePointer>(
|
||||
const_cast<NonConstPointer>(pointer))),
|
||||
extent_(extent),
|
||||
is_residue_tile_(true) {
|
||||
|
||||
|
||||
TensorCoord residue_offset;
|
||||
if (kAdvanceRank) {
|
||||
residue_tile_idx_ =
|
||||
(extent_[kAdvanceRank] - threadblock_offset[kAdvanceRank] - 1) /
|
||||
Shape::kStrided;
|
||||
residue_offset = make_Coord(0, residue_tile_idx_ * Shape::kStrided);
|
||||
} else {
|
||||
residue_tile_idx_ =
|
||||
(extent_[kAdvanceRank] - threadblock_offset[kAdvanceRank] - 1) /
|
||||
Shape::kContiguous;
|
||||
residue_offset = make_Coord(residue_tile_idx_ * Shape::kContiguous, 0);
|
||||
}
|
||||
|
||||
// Per-thread offset in logical coordinates of tensor
|
||||
thread_offset_ = threadblock_offset + residue_offset +
|
||||
ThreadMap::initial_offset(thread_id);
|
||||
|
||||
// update internal pointers
|
||||
Layout layout(params_.stride_);
|
||||
add_pointer_offset(layout(thread_offset_));
|
||||
|
||||
compute_predicates_(false);
|
||||
|
||||
set_iteration_index(0);
|
||||
}
|
||||
|
||||
/// Construct a PredicatedTileAccessIterator2dThreadTile with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile(
|
||||
/// Precomputed parameters object
|
||||
Params const ¶ms,
|
||||
/// Pointer to start of tensor
|
||||
Pointer pointer,
|
||||
/// Extent of tensor
|
||||
TensorCoord extent,
|
||||
///< ID of each participating thread
|
||||
int thread_id)
|
||||
: PredicatedTileAccessIterator2dThreadTile(params, pointer, extent, thread_id,
|
||||
make_Coord(0, 0)) {}
|
||||
|
||||
/// Overrides the internal iteration index
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_iteration_index(int index) {
|
||||
|
||||
int residual = index % (ThreadMap::Iterations::kContiguous * ThreadMap::ThreadAccessShape::kStrided);
|
||||
iteration_strided_ = index / (ThreadMap::Iterations::kContiguous * ThreadMap::ThreadAccessShape::kStrided);
|
||||
|
||||
iteration_contiguous_ = residual / ThreadMap::ThreadAccessShape::kStrided;
|
||||
iteration_thread_ = residual % ThreadMap::ThreadAccessShape::kStrided;
|
||||
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
pointer_ += int(sizeof(Element)) * pointer_offset;
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(
|
||||
TensorCoord const &tile_offset) {
|
||||
if (is_residue_tile_) {
|
||||
TensorCoord residue_offset;
|
||||
if (kAdvanceRank) {
|
||||
residue_offset = TensorCoord(0, residue_tile_idx_ * Shape::kStrided);
|
||||
} else {
|
||||
residue_offset = TensorCoord(residue_tile_idx_ * Shape::kContiguous, 0);
|
||||
}
|
||||
|
||||
thread_offset_ -= residue_offset;
|
||||
|
||||
Layout layout(params_.stride_);
|
||||
add_pointer_offset(-layout(residue_offset));
|
||||
|
||||
compute_predicates_(true);
|
||||
|
||||
if (kAdvanceRank) {
|
||||
pointer_ += params_.inc_advance_ * (tile_offset.strided() - 1);
|
||||
pointer_ += Shape::kContiguous * tile_offset.contiguous();
|
||||
} else {
|
||||
pointer_ += params_.inc_advance_ * (tile_offset.contiguous() - 1);
|
||||
pointer_ += Shape::kStrided * tile_offset.strided();
|
||||
}
|
||||
} else {
|
||||
if (kAdvanceRank) {
|
||||
pointer_ += params_.inc_advance_ * tile_offset.strided();
|
||||
pointer_ += Shape::kContiguous * tile_offset.contiguous();
|
||||
} else {
|
||||
pointer_ += params_.inc_advance_ * tile_offset.contiguous();
|
||||
pointer_ += Shape::kStrided * tile_offset.strided();
|
||||
}
|
||||
}
|
||||
is_residue_tile_ = false;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
AccessType *get() const {
|
||||
|
||||
AccessType *ret_val = reinterpret_cast<AccessType *>(
|
||||
pointer_ + (iteration_thread_ * params_.stride_ + iteration_contiguous_ * ThreadMap::Delta::kContiguous) * int(sizeof(Element)));
|
||||
|
||||
return ret_val;
|
||||
}
|
||||
|
||||
/// Increment and return an instance to self.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile &operator++() {
|
||||
|
||||
iteration_thread_++;
|
||||
|
||||
if (iteration_thread_ < ThreadMap::ThreadAccessShape::kStrided)
|
||||
return *this;
|
||||
|
||||
iteration_thread_ = 0;
|
||||
|
||||
++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) {
|
||||
pointer_ += params_.inc_strided_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
// Enter here only if (iteration_stride_ == ThreadMap::Iteration::kStrided)
|
||||
// which means we enter the next tile.
|
||||
iteration_strided_ = 0;
|
||||
|
||||
// advance to next tile
|
||||
pointer_ += params_.inc_next_;
|
||||
|
||||
// now return to start tile - if the iterator is subsequently advanced, this
|
||||
// subtraction as well as the subsequent integer addition are both elided by
|
||||
// the compiler.
|
||||
pointer_ -= params_.inc_advance_;
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Increment and return an instance to self.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile operator++(int) {
|
||||
PredicatedTileAccessIterator2dThreadTile self(*this);
|
||||
operator++();
|
||||
return self;
|
||||
}
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void clear_mask() {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kPredicateWordCount; ++i) {
|
||||
predicates_[i] = 0u;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void enable_mask() {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kPredicateWordCount; ++i) {
|
||||
predicates_[i] = 0xffffffff;
|
||||
}
|
||||
}
|
||||
|
||||
/// Sets the predicate mask, overriding value stored in predicate iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_mask(Mask const &mask) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kPredicateWordCount; ++i) {
|
||||
predicates_[i] = mask[i];
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
/// Gets the mask
|
||||
CUTLASS_HOST_DEVICE
|
||||
void get_mask(Mask &mask) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kPredicateWordCount; ++i) {
|
||||
mask[i] = predicates_[i];
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns whether access is valid or not
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool valid() {
|
||||
|
||||
int pred_idx =
|
||||
iteration_thread_ +
|
||||
iteration_contiguous_ * ThreadMap::ThreadAccessShape::kStrided +
|
||||
iteration_strided_ * ThreadMap::Iterations::kContiguous * ThreadMap::ThreadAccessShape::kStrided;
|
||||
|
||||
int word_idx = pred_idx / kPredicatesPerWord;
|
||||
int residual = pred_idx % kPredicatesPerWord;
|
||||
int byte_idx = residual / kPredicatesPerByte;
|
||||
int bit_idx = residual % kPredicatesPerByte;
|
||||
|
||||
bool pred = (predicates_[word_idx] & (1u << (byte_idx * 8 + bit_idx))) != 0;
|
||||
|
||||
return pred;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization of PredicatedTileAccessIterator2dThreadTile for pitch-linear data.
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept |
|
||||
/// MaskedTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, typename AccessType_>
|
||||
class PredicatedTileAccessIterator2dThreadTile<Shape_, Element_, layout::ColumnMajor,
|
||||
AdvanceRank, ThreadMap_, AccessType_> {
|
||||
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::ColumnMajor;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
using AccessType = AccessType_;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorView = TensorView<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using Pointer = Element *;
|
||||
using NonConstPointer = typename platform::remove_const<Element>::type *;
|
||||
|
||||
using UnderlyingIterator = PredicatedTileAccessIterator2dThreadTile<
|
||||
layout::PitchLinearShape<Shape::kRow, Shape::kColumn>, Element,
|
||||
layout::PitchLinear, (kAdvanceRank == 0 ? 0 : 1), ThreadMap, AccessType>;
|
||||
|
||||
/// Predicate vector stores mask to guard accesses
|
||||
using Mask = typename UnderlyingIterator::Mask;
|
||||
|
||||
/// Parameters object is precomputed state and is host-constructible
|
||||
class Params {
|
||||
private:
|
||||
friend PredicatedTileAccessIterator2dThreadTile;
|
||||
|
||||
/// Parameters object
|
||||
typename UnderlyingIterator::Params params_;
|
||||
|
||||
public:
|
||||
|
||||
/// Default ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params() { }
|
||||
|
||||
/// Construct the Params object given a pitch-linear tensor's layout
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Layout const &layout)
|
||||
: params_(layout::PitchLinear(layout.stride(0))){};
|
||||
};
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Underlying pitch-linear tile iterator
|
||||
UnderlyingIterator iterator_;
|
||||
|
||||
public:
|
||||
/// Constructs a TileIterator from its precomputed state, threadblock offset,
|
||||
/// and thread ID
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile(
|
||||
///< Precomputed parameters object
|
||||
Params const ¶ms,
|
||||
///< Pointer to start of tensor
|
||||
Pointer pointer,
|
||||
///< Extent of tensor
|
||||
TensorCoord extent,
|
||||
///< ID of each participating thread
|
||||
int thread_id,
|
||||
///< Initial offset of threadblock
|
||||
TensorCoord const &threadblock_offset)
|
||||
: iterator_(params.params_, pointer,
|
||||
layout::PitchLinearCoord(extent.row(), extent.column()),
|
||||
thread_id,
|
||||
layout::PitchLinearCoord(threadblock_offset.row(),
|
||||
threadblock_offset.column())) {}
|
||||
|
||||
/// Construct a PredicatedTileAccessIterator2dThreadTile with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile(
|
||||
Params const ¶ms, ///< Precomputed parameters object
|
||||
Pointer pointer, ///< Pointer to start of tensor
|
||||
TensorCoord extent, ///< Extent of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: PredicatedTileAccessIterator2dThreadTile(params, pointer, extent, thread_id,
|
||||
make_Coord(0, 0)) {}
|
||||
|
||||
/// 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);
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole
|
||||
/// tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_tile_offset(TensorCoord const &tile_offset) {
|
||||
iterator_.add_tile_offset({tile_offset.row(), tile_offset.column()});
|
||||
}
|
||||
|
||||
/// Returns a pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
AccessType *get() const {
|
||||
return reinterpret_cast<AccessType *>(iterator_.get());
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
///
|
||||
/// The first time this method is called, predicates are updated, and the
|
||||
/// iterator's internal pointer is reverted to the first "steady state" tile.
|
||||
/// Subsequent calls are lightweight and must only update the internal
|
||||
/// pointer.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
///
|
||||
/// The first time this method is called, predicates are updated, and the
|
||||
/// iterator's internal pointer is reverted to the first "steady state" tile.
|
||||
/// Subsequent calls are lightweight and must only update the internal
|
||||
/// pointer.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile operator++(int) {
|
||||
PredicatedTileAccessIterator2dThreadTile self(*this);
|
||||
operator++();
|
||||
return self;
|
||||
}
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void clear_mask() { iterator_.clear_mask(); }
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void enable_mask() { iterator_.enable_mask(); }
|
||||
|
||||
/// Sets the predicate mask, overriding value stored in predicate iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_mask(Mask const &mask) { iterator_.set_mask(mask); }
|
||||
|
||||
/// Gets the mask
|
||||
CUTLASS_HOST_DEVICE
|
||||
void get_mask(Mask &mask) { iterator_.get_mask(mask); }
|
||||
|
||||
/// Returns whether access is valid or not
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool valid() {
|
||||
return iterator_.valid();
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization of PredicatedTileAccessIterator2dThreadTile for pitch-linear data.
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept |
|
||||
/// MaskedTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, typename AccessType_>
|
||||
class PredicatedTileAccessIterator2dThreadTile<Shape_, Element_, layout::RowMajor,
|
||||
AdvanceRank, ThreadMap_, AccessType_> {
|
||||
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::RowMajor;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
using AccessType = AccessType_;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorView = TensorView<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using Pointer = Element *;
|
||||
using NonConstPointer = typename platform::remove_const<Element>::type *;
|
||||
|
||||
using UnderlyingIterator = PredicatedTileAccessIterator2dThreadTile<
|
||||
layout::PitchLinearShape<Shape::kColumn, Shape::kRow>, Element,
|
||||
layout::PitchLinear, (kAdvanceRank == 0 ? 1 : 0), ThreadMap, AccessType>;
|
||||
|
||||
/// Predicate vector stores mask to guard accesses
|
||||
using Mask = typename UnderlyingIterator::Mask;
|
||||
|
||||
/// Parameters object is precomputed state and is host-constructible
|
||||
class Params {
|
||||
private:
|
||||
friend PredicatedTileAccessIterator2dThreadTile;
|
||||
|
||||
/// Parameters object
|
||||
typename UnderlyingIterator::Params params_;
|
||||
|
||||
public:
|
||||
|
||||
/// Default ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params() { }
|
||||
|
||||
/// Construct the Params object given a pitch-linear tensor's layout
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Layout const &layout)
|
||||
: params_(layout::PitchLinear(layout.stride(0))){};
|
||||
};
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Underlying pitch-linear tile iterator
|
||||
UnderlyingIterator iterator_;
|
||||
|
||||
public:
|
||||
/// Constructs a TileIterator from its precomputed state, threadblock offset,
|
||||
/// and thread ID
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile(
|
||||
///< Precomputed parameters object
|
||||
Params const ¶ms,
|
||||
///< Pointer to start of tensor
|
||||
Pointer pointer,
|
||||
///< Extent of tensor
|
||||
TensorCoord extent,
|
||||
///< ID of each participating thread
|
||||
int thread_id,
|
||||
///< Initial offset of threadblock
|
||||
TensorCoord const &threadblock_offset)
|
||||
: iterator_(params.params_, pointer,
|
||||
layout::PitchLinearCoord(extent.column(), extent.row()),
|
||||
thread_id,
|
||||
layout::PitchLinearCoord(threadblock_offset.column(),
|
||||
threadblock_offset.row())) {}
|
||||
|
||||
/// Construct a PredicatedTileAccessIterator2dThreadTile with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile(
|
||||
Params const ¶ms, ///< Precomputed parameters object
|
||||
Pointer pointer, ///< Pointer to start of tensor
|
||||
TensorCoord extent, ///< Extent of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: PredicatedTileAccessIterator2dThreadTile(params, pointer, extent, thread_id,
|
||||
make_Coord(0, 0)) {}
|
||||
|
||||
/// 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);
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole
|
||||
/// tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_tile_offset(TensorCoord const &tile_offset) {
|
||||
iterator_.add_tile_offset({tile_offset.column(), tile_offset.row()});
|
||||
}
|
||||
|
||||
/// Returns a pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
AccessType *get() const {
|
||||
return reinterpret_cast<AccessType *>(iterator_.get());
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
///
|
||||
/// The first time this method is called, predicates are updated, and the
|
||||
/// iterator's internal pointer is reverted to the first "steady state" tile.
|
||||
/// Subsequent calls are lightweight and must only update the internal
|
||||
/// pointer.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
///
|
||||
/// The first time this method is called, predicates are updated, and the
|
||||
/// iterator's internal pointer is reverted to the first "steady state" tile.
|
||||
/// Subsequent calls are lightweight and must only update the internal
|
||||
/// pointer.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIterator2dThreadTile operator++(int) {
|
||||
PredicatedTileAccessIterator2dThreadTile self(*this);
|
||||
operator++();
|
||||
return self;
|
||||
}
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void clear_mask() { iterator_.clear_mask(); }
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void enable_mask() { iterator_.enable_mask(); }
|
||||
|
||||
/// Sets the predicate mask, overriding value stored in predicate iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_mask(Mask const &mask) { iterator_.set_mask(mask); }
|
||||
|
||||
/// Gets the mask
|
||||
CUTLASS_HOST_DEVICE
|
||||
void get_mask(Mask &mask) { iterator_.get_mask(mask); }
|
||||
|
||||
/// Returns whether access is valid or not
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool valid() {
|
||||
return iterator_.valid();
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
1112
include/cutlass/transform/threadblock/predicated_tile_iterator.h
Normal file
1112
include/cutlass/transform/threadblock/predicated_tile_iterator.h
Normal file
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,767 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice, this list of
|
||||
* conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright notice, this list of
|
||||
* conditions and the following disclaimer in the documentation and/or other materials
|
||||
* provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
|
||||
* to endorse or promote products derived from this software without specific prior written
|
||||
* permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
|
||||
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
|
||||
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
|
||||
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
|
||||
* STRICT LIABILITY, OR TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates implementing loading of tiles from pitch-linear rank=2 tensors.
|
||||
|
||||
This iterator uses masks to guard out-of-bounds accesses and visits the last "residue" tile
|
||||
first, with the objective of minimizing predicate mask updates during steady-state operation.
|
||||
|
||||
A precomputed "Params" object minimizes the amount of state that must be stored in registers,
|
||||
and integer addition is used to advance the pointer through memory.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/transform/threadblock/predicated_tile_access_iterator_2dthreadtile.h"
|
||||
#include "cutlass/transform/thread/transpose.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace transform {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// PredicatedTileIterator2dThreadTile
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept |
|
||||
/// MaskedTileIteratorConcept
|
||||
///
|
||||
/// Regular tile iterator using a precomputed control structure to minimize register liveness
|
||||
/// and integer arithmetic.
|
||||
///
|
||||
/// Layout is assumed to be invariant at the time the precomputed "Params" object is constructed.
|
||||
///
|
||||
/// Base pointer and tensor extents may be specified at the time the iterator is constructed.
|
||||
/// Subsequently, they are assumed to be immutable.
|
||||
///
|
||||
/// Adding a logical coordinate offset may be performed at the time the iterator is constructed.
|
||||
/// Subsequent additions to logical coordinate offset may be performed but are relatively expensive.
|
||||
///
|
||||
/// Vistitation order is intended to first visit a "residual" tile that may be partially full in
|
||||
/// both the advance dimension and the steady-state dimension. This is assumed to be the last
|
||||
/// tile in the iteration sequence. Advancing an iterator that has just been constructed moves to
|
||||
/// the first tile that is full in the advance dimension and recomputes predicates. Subsequent
|
||||
/// accesses may be performed without updating internal predicates and are efficient in terms of
|
||||
/// live register state and pointer arithmetic instructions.
|
||||
///
|
||||
/// To be efficient, this assumes the iteraor will be dereferenced and advanced at least once
|
||||
/// outside any looping structure to minimize integer arithmetic.
|
||||
///
|
||||
/// Acceses out of bounds are safe so long as `clear_mask()` is called prior to dereferencing
|
||||
/// the iterator.
|
||||
///
|
||||
///
|
||||
/// Example:
|
||||
///
|
||||
/// An efficient pipeline structure may be constructed as follows:
|
||||
///
|
||||
// template <typename Iterator>
|
||||
// __global__ void kernel(
|
||||
// typename Iterator::Params params,
|
||||
// typename Iterator::Element *ptr,
|
||||
// TensorCoord extent) {
|
||||
//
|
||||
// typename Iterator::Fragment fragment;
|
||||
//
|
||||
// TensorCoord threadblock_offset(0, 0);
|
||||
//
|
||||
// Iterator iter(params, ptr, extent, threadIdx.x, threadblock_offsets);
|
||||
//
|
||||
//
|
||||
// fragment = *iter; // load "residue" tile first
|
||||
// ++iter; // advance to first "steady state" tile and update internal masks
|
||||
//
|
||||
//
|
||||
// #pragma unroll
|
||||
// for (int i = Remaining - 1; i >= 0; --i) {
|
||||
//
|
||||
// f(fragment);
|
||||
//
|
||||
// if (!i) {
|
||||
// iter.clear_mask(); // light-weight operation to clear masks - subsequent loads become NO-OPs.
|
||||
// }
|
||||
//
|
||||
// fragment = *iter; // load tile during "steady state" phase
|
||||
// ++iter; // advance to next tile - lightweight due to steady-state masks
|
||||
// }
|
||||
// }
|
||||
//
|
||||
// void host(TensorView<Element, 2, layout::PitchLinear> view) {
|
||||
//
|
||||
// using Iterator = transform::threadblock::PredicatedTileIterator2dThreadTile;
|
||||
//
|
||||
// typename Iterator::Params params(view.layout());
|
||||
//
|
||||
// kernel<Iterator>(params, view.data());
|
||||
// }
|
||||
///
|
||||
///
|
||||
template <
|
||||
typename Shape,
|
||||
typename Element,
|
||||
typename Layout,
|
||||
int AdvanceRank,
|
||||
typename ThreadMap,
|
||||
bool Transpose = false
|
||||
>
|
||||
class PredicatedTileIterator2dThreadTile;
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization of PredicatedTileIterator2dThreadTile for pitch-linear data.
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept |
|
||||
/// MaskedTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank, typename ThreadMap_, bool Transpose_>
|
||||
class PredicatedTileIterator2dThreadTile<Shape_, Element_, layout::PitchLinear, AdvanceRank, ThreadMap_, Transpose_> {
|
||||
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::PitchLinear;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorView = TensorView<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using Pointer = Element *;
|
||||
using NonConstPointer = typename platform::remove_const<Element>::type *;
|
||||
|
||||
/// Type used for internal memory accesses
|
||||
/// extra set of parenthesis is needed for VS compiler
|
||||
struct alignas((ThreadMap::kElementsPerAccess * sizeof_bits<Element>::value /
|
||||
8)) AccessType {
|
||||
|
||||
Array<Element, ThreadMap::kElementsPerAccess> storage;
|
||||
|
||||
static int const kElements = ThreadMap::kElementsPerAccess;
|
||||
};
|
||||
|
||||
/// Optinally this fragment can be 4x4 transposed
|
||||
using Transform = thread::Transpose< ThreadMap::Iterations::kCount * ThreadMap::ThreadAccessShape::kCount , layout::PitchLinearShape<4,4>, Element>;
|
||||
static bool const transpose = Transpose_;
|
||||
|
||||
/// Underlying iterator to compute the addresses
|
||||
using TileAccessIterator =
|
||||
PredicatedTileAccessIterator2dThreadTile<Shape, Element, Layout, kAdvanceRank,
|
||||
ThreadMap, AccessType>;
|
||||
|
||||
/// Fragment object to be loaded or stored
|
||||
using Fragment = cutlass::Array<Element, ThreadMap::Iterations::kCount *
|
||||
ThreadMap::ThreadAccessShape::kCount>;
|
||||
|
||||
/// Predicate vector stores mask to guard accesses
|
||||
using Mask = typename TileAccessIterator::Mask;
|
||||
|
||||
/// Parameters object is precomputed state and is host-constructible
|
||||
class Params {
|
||||
public:
|
||||
friend PredicatedTileIterator2dThreadTile;
|
||||
|
||||
private:
|
||||
/// Parameters object
|
||||
typename TileAccessIterator::Params params_;
|
||||
|
||||
public:
|
||||
/// Construct the Params object given a pitch-linear tensor's layout
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Layout const &layout) : params_(layout) { }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params() { }
|
||||
};
|
||||
|
||||
private:
|
||||
/// Internal pointer type permits fast address arithmetic
|
||||
using BytePointer = char *;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Data member to the tile access iterator
|
||||
TileAccessIterator address_iterator_;
|
||||
|
||||
public:
|
||||
/// Constructs a TileIterator from its precomputed state, threadblock offset,
|
||||
/// and thread ID
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile(
|
||||
/// Precomputed parameters object
|
||||
Params const ¶ms,
|
||||
/// Pointer to start of tensor
|
||||
Pointer pointer,
|
||||
/// Extent of tensor
|
||||
TensorCoord extent,
|
||||
/// ID of each participating thread
|
||||
int thread_id,
|
||||
/// Initial offset of threadblock
|
||||
TensorCoord const &threadblock_offset)
|
||||
: address_iterator_(params.params_, pointer, extent, thread_id,
|
||||
threadblock_offset) {}
|
||||
|
||||
/// Construct a PredicatedTileIterator2dThreadTile with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile(
|
||||
Params const ¶ms, ///< Precomputed parameters object
|
||||
Pointer pointer, ///< Pointer to start of tensor
|
||||
TensorCoord extent, ///< Extent of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: PredicatedTileIterator2dThreadTile(params, pointer, extent, thread_id,
|
||||
make_Coord(0, 0)) {}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
address_iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
///
|
||||
/// The first time this method is called, predicates are updated, and the
|
||||
/// iterator's internal pointer is reverted to the first "steady state" tile.
|
||||
/// Subsequent calls are lightweight and must only update the internal
|
||||
/// pointer.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile &operator++() {
|
||||
if (kAdvanceRank)
|
||||
address_iterator_.add_tile_offset({0, 1});
|
||||
else
|
||||
address_iterator_.add_tile_offset({1, 0});
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
///
|
||||
/// The first time this method is called, predicates are updated, and the
|
||||
/// iterator's internal pointer is reverted to the first "steady state" tile.
|
||||
/// Subsequent calls are lightweight and must only update the internal
|
||||
/// pointer.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile operator++(int) {
|
||||
PredicatedTileIterator2dThreadTile self(*this);
|
||||
operator++();
|
||||
return self;
|
||||
}
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void clear_mask() { address_iterator_.clear_mask(); }
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void enable_mask() { address_iterator_.enable_mask(); }
|
||||
|
||||
/// Sets the predicate mask, overriding value stored in predicate iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_mask(Mask const &mask) { address_iterator_.set_mask(mask); }
|
||||
|
||||
/// Gets the mask
|
||||
CUTLASS_HOST_DEVICE
|
||||
void get_mask(Mask &mask) { address_iterator_.get_mask(mask); }
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ts = 0; ts < ThreadMap::ThreadAccessShape::kStrided; ts++){
|
||||
|
||||
int access_idx = ts + c * ThreadMap::ThreadAccessShape::kStrided + \
|
||||
s * ThreadMap::Iterations::kContiguous * ThreadMap::ThreadAccessShape::kStrided;
|
||||
|
||||
address_iterator_.set_iteration_index(access_idx);
|
||||
if (address_iterator_.valid()) {
|
||||
|
||||
frag_ptr[access_idx] =
|
||||
*(address_iterator_.get() + pointer_offset);
|
||||
}
|
||||
|
||||
++address_iterator_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (transpose) {
|
||||
Transform t;
|
||||
t.transform(frag, frag);
|
||||
}
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) { load_with_pointer_offset(frag, 0); }
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
|
||||
AccessType const *frag_ptr = reinterpret_cast<AccessType const *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ts = 0; ts < ThreadMap::ThreadAccessShape::kStrided; ts++){
|
||||
|
||||
int access_idx = ts + c * ThreadMap::ThreadAccessShape::kStrided + \
|
||||
s * ThreadMap::Iterations::kContiguous * ThreadMap::ThreadAccessShape::kStrided;
|
||||
|
||||
address_iterator_.set_iteration_index(access_idx);
|
||||
if (address_iterator_.valid()) {
|
||||
*(address_iterator_.get() + pointer_offset) = frag_ptr[access_idx];
|
||||
}
|
||||
++address_iterator_;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) { store_with_pointer_offset(frag, 0); }
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization of PredicatedTileIterator2dThreadTile for pitch-linear data.
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept |
|
||||
/// MaskedTileIteratorConcept
|
||||
///
|
||||
template <
|
||||
typename Shape_,
|
||||
typename Element_,
|
||||
int AdvanceRank,
|
||||
typename ThreadMap_,
|
||||
bool Transpose_
|
||||
>
|
||||
class PredicatedTileIterator2dThreadTile<Shape_, Element_, layout::ColumnMajor, AdvanceRank, ThreadMap_, Transpose_> {
|
||||
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::ColumnMajor;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
static bool const Transpose = Transpose_;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorView = TensorView<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using Pointer = Element *;
|
||||
using NonConstPointer = typename platform::remove_const<Element>::type *;
|
||||
|
||||
using UnderlyingIterator = PredicatedTileIterator2dThreadTile<
|
||||
layout::PitchLinearShape<Shape::kRow, Shape::kColumn>,
|
||||
Element,
|
||||
layout::PitchLinear,
|
||||
(kAdvanceRank == 0 ? 0 : 1),
|
||||
ThreadMap,
|
||||
Transpose
|
||||
>;
|
||||
|
||||
using AccessType = typename UnderlyingIterator::AccessType;
|
||||
|
||||
/// Fragment object to be loaded or stored
|
||||
using Fragment = cutlass::Array<Element, ThreadMap::Iterations::kCount * ThreadMap::ThreadAccessShape::kCount>;
|
||||
|
||||
/// Predicate vector stores mask to guard accesses
|
||||
using Mask = typename UnderlyingIterator::Mask;
|
||||
|
||||
/// Parameters object is precomputed state and is host-constructible
|
||||
class Params {
|
||||
private:
|
||||
|
||||
friend PredicatedTileIterator2dThreadTile;
|
||||
|
||||
/// Parameters object
|
||||
typename UnderlyingIterator::Params params_;
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params() { }
|
||||
|
||||
/// Construct the Params object given a pitch-linear tensor's layout
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Layout const &layout): params_(layout::PitchLinear(layout.stride(0))) {
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Underlying pitch-linear tile iterator
|
||||
UnderlyingIterator iterator_;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructs a TileIterator from its precomputed state, threadblock offset, and thread ID
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile(
|
||||
Params const ¶ms, ///< Precomputed parameters object
|
||||
Pointer pointer, ///< Pointer to start of tensor
|
||||
TensorCoord extent, ///< Extent of tensor
|
||||
int thread_id, ///< ID of each participating thread
|
||||
TensorCoord const &threadblock_offset ///< Initial offset of threadblock
|
||||
):
|
||||
iterator_(
|
||||
params.params_,
|
||||
pointer,
|
||||
layout::PitchLinearCoord(extent.row(), extent.column()),
|
||||
thread_id,
|
||||
layout::PitchLinearCoord(threadblock_offset.row(), threadblock_offset.column())
|
||||
) { }
|
||||
|
||||
/// Construct a PredicatedTileIterator2dThreadTile with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile(
|
||||
Params const ¶ms, ///< Precomputed parameters object
|
||||
Pointer pointer, ///< Pointer to start of tensor
|
||||
TensorCoord extent, ///< Extent of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
): PredicatedTileIterator2dThreadTile(params, pointer, extent, thread_id, make_Coord(0, 0)) { }
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
///
|
||||
/// The first time this method is called, predicates are updated, and the iterator's
|
||||
/// internal pointer is reverted to the first "steady state" tile. Subsequent calls
|
||||
/// are lightweight and must only update the internal pointer.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
///
|
||||
/// The first time this method is called, predicates are updated, and the iterator's
|
||||
/// internal pointer is reverted to the first "steady state" tile. Subsequent calls
|
||||
/// are lightweight and must only update the internal pointer.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile operator++(int) {
|
||||
PredicatedTileIterator2dThreadTile self(*this);
|
||||
operator++();
|
||||
return self;
|
||||
}
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void clear_mask() {
|
||||
iterator_.clear_mask();
|
||||
}
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void enable_mask() {
|
||||
iterator_.enable_mask();
|
||||
}
|
||||
|
||||
/// Sets the predicate mask, overriding value stored in predicate iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_mask(Mask const &mask) {
|
||||
iterator_.set_mask(mask);
|
||||
}
|
||||
|
||||
/// Gets the mask
|
||||
CUTLASS_HOST_DEVICE
|
||||
void get_mask(Mask &mask) {
|
||||
iterator_.get_mask(mask);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization of PredicatedTileIterator2dThreadTile for pitch-linear data.
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept |
|
||||
/// MaskedTileIteratorConcept
|
||||
///
|
||||
template <
|
||||
typename Shape_,
|
||||
typename Element_,
|
||||
int AdvanceRank,
|
||||
typename ThreadMap_,
|
||||
bool Transpose_
|
||||
>
|
||||
class PredicatedTileIterator2dThreadTile<Shape_, Element_, layout::RowMajor, AdvanceRank, ThreadMap_, Transpose_> {
|
||||
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::RowMajor;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
static bool const Transpose = Transpose_;
|
||||
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
using TensorView = TensorView<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
|
||||
using Pointer = Element *;
|
||||
using NonConstPointer = typename platform::remove_const<Element>::type *;
|
||||
|
||||
using UnderlyingIterator = PredicatedTileIterator2dThreadTile<
|
||||
layout::PitchLinearShape<Shape::kColumn, Shape::kRow>,
|
||||
Element,
|
||||
layout::PitchLinear,
|
||||
(kAdvanceRank == 0 ? 1 : 0),
|
||||
ThreadMap,
|
||||
Transpose
|
||||
>;
|
||||
|
||||
using AccessType = typename UnderlyingIterator::AccessType;
|
||||
|
||||
/// Fragment object to be loaded or stored
|
||||
using Fragment = cutlass::Array<Element, ThreadMap::Iterations::kCount * ThreadMap::ThreadAccessShape::kCount>;
|
||||
|
||||
/// Predicate vector stores mask to guard accesses
|
||||
using Mask = typename UnderlyingIterator::Mask;
|
||||
|
||||
/// Parameters object is precomputed state and is host-constructible
|
||||
class Params {
|
||||
private:
|
||||
|
||||
friend PredicatedTileIterator2dThreadTile;
|
||||
|
||||
/// Parameters object
|
||||
typename UnderlyingIterator::Params params_;
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params() { }
|
||||
|
||||
/// Construct the Params object given a pitch-linear tensor's layout
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Layout const &layout): params_(layout::PitchLinear(layout.stride(0))) {
|
||||
|
||||
};
|
||||
};
|
||||
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Underlying pitch-linear tile iterator
|
||||
UnderlyingIterator iterator_;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructs a TileIterator from its precomputed state, threadblock offset, and thread ID
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile(
|
||||
Params const ¶ms, ///< Precomputed parameters object
|
||||
Pointer pointer, ///< Pointer to start of tensor
|
||||
TensorCoord extent, ///< Extent of tensor
|
||||
int thread_id, ///< ID of each participating thread
|
||||
TensorCoord const &threadblock_offset ///< Initial offset of threadblock
|
||||
):
|
||||
iterator_(
|
||||
params.params_,
|
||||
pointer,
|
||||
layout::PitchLinearCoord(extent.column(), extent.row()),
|
||||
thread_id,
|
||||
layout::PitchLinearCoord(threadblock_offset.column(), threadblock_offset.row())
|
||||
) { }
|
||||
|
||||
/// Construct a PredicatedTileIterator2dThreadTile with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile(
|
||||
Params const ¶ms, ///< Precomputed parameters object
|
||||
Pointer pointer, ///< Pointer to start of tensor
|
||||
TensorCoord extent, ///< Extent of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
): PredicatedTileIterator2dThreadTile(params, pointer, extent, thread_id, make_Coord(0, 0)) { }
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
///
|
||||
/// The first time this method is called, predicates are updated, and the iterator's
|
||||
/// internal pointer is reverted to the first "steady state" tile. Subsequent calls
|
||||
/// are lightweight and must only update the internal pointer.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
///
|
||||
/// The first time this method is called, predicates are updated, and the iterator's
|
||||
/// internal pointer is reverted to the first "steady state" tile. Subsequent calls
|
||||
/// are lightweight and must only update the internal pointer.
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileIterator2dThreadTile operator++(int) {
|
||||
PredicatedTileIterator2dThreadTile self(*this);
|
||||
operator++();
|
||||
return self;
|
||||
}
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void clear_mask() {
|
||||
iterator_.clear_mask();
|
||||
}
|
||||
|
||||
/// Clears the predicate set efficiently
|
||||
CUTLASS_HOST_DEVICE
|
||||
void enable_mask() {
|
||||
iterator_.enable_mask();
|
||||
}
|
||||
|
||||
/// Sets the predicate mask, overriding value stored in predicate iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_mask(Mask const &mask) {
|
||||
iterator_.set_mask(mask);
|
||||
}
|
||||
|
||||
/// Gets the mask
|
||||
CUTLASS_HOST_DEVICE
|
||||
void get_mask(Mask &mask) {
|
||||
iterator_.get_mask(mask);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,54 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
*modification, are permitted provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice,
|
||||
*this list of conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright
|
||||
*notice, this list of conditions and the following disclaimer in the
|
||||
*documentation and/or other materials provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its
|
||||
*contributors may be used to endorse or promote products derived from this
|
||||
*software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
*AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
*IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
*DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE FOR ANY DIRECT,
|
||||
*INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE,
|
||||
*DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY
|
||||
*OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TOR (INCLUDING
|
||||
*NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE,
|
||||
*EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates implementing the address computation of storing of tiles
|
||||
from pitch-linear rank=2 tensors.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace transform {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Shape, typename Element, typename Layout, int AdvanceRank,
|
||||
typename ThreadMap,
|
||||
int Alignment =
|
||||
sizeof_bits<Element>::value* ThreadMap::kElementsPerAccess / 8>
|
||||
class RegularTileAccessIterator;
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
@@ -0,0 +1,393 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice, this list of
|
||||
* conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright notice, this list of
|
||||
* conditions and the following disclaimer in the documentation and/or other materials
|
||||
* provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
|
||||
* to endorse or promote products derived from this software without specific prior written
|
||||
* permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
|
||||
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
|
||||
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
|
||||
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
|
||||
* STRICT LIABILITY, OR TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates implementing computing the addresses of storing of tiles
|
||||
from pitch-linear rank=2 tensors.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/layout/pitch_linear.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
|
||||
#include "cutlass/transform/threadblock/regular_tile_access_iterator.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace transform {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for congruous arrangements for TensorOps
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment>
|
||||
class RegularTileAccessIterator<
|
||||
Shape_, Element_,
|
||||
layout::PitchLinear,
|
||||
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::PitchLinear;
|
||||
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_;
|
||||
|
||||
/// Element type per access
|
||||
using AccessType = Array<Element, ThreadMap::kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Stride value
|
||||
Index stride_;
|
||||
|
||||
/// 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
|
||||
)
|
||||
: stride_(ref.stride(0) / ThreadMap::kElementsPerAccess),
|
||||
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_DEVICE
|
||||
AccessType *get() const {
|
||||
|
||||
AccessType *access_ptr = pointer_;
|
||||
|
||||
int access_offset = iteration_strided_ * ThreadMap::Delta::kStrided * stride_ +
|
||||
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_stride_ == 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::kContiguous +
|
||||
coord.strided() * Shape::kStrided * stride_ *
|
||||
ThreadMap::kElementsPerAccess);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for column major layouts
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment>
|
||||
class RegularTileAccessIterator<
|
||||
Shape_, Element_,
|
||||
layout::ColumnMajor,
|
||||
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::ColumnMajor;
|
||||
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<
|
||||
layout::PitchLinearShape<Shape::kRow, Shape::kColumn>, Element,
|
||||
layout::PitchLinear,
|
||||
(kAdvanceRank == 0 ? 0 : 1),
|
||||
ThreadMap_>;
|
||||
|
||||
using AccessType = typename UnderlyingIterator::AccessType;
|
||||
|
||||
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 reinterpret_cast<AccessType *>(iterator_.get());
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.row(), coord.column()});
|
||||
}
|
||||
|
||||
/// 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;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for row major layouts
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment>
|
||||
class RegularTileAccessIterator<
|
||||
Shape_, Element_,
|
||||
layout::RowMajor,
|
||||
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::RowMajor;
|
||||
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<
|
||||
layout::PitchLinearShape<Shape::kColumn, Shape::kRow>, Element,
|
||||
layout::PitchLinear,
|
||||
(kAdvanceRank == 0 ? 1 : 0),
|
||||
ThreadMap_>;
|
||||
|
||||
using AccessType = typename UnderlyingIterator::AccessType;
|
||||
|
||||
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 reinterpret_cast<AccessType *>(iterator_.get());
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.column(), coord.row()});
|
||||
}
|
||||
|
||||
/// 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;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,805 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice, this list of
|
||||
* conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright notice, this list of
|
||||
* conditions and the following disclaimer in the documentation and/or other materials
|
||||
* provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
|
||||
* to endorse or promote products derived from this software without specific prior written
|
||||
* permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
|
||||
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
|
||||
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
|
||||
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
|
||||
* STRICT LIABILITY, OR TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates implementing computing the addresses of storing of tiles
|
||||
from pitch-linear rank=2 tensors.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/layout/pitch_linear.h"
|
||||
#include "cutlass/layout/tensor_op_multiplicand_sm75.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/transform/threadblock/regular_tile_access_iterator.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace transform {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for congruous arrangements for TensorOps
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment>
|
||||
class RegularTileAccessIterator<
|
||||
Shape_, Element_,
|
||||
layout::TensorOpMultiplicandCongruous<sizeof_bits<Element_>::value,
|
||||
int(128 / sizeof(Element_))>,
|
||||
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::TensorOpMultiplicandCongruous<sizeof_bits<Element_>::value,
|
||||
int(128 / sizeof(Element_))>;
|
||||
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");
|
||||
|
||||
///< Number of pointers
|
||||
static int const kPointerCount =
|
||||
(ThreadMap::Iterations::kStrided > 1 ? 2 : 1);
|
||||
};
|
||||
|
||||
/// Element type per access
|
||||
using AccessType = Array<Element, Layout::kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Stride value
|
||||
Index stride_;
|
||||
|
||||
/// Internal pointer to first access of tile
|
||||
AccessType *pointer_[Detail::kPointerCount];
|
||||
|
||||
/// 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
|
||||
)
|
||||
: stride_(ref.stride(0) / Layout::kElementsPerAccess),
|
||||
byte_offset_(0) {
|
||||
layout::PitchLinearCoord thread_offset_base =
|
||||
ThreadMap::initial_offset(thread_id);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < Detail::kPointerCount; ++i) {
|
||||
// This is the offset of a thread within a threadblock tile for a specific
|
||||
// pointer (units of elements)
|
||||
layout::PitchLinearCoord thread_offset_in_threadblock_tile =
|
||||
thread_offset_base +
|
||||
layout::PitchLinearCoord{
|
||||
0, ThreadMap::Detail::WarpThreadArrangement::kStrided * i};
|
||||
|
||||
// initialize pointer
|
||||
pointer_[i] = reinterpret_cast<AccessType *>(
|
||||
ref.data() + ref.offset(thread_offset_in_threadblock_tile));
|
||||
}
|
||||
|
||||
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_[iteration_strided_ & 1];
|
||||
int stride_idx = (iteration_strided_ & ~1);
|
||||
|
||||
int access_offset = stride_idx * ThreadMap::Delta::kStrided * stride_ +
|
||||
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::kContiguous +
|
||||
coord.strided() * Shape::kStrided * stride_ *
|
||||
Layout::kElementsPerAccess);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile Iterator specialized for column-major congruous TensorOp formats.
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment>
|
||||
class RegularTileAccessIterator<
|
||||
Shape_, Element_,
|
||||
layout::ColumnMajorTensorOpMultiplicandCongruous<
|
||||
sizeof_bits<Element_>::value, int(128 / sizeof(Element_))>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for column-major iterator may along advance along the "
|
||||
"columns(rank=0) or rows(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::ColumnMajorTensorOpMultiplicandCongruous<
|
||||
sizeof_bits<Element_>::value, int(128 / sizeof(Element_))>;
|
||||
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<
|
||||
layout::PitchLinearShape<Shape::kRow, Shape::kColumn>, Element,
|
||||
layout::TensorOpMultiplicandCongruous<sizeof_bits<Element_>::value,
|
||||
int(128 / sizeof(Element_))>,
|
||||
(kAdvanceRank == 0 ? 0 : 1), ThreadMap_>;
|
||||
|
||||
using AccessType = typename UnderlyingIterator::AccessType;
|
||||
|
||||
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 reinterpret_cast<AccessType *>(iterator_.get());
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.row(), coord.column()});
|
||||
}
|
||||
|
||||
/// 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;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile Iterator specialized for row-major congruous TensorOp formats.
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment>
|
||||
class RegularTileAccessIterator<
|
||||
Shape_, Element_,
|
||||
layout::RowMajorTensorOpMultiplicandCongruous<sizeof_bits<Element_>::value,
|
||||
int(128 / sizeof(Element_))>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for row-major iterator may along advance along the "
|
||||
"columns(rank=0) or rows(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::RowMajorTensorOpMultiplicandCongruous<
|
||||
sizeof_bits<Element_>::value, int(128 / sizeof(Element_))>;
|
||||
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<
|
||||
layout::PitchLinearShape<Shape::kColumn, Shape::kRow>, Element,
|
||||
layout::TensorOpMultiplicandCongruous<sizeof_bits<Element_>::value,
|
||||
int(128 / sizeof(Element_))>,
|
||||
(kAdvanceRank == 0 ? 1 : 0), ThreadMap_>;
|
||||
|
||||
using AccessType = typename UnderlyingIterator::AccessType;
|
||||
|
||||
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 reinterpret_cast<AccessType *>(iterator_.get());
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.column(), coord.row()});
|
||||
}
|
||||
|
||||
/// 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;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for crosswise arrangements for TensorOps
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment, int Crosswise>
|
||||
class RegularTileAccessIterator<Shape_, Element_,
|
||||
layout::TensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<Element_>::value, Crosswise>,
|
||||
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::TensorOpMultiplicandCrosswise<sizeof_bits<Element_>::value,
|
||||
Crosswise>;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
static int const kAlignment = Alignment;
|
||||
static int const kCrosswise = Crosswise;
|
||||
|
||||
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");
|
||||
|
||||
/// Number of pointers
|
||||
///
|
||||
/// Note:TN kblock32 layouts only needs 1 pointer, but strangely
|
||||
/// reducing pointer count hurts perfomrnace
|
||||
static int const kPointerCount =
|
||||
(ThreadMap::Iterations::kStrided > 1 ? 2 : 1);
|
||||
};
|
||||
|
||||
/// Element type per access
|
||||
using AccessType = Array<Element, Layout::kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Total number of sections. The memory is divided into stages. One stage
|
||||
/// can store one tile. Stage is divided into sections. Interleaved layout
|
||||
/// can have multiple sections in a stage. The rest layout only has one section
|
||||
/// in a stage.
|
||||
int sections_;
|
||||
|
||||
/// Sections that a stage has
|
||||
int sections_per_stage_;
|
||||
|
||||
/// Stride value
|
||||
Index stride_;
|
||||
|
||||
/// Internal pointer to first access of tile
|
||||
AccessType *pointer_[Detail::kPointerCount];
|
||||
|
||||
/// 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
|
||||
)
|
||||
: sections_(ref.stride(0) / kCrosswise),
|
||||
sections_per_stage_(Shape::kContiguous / kCrosswise),
|
||||
// stride_ = kCrosswise x sections_ x kFactor
|
||||
stride_(ref.stride(0) * Layout::kFactor / Layout::kElementsPerAccess),
|
||||
byte_offset_(0) {
|
||||
layout::PitchLinearCoord thread_offset_base =
|
||||
ThreadMap::initial_offset(thread_id);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < Detail::kPointerCount; ++i) {
|
||||
// This is the offset of a thread within a threadblock tile for a specific
|
||||
// pointer (units of elements)
|
||||
layout::PitchLinearCoord thread_offset_in_threadblock_tile =
|
||||
thread_offset_base +
|
||||
layout::PitchLinearCoord{
|
||||
0, ThreadMap::Detail::WarpThreadArrangement::kStrided * i};
|
||||
// initialize pointer
|
||||
pointer_[i] = reinterpret_cast<AccessType *>(ref.data()) +
|
||||
ref.offset(thread_offset_in_threadblock_tile) /
|
||||
Layout::kElementsPerAccess;
|
||||
}
|
||||
|
||||
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_bits<Element>::value / 8;
|
||||
}
|
||||
|
||||
/// Returns a pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
AccessType *get() const {
|
||||
AccessType *access_ptr = pointer_[iteration_strided_ & 1];
|
||||
int stride_idx = (iteration_strided_ & ~1);
|
||||
|
||||
int access_offset =
|
||||
stride_idx * ThreadMap::Delta::kStrided * stride_ / Layout::kFactor +
|
||||
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 section.
|
||||
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() * sections_per_stage_ * stride_ *
|
||||
ThreadMap::kElementsPerAccess / sections_ +
|
||||
coord.strided() * Shape::kStrided * stride_ *
|
||||
Layout::kElementsPerAccess);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile Iterator specialized for column-major crosswise TensorOp formats.
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment, int Crosswise>
|
||||
class RegularTileAccessIterator<
|
||||
Shape_, Element_,
|
||||
layout::ColumnMajorTensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<Element_>::value, Crosswise>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for column-major iterator may along advance along the "
|
||||
"columns(rank=0) or rows(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::ColumnMajorTensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<Element_>::value, Crosswise>;
|
||||
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<
|
||||
layout::PitchLinearShape<Shape::kRow, Shape::kColumn>, Element,
|
||||
layout::TensorOpMultiplicandCrosswise<sizeof_bits<Element_>::value,
|
||||
Crosswise>,
|
||||
(kAdvanceRank == 0 ? 0 : 1), ThreadMap_>;
|
||||
|
||||
using AccessType = typename UnderlyingIterator::AccessType;
|
||||
|
||||
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 reinterpret_cast<AccessType *>(iterator_.get());
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.row(), coord.column()});
|
||||
}
|
||||
|
||||
/// 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;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile Iterator specialized for row-major crosswise TensorOp formats.
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment, int Crosswise>
|
||||
class RegularTileAccessIterator<Shape_, Element_,
|
||||
layout::RowMajorTensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<Element_>::value, Crosswise>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for row-major iterator may along advance along the "
|
||||
"columns(rank=0) or rows(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::RowMajorTensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<Element_>::value, Crosswise>;
|
||||
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<
|
||||
layout::PitchLinearShape<Shape::kColumn, Shape::kRow>, Element,
|
||||
layout::TensorOpMultiplicandCrosswise<sizeof_bits<Element_>::value,
|
||||
Crosswise>,
|
||||
(kAdvanceRank == 0 ? 1 : 0), ThreadMap_>;
|
||||
|
||||
using AccessType = typename UnderlyingIterator::AccessType;
|
||||
|
||||
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 reinterpret_cast<AccessType *>(iterator_.get());
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.column(), coord.row()});
|
||||
}
|
||||
|
||||
/// 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;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,56 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice, this list of
|
||||
* conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright notice, this list of
|
||||
* conditions and the following disclaimer in the documentation and/or other materials
|
||||
* provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
|
||||
* to endorse or promote products derived from this software without specific prior written
|
||||
* permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
|
||||
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
|
||||
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
|
||||
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
|
||||
* STRICT LIABILITY, OR TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates implementing storing of tiles from pitch-linear rank=2 tensors.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace transform {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename Shape,
|
||||
typename Element,
|
||||
typename Layout,
|
||||
int AdvanceRank,
|
||||
typename ThreadMap,
|
||||
int Alignment = sizeof_bits<Element>::value * ThreadMap::kElementsPerAccess / 8
|
||||
>
|
||||
class RegularTileIterator;
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
@@ -0,0 +1,485 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice, this list of
|
||||
* conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright notice, this list of
|
||||
* conditions and the following disclaimer in the documentation and/or other materials
|
||||
* provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
|
||||
* to endorse or promote products derived from this software without specific prior written
|
||||
* permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
|
||||
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
|
||||
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
|
||||
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
|
||||
* STRICT LIABILITY, OR TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates implementing loading of tiles from pitch-linear rank=2 tensors.
|
||||
|
||||
This iterator uses masks to guard out-of-bounds accesses and visits the last "residue" tile
|
||||
first, with the objective of minimizing predicate mask updates during steady-state operation.
|
||||
|
||||
A precomputed "Params" object minimizes the amount of state that must be stored in registers,
|
||||
and integer addition is used to advance the pointer through memory.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/layout/pitch_linear.h"
|
||||
|
||||
#include "regular_tile_iterator.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace transform {
|
||||
namespace threadblock {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Regular tile iterator specialized for pitch-linear
|
||||
template <
|
||||
typename Shape_,
|
||||
typename Element_,
|
||||
int AdvanceRank,
|
||||
typename ThreadMap_,
|
||||
int Alignment
|
||||
>
|
||||
class RegularTileIterator<Shape_, Element_, layout::PitchLinear, AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::PitchLinear;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
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 Fragment = Array<Element, ThreadMap::Iterations::kCount * ThreadMap::kElementsPerAccess>;
|
||||
|
||||
static_assert(kAdvanceRank == 0 || kAdvanceRank == 1,
|
||||
"Advance rank may only be along the contiguous or strided dimensions.");
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Types
|
||||
//
|
||||
|
||||
using AccessType = AlignedArray<Element, ThreadMap::kElementsPerAccess, kAlignment>;
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Pointer to memory
|
||||
uint8_t *pointer_;
|
||||
|
||||
/// Stride quantity
|
||||
Index stride_;
|
||||
|
||||
/// Amount to increment pointer along strided dimension
|
||||
Index increment_strided_;
|
||||
|
||||
/// Amount to advance pointer between tiles
|
||||
Index increment_advance_;
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator(): pointer_(nullptr), increment_strided_(0), increment_advance_(0) { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator(
|
||||
TensorRef const &ref,
|
||||
int thread_idx
|
||||
):
|
||||
pointer_(reinterpret_cast<uint8_t *>(ref.data()) + (ref.offset(ThreadMap::initial_offset(thread_idx)) * sizeof_bits<Element>::value / 8)) {
|
||||
|
||||
stride_ = ref.stride()[0];
|
||||
increment_strided_ = (ref.stride()[0] * sizeof_bits<Element>::value) * ThreadMap::Delta::kStrided / 8;
|
||||
|
||||
increment_advance_ =
|
||||
(kAdvanceRank == 0 ?
|
||||
Shape::kContiguous * sizeof_bits<Element>::value / 8 :
|
||||
Shape::kStrided * (ref.stride()[0] * sizeof_bits<Element>::value / 8));
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
uint8_t const *byte_pointer = pointer_ + pointer_offset * sizeof_bits<Element>::value / 8;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
|
||||
AccessType const *access_ptr = reinterpret_cast<AccessType const *>(byte_pointer);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
|
||||
int idx = c + s * ThreadMap::Iterations::kContiguous;
|
||||
frag_ptr[idx] = access_ptr[c * ThreadMap::Delta::kContiguous];
|
||||
}
|
||||
|
||||
if (s + 1 < ThreadMap::Iterations::kStrided) {
|
||||
byte_pointer += increment_strided_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag, TensorCoord const & tile_offset) {
|
||||
load_with_pointer_offset(
|
||||
frag,
|
||||
tile_offset.contiguous() * Shape::kContiguous / ThreadMap::kElementsPerAccess +
|
||||
tile_offset.strided() * Shape::kStrided * stride_
|
||||
);
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
|
||||
AccessType const *frag_ptr = reinterpret_cast<AccessType const*>(&frag);
|
||||
uint8_t *byte_pointer = pointer_ + pointer_offset * sizeof_bits<Element>::value / 8;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
|
||||
AccessType *access_ptr = reinterpret_cast<AccessType *>(byte_pointer);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
|
||||
int idx = c + s * ThreadMap::Iterations::kContiguous;
|
||||
access_ptr[c * ThreadMap::Delta::kContiguous] = frag_ptr[idx];
|
||||
}
|
||||
|
||||
if (s + 1 < ThreadMap::Iterations::kStrided) {
|
||||
byte_pointer += increment_strided_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag, TensorCoord const & tile_offset) {
|
||||
store_with_pointer_offset(
|
||||
frag,
|
||||
tile_offset.contiguous() * Shape::kContiguous + tile_offset.strided() * Shape::kStrided * stride_
|
||||
);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator++() {
|
||||
pointer_ += increment_advance_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator--() {
|
||||
pointer_ -= increment_advance_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
pointer_ += pointer_offset;
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
int offset = sizeof_bits<Element>::value *
|
||||
(coord.contiguous() * Shape::kContiguous + coord.strided() * Shape::kStrided * stride_) / 8;
|
||||
add_pointer_offset(offset);
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Regular tile iterator specialized for pitch-linear
|
||||
template <
|
||||
typename Shape_,
|
||||
typename Element_,
|
||||
int AdvanceRank,
|
||||
typename ThreadMap_,
|
||||
int Alignment
|
||||
>
|
||||
class RegularTileIterator<Shape_, Element_, layout::RowMajor, AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::RowMajor;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
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 Fragment = Array<Element, ThreadMap::Iterations::kCount * ThreadMap::kElementsPerAccess>;
|
||||
|
||||
using Underlying = RegularTileIterator<
|
||||
layout::PitchLinearShape<Shape::kColumn, Shape::kRow>,
|
||||
Element,
|
||||
layout::PitchLinear,
|
||||
(kAdvanceRank == 0 ? 1 : 0),
|
||||
ThreadMap,
|
||||
kAlignment
|
||||
>;
|
||||
|
||||
static_assert(kAdvanceRank == 0 || kAdvanceRank == 1,
|
||||
"Advance rank may only be along the row or column dimensions.");
|
||||
|
||||
private:
|
||||
|
||||
Underlying iterator_;
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator() { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator(
|
||||
TensorRef const &ref,
|
||||
int thread_idx
|
||||
):
|
||||
iterator_({ref.data(), ref.stride()}, thread_idx) {
|
||||
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag, TensorCoord const & tile_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, {tile_offset.column(), tile_offset.row()});
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
iterator_.load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag, TensorCoord const & tile_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, {tile_offset.column(), tile_offset.row()});
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
iterator_.store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator--() {
|
||||
--iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.column(), coord.row()});
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Regular tile iterator specialized for pitch-linear
|
||||
template <
|
||||
typename Shape_,
|
||||
typename Element_,
|
||||
int AdvanceRank,
|
||||
typename ThreadMap_,
|
||||
int Alignment
|
||||
>
|
||||
class RegularTileIterator<Shape_, Element_, layout::ColumnMajor, AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::ColumnMajor;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
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 Fragment = Array<Element, ThreadMap::Iterations::kCount * ThreadMap::kElementsPerAccess>;
|
||||
|
||||
using Underlying = RegularTileIterator<
|
||||
layout::PitchLinearShape<Shape::kRow, Shape::kColumn>,
|
||||
Element,
|
||||
layout::PitchLinear,
|
||||
(kAdvanceRank == 0 ? 0 : 1),
|
||||
ThreadMap
|
||||
>;
|
||||
|
||||
static_assert(kAdvanceRank == 0 || kAdvanceRank == 1,
|
||||
"Advance rank may only be along the row or column dimensions.");
|
||||
|
||||
private:
|
||||
|
||||
Underlying iterator_;
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator() { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator(
|
||||
TensorRef const &ref,
|
||||
int thread_idx
|
||||
):
|
||||
iterator_({ref.data(), ref.stride()}, thread_idx) {
|
||||
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag, TensorCoord const & tile_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, {tile_offset.row(), tile_offset.column()});
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
iterator_.load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag, TensorCoord const & tile_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, {tile_offset.row(), tile_offset.column()});
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
iterator_.store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator--() {
|
||||
--iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.row(), coord.column()});
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
|
||||
@@ -0,0 +1,502 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice, this list of
|
||||
* conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright notice, this list of
|
||||
* conditions and the following disclaimer in the documentation and/or other materials
|
||||
* provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
|
||||
* to endorse or promote products derived from this software without specific prior written
|
||||
* permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
|
||||
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
|
||||
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
|
||||
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
|
||||
* STRICT LIABILITY, OR TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates implementing loading of tiles from pitch-linear rank=2 tensors.
|
||||
|
||||
This iterator uses masks to guard out-of-bounds accesses and visits the last "residue" tile
|
||||
first, with the objective of minimizing predicate mask updates during steady-state operation.
|
||||
|
||||
A precomputed "Params" object minimizes the amount of state that must be stored in registers,
|
||||
and integer addition is used to advance the pointer through memory.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/layout/pitch_linear.h"
|
||||
|
||||
#include "regular_tile_iterator.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace transform {
|
||||
namespace threadblock {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
template <
|
||||
typename Shape,
|
||||
typename Element,
|
||||
typename Layout,
|
||||
int AdvanceRank,
|
||||
typename ThreadMap,
|
||||
int Alignment = sizeof_bits<Element>::value * ThreadMap::kElementsPerAccess / 8
|
||||
>
|
||||
class RegularTileIterator2dThreadTile;
|
||||
|
||||
|
||||
/// Regular tile iterator specialized for pitch-linear + 2d thread-tiled threadmapping
|
||||
template <
|
||||
typename Shape_,
|
||||
typename Element_,
|
||||
int AdvanceRank,
|
||||
typename ThreadMap_,
|
||||
int Alignment
|
||||
>
|
||||
class RegularTileIterator2dThreadTile<Shape_, Element_, layout::PitchLinear, AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::PitchLinear;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
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 Fragment = Array<Element, ThreadMap::Iterations::kCount * ThreadMap::ThreadAccessShape::kCount>;
|
||||
|
||||
static_assert(kAdvanceRank == 0 || kAdvanceRank == 1,
|
||||
"Advance rank may only be along the contiguous or strided dimensions.");
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Types
|
||||
//
|
||||
|
||||
using AccessType = AlignedArray<Element, ThreadMap::ThreadAccessShape::kCount, kAlignment>;
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Pointer to memory
|
||||
uint8_t *pointer_;
|
||||
|
||||
/// Stride quantity
|
||||
Index stride_;
|
||||
|
||||
/// Amount to increment pointer along strided dimension
|
||||
Index increment_strided_;
|
||||
|
||||
/// Amount to advance pointer between tiles
|
||||
Index increment_advance_;
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator2dThreadTile(): pointer_(nullptr), increment_strided_(0), increment_advance_(0) { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator2dThreadTile(
|
||||
TensorRef const &ref,
|
||||
int thread_idx,
|
||||
int interleave
|
||||
){
|
||||
|
||||
TensorCoord t = ThreadMap::initial_offset(thread_idx);
|
||||
long int offset = t[0] * interleave + t[1] * ref.stride()[0]/interleave;
|
||||
pointer_ = reinterpret_cast<uint8_t *>(ref.data() + offset);
|
||||
|
||||
stride_ = ref.stride()[0] / interleave;
|
||||
increment_strided_ = (ref.stride()[0] * sizeof_bits<Element>::value / 8) * ThreadMap::Delta::kStrided / interleave;
|
||||
|
||||
increment_advance_ =
|
||||
(kAdvanceRank == 0 ?
|
||||
Shape::kContiguous * sizeof_bits<Element>::value / 8 :
|
||||
Shape::kStrided * (ref.stride()[0] * sizeof_bits<Element>::value / 8) / interleave);
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
uint8_t const *byte_pointer = pointer_ + pointer_offset * sizeof_bits<Element>::value / 8;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
|
||||
AccessType const *access_ptr = reinterpret_cast<AccessType const *>(byte_pointer);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
|
||||
int idx = c + s * ThreadMap::Iterations::kContiguous;
|
||||
frag_ptr[idx] = access_ptr[c * ThreadMap::Delta::kContiguous / ThreadMap::ThreadAccessShape::kStrided];
|
||||
}
|
||||
|
||||
if (s + 1 < ThreadMap::Iterations::kStrided) {
|
||||
byte_pointer += increment_strided_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag, TensorCoord const & tile_offset) {
|
||||
load_with_pointer_offset(
|
||||
frag,
|
||||
tile_offset.contiguous() * Shape::kContiguous / ThreadMap::kElementsPerAccess +
|
||||
tile_offset.strided() * Shape::kStrided * stride_
|
||||
);
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
|
||||
AccessType const *frag_ptr = reinterpret_cast<AccessType const*>(&frag);
|
||||
uint8_t *byte_pointer = pointer_ + pointer_offset * sizeof_bits<Element>::value / 8;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
|
||||
AccessType *access_ptr = reinterpret_cast<AccessType *>(byte_pointer);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
|
||||
int idx = c + s * ThreadMap::Iterations::kContiguous;
|
||||
access_ptr[c * ThreadMap::Delta::kContiguous / ThreadMap::ThreadAccessShape::kStrided] = frag_ptr[idx];
|
||||
}
|
||||
|
||||
if (s + 1 < ThreadMap::Iterations::kStrided) {
|
||||
byte_pointer += increment_strided_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag, TensorCoord const & tile_offset) {
|
||||
store_with_pointer_offset(
|
||||
frag,
|
||||
tile_offset.contiguous() * Shape::kContiguous + tile_offset.strided() * Shape::kStrided * stride_
|
||||
);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator2dThreadTile &operator++() {
|
||||
pointer_ += increment_advance_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator2dThreadTile &operator--() {
|
||||
pointer_ -= increment_advance_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
pointer_ += pointer_offset;
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
int offset = sizeof_bits<Element>::value *
|
||||
(coord.contiguous() * Shape::kContiguous + coord.strided() * Shape::kStrided * stride_) / 8;
|
||||
add_pointer_offset(offset);
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Regular tile iterator specialized for interleaved layout + 2d thread-tiled threadmapping
|
||||
template <
|
||||
typename Shape_,
|
||||
typename Element_,
|
||||
int AdvanceRank,
|
||||
typename ThreadMap_,
|
||||
int Alignment
|
||||
>
|
||||
class RegularTileIterator2dThreadTile<Shape_, Element_, layout::RowMajorInterleaved<4>, AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::RowMajorInterleaved<4>;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
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 Fragment = Array<Element, ThreadMap::Iterations::kCount * ThreadMap::ThreadAccessShape::kCount>;
|
||||
|
||||
using Underlying = RegularTileIterator2dThreadTile<
|
||||
layout::PitchLinearShape<Shape::kColumn, Shape::kRow>,
|
||||
Element,
|
||||
layout::PitchLinear,
|
||||
(kAdvanceRank == 0 ? 1 : 0),
|
||||
ThreadMap,
|
||||
kAlignment
|
||||
>;
|
||||
|
||||
static_assert(kAdvanceRank == 0 || kAdvanceRank == 1,
|
||||
"Advance rank may only be along the row or column dimensions.");
|
||||
|
||||
private:
|
||||
|
||||
Underlying iterator_;
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator2dThreadTile() { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator2dThreadTile(
|
||||
TensorRef const &ref,
|
||||
int thread_idx
|
||||
):
|
||||
iterator_({ref.data(), ref.stride()}, thread_idx, 4) {
|
||||
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag, TensorCoord const & tile_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, {tile_offset.column(), tile_offset.row()});
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
iterator_.load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag, TensorCoord const & tile_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, {tile_offset.column(), tile_offset.row()});
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
iterator_.store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator2dThreadTile &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator2dThreadTile &operator--() {
|
||||
--iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.column(), coord.row()});
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Regular tile iterator specialized for interleaved layout + 2d thread-tiled threadmapping
|
||||
template <
|
||||
typename Shape_,
|
||||
typename Element_,
|
||||
int AdvanceRank,
|
||||
typename ThreadMap_,
|
||||
int Alignment
|
||||
>
|
||||
class RegularTileIterator2dThreadTile<Shape_, Element_, layout::ColumnMajorInterleaved<4>, AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::ColumnMajorInterleaved<4>;
|
||||
static int const kAdvanceRank = AdvanceRank;
|
||||
using ThreadMap = ThreadMap_;
|
||||
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 Fragment = Array<Element, ThreadMap::Iterations::kCount * ThreadMap::ThreadAccessShape::kCount>;
|
||||
using PitchLinearThreadMap = PitchLinearStripminedThreadMap< layout::PitchLinearShape<Shape::kRow, Shape::kColumn>,
|
||||
ThreadMap::kThreads, ThreadMap::ThreadAccessShape::kCount >;
|
||||
|
||||
|
||||
using Underlying = RegularTileIterator2dThreadTile<
|
||||
layout::PitchLinearShape<Shape::kRow, Shape::kColumn>,
|
||||
Element,
|
||||
layout::PitchLinear,
|
||||
(kAdvanceRank == 0 ? 0 : 1),
|
||||
ThreadMap
|
||||
>;
|
||||
|
||||
static_assert(kAdvanceRank == 0 || kAdvanceRank == 1,
|
||||
"Advance rank may only be along the row or column dimensions.");
|
||||
|
||||
private:
|
||||
|
||||
Underlying iterator_;
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator2dThreadTile() { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
RegularTileIterator2dThreadTile(
|
||||
TensorRef const &ref,
|
||||
int thread_idx
|
||||
):
|
||||
iterator_({ref.data(), ref.stride()}, thread_idx, 4) {
|
||||
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag, TensorCoord const & tile_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, {tile_offset.row(), tile_offset.column()});
|
||||
}
|
||||
|
||||
/// Loads a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
iterator_.load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag, TensorCoord const & tile_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, {tile_offset.row(), tile_offset.column()});
|
||||
}
|
||||
|
||||
/// Stores a fragment
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
iterator_.store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator2dThreadTile &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the pointer
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator2dThreadTile &operator--() {
|
||||
--iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.row(), coord.column()});
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
|
||||
@@ -0,0 +1,808 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice, this list of
|
||||
* conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright notice, this list of
|
||||
* conditions and the following disclaimer in the documentation and/or other materials
|
||||
* provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
|
||||
* to endorse or promote products derived from this software without specific prior written
|
||||
* permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
|
||||
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
|
||||
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
|
||||
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
|
||||
* STRICT LIABILITY, OR TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates implementing storing of tiles from pitch-linear rank=2 tensors.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/transform/threadblock/regular_tile_iterator.h"
|
||||
#include "cutlass/transform/threadblock/regular_tile_access_iterator_tensor_op.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace transform {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for congruous arrangements for TensorOps
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment>
|
||||
class RegularTileIterator<
|
||||
Shape_, Element_,
|
||||
layout::TensorOpMultiplicandCongruous<sizeof_bits<Element_>::value,
|
||||
int(128 / sizeof(Element_))>,
|
||||
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::TensorOpMultiplicandCongruous<sizeof_bits<Element_>::value,
|
||||
int(128 / sizeof(Element))>;
|
||||
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>;
|
||||
|
||||
public:
|
||||
|
||||
/// Fragment object to be loaded or stored
|
||||
using Fragment = Array<Element, ThreadMap::Iterations::kCount * Layout::kElementsPerAccess>;
|
||||
|
||||
/// Underlying iterator to compute the addresses
|
||||
using TileAccessIterator = RegularTileAccessIterator<Shape, Element, Layout,
|
||||
kAdvanceRank, ThreadMap>;
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Data member to the tile access iterator
|
||||
TileAccessIterator address_iterator_;
|
||||
|
||||
public:
|
||||
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator(TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: address_iterator_(ref, thread_id) {}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
address_iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator++() {
|
||||
address_iterator_.add_tile_offset({0, 1});
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator operator++(int) {
|
||||
RegularTileIterator prev(*this);
|
||||
this->operator++();
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
address_iterator_.add_tile_offset(coord);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
address_iterator_.set_iteration_index(0);
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
int access_idx = c + s * ThreadMap::Iterations::kContiguous;
|
||||
frag_ptr[access_idx] = *(address_iterator_.get() + pointer_offset);
|
||||
++address_iterator_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
address_iterator_.set_iteration_index(0);
|
||||
AccessType const *frag_ptr = reinterpret_cast<AccessType const *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
int access_idx = c + s * ThreadMap::Iterations::kContiguous;
|
||||
*(address_iterator_.get() + pointer_offset) = frag_ptr[access_idx];
|
||||
++address_iterator_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile Iterator specialized for column-major congruous TensorOp formats.
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment>
|
||||
class RegularTileIterator<
|
||||
Shape_, Element_,
|
||||
layout::ColumnMajorTensorOpMultiplicandCongruous<
|
||||
sizeof_bits<Element_>::value, int(128 / sizeof(Element_))>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
|
||||
static_assert(AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for column-major iterator may along advance along the "
|
||||
"columns(rank=0) or rows(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::ColumnMajorTensorOpMultiplicandCongruous<
|
||||
sizeof_bits<Element_>::value, int(128 / sizeof(Element))>;
|
||||
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 = RegularTileIterator<
|
||||
layout::PitchLinearShape<Shape::kRow, Shape::kColumn>, Element,
|
||||
layout::TensorOpMultiplicandCongruous<sizeof_bits<Element_>::value,
|
||||
int(128 / sizeof(Element))>,
|
||||
(kAdvanceRank == 0 ? 0 : 1), ThreadMap_>;
|
||||
|
||||
public:
|
||||
|
||||
/// Fragment object to be loaded or stored
|
||||
using Fragment = Array<Element, UnderlyingIterator::Fragment::kElements>;
|
||||
|
||||
private:
|
||||
|
||||
/// Underlying iterator
|
||||
UnderlyingIterator iterator_;
|
||||
|
||||
public:
|
||||
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator(
|
||||
TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
): iterator_({ref.data(), ref.stride()}, thread_id) {
|
||||
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.row(), coord.column()});
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator operator++(int) {
|
||||
RegularTileIterator prev(*this);
|
||||
++iterator_;
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(
|
||||
Fragment const &frag,
|
||||
Index pointer_offset) {
|
||||
|
||||
iterator_.store_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile Iterator specialized for row-major congruous TensorOp formats.
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment>
|
||||
class RegularTileIterator<
|
||||
Shape_, Element_,
|
||||
layout::RowMajorTensorOpMultiplicandCongruous<sizeof_bits<Element_>::value,
|
||||
int(128 / sizeof(Element_))>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
|
||||
static_assert(AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for row-major iterator may along advance along the "
|
||||
"columns(rank=0) or rows(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::RowMajorTensorOpMultiplicandCongruous<
|
||||
sizeof_bits<Element_>::value, int(128 / sizeof(Element))>;
|
||||
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 = RegularTileIterator<
|
||||
layout::PitchLinearShape<Shape::kColumn, Shape::kRow>, Element,
|
||||
layout::TensorOpMultiplicandCongruous<sizeof_bits<Element_>::value,
|
||||
int(128 / sizeof(Element))>,
|
||||
(kAdvanceRank == 0 ? 1 : 0), ThreadMap_>;
|
||||
|
||||
public:
|
||||
|
||||
/// Fragment object to be loaded or stored
|
||||
using Fragment = Array<Element, UnderlyingIterator::Fragment::kElements>;
|
||||
|
||||
private:
|
||||
|
||||
/// Underlying iterator
|
||||
UnderlyingIterator iterator_;
|
||||
|
||||
public:
|
||||
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator(
|
||||
TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
): iterator_({ref.data(), ref.stride()}, thread_id) {
|
||||
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.column(), coord.row()});
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator++() {
|
||||
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator operator++(int) {
|
||||
|
||||
RegularTileIterator prev(*this);
|
||||
++iterator_;
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) {
|
||||
load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(
|
||||
Fragment const &frag,
|
||||
Index pointer_offset) {
|
||||
|
||||
iterator_.store_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) {
|
||||
store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for crosswise arrangements for TensorOps
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment, int Crosswise>
|
||||
class RegularTileIterator<Shape_, Element_,
|
||||
layout::TensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<Element_>::value, Crosswise>,
|
||||
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::TensorOpMultiplicandCrosswise<sizeof_bits<Element_>::value,
|
||||
Crosswise>;
|
||||
|
||||
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>;
|
||||
|
||||
public:
|
||||
/// Fragment object to be loaded or stored
|
||||
using Fragment =
|
||||
Array<Element, ThreadMap::Iterations::kCount * Layout::kElementsPerAccess>;
|
||||
|
||||
/// Underlying iterator to compute the addresses
|
||||
using TileAccessIterator = RegularTileAccessIterator<Shape, Element, Layout,
|
||||
kAdvanceRank, ThreadMap>;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Data member to the tile access iterator
|
||||
TileAccessIterator address_iterator_;
|
||||
|
||||
public:
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator(TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: address_iterator_(ref, thread_id) {}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
address_iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator++() {
|
||||
address_iterator_.add_tile_offset({1, 0});
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator operator++(int) {
|
||||
RegularTileIterator prev(*this);
|
||||
this->operator++();
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
address_iterator_.add_tile_offset(coord);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
address_iterator_.set_iteration_index(0);
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
int access_idx = c + s * ThreadMap::Iterations::kContiguous;
|
||||
frag_ptr[access_idx] = *(address_iterator_.get() + pointer_offset);
|
||||
++address_iterator_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) { load_with_pointer_offset(frag, 0); }
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
address_iterator_.set_iteration_index(0);
|
||||
AccessType const *frag_ptr = reinterpret_cast<AccessType const *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < ThreadMap::Iterations::kContiguous; ++c) {
|
||||
int access_idx = c + s * ThreadMap::Iterations::kContiguous;
|
||||
*(address_iterator_.get() + pointer_offset) = frag_ptr[access_idx];
|
||||
++address_iterator_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) { store_with_pointer_offset(frag, 0); }
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile Iterator specialized for column-major crosswise TensorOp formats.
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment, int Crosswise>
|
||||
class RegularTileIterator<Shape_, Element_,
|
||||
layout::ColumnMajorTensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<Element_>::value, Crosswise>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for column-major iterator may along advance along the "
|
||||
"columns(rank=0) or rows(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::ColumnMajorTensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<Element_>::value, Crosswise>;
|
||||
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 = RegularTileIterator<
|
||||
layout::PitchLinearShape<Shape::kRow, Shape::kColumn>, Element,
|
||||
layout::TensorOpMultiplicandCrosswise<sizeof_bits<Element_>::value,
|
||||
Crosswise>,
|
||||
(kAdvanceRank == 0 ? 0 : 1), ThreadMap_>;
|
||||
|
||||
public:
|
||||
/// Fragment object to be loaded or stored
|
||||
using Fragment = Array<Element, UnderlyingIterator::Fragment::kElements>;
|
||||
|
||||
private:
|
||||
/// Underlying iterator
|
||||
UnderlyingIterator iterator_;
|
||||
|
||||
public:
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator(TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: iterator_({ref.data(), ref.stride()}, thread_id) {}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.row(), coord.column()});
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator operator++(int) {
|
||||
RegularTileIterator prev(*this);
|
||||
++iterator_;
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) { load_with_pointer_offset(frag, 0); }
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) { store_with_pointer_offset(frag, 0); }
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile Iterator specialized for row-major crosswise TensorOp formats.
|
||||
///
|
||||
///
|
||||
/// Satisfies: ForwardTileIteratorConcept |
|
||||
/// ReadableContiguousTileIteratorConcept |
|
||||
/// WriteableContiguousTileIteratorConcept
|
||||
///
|
||||
template <typename Shape_, typename Element_, int AdvanceRank,
|
||||
typename ThreadMap_, int Alignment, int Crosswise>
|
||||
class RegularTileIterator<Shape_, Element_,
|
||||
layout::RowMajorTensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<Element_>::value, Crosswise>,
|
||||
AdvanceRank, ThreadMap_, Alignment> {
|
||||
public:
|
||||
static_assert(
|
||||
AdvanceRank == 0 || AdvanceRank == 1,
|
||||
"Specialization for row-major iterator may along advance along the "
|
||||
"columns(rank=0) or rows(rank=1) dimension.");
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = layout::RowMajorTensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<Element_>::value, Crosswise>;
|
||||
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 = RegularTileIterator<
|
||||
layout::PitchLinearShape<Shape::kColumn, Shape::kRow>, Element,
|
||||
layout::TensorOpMultiplicandCrosswise<sizeof_bits<Element_>::value,
|
||||
Crosswise>,
|
||||
(kAdvanceRank == 0 ? 1 : 0), ThreadMap_>;
|
||||
|
||||
public:
|
||||
/// Fragment object to be loaded or stored
|
||||
using Fragment = Array<Element, UnderlyingIterator::Fragment::kElements>;
|
||||
|
||||
private:
|
||||
/// Underlying iterator
|
||||
UnderlyingIterator iterator_;
|
||||
|
||||
public:
|
||||
/// Construct a TileIterator with zero threadblock offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator(TensorRef ref, ///< Pointer to start of tensor
|
||||
int thread_id ///< ID of each participating thread
|
||||
)
|
||||
: iterator_({ref.data(), ref.stride()}, thread_id) {}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
iterator_.add_pointer_offset(pointer_offset);
|
||||
}
|
||||
|
||||
/// Adds a tile offset
|
||||
CUTLASS_DEVICE
|
||||
void add_tile_offset(TensorCoord const &coord) {
|
||||
iterator_.add_tile_offset({coord.column(), coord.row()});
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator &operator++() {
|
||||
++iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances to the next tile in memory.
|
||||
CUTLASS_HOST_DEVICE
|
||||
RegularTileIterator operator++(int) {
|
||||
RegularTileIterator prev(*this);
|
||||
++iterator_;
|
||||
|
||||
return prev;
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) {
|
||||
iterator_.load_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory
|
||||
CUTLASS_DEVICE
|
||||
void load(Fragment &frag) { load_with_pointer_offset(frag, 0); }
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) {
|
||||
iterator_.store_with_pointer_offset(frag, pointer_offset);
|
||||
}
|
||||
|
||||
/// Store a fragment to memory
|
||||
CUTLASS_DEVICE
|
||||
void store(Fragment const &frag) { store_with_pointer_offset(frag, 0); }
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
} // namespace threadblock
|
||||
} // namespace transform
|
||||
} // namespace cutlass
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user