CUTLASS 2.1 (#83)
CUTLASS 2.1 contributes: - BLAS-style host-side API added to CUTLASS Library - Planar Complex GEMM kernels targeting Volta and Turing Tensor Cores - Minor enhancements and bug fixes
This commit is contained in:
@@ -114,3 +114,6 @@ struct DefaultMmaTensorOp {
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
@@ -1,351 +0,0 @@
|
||||
/***************************************************************************************************
|
||||
* 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 warp-level matrix multiply-accumulate operations targeting
|
||||
Tensor Cores.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/complex.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/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"
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace warp {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape_,
|
||||
/// Data type of A elements
|
||||
typename RealElementA,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA_,
|
||||
/// Data type of B elements
|
||||
typename RealElementB,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB_,
|
||||
/// Element type of C matrix
|
||||
typename RealElementC,
|
||||
/// Layout of C matrix (concept: MatrixLayout)
|
||||
typename LayoutC_,
|
||||
/// Policy describing warp-level MmaTensorOp (concept: MmaTensorOp policy)
|
||||
typename Policy_,
|
||||
/// Complex transform on A operand
|
||||
ComplexTransform TransformA = ComplexTransform::kNone,
|
||||
/// Complex transform on B operand
|
||||
ComplexTransform TransformB = ComplexTransform::kNone,
|
||||
/// Used for partial specialization
|
||||
typename Enable = bool
|
||||
>
|
||||
class MmaComplexTensorOp;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for complex*complex+complex => complex using real-valued TensorOps
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape_,
|
||||
/// Data type of A elements
|
||||
typename RealElementA,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA_,
|
||||
/// Data type of B elements
|
||||
typename RealElementB,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB_,
|
||||
/// Element type of C matrix
|
||||
typename RealElementC,
|
||||
/// Layout of C matrix (concept: MatrixLayout)
|
||||
typename LayoutC_,
|
||||
/// Policy describing warp-level MmaTensorOp (concept: MmaTensorOp policy)
|
||||
typename Policy_,
|
||||
/// Complex transform on A operand
|
||||
ComplexTransform TransformA,
|
||||
/// Complex transform on B operand
|
||||
ComplexTransform TransformB,
|
||||
/// Used for partial specialization
|
||||
typename Enable
|
||||
>
|
||||
class MmaComplexTensorOp<
|
||||
Shape_,
|
||||
complex<RealElementA>,
|
||||
LayoutA_,
|
||||
complex<RealElementB>,
|
||||
LayoutB_,
|
||||
complex<RealElementC>,
|
||||
LayoutC_,
|
||||
Policy_,
|
||||
TransformA,
|
||||
TransformB,
|
||||
Enable> {
|
||||
public:
|
||||
/// Shape of warp-level matrix operation (concept: GemmShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Data type of multiplicand A
|
||||
using ElementA = complex<RealElementA>;
|
||||
|
||||
/// Layout of multiplicand A
|
||||
using LayoutA = LayoutA_;
|
||||
|
||||
/// Data type of multiplicand B
|
||||
using ElementB = complex<RealElementB>;
|
||||
|
||||
/// Layout of multiplicand B
|
||||
using LayoutB = LayoutB_;
|
||||
|
||||
/// Data type of accumulator matrix C
|
||||
using ElementC = complex<RealElementC>;
|
||||
|
||||
/// Layout of accumulator matrix C
|
||||
using LayoutC = LayoutC_;
|
||||
|
||||
/// Shape of the warp in units of thread (concept: MmaLanePolicySimt)
|
||||
using Policy = Policy_;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
|
||||
/// Complex transform on B operand
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Number of threads participating in warp-level matrix product
|
||||
static int const kThreadCount = 32;
|
||||
|
||||
public:
|
||||
|
||||
/// Iterates over the A operand in memory
|
||||
using IteratorA = MmaTensorOpMultiplicandTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kK>,
|
||||
Operand::kA,
|
||||
ElementA,
|
||||
LayoutA,
|
||||
MatrixShape<Policy::Operator::Shape::kM, Policy::Operator::Shape::kK>,
|
||||
Policy::OpDelta::kRow,
|
||||
32,
|
||||
1
|
||||
>;
|
||||
|
||||
/// Storage for A tile
|
||||
using FragmentA = typename IteratorA::Fragment;
|
||||
|
||||
/// Iterates over the B operand in memory
|
||||
using IteratorB = MmaTensorOpMultiplicandTileIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>,
|
||||
Operand::kB,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
MatrixShape<Policy::Operator::Shape::kK, Policy::Operator::Shape::kN>,
|
||||
Policy::OpDelta::kColumn,
|
||||
32,
|
||||
1
|
||||
>;
|
||||
|
||||
/// Storage for B tile
|
||||
using FragmentB = typename IteratorB::Fragment;
|
||||
|
||||
|
||||
static_assert(
|
||||
!(Shape::kM % Policy::Operator::Shape::kM) &&
|
||||
!(Shape::kN % Policy::Operator::Shape::kN),
|
||||
"Shape of warp-level Mma must be divisible by operator shape.");
|
||||
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = MatrixShape<
|
||||
Shape::kM / Policy::Operator::Shape::kM,
|
||||
Shape::kN / Policy::Operator::Shape::kN
|
||||
>;
|
||||
|
||||
/// Iterates over the C operand in memory
|
||||
using IteratorC = MmaTensorOpAccumulatorTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kN>,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
typename Policy::Operator::Shape,
|
||||
typename Policy::OpDelta>;
|
||||
|
||||
/// Storage for C tile, the accumulator. Note, regardless of multiplicand type, this
|
||||
/// storage arrangement is to be considered 'planar complex' in the sense that all real-valued
|
||||
/// parts are stored consecutively followed by all imaginary parts. This matches the structure
|
||||
/// of Tensor Cores which are always real-valued matrix multiplies.
|
||||
using FragmentC = typename IteratorC::Fragment;
|
||||
|
||||
static_assert(
|
||||
FragmentC::kElements == 2 * MmaIterations::kCount * Policy::Operator::FragmentC::kElements,
|
||||
"Unexpected planar complex fragment length.");
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Underlying real-valued matrix multiply operator (concept: arch::Mma)
|
||||
typename Policy::Operator mma;
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_DEVICE
|
||||
MmaComplexTensorOp() {}
|
||||
|
||||
/// Performs a warp-level matrix multiply-accumulate operation
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
FragmentC &D,
|
||||
FragmentA const &A,
|
||||
FragmentB const &B,
|
||||
FragmentC const &C) const {
|
||||
|
||||
// Alias types for underlying real-valued matrix multiply operator
|
||||
using MmaOperandA = typename Policy::Operator::FragmentA;
|
||||
using MmaOperandB = typename Policy::Operator::FragmentB;
|
||||
using MmaOperandC = typename Policy::Operator::FragmentC;
|
||||
|
||||
static_assert(MmaOperandA::kElements == 1,
|
||||
"This implementation only supports math instructions in which exactly one element is needed for the A operand."
|
||||
"We can geneneralize later.");
|
||||
|
||||
static_assert(MmaOperandB::kElements == 1,
|
||||
"This implementation only supports math instructions in which exactly one element is needed for the A operand."
|
||||
"We can geneneralize later.");
|
||||
|
||||
D = C;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int m = 0; m < MmaIterations::kRow; ++m) {
|
||||
|
||||
// mma(accum.real(), a.real(), b.real(), accum.real());
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < MmaIterations::kColumn; ++n) {
|
||||
|
||||
// Pack operands together. This may result in actual MOVs
|
||||
MmaOperandA operand_A;
|
||||
MmaOperandB operand_B;
|
||||
|
||||
operand_A[0] = A[m].real();
|
||||
operand_B[0] = B[n].real();
|
||||
|
||||
// Real-valued accumulator part
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow);
|
||||
|
||||
mma(*accum, operand_A, operand_B, *accum);
|
||||
}
|
||||
|
||||
// mma(accum.imag(), a.real(), b.imag(), accum.imag());
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = MmaIterations::kColumn - 1; n >= 0; --n) {
|
||||
|
||||
// Pack operands together. This may result in actual MOVs
|
||||
MmaOperandA operand_A;
|
||||
MmaOperandB operand_B;
|
||||
|
||||
operand_A[0] = A[m].real();
|
||||
operand_B[0] = (kTransformB == ComplexTransform::kConjugate ? -B[n].imag() : B[n].imag());
|
||||
|
||||
// Complex-valued accumulator part
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow) + MmaIterations::kCount;
|
||||
|
||||
mma(*accum, operand_A, operand_B, *accum);
|
||||
}
|
||||
|
||||
// mma(accum.real(), -a.imag(), b.imag(), accum.real())
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < MmaIterations::kColumn; ++n) {
|
||||
|
||||
// Pack operands together. This may result in actual MOVs
|
||||
MmaOperandA operand_A;
|
||||
MmaOperandB operand_B;
|
||||
|
||||
// A imaginary part is intentionally negated
|
||||
operand_A[0] = (kTransformA == ComplexTransform::kConjugate ? A[m].imag() : -A[m].imag());
|
||||
operand_B[0] = (kTransformB == ComplexTransform::kConjugate ? -B[n].imag() : B[n].imag());
|
||||
|
||||
// Complex-valued accumulator part
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow);
|
||||
|
||||
mma(*accum, operand_A, operand_B, *accum);
|
||||
}
|
||||
|
||||
// mma(accum.imag(), a.imag(), b.real(), accum.imag())
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = MmaIterations::kColumn - 1; n >= 0; --n) {
|
||||
|
||||
// Pack operands together. This may result in actual MOVs
|
||||
MmaOperandA operand_A;
|
||||
MmaOperandB operand_B;
|
||||
|
||||
operand_A[0] = (kTransformA == ComplexTransform::kConjugate ? -A[m].imag() : A[m].imag());
|
||||
operand_B[0] = B[n].real();
|
||||
|
||||
// Real-valued accumulator part
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow) + MmaIterations::kCount;
|
||||
|
||||
mma(*accum, operand_A, operand_B, *accum);
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// TODO - partial specializations of real*complex and complex*real
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,176 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2020, 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.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/complex.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
|
||||
#include "cutlass/array_planar_complex.h"
|
||||
#include "cutlass/gemm/warp/tile_iterator_planar_complex.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace warp {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Underlying real-valued warp-level matrix multiply
|
||||
typename Operator_,
|
||||
/// Transformation applied to A operand (typically folded into math instruction)
|
||||
ComplexTransform TransformA = ComplexTransform::kNone,
|
||||
/// Transformation applied to B operand (typically folded into math instruction)
|
||||
ComplexTransform TransformB = ComplexTransform::kNone
|
||||
>
|
||||
class MmaPlanarComplex {
|
||||
public:
|
||||
|
||||
/// Underlying real-valued warp-level matrix multiply
|
||||
using Operator = Operator_;
|
||||
|
||||
/// Shape of warp-level matrix multipy
|
||||
using Shape = typename Operator::Shape;
|
||||
|
||||
/// Transformation applied to A operand (typically folded into math instruction)
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
|
||||
/// Transformation applied to B operand (typically folded into math instruction)
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
|
||||
/// Fragment of elements
|
||||
using FragmentA = ArrayPlanarComplex<typename Operator::ElementA, Operator::FragmentA::kElements>;
|
||||
|
||||
/// Iterator into planar complex
|
||||
using IteratorA = TileIteratorPlanarComplex<typename Operator::IteratorA>;
|
||||
|
||||
/// Layout in memory of the A operand
|
||||
using LayoutA = typename Operator::LayoutA;
|
||||
|
||||
using FragmentB = ArrayPlanarComplex<typename Operator::ElementB, Operator::FragmentB::kElements>;
|
||||
|
||||
/// Iterator into planar complex
|
||||
using IteratorB = TileIteratorPlanarComplex<typename Operator::IteratorB>;
|
||||
|
||||
/// Layout in memory of the B operand
|
||||
using LayoutB = typename Operator::LayoutB;
|
||||
|
||||
/// Tile iterator for accumulator
|
||||
using IteratorC = TileIteratorPlanarComplex<typename Operator::IteratorC>;
|
||||
|
||||
/// Accumulator fragment
|
||||
using FragmentC = ArrayPlanarComplex<typename Operator::ElementC, Operator::FragmentC::kElements>;
|
||||
|
||||
/// Layout of accumulator fragment in memory
|
||||
using LayoutC = typename Operator::LayoutC;
|
||||
|
||||
private:
|
||||
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = MatrixShape<
|
||||
Operator::Shape::kM / Operator::Policy::Operator::Shape::kM,
|
||||
Operator::Shape::kN / Operator::Policy::Operator::Shape::kN
|
||||
>;
|
||||
|
||||
public:
|
||||
/// Ctor
|
||||
CUTLASS_DEVICE
|
||||
MmaPlanarComplex() {}
|
||||
|
||||
/// Performs a warp-level matrix multiply-accumulate operation
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
FragmentC &D,
|
||||
FragmentA const &A_in,
|
||||
FragmentB const &B_in,
|
||||
FragmentC const &C) const {
|
||||
|
||||
D.real = C.real;
|
||||
D.imag = C.imag;
|
||||
|
||||
//
|
||||
// Transform fragments based on conjugate operations.
|
||||
//
|
||||
|
||||
negate<typename FragmentA::ArrayReal> neg_A;
|
||||
|
||||
FragmentA frag_A;
|
||||
frag_A.real = A_in.real;
|
||||
|
||||
if (kTransformA == ComplexTransform::kConjugate) {
|
||||
frag_A.imag = neg_A(frag_A.imag);
|
||||
}
|
||||
else {
|
||||
frag_A.imag = frag_A.imag;
|
||||
}
|
||||
|
||||
FragmentB frag_B;
|
||||
frag_B.real = B_in.real;
|
||||
|
||||
if (kTransformB == ComplexTransform::kConjugate) {
|
||||
negate<typename FragmentB::ArrayReal> neg;
|
||||
frag_B.imag = neg(frag_B.imag);
|
||||
}
|
||||
else {
|
||||
frag_B.imag = frag_B.imag;
|
||||
}
|
||||
|
||||
//
|
||||
// Accumulated real-valued matrix multiplies
|
||||
//
|
||||
|
||||
Operator real_mma;
|
||||
|
||||
// D.i += A.i * B.r
|
||||
real_mma(D.imag, frag_A.imag, frag_B.real, D.imag);
|
||||
|
||||
// D.r += A.r * B.r
|
||||
real_mma(D.real, frag_A.real, frag_B.real, D.real);
|
||||
|
||||
// D.i += A.r * B.i
|
||||
real_mma(D.imag, frag_A.real, frag_B.imag, D.imag);
|
||||
|
||||
// D.r += -A.i * B.i
|
||||
frag_A.imag = neg_A(frag_A.imag);
|
||||
real_mma(D.real, frag_A.imag, frag_B.imag, D.real);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -100,6 +100,16 @@ public:
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassSimt;
|
||||
|
||||
/// Hard-coded for now
|
||||
using ArchTag = arch::Sm50;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = ComplexTransform::kNone;
|
||||
|
||||
/// Complex transform on B operand
|
||||
static ComplexTransform const kTransformB = ComplexTransform::kNone;
|
||||
|
||||
/// Layout of threads
|
||||
using ThreadLayoutA = typename platform::conditional< platform::is_same< layout::ColumnMajorInterleaved<4>, LayoutA >::value,
|
||||
layout::ColumnMajor,
|
||||
typename platform::conditional < platform::is_same< layout::RowMajorInterleaved<4>, LayoutA >::value,
|
||||
@@ -153,6 +163,9 @@ public:
|
||||
/// Storage for A tile
|
||||
using FragmentA = typename IteratorA::Fragment;
|
||||
|
||||
/// Storage for transformed A tile
|
||||
using TransformedFragmentA = FragmentA;
|
||||
|
||||
/// Iterates over the B operand in memory
|
||||
using IteratorB = MmaSimtTileIterator<
|
||||
MatrixShape<Policy::LaneMmaShape::kK, Shape::kN>,
|
||||
@@ -167,6 +180,9 @@ public:
|
||||
/// Storage for B tile
|
||||
using FragmentB = typename IteratorB::Fragment;
|
||||
|
||||
/// Storage for transformed A tile
|
||||
using TransformedFragmentB = FragmentB;
|
||||
|
||||
/// Iterates over the C operand in memory
|
||||
using IteratorC = MmaSimtTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kN>,
|
||||
@@ -201,6 +217,15 @@ public:
|
||||
|
||||
mma(d, a, b, c);
|
||||
}
|
||||
|
||||
/// 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 {
|
||||
//TODO: Implement this
|
||||
dst_A = A;
|
||||
dst_B = B;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -31,7 +31,9 @@
|
||||
|
||||
#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"
|
||||
|
||||
@@ -51,6 +53,60 @@ namespace warp {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
|
||||
template <typename T, typename S, int N, FloatRoundStyle Round>
|
||||
struct ConvertAndPack {
|
||||
|
||||
using Converter = NumericArrayConverter<T, S, N, Round>;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> operator()(Array<S, N> const &source) {
|
||||
Converter converter;
|
||||
|
||||
return converter(source);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T, int N, FloatRoundStyle Round>
|
||||
struct ConvertAndPack<T, T, N, Round> {
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> operator()(Array<T, N> const &source) {
|
||||
return source;
|
||||
}
|
||||
};
|
||||
|
||||
template <int N, FloatRoundStyle Round>
|
||||
struct ConvertAndPack<half_t, float, N, Round> {
|
||||
|
||||
using Converter = NumericArrayConverter<half_t, float, N, Round>;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<half_t, N> operator()(Array<float, N> const &source) {
|
||||
Converter converter;
|
||||
|
||||
Array<float, N> tmp;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < N; ++i) {
|
||||
int idx = (((i << 1) & 2) | ((i >> 1) & 1) | (i & 0xfffffffc));
|
||||
tmp[i] = source[idx];
|
||||
}
|
||||
|
||||
return converter(tmp);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Structure to compute the matrix product targeting CUDA cores and SIMT math instructions.
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
@@ -105,9 +161,18 @@ public:
|
||||
/// Shape of the warp in units of thread (concept: MmaLanePolicySimt)
|
||||
using Policy = Policy_;
|
||||
|
||||
/// Architecture tag from underlying instruction
|
||||
using ArchTag = typename Policy::Operator::ArchTag;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// 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;
|
||||
|
||||
@@ -128,6 +193,10 @@ public:
|
||||
/// Storage for A tile
|
||||
using FragmentA = typename IteratorA::Fragment;
|
||||
|
||||
/// Storage for transformed A tile
|
||||
using TransformedFragmentA =
|
||||
Array<typename Policy::Operator::ElementA, FragmentA::kElements>;
|
||||
|
||||
/// Iterates over the B operand in memory
|
||||
using IteratorB = MmaTensorOpMultiplicandTileIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>, Operand::kB, ElementB, LayoutB,
|
||||
@@ -137,6 +206,10 @@ public:
|
||||
/// Storage for B tile
|
||||
using FragmentB = typename IteratorB::Fragment;
|
||||
|
||||
/// Storage for transformed B tile
|
||||
using TransformedFragmentB =
|
||||
Array<typename Policy::Operator::ElementB, FragmentB::kElements>;
|
||||
|
||||
/// Iterates over the C operand in memory
|
||||
using IteratorC = MmaTensorOpAccumulatorTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kN>, ElementC, LayoutC,
|
||||
@@ -179,8 +252,8 @@ public:
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
FragmentC &D,
|
||||
FragmentA const &A,
|
||||
FragmentB const &B,
|
||||
TransformedFragmentA const &A,
|
||||
TransformedFragmentB const &B,
|
||||
FragmentC const &C,
|
||||
int const &partitionN_idx = 0) const {
|
||||
|
||||
@@ -221,6 +294,44 @@ public:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 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 {
|
||||
bool midway_depstage =
|
||||
!(platform::is_same<typename Policy::Operator::ElementA,
|
||||
ElementA>::value &&
|
||||
platform::is_same<typename Policy::Operator::ElementB,
|
||||
ElementB>::value);
|
||||
|
||||
//
|
||||
// Define conversions from source type to instruction type
|
||||
//
|
||||
FloatRoundStyle const kRoundA =
|
||||
PreferredRoundingMode<typename Policy::Operator::ElementA,
|
||||
ElementA>::kRound;
|
||||
FloatRoundStyle const kRoundB =
|
||||
PreferredRoundingMode<typename Policy::Operator::ElementB,
|
||||
ElementB>::kRound;
|
||||
detail::ConvertAndPack<typename Policy::Operator::ElementA, ElementA,
|
||||
FragmentA::kElements, kRoundA>
|
||||
convert_A;
|
||||
NumericArrayConverter<typename Policy::Operator::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 Policy::Operator::ElementB, FragmentB::kElements / 2> *
|
||||
ptr_dst_B = reinterpret_cast<Array<typename Policy::Operator::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]);
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -228,3 +339,5 @@ public:
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -103,6 +103,15 @@ public:
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Architecture tag
|
||||
using ArchTag = arch::Sm70;
|
||||
|
||||
/// 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;
|
||||
|
||||
|
||||
@@ -199,7 +199,8 @@ public:
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads>;
|
||||
using Fragment =
|
||||
Array<Element, Shape::kContiguous * InstructionShape::kStrided / kThreads>;
|
||||
|
||||
private:
|
||||
|
||||
@@ -516,7 +517,7 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads>;
|
||||
using Fragment = typename Base::Fragment;
|
||||
|
||||
private:
|
||||
|
||||
@@ -747,7 +748,7 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads>;
|
||||
using Fragment = typename Base::Fragment;
|
||||
|
||||
private:
|
||||
|
||||
@@ -1023,7 +1024,8 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads>;
|
||||
using Fragment = Array<Element, Shape::kStrided *
|
||||
InstructionShape::kContiguous / kThreads>;
|
||||
|
||||
private:
|
||||
|
||||
@@ -1151,7 +1153,8 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
int k_groups_delta = tile_offset.contiguous() % Policy::kGroupsPerTile;
|
||||
|
||||
byte_offset_ ^= k_groups_delta * sizeof_bits<Element>::value *
|
||||
Layout::kElementsPerAccess / 8;
|
||||
Layout::kElementsPerAccess *
|
||||
Policy::LdsmShape::kContiguous / 8;
|
||||
pointer_ +=
|
||||
tile_offset.strided() * stride_ * Shape::kStrided / Layout::kFactor +
|
||||
whole_tiles * stride_ / sections_;
|
||||
@@ -1406,7 +1409,7 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads>;
|
||||
using Fragment = typename Base::Fragment;
|
||||
|
||||
private:
|
||||
/// Underlying tile iterator
|
||||
@@ -1636,7 +1639,7 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads>;
|
||||
using Fragment = typename Base::Fragment;
|
||||
|
||||
private:
|
||||
/// Underlying tile iterator
|
||||
|
||||
@@ -165,7 +165,8 @@ public:
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads * 2>;
|
||||
using Fragment = Array<Element, Shape::kContiguous *
|
||||
InstructionShape::kStrided / kThreads * 2>;
|
||||
|
||||
private:
|
||||
|
||||
@@ -473,7 +474,8 @@ public:
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile, needs on more time number of registers
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads * 2>;
|
||||
using Fragment = Array<Element, Shape::kContiguous *
|
||||
InstructionShape::kStrided / kThreads * 2>;
|
||||
|
||||
private:
|
||||
|
||||
@@ -738,7 +740,7 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads * 2>;
|
||||
using Fragment = typename Base::Fragment;
|
||||
|
||||
private:
|
||||
|
||||
@@ -962,7 +964,7 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads * 2>;
|
||||
using Fragment = typename Base::Fragment;
|
||||
|
||||
private:
|
||||
|
||||
@@ -1557,7 +1559,9 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads * 2>;
|
||||
using Fragment =
|
||||
Array<Element,
|
||||
Shape::kStrided * InstructionShape::kContiguous / kThreads * 2>;
|
||||
|
||||
private:
|
||||
|
||||
@@ -1869,7 +1873,7 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads * 2>;
|
||||
using Fragment = typename Base::Fragment;
|
||||
|
||||
private:
|
||||
/// Underlying tile iterator
|
||||
@@ -2097,7 +2101,7 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads * 2>;
|
||||
using Fragment = typename Base::Fragment;
|
||||
|
||||
private:
|
||||
/// Underlying tile iterator
|
||||
|
||||
@@ -106,6 +106,15 @@ public:
|
||||
/// Shape of the warp in units of thread (concept: MmaTensorOpPolicy)
|
||||
using Policy = Policy_;
|
||||
|
||||
/// Underlying architecture tag
|
||||
using ArchTag = typename Policy::Operator::ArchTag;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = ComplexTransform::kNone;
|
||||
|
||||
/// Complex transform on B operand
|
||||
static ComplexTransform const kTransformB = ComplexTransform::kNone;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
@@ -193,7 +202,6 @@ public:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -0,0 +1,244 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2020, 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.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
|
||||
#include "cutlass/array_planar_complex.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace warp {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename TileIterator_>
|
||||
class TileIteratorPlanarComplex {
|
||||
public:
|
||||
|
||||
/// Underlying iterator over real-valued tiles
|
||||
using TileIterator = TileIterator_;
|
||||
|
||||
/// Underlying element type
|
||||
using Element = typename TileIterator::Element;
|
||||
|
||||
/// Underlying layout type
|
||||
using Layout = typename TileIterator::Layout;
|
||||
|
||||
/// TensorRef type for loading element from a tensor
|
||||
using TensorRef = typename TileIterator::TensorRef;
|
||||
|
||||
/// 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;
|
||||
|
||||
/// Planar complex fragment
|
||||
using Fragment = ArrayPlanarComplex<Element, TileIterator::Fragment::kElements>;
|
||||
|
||||
public:
|
||||
|
||||
/// Underlying tile iterator
|
||||
TileIterator tile_iterator_;
|
||||
|
||||
/// Offset (in units of bytes) to the imaginary part of the planar complex matrix
|
||||
LongIndex imaginary_offset_;
|
||||
|
||||
public:
|
||||
/// Default ctor constructs null iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
TileIteratorPlanarComplex(): imaginary_offset_(0) { }
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_DEVICE
|
||||
TileIteratorPlanarComplex(
|
||||
TensorRef const &ref,
|
||||
int lane_id,
|
||||
LongIndex imaginary_offset
|
||||
):
|
||||
tile_iterator_(ref, lane_id),
|
||||
imaginary_offset_((imaginary_offset * sizeof_bits<Element>::value) / 8) { }
|
||||
|
||||
|
||||
/// Adds a pointer offset to internal pointer(s) to advance through memory
|
||||
CUTLASS_DEVICE
|
||||
TileIteratorPlanarComplex &add_pointer_offset(LongIndex offset) {
|
||||
|
||||
tile_iterator_.add_pointer_offset(offset);
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
TileIteratorPlanarComplex &add_tile_offset(TensorCoord const &tile_offset) {
|
||||
|
||||
tile_iterator_.add_tile_offset(tile_offset);
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_DEVICE
|
||||
TileIteratorPlanarComplex & operator++() {
|
||||
++tile_iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
//
|
||||
// WIP
|
||||
//
|
||||
|
||||
/// Advances the iterator along the opposite of the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
TileIteratorPlanarComplex & operator--() {
|
||||
--tile_iterator_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
///< advances in units of whole tiles along the logical coordinate space of the tensor
|
||||
CUTLASS_DEVICE
|
||||
TileIteratorPlanarComplex & operator+=(TensorCoord const &tile_offset) {
|
||||
tile_iterator_.add_tile_offset(tile_offset);
|
||||
return *this;
|
||||
}
|
||||
|
||||
///< advances in units of whole tiles along the logical coordinate space of the tensor
|
||||
CUTLASS_DEVICE
|
||||
TileIteratorPlanarComplex & operator-=(TensorCoord const &tile_offset) {
|
||||
tile_iterator_.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 {
|
||||
|
||||
tile_iterator_.load_with_byte_offset(frag.real, 0);
|
||||
tile_iterator_.load_with_byte_offset(frag.imag, imaginary_offset_);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with additional logical offset
|
||||
CUTLASS_DEVICE
|
||||
void load_with_byte_offset(
|
||||
/// fragment to load from the tensor
|
||||
Fragment &frag,
|
||||
/// loads a tile with a linear offset in units of bytes
|
||||
Index byte_offset) const {
|
||||
|
||||
tile_iterator_.load_with_byte_offset(frag.real, byte_offset);
|
||||
tile_iterator_.load_with_byte_offset(frag.imag, byte_offset + imaginary_offset_);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with additional logical offset
|
||||
CUTLASS_DEVICE
|
||||
void load_with_pointer_offset(
|
||||
/// fragment to load from the tensor
|
||||
Fragment &frag,
|
||||
/// loads a tile with a linear offset
|
||||
Index pointer_offset) const {
|
||||
|
||||
Index byte_offset = (pointer_offset * sizeof_bits<Element>::value)/8;
|
||||
|
||||
tile_iterator_.load_with_byte_offset(frag.real, byte_offset);
|
||||
tile_iterator_.load_with_byte_offset(frag.imag, byte_offset + imaginary_offset_);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void load(
|
||||
/// fragment to load from the tensor
|
||||
Fragment &frag,
|
||||
/// loads a tile with a logical offset in units of whole tiles
|
||||
TensorCoord const &tile_offset) const {
|
||||
|
||||
tile_iterator_.load_with_byte_offset(frag.real, tile_offset, 0);
|
||||
tile_iterator_.load_with_byte_offset(frag.imag, tile_offset, imaginary_offset_);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void load(
|
||||
/// fragment to load from the tensor
|
||||
Fragment &frag,
|
||||
/// loads a tile with a logical offset in units of whole tiles
|
||||
TensorCoord const &tile_offset,
|
||||
/// loads a tile with a logical offset AND a pointer offset
|
||||
Index pointer_offset) const {
|
||||
|
||||
Index byte_offset = (pointer_offset * sizeof_bits<Element>::value)/8;
|
||||
|
||||
tile_iterator_.load_with_byte_offset(frag.real, tile_offset, byte_offset);
|
||||
tile_iterator_.load_with_byte_offset(frag.real, tile_offset, byte_offset + imaginary_offset_);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void load_with_byte_offset(
|
||||
/// fragment to load from the tensor
|
||||
Fragment &frag,
|
||||
/// loads a tile with a logical offset in units of whole tiles
|
||||
TensorCoord const &tile_offset,
|
||||
/// loads a tile with a logical offset AND a pointer offset
|
||||
Index byte_offset) const {
|
||||
|
||||
tile_iterator_.load_with_byte_offset(frag.real, tile_offset, byte_offset);
|
||||
tile_iterator_.load_with_byte_offset(frag.imag, tile_offset, byte_offset + imaginary_offset_);
|
||||
}
|
||||
|
||||
/// Notify the iterator which k-group it is currently pointing to.
|
||||
///
|
||||
/// This does not advance the iterator. Rather, it overrides its internal
|
||||
/// tracking with constant-valued k-group index to enable the compiler to
|
||||
/// fold constants and achieve more efficient code.
|
||||
///
|
||||
/// This is used by some nontrivial permuted layouts.
|
||||
CUTLASS_DEVICE
|
||||
void set_kgroup_index(int k_group) {
|
||||
tile_iterator_.set_kgroup_index(k_group);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
Reference in New Issue
Block a user