@@ -0,0 +1,84 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2021, 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 Default warp-level GEMM operators selected by data type, size, and layouts of operands.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/warp/mma_with_reduction_tensor_op.h"
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace warp {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename WarpShape_,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// Data type of A elements
|
||||
typename ElementA,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA,
|
||||
/// Data type of B elements
|
||||
typename ElementB,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB,
|
||||
/// Element type of C matrix
|
||||
typename ElementC,
|
||||
/// Layout of C matrix (concept: MatrixLayout)
|
||||
typename LayoutC,
|
||||
/// Operator describing the tensor operation
|
||||
typename Operator_ = arch::OpMultiplyAdd,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK = 1,
|
||||
/// Store the accumulators in row major or column major. Row major is used
|
||||
/// when output layout is interleaved.
|
||||
bool AccumulatorsInRowMajor = false>
|
||||
struct DefaultMmaWithReductionTensorOp {
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<InstructionShape_, 32, ElementA,
|
||||
cutlass::layout::RowMajor, ElementB,
|
||||
cutlass::layout::ColumnMajor, ElementC,
|
||||
cutlass::layout::RowMajor, Operator_>,
|
||||
cutlass::MatrixShape<1, 1> >;
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using Type = cutlass::gemm::warp::MmaWithReductionTensorOp<
|
||||
WarpShape_, ElementA, LayoutA, ElementB, LayoutB, ElementC, LayoutC,
|
||||
Policy, PartitionsK, AccumulatorsInRowMajor>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -326,6 +326,9 @@ public:
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename ArchMmaOperator::Shape;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
|
||||
@@ -618,6 +621,9 @@ public:
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
|
||||
|
||||
@@ -121,6 +121,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -162,7 +165,7 @@ public:
|
||||
private:
|
||||
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_;
|
||||
@@ -395,6 +398,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -619,6 +625,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -835,6 +844,9 @@ class MmaTensorOpAccumulatorTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -1159,6 +1171,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -1200,7 +1215,7 @@ public:
|
||||
private:
|
||||
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_;
|
||||
@@ -1441,6 +1456,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -1666,6 +1684,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -1901,6 +1922,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -1946,7 +1970,7 @@ public:
|
||||
private:
|
||||
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_;
|
||||
@@ -2207,6 +2231,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -2249,7 +2276,7 @@ public:
|
||||
private:
|
||||
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_;
|
||||
@@ -2305,6 +2332,18 @@ public:
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator &add_tile_offset_negative(TensorCoord const &tile_offset) {
|
||||
|
||||
add_tile_offset(tile_offset);
|
||||
|
||||
if (k_group_idx_ & 1)
|
||||
byte_offset_ ^= 0x40;
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator & operator++() {
|
||||
|
||||
@@ -159,6 +159,9 @@ public:
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
|
||||
|
||||
@@ -154,6 +154,9 @@ public:
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename ThreadMma::ArchMmaOperator;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
/// Shape of the underlying instruction
|
||||
using InstructionShape = GemmShape<1,1,use_dp4a ? 4 : 1>;
|
||||
|
||||
|
||||
@@ -354,18 +354,27 @@ private:
|
||||
/// Internal reference
|
||||
cutlass::TensorRef<Element, layout::RowMajor> ref_;
|
||||
|
||||
/// Extent of tensor
|
||||
MatrixCoord extent_;
|
||||
|
||||
/// Origin
|
||||
MatrixCoord origin_;
|
||||
|
||||
/// Used to conditionally enable extents checking
|
||||
bool divisible_;
|
||||
|
||||
public:
|
||||
|
||||
/// Default ctor constructs null iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaSimtTileIterator() { }
|
||||
MmaSimtTileIterator() : divisible_(true) { }
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaSimtTileIterator(
|
||||
TensorRef ref,
|
||||
int lane_id
|
||||
) {
|
||||
) : extent_(Shape::kRow, Shape::kColumn), divisible_ (true) {
|
||||
|
||||
// compute offset based on thread ID and lane layout
|
||||
typename Policy::LaneLayout lane_layout = Policy::get_lane_layout();
|
||||
@@ -373,12 +382,35 @@ public:
|
||||
MatrixCoord lane_offset = lane_layout.inverse(lane_id) *
|
||||
MatrixCoord(Policy::LaneMmaShape::kM, 0);
|
||||
|
||||
origin_ = lane_offset;
|
||||
|
||||
ref.add_coord_offset(lane_offset);
|
||||
|
||||
ref_.reset(ref.data(), ref.stride(0));
|
||||
|
||||
}
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaSimtTileIterator(
|
||||
TensorRef ref,
|
||||
TensorCoord extent,
|
||||
int lane_id
|
||||
) : extent_(extent), divisible_ (false) {
|
||||
|
||||
// compute offset based on thread ID and lane layout
|
||||
typename Policy::LaneLayout lane_layout = Policy::get_lane_layout();
|
||||
|
||||
MatrixCoord lane_offset = lane_layout.inverse(lane_id) *
|
||||
MatrixCoord(Policy::LaneMmaShape::kM, 0);
|
||||
|
||||
origin_ = lane_offset;
|
||||
|
||||
ref.add_coord_offset(lane_offset);
|
||||
|
||||
ref_.reset(ref.data(), ref.stride(0));
|
||||
|
||||
}
|
||||
|
||||
/// Adds a pointer offset to internal pointer(s) to advance through memory
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -391,9 +423,13 @@ public:
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaSimtTileIterator &add_tile_offset(TensorCoord const &coord) {
|
||||
|
||||
ref_.add_coord_offset({
|
||||
TensorCoord coord_offset(
|
||||
coord.row() * Shape::kRow,
|
||||
coord.column() * Shape::kColumn});
|
||||
coord.column() * Shape::kColumn);
|
||||
|
||||
origin_ += coord_offset;
|
||||
|
||||
ref_.add_coord_offset(coord_offset);
|
||||
|
||||
return *this;
|
||||
}
|
||||
@@ -426,11 +462,21 @@ public:
|
||||
for (int m = 0; m < Iterations::kRow; ++m) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < Policy::LaneMmaShape::kM; i++) {
|
||||
|
||||
frag[m * Policy::LaneMmaShape::kM + i + k * Iterations::kRow] =
|
||||
*(ref_.data() +
|
||||
ref_.offset({m * Policy::WarpShape::kRow * Policy::LaneMmaShape::kM + i, k}) +
|
||||
pointer_offset);
|
||||
|
||||
MatrixCoord offset(m * Policy::WarpShape::kRow * Policy::LaneMmaShape::kM + i, k);
|
||||
|
||||
MatrixCoord access_coord = origin_ + offset;
|
||||
|
||||
int frag_idx = m * Policy::LaneMmaShape::kM + i + k * Iterations::kRow;
|
||||
|
||||
if (divisible_ ||
|
||||
(access_coord.row() < extent_.row() && access_coord.column() < extent_.column())) {
|
||||
|
||||
frag[frag_idx] = *(ref_.data() + ref_.offset(offset) + pointer_offset);
|
||||
}
|
||||
else {
|
||||
frag[frag_idx] = Element();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -765,18 +811,27 @@ private:
|
||||
/// Internal reference
|
||||
cutlass::TensorRef<Element, layout::ColumnMajor> ref_;
|
||||
|
||||
/// Extent of tensor
|
||||
MatrixCoord extent_;
|
||||
|
||||
/// Origin
|
||||
MatrixCoord origin_;
|
||||
|
||||
/// Used to conditionally enable extents checking
|
||||
bool divisible_;
|
||||
|
||||
public:
|
||||
|
||||
/// Default ctor constructs null iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaSimtTileIterator() { }
|
||||
MmaSimtTileIterator(): divisible_(true) { }
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaSimtTileIterator(
|
||||
TensorRef ref,
|
||||
int lane_id
|
||||
) {
|
||||
): extent_(Shape::kRow, Shape::kColumn), divisible_(true) {
|
||||
|
||||
// compute offset based on thread ID and lane layout
|
||||
typename Policy::LaneLayout lane_layout = Policy::get_lane_layout();
|
||||
@@ -784,11 +839,34 @@ public:
|
||||
MatrixCoord lane_offset = lane_layout.inverse(lane_id) *
|
||||
MatrixCoord(0, Policy::LaneMmaShape::kN);
|
||||
|
||||
origin_ = lane_offset;
|
||||
|
||||
ref.add_coord_offset(lane_offset);
|
||||
|
||||
ref_.reset(ref.data(), ref.stride(0));
|
||||
}
|
||||
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaSimtTileIterator(
|
||||
TensorRef ref,
|
||||
TensorCoord extent,
|
||||
int lane_id
|
||||
): extent_(extent), divisible_(false) {
|
||||
|
||||
// compute offset based on thread ID and lane layout
|
||||
typename Policy::LaneLayout lane_layout = Policy::get_lane_layout();
|
||||
|
||||
MatrixCoord lane_offset = lane_layout.inverse(lane_id) *
|
||||
MatrixCoord(0, Policy::LaneMmaShape::kN);
|
||||
|
||||
origin_ = lane_offset;
|
||||
|
||||
ref.add_coord_offset(lane_offset);
|
||||
|
||||
ref_.reset(ref.data(), ref.stride(0));
|
||||
}
|
||||
|
||||
/// Adds a pointer offset to internal pointer(s) to advance through memory
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaSimtTileIterator &add_pointer_offset(LongIndex offset) {
|
||||
@@ -800,9 +878,13 @@ public:
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaSimtTileIterator &add_tile_offset(TensorCoord const &coord) {
|
||||
|
||||
ref_.add_coord_offset({
|
||||
TensorCoord coord_offset(
|
||||
coord.row() * Shape::kRow,
|
||||
coord.column() * Shape::kColumn});
|
||||
coord.column() * Shape::kColumn);
|
||||
|
||||
origin_ += coord_offset;
|
||||
|
||||
ref_.add_coord_offset(coord_offset);
|
||||
|
||||
return *this;
|
||||
}
|
||||
@@ -835,10 +917,21 @@ public:
|
||||
for (int n = 0; n < Iterations::kColumn; ++n) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < Policy::LaneMmaShape::kN; ++i) {
|
||||
frag[n * Policy::LaneMmaShape::kN + i + k * Iterations::kColumn] =
|
||||
*(ref_.data() +
|
||||
ref_.offset({k, n * Policy::WarpShape::kColumn * Policy::LaneMmaShape::kN + i}) +
|
||||
pointer_offset);
|
||||
|
||||
MatrixCoord offset(k, n * Policy::WarpShape::kColumn * Policy::LaneMmaShape::kN + i);
|
||||
|
||||
MatrixCoord access_coord = origin_ + offset;
|
||||
|
||||
int frag_idx = n * Policy::LaneMmaShape::kN + i + k * Iterations::kColumn;
|
||||
|
||||
if (divisible_ ||
|
||||
(access_coord.row() < extent_.row() && access_coord.column() < extent_.column())) {
|
||||
|
||||
frag[frag_idx] = *(ref_.data() + ref_.offset(offset) + pointer_offset);
|
||||
}
|
||||
else {
|
||||
frag[frag_idx] = Element();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -119,6 +119,9 @@ public:
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename Base::ArchMmaOperator;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
/// Architecture tag from underlying instruction
|
||||
using ArchTag = typename Base::ArchTag;
|
||||
|
||||
|
||||
@@ -187,6 +187,9 @@ public:
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename Policy::Operator;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
/// Architecture tag from underlying instruction
|
||||
using ArchTag = typename ArchMmaOperator::ArchTag;
|
||||
|
||||
@@ -400,3 +403,4 @@ public:
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
@@ -90,9 +90,6 @@ class MmaTensorOpFragmentIterator<Shape_, AccumulatorShape_, KBlocksColumn_, Ele
|
||||
/// Output operation on fragment
|
||||
using OutputOp = OutputOp_;
|
||||
|
||||
/// Whether beta is zero
|
||||
static bool const IsBetaZero = true;
|
||||
|
||||
/// Number of participating threads
|
||||
static int const kThreads = 32;
|
||||
|
||||
@@ -274,9 +271,6 @@ class MmaTensorOpFragmentIterator<Shape_, AccumulatorShape_, KBlocksColumn_, Ele
|
||||
/// Output operation on fragment
|
||||
using OutputOp = OutputOp_;
|
||||
|
||||
/// Whether beta is zero
|
||||
static bool const IsBetaZero = true;
|
||||
|
||||
/// Number of participating threads
|
||||
static int const kThreads = 32;
|
||||
|
||||
|
||||
@@ -109,6 +109,9 @@ public:
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename Policy::Operator;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
/// Underlying instruction shape
|
||||
using InstructionShape = typename ArchMmaOperator::Shape;
|
||||
|
||||
|
||||
@@ -140,6 +140,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -205,7 +208,7 @@ public:
|
||||
private:
|
||||
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_[kPointerCount];
|
||||
@@ -535,6 +538,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -607,7 +613,7 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
|
||||
private:
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_[kPointerCount];
|
||||
@@ -892,6 +898,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -1363,6 +1372,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -1430,7 +1442,7 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
int sections_;
|
||||
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_;
|
||||
@@ -2643,6 +2655,307 @@ public:
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// This tile iterator is specialized for 32-thread TensorOps. It is used to load or store
|
||||
/// accumulators from memory and is agnostic to layout.
|
||||
///
|
||||
/// This iterator is not tested.
|
||||
///
|
||||
/// Satisfies:
|
||||
/// ReadableRandomAccessContiguousTileIteratorConcept |
|
||||
/// WriteableRandomAccessContiguousTileIteratorConcept
|
||||
///
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Element type
|
||||
typename Element_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Interval between adjacent *MMA instructions (in units of MMA
|
||||
/// instructions, concept: MatrixShape)
|
||||
typename OpDelta_>
|
||||
class MmaTensorOpAccumulatorTileIterator<
|
||||
Shape_, Element_, cutlass::layout::AffineRankN<2>, InstructionShape_, OpDelta_> {
|
||||
public:
|
||||
|
||||
/// Shape of tile to load (concept: MatrixShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Operand tag
|
||||
static Operand const kOperand = Operand::kC;
|
||||
|
||||
/// Element type
|
||||
using Element = Element_;
|
||||
|
||||
/// Layout of source tile
|
||||
using Layout = cutlass::layout::RowMajor;
|
||||
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
using InstructionShape = InstructionShape_;
|
||||
|
||||
/// Delta between *MMA operations (in units of *MMA operations, concept: MatrixShape)
|
||||
using OpDelta = OpDelta_;
|
||||
|
||||
/// Number of participating threads
|
||||
static int const kThreads = 32;
|
||||
|
||||
/// TensorRef type for loading element from a tensor
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
|
||||
/// Index type
|
||||
using Index = typename TensorRef::Index;
|
||||
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
/// Internal structure of iterator - made public to enable introspection
|
||||
struct Policy {
|
||||
static bool const kDivisible =
|
||||
!(Shape::kRow % InstructionShape::kM) &&
|
||||
!(Shape::kColumn % InstructionShape::kN);
|
||||
|
||||
static_assert(platform::is_same<TensorCoord, MatrixCoord>::value,
|
||||
"Layouts must be defined for logical MatrixCoord coordinate space.");
|
||||
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = MatrixShape<
|
||||
(Shape::kRow + InstructionShape::kM - 1) / InstructionShape::kM,
|
||||
(Shape::kColumn + InstructionShape::kN - 1) / InstructionShape::kN
|
||||
>;
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
// Assume accumulator tile is an arrangement of 8-by-8 tiles replicated over the entire
|
||||
// shape, with each quad mapped to one row and each thread mapped to 1/4 of the elements
|
||||
// of that row. The accumulators within one row are assumed to be consecutive.
|
||||
static int const kElementsPerAccess = InstructionShape::kN / 4;
|
||||
static int const kRowsPerTile = 8;
|
||||
static int const kAccumulatorRows = InstructionShape::kM / kRowsPerTile;
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Derived quantities
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<
|
||||
Element,
|
||||
Policy::MmaIterations::kCount * InstructionShape::kMN / kThreads>;
|
||||
|
||||
private:
|
||||
|
||||
/// Reference to output tensor
|
||||
TensorRef ref_;
|
||||
|
||||
public:
|
||||
|
||||
/// Default ctor constructs null iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator() { }
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator(
|
||||
TensorRef const &ref,
|
||||
int lane_id
|
||||
):
|
||||
ref_(ref) {
|
||||
|
||||
int quad = (lane_id >> 2);
|
||||
int lane_in_quad = (lane_id & 3);
|
||||
|
||||
MatrixCoord lane_offset(quad, lane_in_quad * kElementsPerAccess);
|
||||
|
||||
ref_.add_coord_offset(lane_offset);
|
||||
}
|
||||
|
||||
/// Adds a pointer offset to internal pointer(s) to advance through memory
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator &add_pointer_offset(LongIndex offset) {
|
||||
ref_.add_pointer_offset(offset);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator &add_tile_offset(TensorCoord const &tile_offset) {
|
||||
|
||||
ref_.add_coord_offset(tile_offset * make_Coord(Shape::kRow, Shape::kColumn));
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator & operator++() {
|
||||
// deliberate no-op
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator & operator--() {
|
||||
// deliberate no-op
|
||||
return *this;
|
||||
}
|
||||
|
||||
///< advances in units of whole tiles along the logical coordinate space of the tensor
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator & operator+=(TensorCoord const &tile_offset) {
|
||||
add_tile_offset(tile_offset);
|
||||
return *this;
|
||||
}
|
||||
|
||||
///< advances in units of whole tiles along the logical coordinate space of the tensor
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator & operator-=(TensorCoord const &tile_offset) {
|
||||
add_tile_offset(-tile_offset);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory at the location pointed to by the iterator.
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) const {
|
||||
load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with additional logical offset
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(
|
||||
Fragment &frag, ///< fragment to load from the tensor
|
||||
Index pointer_offset) const { ///< loads a tile with a linear offset
|
||||
|
||||
TensorRef offset_ref(ref_);
|
||||
offset_ref.add_pointer_offset(pointer_offset);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_n = 0; mma_n < Policy::MmaIterations::kColumn; ++mma_n) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_m = 0; mma_m < Policy::MmaIterations::kRow; ++mma_m) {
|
||||
|
||||
int mma_accum_start = kAccumulatorRows * kElementsPerAccess *
|
||||
(mma_n * Policy::MmaIterations::kRow + mma_m);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int row = 0; row < kAccumulatorRows; ++row) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int col = 0; col < kElementsPerAccess; ++col) {
|
||||
int accum_m = mma_m * InstructionShape::kM * OpDelta::kRow +
|
||||
row * kRowsPerTile;
|
||||
int accum_n = mma_n * InstructionShape::kN * OpDelta::kColumn + col;
|
||||
|
||||
frag[mma_accum_start + row * kElementsPerAccess + col] = offset_ref.at({accum_m, accum_n});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with additional logical offset
|
||||
CUTLASS_DEVICE
|
||||
void load_with_byte_offset(
|
||||
Fragment &frag, ///< fragment to load from the tensor
|
||||
Index byte_offset) const { ///< loads a tile with a linear offset
|
||||
|
||||
load_with_pointer_offset(byte_offset / sizeof(Element));
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void load(
|
||||
Fragment &frag, ///< fragment to load from the tensor
|
||||
TensorCoord const &tile_offset) const { ///< loads a tile with a logical offset in units of whole tiles
|
||||
|
||||
load(frag, tile_offset, 0);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void load(
|
||||
Fragment &frag, ///< fragment to load from the tensor
|
||||
TensorCoord const &tile_offset, ///< loads a tile with a logical offset in units of whole tiles
|
||||
Index pointer_offset) const { ///< loads a tile with a logical offset AND a pointer offset
|
||||
|
||||
load_with_pointer_offset(frag, ref_.offset(tile_offset) + pointer_offset);
|
||||
}
|
||||
|
||||
/// Stores a fragment to memory
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag) const {
|
||||
store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Stores a fragment to memory with additional pointer offset
|
||||
CUTLASS_DEVICE
|
||||
void store_with_pointer_offset(
|
||||
Fragment const &frag, ///< fragment to store from the tensor
|
||||
Index pointer_offset) const { ///< store a tile with a linear offset
|
||||
|
||||
TensorRef offset_ref(ref_);
|
||||
offset_ref.add_pointer_offset(pointer_offset);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_n = 0; mma_n < Policy::MmaIterations::kColumn; ++mma_n) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_m = 0; mma_m < Policy::MmaIterations::kRow; ++mma_m) {
|
||||
|
||||
int mma_accum_start = kAccumulatorRows * kElementsPerAccess *
|
||||
(mma_n * Policy::MmaIterations::kRow + mma_m);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int row = 0; row < kAccumulatorRows; ++row) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int col = 0; col < kElementsPerAccess; ++col) {
|
||||
int accum_m = mma_m * InstructionShape::kM * OpDelta::kRow +
|
||||
row * kRowsPerTile;
|
||||
int accum_n = mma_n * InstructionShape::kN * OpDelta::kColumn + col;
|
||||
int idx = mma_accum_start + row * kElementsPerAccess + col;
|
||||
|
||||
offset_ref.at({accum_m, accum_n}) = frag[idx];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Stores a fragment to memory with additional pointer offset
|
||||
CUTLASS_DEVICE
|
||||
void store_with_byte_offset(
|
||||
Fragment const &frag, ///< fragment to store from the tensor
|
||||
Index byte_offset) const { ///< store a tile with a linear offset
|
||||
|
||||
store_with_pointer_offset(byte_offset / sizeof(Element));
|
||||
}
|
||||
|
||||
/// Stores a fragment to memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void store(
|
||||
Fragment &frag, ///< fragment to store to the tensor
|
||||
TensorCoord const &tile_offset) const { ///< stores a tile with a logical offset in units of whole tiles
|
||||
|
||||
store(frag, tile_offset, 0);
|
||||
}
|
||||
|
||||
/// Stores a fragment from memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void store(
|
||||
/// fragment to store to the tensor
|
||||
Fragment const &frag,
|
||||
/// stores a tile with a logical offset in units of whole tiles
|
||||
TensorCoord const &tile_offset,
|
||||
/// stores a tile with a logical offset AND a pointer offset
|
||||
Index pointer_offset) const {
|
||||
store_with_pointer_offset(frag, ref_.offset(tile_offset) + pointer_offset);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// This tile iterator is specialized for 32-thread TensorOps. It is used to load or store
|
||||
/// accumulators from memory and is agnostic to layout. It could be faster if it assumed row-major
|
||||
/// accumulator layout.
|
||||
@@ -3289,6 +3602,9 @@ class MmaTensorOpAccumulatorTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
|
||||
@@ -123,6 +123,9 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -171,7 +174,7 @@ public:
|
||||
private:
|
||||
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_[kPointerCount];
|
||||
@@ -436,6 +439,9 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -480,7 +486,7 @@ public:
|
||||
private:
|
||||
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_;
|
||||
@@ -1526,6 +1532,9 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -1566,7 +1575,7 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
private:
|
||||
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_;
|
||||
|
||||
@@ -121,6 +121,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -166,7 +169,7 @@ public:
|
||||
private:
|
||||
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_;
|
||||
@@ -877,6 +880,9 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Long Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -919,7 +925,7 @@ public:
|
||||
private:
|
||||
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_;
|
||||
@@ -982,6 +988,16 @@ public:
|
||||
return *this;
|
||||
}
|
||||
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator &add_tile_offset_negative(TensorCoord const &tile_offset) {
|
||||
|
||||
add_tile_offset(tile_offset); // TODO fix this if it becomes an issue during warp it reset
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator & operator++() {
|
||||
@@ -1237,6 +1253,15 @@ public:
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator &add_tile_offset_negative(TensorCoord const &tile_offset) {
|
||||
|
||||
iterator_.add_tile_offset_negative({tile_offset.column(), tile_offset.row()});
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator & operator++() {
|
||||
@@ -1461,6 +1486,15 @@ public:
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator &add_tile_offset_negative(TensorCoord const &tile_offset) {
|
||||
|
||||
iterator_.add_tile_offset_negative({tile_offset.row(), tile_offset.column()});
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator & operator++() {
|
||||
|
||||
@@ -130,6 +130,9 @@ class MmaTensorOpWmmaMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Stride Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -180,7 +183,7 @@ private:
|
||||
Index byte_offset_;
|
||||
|
||||
/// Stride in units of number of elements
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Layout of shared memory
|
||||
Layout layout_;
|
||||
@@ -375,6 +378,9 @@ class MmaTensorOpWmmaMultiplicandTileIterator<
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Stride Index type
|
||||
using StrideIndex = typename TensorRef::Layout::Stride::Index;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
@@ -425,7 +431,7 @@ private:
|
||||
Index byte_offset_;
|
||||
|
||||
/// Stride in units of number of elements
|
||||
Index stride_;
|
||||
StrideIndex stride_;
|
||||
|
||||
/// Layout of shared memory
|
||||
Layout layout_;
|
||||
|
||||
@@ -109,6 +109,12 @@ public:
|
||||
/// Underlying instruction shape
|
||||
using InstructionShape = typename Policy::Operator::Shape;
|
||||
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename Policy::Operator;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
/// Underlying architecture tag
|
||||
using ArchTag = typename Policy::Operator::ArchTag;
|
||||
|
||||
|
||||
@@ -0,0 +1,405 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2021, 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 warp-level matrix multiply-accumulate operations targeting
|
||||
Tensor Cores.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/platform/platform.h"
|
||||
|
||||
#include "cutlass/numeric_conversion.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
|
||||
#include "cutlass/arch/memory_sm75.h"
|
||||
#include "cutlass/arch/mma_sm75.h"
|
||||
#include "cutlass/arch/mma_sm80.h"
|
||||
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/gemm/warp/mma.h"
|
||||
|
||||
#include "cutlass/gemm/warp/mma_tensor_op_policy.h"
|
||||
#include "cutlass/gemm/warp/mma_tensor_op.h"
|
||||
#include "cutlass/gemm/warp/mma_tensor_op_tile_iterator.h"
|
||||
#include "cutlass/gemm/warp/mma_tensor_op_tile_iterator_sm80.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace warp {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Structure to compute the matrix product targeting CUDA cores and SIMT math instructions.
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape_,
|
||||
/// Data type of A elements
|
||||
typename ElementA_,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA_,
|
||||
/// Data type of B elements
|
||||
typename ElementB_,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB_,
|
||||
/// Element type of C matrix
|
||||
typename ElementC_,
|
||||
/// Layout of C matrix (concept: MatrixLayout)
|
||||
typename LayoutC_,
|
||||
/// Policy describing warp-level MmaTensorOp (concept: MmaTensorOp policy)
|
||||
typename Policy_,
|
||||
///
|
||||
bool ReduceKForA_,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK_ = 1,
|
||||
/// Store the accumulators in row major or column major. Row major is used
|
||||
/// when output layout is interleaved.
|
||||
bool AccumulatorsInRowMajor = false,
|
||||
/// Used for partial specialization
|
||||
typename Enable = bool
|
||||
>
|
||||
class MmaWithReductionTensorOp {
|
||||
public:
|
||||
/// Shape of warp-level matrix operation (concept: GemmShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Data type of multiplicand A
|
||||
using ElementA = ElementA_;
|
||||
|
||||
/// Layout of multiplicand A
|
||||
using LayoutA = LayoutA_;
|
||||
|
||||
/// Data type of multiplicand B
|
||||
using ElementB = ElementB_;
|
||||
|
||||
/// Layout of multiplicand B
|
||||
using LayoutB = LayoutB_;
|
||||
|
||||
/// Data type of accumulator matrix C
|
||||
using ElementC = ElementC_;
|
||||
|
||||
/// Layout of accumulator matrix C
|
||||
using LayoutC = LayoutC_;
|
||||
|
||||
/// Shape of the warp in units of thread (concept: MmaLanePolicySimt)
|
||||
using Policy = Policy_;
|
||||
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename Policy::Operator;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
/// Architecture tag from underlying instruction
|
||||
using ArchTag = typename ArchMmaOperator::ArchTag;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename ArchMmaOperator::Shape;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = ComplexTransform::kNone;
|
||||
|
||||
/// Complex transform on B operand
|
||||
static ComplexTransform const kTransformB = ComplexTransform::kNone;
|
||||
|
||||
/// Number of threads participating in warp-level matrix product
|
||||
static int const kThreadCount = 32;
|
||||
|
||||
/// Number of partitions along K dimension
|
||||
static int const kPartitionsK = PartitionsK_;
|
||||
|
||||
static bool const kReduceKForA = ReduceKForA_;
|
||||
|
||||
public:
|
||||
|
||||
/// Iterates over the A operand in memory
|
||||
using IteratorA = MmaTensorOpMultiplicandTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kK>, Operand::kA, ElementA, LayoutA,
|
||||
MatrixShape<ArchMmaOperator::Shape::kM, ArchMmaOperator::Shape::kK>,
|
||||
Policy::OpDelta::kRow, kThreadCount, kPartitionsK>;
|
||||
|
||||
/// Storage for A tile
|
||||
using FragmentA = typename IteratorA::Fragment;
|
||||
|
||||
/// Storage for transformed A tile
|
||||
using TransformedFragmentA =
|
||||
Array<typename ArchMmaOperator::ElementA, FragmentA::kElements>;
|
||||
|
||||
/// Iterates over the B operand in memory
|
||||
using IteratorB = MmaTensorOpMultiplicandTileIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>, Operand::kB, ElementB, LayoutB,
|
||||
MatrixShape<ArchMmaOperator::Shape::kK, ArchMmaOperator::Shape::kN>,
|
||||
Policy::OpDelta::kRow, kThreadCount, kPartitionsK>;
|
||||
|
||||
/// Storage for B tile
|
||||
using FragmentB = typename IteratorB::Fragment;
|
||||
|
||||
/// Storage for transformed B tile
|
||||
using TransformedFragmentB =
|
||||
Array<typename ArchMmaOperator::ElementB, FragmentB::kElements>;
|
||||
|
||||
/// Iterates over the C operand in memory
|
||||
using IteratorC = MmaTensorOpAccumulatorTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kN>, ElementC, LayoutC,
|
||||
typename ArchMmaOperator::Shape, typename Policy::OpDelta>;
|
||||
|
||||
/// Storage for C tile
|
||||
using FragmentC = typename IteratorC::Fragment;
|
||||
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = MatrixShape<
|
||||
(Shape::kM + ArchMmaOperator::Shape::kM - 1) / ArchMmaOperator::Shape::kM,
|
||||
(Shape::kN + ArchMmaOperator::Shape::kN - 1) / ArchMmaOperator::Shape::kN
|
||||
>;
|
||||
|
||||
using FragmentReduction = Array<ElementC, kReduceKForA ? (Shape::kM / 8) : (Shape::kN / 8)>;
|
||||
|
||||
public:
|
||||
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
ArchMmaOperator mma;
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_DEVICE
|
||||
MmaWithReductionTensorOp() {}
|
||||
|
||||
/// Performs a warp-level matrix multiply-accumulate operation
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
FragmentC &D,
|
||||
TransformedFragmentA const &A,
|
||||
TransformedFragmentB const &B,
|
||||
FragmentC const &C,
|
||||
FragmentReduction &gemm_k_reduction
|
||||
) const {
|
||||
|
||||
using MmaOperandA = typename ArchMmaOperator::FragmentA;
|
||||
using MmaOperandB = typename ArchMmaOperator::FragmentB;
|
||||
using MmaOperandC = typename ArchMmaOperator::FragmentC;
|
||||
|
||||
D = C;
|
||||
|
||||
MmaOperandA const *ptr_A = reinterpret_cast<MmaOperandA const *>(&A);
|
||||
MmaOperandB const *ptr_B = reinterpret_cast<MmaOperandB const *>(&B);
|
||||
MmaOperandC *ptr_D = reinterpret_cast<MmaOperandC *>(&D);
|
||||
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ < 800)
|
||||
// Serpentine visitation order maximizing reuse of Rb
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < MmaIterations::kColumn; ++n) {
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int m = 0; m < MmaIterations::kRow; ++m) {
|
||||
|
||||
int m_serpentine = ((n % 2) ? (MmaIterations::kRow - 1 - m) : m);
|
||||
|
||||
if (AccumulatorsInRowMajor) { // matrix B is reordered
|
||||
mma(
|
||||
ptr_D[n + m_serpentine * MmaIterations::kColumn],
|
||||
ptr_A[m_serpentine],
|
||||
ptr_B[n],
|
||||
ptr_D[n + m_serpentine * MmaIterations::kColumn]);
|
||||
} else {
|
||||
mma(
|
||||
ptr_D[m_serpentine + n * MmaIterations::kRow],
|
||||
ptr_A[m_serpentine],
|
||||
ptr_B[n],
|
||||
ptr_D[m_serpentine + n * MmaIterations::kRow]);
|
||||
}
|
||||
}
|
||||
}
|
||||
#elif defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
|
||||
// Serpentine visitation order maximizing reuse of Ra
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int m = 0; m < MmaIterations::kRow; ++m) {
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < MmaIterations::kColumn; ++n) {
|
||||
|
||||
int n_serpentine = ((m % 2) ? (MmaIterations::kColumn - 1 - n) : n);
|
||||
|
||||
if (AccumulatorsInRowMajor) { // matrix B is reordered
|
||||
mma(
|
||||
ptr_D[n_serpentine + m * MmaIterations::kColumn],
|
||||
ptr_A[m],
|
||||
ptr_B[n_serpentine],
|
||||
ptr_D[n_serpentine + m * MmaIterations::kColumn]);
|
||||
} else {
|
||||
mma(ptr_D[m + n_serpentine * MmaIterations::kRow],
|
||||
ptr_A[m],
|
||||
ptr_B[n_serpentine],
|
||||
ptr_D[m + n_serpentine * MmaIterations::kRow]);
|
||||
|
||||
if (!kReduceKForA && m == 0) {
|
||||
// gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4]);
|
||||
// gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4 + 1]);
|
||||
// gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4 + 2]);
|
||||
// gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4 + 3]);
|
||||
|
||||
uint32_t const *tmp = reinterpret_cast<uint32_t const *>(&B);
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
" .reg .f16 low, high;\n\t"
|
||||
" .reg .f32 tmp;\n\t"
|
||||
" mov.b32 {low, high}, %1;\n\t"
|
||||
" cvt.f32.f16 tmp, low;\n\t"
|
||||
" add.f32 %0, tmp, %0;\n\t"
|
||||
" cvt.f32.f16 tmp, high;\n\t"
|
||||
" add.f32 %0, tmp, %0;\n\t"
|
||||
" mov.b32 {low, high}, %2;\n\t"
|
||||
" cvt.f32.f16 tmp, low;\n\t"
|
||||
" add.f32 %0, tmp, %0;\n\t"
|
||||
" cvt.f32.f16 tmp, high;\n\t"
|
||||
" add.f32 %0, tmp, %0;\n\t"
|
||||
"}\n\t"
|
||||
: "+f"(gemm_k_reduction[n_serpentine])
|
||||
: "r"(tmp[n_serpentine * 2]), "r"(tmp[n_serpentine * 2 + 1]));
|
||||
}
|
||||
}
|
||||
|
||||
if (kReduceKForA && (n == 0)) {
|
||||
// gemm_k_reduction[m * 2] += float(A[m * 8]);
|
||||
// gemm_k_reduction[m * 2] += float(A[m * 8 + 1]);
|
||||
// gemm_k_reduction[m * 2] += float(A[m * 8 + 4]);
|
||||
// gemm_k_reduction[m * 2] += float(A[m * 8 + 5]);
|
||||
//
|
||||
// gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 2]);
|
||||
// gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 3]);
|
||||
// gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 6]);
|
||||
// gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 7]);
|
||||
|
||||
uint32_t const *tmp = reinterpret_cast<uint32_t const *>(&A);
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
" .reg .f16 low, high;\n\t"
|
||||
" .reg .f32 tmp;\n\t"
|
||||
" mov.b32 {low, high}, %2;\n\t"
|
||||
" cvt.f32.f16 tmp, low;\n\t"
|
||||
" add.f32 %0, tmp, %0;\n\t"
|
||||
" cvt.f32.f16 tmp, high;\n\t"
|
||||
" add.f32 %0, tmp, %0;\n\t"
|
||||
" mov.b32 {low, high}, %3;\n\t"
|
||||
" cvt.f32.f16 tmp, low;\n\t"
|
||||
" add.f32 %1, tmp, %1;\n\t"
|
||||
" cvt.f32.f16 tmp, high;\n\t"
|
||||
" add.f32 %1, tmp, %1;\n\t"
|
||||
" mov.b32 {low, high}, %4;\n\t"
|
||||
" cvt.f32.f16 tmp, low;\n\t"
|
||||
" add.f32 %0, tmp, %0;\n\t"
|
||||
" cvt.f32.f16 tmp, high;\n\t"
|
||||
" add.f32 %0, tmp, %0;\n\t"
|
||||
" mov.b32 {low, high}, %5;\n\t"
|
||||
" cvt.f32.f16 tmp, low;\n\t"
|
||||
" add.f32 %1, tmp, %1;\n\t"
|
||||
" cvt.f32.f16 tmp, high;\n\t"
|
||||
" add.f32 %1, tmp, %1;\n\t"
|
||||
"}\n\t"
|
||||
: "+f"(gemm_k_reduction[m * 2]), "+f"(gemm_k_reduction[m * 2 + 1])
|
||||
: "r"(tmp[m * 4]), "r"(tmp[m * 4 + 1]),"r"(tmp[m * 4 + 2]), "r"(tmp[m * 4 + 3]));
|
||||
}
|
||||
}
|
||||
}
|
||||
#else
|
||||
assert(0);
|
||||
#endif
|
||||
}
|
||||
|
||||
/// Transform the mma operands to the required types
|
||||
CUTLASS_DEVICE
|
||||
void transform(TransformedFragmentA &dst_A, TransformedFragmentB &dst_B,
|
||||
FragmentA const &A, FragmentB const &B) const {
|
||||
|
||||
//
|
||||
// Define conversions from source type to instruction type
|
||||
//
|
||||
FloatRoundStyle const kRoundA =
|
||||
PreferredRoundingMode<typename ArchMmaOperator::ElementA,
|
||||
ElementA>::kRound;
|
||||
FloatRoundStyle const kRoundB =
|
||||
PreferredRoundingMode<typename ArchMmaOperator::ElementB,
|
||||
ElementB>::kRound;
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ < 800)
|
||||
detail::ConvertAndPack<typename ArchMmaOperator::ElementA, ElementA,
|
||||
FragmentA::kElements, kRoundA>
|
||||
convert_A;
|
||||
NumericArrayConverter<typename ArchMmaOperator::ElementB, ElementB,
|
||||
FragmentB::kElements / 2, kRoundB>
|
||||
convert_B;
|
||||
Array<ElementB, FragmentB::kElements / 2> const *ptr_B =
|
||||
reinterpret_cast<Array<ElementB, FragmentB::kElements / 2> const *>(&B);
|
||||
Array<typename ArchMmaOperator::ElementB, FragmentB::kElements / 2> *
|
||||
ptr_dst_B = reinterpret_cast<Array<typename ArchMmaOperator::ElementB,
|
||||
FragmentB::kElements / 2> *>(&dst_B);
|
||||
|
||||
dst_A = convert_A(A);
|
||||
|
||||
ptr_dst_B[0] = convert_B(ptr_B[0]);
|
||||
ptr_dst_B[1] = convert_B(ptr_B[1]);
|
||||
|
||||
#elif defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
|
||||
detail::ConvertAndPack<typename ArchMmaOperator::ElementA, ElementA,
|
||||
FragmentA::kElements / 2, kRoundA>
|
||||
convert_A;
|
||||
NumericArrayConverter<typename ArchMmaOperator::ElementB, ElementB,
|
||||
FragmentB::kElements, kRoundB>
|
||||
convert_B;
|
||||
Array<ElementA, FragmentA::kElements / 2> const *ptr_A =
|
||||
reinterpret_cast<Array<ElementA, FragmentA::kElements / 2> const *>(&A);
|
||||
Array<typename ArchMmaOperator::ElementA, FragmentA::kElements / 2> *
|
||||
ptr_dst_A = reinterpret_cast<Array<typename ArchMmaOperator::ElementA,
|
||||
FragmentA::kElements / 2> *>(&dst_A);
|
||||
|
||||
dst_B = convert_B(B);
|
||||
|
||||
ptr_dst_A[0] = convert_A(ptr_A[0]);
|
||||
ptr_dst_A[1] = convert_A(ptr_A[1]);
|
||||
#else
|
||||
assert(0);
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
Reference in New Issue
Block a user