releaase 2.11 (#703)
This commit is contained in:
@@ -214,7 +214,7 @@ struct DefaultMmaComplexTensorOp<
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Partial specialization - input and output types are complex<float>*complex<float>
|
||||
// Use TF32 tensor operation internally
|
||||
// 4 real-valued MMA.1688.F32.TF32 operations on TF32
|
||||
// 4 real-valued mma.sync.aligned.m16n8k8.f32.tf32.tf32.f32 operations on TF32
|
||||
// A = (ar + j ai), B (br +j bi), D = AB
|
||||
// D = dr + j di = (ar*br - ai*bi) + j (ar*bi + ai*br)
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -246,7 +246,7 @@ struct DefaultMmaComplexTensorOp<
|
||||
TransformB,
|
||||
arch::OpMultiplyAddComplex> {
|
||||
|
||||
// Complex floating point tensor operation use MMA.1688.F32.TF32 mma instruction
|
||||
// Complex floating point tensor operation use mma.sync.aligned.m16n8k8.f32.tf32.tf32.f32 mma instruction
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
InstructionShape_,
|
||||
@@ -278,7 +278,7 @@ struct DefaultMmaComplexTensorOp<
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Partial specialization - input and output types are complex<float>*complex<float>
|
||||
// Use BF16 tensor operation internally
|
||||
// 4 real-valued MMA.1688.F32.BF16 operations on BF16
|
||||
// 4 real-valued mma.sync.aligned.m16n8k8.f32.bf16.bf16.f32 operations on BF16
|
||||
// A = (ar + j ai), B (br +j bi), D = AB
|
||||
// D = dr + j di = (ar*br - ai*bi) + j (ar*bi + ai*br)
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -310,7 +310,7 @@ struct DefaultMmaComplexTensorOp<
|
||||
TransformB,
|
||||
arch::OpMultiplyAddFastBF16> {
|
||||
|
||||
// Complex floating point tensor operation use MMA.1688.F32.BF16 mma instruction
|
||||
// Complex floating point tensor operation use mma.sync.aligned.m16n8k8.f32.bf16.bf16.f32 mma instruction
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
InstructionShape_,
|
||||
@@ -342,7 +342,7 @@ struct DefaultMmaComplexTensorOp<
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Partial specialization - input and output types are complex<float>*complex<float>
|
||||
// Use F16 tensor operation internally
|
||||
// 4 real-valued MMA.1688.F32.F16 operations on F16
|
||||
// 4 real-valued mma.sync.aligned.m16n8k8.f32.f16.f16.f32 operations on F16
|
||||
// A = (ar + j ai), B (br +j bi), D = AB
|
||||
// D = dr + j di = (ar*br - ai*bi) + j (ar*bi + ai*br)
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -374,7 +374,7 @@ struct DefaultMmaComplexTensorOp<
|
||||
TransformB,
|
||||
arch::OpMultiplyAddFastF16> {
|
||||
|
||||
// Complex floating point tensor operation use MMA.1688.F32.F16 mma instruction
|
||||
// Complex floating point tensor operation use mma.sync.aligned.m16n8k8.f32.f16.f16.f32 mma instruction
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
InstructionShape_,
|
||||
@@ -407,7 +407,7 @@ struct DefaultMmaComplexTensorOp<
|
||||
/// 3xTF32 or 4xTF32 (fast and accurate complex<float> operation)
|
||||
/// Partial specialization - input and output types are complex<float> * complex<float>
|
||||
// Use 3xTF32 or 4xTF32 tensor operation internally
|
||||
// 4 real-valued MMA.1688.F32.TF32 operations on TF32
|
||||
// 4 real-valued mma.sync.aligned.m16n8k8.f32.tf32.tf32.f32 operations on TF32
|
||||
// A = (ar + j ai), B (br +j bi), D = AB
|
||||
// D = dr + j di = 3x[(ar*br - ai*bi) + j (ar*bi + ai*br)]
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -441,7 +441,7 @@ struct DefaultMmaComplexTensorOp<
|
||||
TransformB,
|
||||
arch::OpMultiplyAddComplexFastF32> {
|
||||
|
||||
// Complex floating point tensor operation use MMA.1688.F32.TF32 mma instruction
|
||||
// Complex floating point tensor operation use mma.sync.aligned.m16n8k8.f32.tf32.tf32.f32 mma instruction
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
InstructionShape_,
|
||||
@@ -470,6 +470,143 @@ struct DefaultMmaComplexTensorOp<
|
||||
TransformB>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for complex<double>*complex<double> case
|
||||
// 4 real-valued mma.sync.aligned.m16n8k4.f64.f64.f64.f64 operations
|
||||
// A = (ar + j ai), B (br +j bi), D = AB
|
||||
// D = dr + j di = (ar*br - ai*bi) + j (ar*bi + ai*br)
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename WarpShape_,
|
||||
/// Real-valued underlying type of complex-valued A operand
|
||||
typename RealElementA,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA,
|
||||
/// Real-valued underlying type of complex-valued B operand
|
||||
typename RealElementB,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB,
|
||||
/// Real-valued underlying type of complex-valued C operand
|
||||
typename RealElementC,
|
||||
/// Layout of C matrix (concept: MatrixLayout)
|
||||
typename LayoutC,
|
||||
/// Complex transform on A operand
|
||||
ComplexTransform TransformA,
|
||||
/// Complex transform on B operand
|
||||
ComplexTransform TransformB>
|
||||
struct DefaultMmaComplexTensorOp<
|
||||
WarpShape_,
|
||||
GemmShape<16, 8, 4>,
|
||||
complex<RealElementA>,
|
||||
LayoutA,
|
||||
complex<RealElementB>,
|
||||
LayoutB,
|
||||
complex<RealElementC>,
|
||||
LayoutC,
|
||||
TransformA,
|
||||
TransformB,
|
||||
arch::OpMultiplyAddComplex> {
|
||||
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
GemmShape<16, 8, 4>,
|
||||
32,
|
||||
RealElementA,
|
||||
cutlass::layout::RowMajor,
|
||||
RealElementB,
|
||||
cutlass::layout::ColumnMajor,
|
||||
RealElementC,
|
||||
cutlass::layout::RowMajor,
|
||||
arch::OpMultiplyAdd>,
|
||||
cutlass::MatrixShape<1, 1>
|
||||
>;
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using Type = cutlass::gemm::warp::MmaComplexTensorOp<
|
||||
WarpShape_,
|
||||
complex<RealElementA>,
|
||||
LayoutA,
|
||||
complex<RealElementB>,
|
||||
LayoutB,
|
||||
complex<RealElementC>,
|
||||
LayoutC,
|
||||
Policy,
|
||||
TransformA,
|
||||
TransformB,
|
||||
true>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Partial specialization for complex<T>*complex<T> case using GaussianComplex operation
|
||||
// 3 real-valued mma.sync.aligned.m16n8k4.f64.f64.f64.f64 operations
|
||||
// A = (ar + j ai), B = (br +j bi), D = AB
|
||||
// P1 = (ar + ai) * br, P2 = - ar * (br - bi), P3 = ai * (br + bi)
|
||||
// D = dr + j di = (P1 - P3) + j (P1 + P2)
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename WarpShape_,
|
||||
/// Real-valued underlying type of complex-valued A operand
|
||||
typename RealElementA,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA,
|
||||
/// Real-valued underlying type of complex-valued B operand
|
||||
typename RealElementB,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB,
|
||||
/// Real-valued underlying type of complex-valued C operand
|
||||
typename RealElementC,
|
||||
/// Layout of C matrix (concept: MatrixLayout)
|
||||
typename LayoutC,
|
||||
/// Complex transform on A operand
|
||||
ComplexTransform TransformA,
|
||||
/// Complex transform on B operand
|
||||
ComplexTransform TransformB>
|
||||
struct DefaultMmaComplexTensorOp<
|
||||
WarpShape_,
|
||||
GemmShape<16, 8, 4>,
|
||||
complex<RealElementA>,
|
||||
LayoutA,
|
||||
complex<RealElementB>,
|
||||
LayoutB,
|
||||
complex<RealElementC>,
|
||||
LayoutC,
|
||||
TransformA,
|
||||
TransformB,
|
||||
arch::OpMultiplyAddGaussianComplex> {
|
||||
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
GemmShape<16, 8, 4>,
|
||||
32,
|
||||
RealElementA,
|
||||
cutlass::layout::RowMajor,
|
||||
RealElementB,
|
||||
cutlass::layout::ColumnMajor,
|
||||
RealElementC,
|
||||
cutlass::layout::RowMajor,
|
||||
arch::OpMultiplyAdd>,
|
||||
cutlass::MatrixShape<1, 1>
|
||||
>;
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using Type = cutlass::gemm::warp::MmaGaussianComplexTensorOp<
|
||||
WarpShape_,
|
||||
complex<RealElementA>,
|
||||
LayoutA,
|
||||
complex<RealElementB>,
|
||||
LayoutB,
|
||||
complex<RealElementC>,
|
||||
LayoutC,
|
||||
Policy,
|
||||
TransformA,
|
||||
TransformB,
|
||||
true>;
|
||||
};
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
@@ -46,6 +46,8 @@
|
||||
#include "cutlass/arch/memory_sm75.h"
|
||||
#include "cutlass/arch/mma_sm75.h"
|
||||
#include "cutlass/arch/mma_sm80.h"
|
||||
#include "cutlass/arch/mma_sm90.h"
|
||||
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/gemm/warp/mma.h"
|
||||
|
||||
@@ -545,7 +547,7 @@ public:
|
||||
/// Partial specialization for complex*complex+complex => complex:
|
||||
// Operands data type: complex<float>
|
||||
// Rounding: float -> tfloat32_t (round half_ulp_truncate nearest)
|
||||
// Math instruction: MMA.1688.F32.TF32
|
||||
// Math instruction: mma.sync.aligned.m16n8k8.f32.tf32.tf32.f32
|
||||
// Output data type: complex<float>
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -733,7 +735,7 @@ public:
|
||||
using MmaOperandC = typename ArchMmaOperator::FragmentC;
|
||||
|
||||
static_assert(platform::is_same<cutlass::gemm::GemmShape<16, 8, 8>, typename ArchMmaOperator::Shape>::value,
|
||||
"This implementation only supports MMA.1688 math instructions.");
|
||||
"This implementation only supports mma.m16n8k8 math instructions.");
|
||||
|
||||
static_assert(InstMmaOperandA::kElements == 4,
|
||||
"This implementation only supports math instructions in which exactly four element is needed for the A operand."
|
||||
@@ -846,6 +848,312 @@ public:
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Partial specialization for complex*complex+complex => complex:
|
||||
// Operands data type: complex<double>
|
||||
// Math instruction: mma.sync.aligned.m16n8k4.f64.f64.f64.f64
|
||||
// Output data type: complex<double>
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape_,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA_,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB_,
|
||||
/// 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
|
||||
>
|
||||
class MmaComplexTensorOp<
|
||||
Shape_,
|
||||
complex<double>,
|
||||
LayoutA_,
|
||||
complex<double>,
|
||||
LayoutB_,
|
||||
complex<double>,
|
||||
LayoutC_,
|
||||
Policy_,
|
||||
TransformA,
|
||||
TransformB,
|
||||
true> {
|
||||
public:
|
||||
/// Shape of warp-level matrix operation (concept: GemmShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Data type of members of complex multiplicand A
|
||||
using RealElementA = double;
|
||||
|
||||
/// Data type of multiplicand A
|
||||
using ElementA = complex<RealElementA>;
|
||||
|
||||
/// Layout of multiplicand A
|
||||
using LayoutA = LayoutA_;
|
||||
|
||||
/// Data type of members of complex multiplicand B
|
||||
using RealElementB = double;
|
||||
|
||||
/// Data type of multiplicand B
|
||||
using ElementB = complex<RealElementB>;
|
||||
|
||||
/// Layout of multiplicand B
|
||||
using LayoutB = LayoutB_;
|
||||
|
||||
/// Data type of members of complex accumulator matrix C
|
||||
using RealElementC = double;
|
||||
|
||||
/// 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: MmaLanePolicyTensorOp)
|
||||
using Policy = Policy_;
|
||||
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename Policy::Operator;
|
||||
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename ArchMmaOperator::Shape;
|
||||
|
||||
/// Underlying arch tag
|
||||
using ArchTag = typename ArchMmaOperator::ArchTag;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = typename arch::OpMultiplyAddComplex;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
|
||||
/// Complex transform on B operand
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
|
||||
/// 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<ArchMmaOperator::Shape::kM, ArchMmaOperator::Shape::kK>,
|
||||
Policy::OpDelta::kRow,
|
||||
32,
|
||||
1
|
||||
>;
|
||||
|
||||
/// 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 = MmaTensorOpMultiplicandTileIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>,
|
||||
Operand::kB,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
MatrixShape<ArchMmaOperator::Shape::kK, ArchMmaOperator::Shape::kN>,
|
||||
Policy::OpDelta::kColumn,
|
||||
32,
|
||||
1
|
||||
>;
|
||||
|
||||
/// Storage for B tile
|
||||
using FragmentB = typename IteratorB::Fragment;
|
||||
|
||||
/// Storage for transformed B tile
|
||||
using TransformedFragmentB = FragmentB;
|
||||
|
||||
static_assert(
|
||||
!(Shape::kM % ArchMmaOperator::Shape::kM) &&
|
||||
!(Shape::kN % ArchMmaOperator::Shape::kN),
|
||||
"Shape of warp-level Mma must be divisible by operator shape.");
|
||||
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = MatrixShape<
|
||||
Shape::kM / ArchMmaOperator::Shape::kM,
|
||||
Shape::kN / ArchMmaOperator::Shape::kN
|
||||
>;
|
||||
|
||||
/// 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, 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 * ArchMmaOperator::FragmentC::kElements,
|
||||
"Unexpected planar complex fragment length.");
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Underlying real-valued matrix multiply operator (concept: arch::Mma)
|
||||
ArchMmaOperator 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 ArchMmaOperator::FragmentA;
|
||||
using MmaOperandB = typename ArchMmaOperator::FragmentB;
|
||||
using MmaOperandC = typename ArchMmaOperator::FragmentC;
|
||||
|
||||
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;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mk = 0; mk < MmaOperandA::kElements; ++mk)
|
||||
operand_A[mk] = A[m*MmaOperandA::kElements + mk].real();
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int nk = 0; nk < MmaOperandB::kElements; ++nk)
|
||||
operand_B[nk] = B[n*MmaOperandB::kElements + nk].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;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mk = 0; mk < MmaOperandA::kElements; ++mk)
|
||||
operand_A[mk] = A[m*MmaOperandA::kElements + mk].real();
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int nk = 0; nk < MmaOperandB::kElements; ++nk)
|
||||
operand_B[nk] = (kTransformB == ComplexTransform::kConjugate ?
|
||||
-B[n*MmaOperandB::kElements + nk].imag() : B[n*MmaOperandB::kElements + nk].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
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mk = 0; mk < MmaOperandA::kElements; ++mk)
|
||||
operand_A[mk] = (kTransformA == ComplexTransform::kConjugate ?
|
||||
A[m*MmaOperandA::kElements + mk].imag() : -A[m*MmaOperandA::kElements + mk].imag());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int nk = 0; nk < MmaOperandB::kElements; ++nk)
|
||||
operand_B[nk] = (kTransformB == ComplexTransform::kConjugate ?
|
||||
-B[n*MmaOperandB::kElements + nk].imag() : B[n*MmaOperandB::kElements + nk].imag());
|
||||
|
||||
// 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.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;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mk = 0; mk < MmaOperandA::kElements; ++mk)
|
||||
operand_A[mk] = (kTransformA == ComplexTransform::kConjugate ?
|
||||
-A[m*MmaOperandA::kElements + mk].imag() : A[m*MmaOperandA::kElements + mk].imag());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int nk = 0; nk < MmaOperandB::kElements; ++nk)
|
||||
operand_B[nk] = B[n*MmaOperandB::kElements + nk].real();
|
||||
|
||||
// Complex-valued accumulator part
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow) + MmaIterations::kCount;
|
||||
|
||||
mma(*accum, operand_A, operand_B, *accum);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 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 {
|
||||
dst_A = A;
|
||||
dst_B = B;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// TODO - partial specializations of real*complex and complex*real
|
||||
|
||||
@@ -301,7 +301,7 @@ class MmaComplexTensorOpFastF32;
|
||||
/// Partial specialization for complex*complex+complex => complex:
|
||||
// Operands data type: complex<float>
|
||||
// Rounding: float -> tfloat32_t (round half_ulp_truncate nearest)
|
||||
// Math instruction: MMA.1688.F32.TF32
|
||||
// Math instruction: mma.sync.aligned.m16n8k8.f32.tf32.tf32.f32
|
||||
// Output data type: complex<float>
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -497,7 +497,7 @@ public:
|
||||
using MmaOperandC = typename ArchMmaOperator::FragmentC;
|
||||
|
||||
static_assert(platform::is_same<cutlass::gemm::GemmShape<16, 8, 8>, typename ArchMmaOperator::Shape>::value,
|
||||
"This implementation only supports MMA.1688 math instructions.");
|
||||
"This implementation only supports mma.m16n8k8 math instructions.");
|
||||
|
||||
static_assert(InstMmaOperandA::kElements == 4,
|
||||
"This implementation only supports math instructions in which exactly four element is needed for the A operand."
|
||||
|
||||
@@ -84,6 +84,8 @@ template <
|
||||
ComplexTransform TransformA = ComplexTransform::kNone,
|
||||
/// Complex transform on B operand
|
||||
ComplexTransform TransformB = ComplexTransform::kNone,
|
||||
/// Do source operands need more than one elements
|
||||
bool GeneralizedOperatorElements = false,
|
||||
/// Used for partial specialization
|
||||
typename Enable = bool
|
||||
>
|
||||
@@ -112,9 +114,7 @@ template <
|
||||
/// Complex transform on A operand
|
||||
ComplexTransform TransformA,
|
||||
/// Complex transform on B operand
|
||||
ComplexTransform TransformB,
|
||||
/// Used for partial specialization
|
||||
typename Enable
|
||||
ComplexTransform TransformB
|
||||
>
|
||||
class MmaGaussianComplexTensorOp<
|
||||
Shape_,
|
||||
@@ -126,8 +126,7 @@ class MmaGaussianComplexTensorOp<
|
||||
LayoutC_,
|
||||
Policy_,
|
||||
TransformA,
|
||||
TransformB,
|
||||
Enable> {
|
||||
TransformB> {
|
||||
public:
|
||||
/// Shape of warp-level matrix operation (concept: GemmShape)
|
||||
using Shape = Shape_;
|
||||
@@ -359,6 +358,282 @@ public:
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// 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
|
||||
>
|
||||
class MmaGaussianComplexTensorOp<
|
||||
Shape_,
|
||||
complex<RealElementA>,
|
||||
LayoutA_,
|
||||
complex<RealElementB>,
|
||||
LayoutB_,
|
||||
complex<RealElementC>,
|
||||
LayoutC_,
|
||||
Policy_,
|
||||
TransformA,
|
||||
TransformB,
|
||||
true> {
|
||||
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_;
|
||||
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename Policy::Operator;
|
||||
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename ArchMmaOperator::Shape;
|
||||
|
||||
/// Underlying arch tag
|
||||
using ArchTag = typename ArchMmaOperator::ArchTag;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = arch::OpMultiplyAddGaussianComplex;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
|
||||
/// Complex transform on B operand
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
|
||||
|
||||
/// 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<ArchMmaOperator::Shape::kM, ArchMmaOperator::Shape::kK>,
|
||||
Policy::OpDelta::kRow,
|
||||
32,
|
||||
1
|
||||
>;
|
||||
|
||||
/// 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 = MmaTensorOpMultiplicandTileIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>,
|
||||
Operand::kB,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
MatrixShape<ArchMmaOperator::Shape::kK, ArchMmaOperator::Shape::kN>,
|
||||
Policy::OpDelta::kColumn,
|
||||
32,
|
||||
1
|
||||
>;
|
||||
|
||||
/// Storage for B tile
|
||||
using FragmentB = typename IteratorB::Fragment;
|
||||
|
||||
/// Storage for transformed B tile
|
||||
using TransformedFragmentB = FragmentB;
|
||||
|
||||
static_assert(
|
||||
!(Shape::kM % ArchMmaOperator::Shape::kM) &&
|
||||
!(Shape::kN % ArchMmaOperator::Shape::kN),
|
||||
"Shape of warp-level Mma must be divisible by operator shape.");
|
||||
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = MatrixShape<
|
||||
Shape::kM / ArchMmaOperator::Shape::kM,
|
||||
Shape::kN / ArchMmaOperator::Shape::kN
|
||||
>;
|
||||
|
||||
/// Iterates over the C operand in memory
|
||||
using IteratorC = MmaTensorOpGaussianComplexAccumulatorTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kN>,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
typename ArchMmaOperator::Shape,
|
||||
typename Policy::OpDelta>;
|
||||
|
||||
/// Storage for C tile, the accumulator. Note, regardless of multiplicand type, this
|
||||
/// storage arrangement is to be considered 'gaussian complex' in the sense that the accumulation is
|
||||
/// done in three parts namely part1, part2, and part3. The parts 1, 2, and 3 are stored consecutively
|
||||
/// in InteratorC::Frament. This matches the structure of Tensor Cores which are always real-valued matrix multiplies.
|
||||
using FragmentC = typename IteratorC::Fragment;
|
||||
|
||||
static_assert(
|
||||
FragmentC::kElements == 3 * MmaIterations::kCount * ArchMmaOperator::FragmentC::kElements,
|
||||
"Unexpected gaussian complex fragment length.");
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Underlying real-valued matrix multiply operator (concept: arch::Mma)
|
||||
ArchMmaOperator mma;
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_DEVICE
|
||||
MmaGaussianComplexTensorOp() {}
|
||||
|
||||
/// 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 ArchMmaOperator::FragmentA;
|
||||
using MmaOperandB = typename ArchMmaOperator::FragmentB;
|
||||
using MmaOperandC = typename ArchMmaOperator::FragmentC;
|
||||
|
||||
D = C;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int m = 0; m < MmaIterations::kRow; ++m) {
|
||||
|
||||
// mma(accum.part1(), (a.real() + a.imag()), b.real(), accum.part1());
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < MmaIterations::kColumn; ++n) {
|
||||
|
||||
// Pack operands together. This may result in actual MOVs
|
||||
MmaOperandA operand_Asum;
|
||||
MmaOperandB operand_Br;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mk = 0; mk < MmaOperandA::kElements; ++mk)
|
||||
operand_Asum[mk] = A[m*MmaOperandA::kElements + mk].real() + ((kTransformA == ComplexTransform::kConjugate) ?
|
||||
-A[m*MmaOperandA::kElements + mk].imag() : +A[m*MmaOperandA::kElements + mk].imag());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int nk = 0; nk < MmaOperandB::kElements; ++nk)
|
||||
operand_Br[nk] = B[n*MmaOperandB::kElements + nk].real();
|
||||
|
||||
// accumulator part1
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow);
|
||||
|
||||
mma(*accum, operand_Asum, operand_Br, *accum);
|
||||
}
|
||||
|
||||
// mma(accum.part2(), -a.real(), (b.real() - b.imag()), accum.part2());
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = MmaIterations::kColumn - 1; n >= 0; --n) {
|
||||
|
||||
// Pack operands together. This may result in actual MOVs
|
||||
MmaOperandA operand_Ar;
|
||||
MmaOperandB operand_Bdiff;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mk = 0; mk < MmaOperandA::kElements; ++mk)
|
||||
operand_Ar[mk] = -A[m*MmaOperandA::kElements + mk].real();
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int nk = 0; nk < MmaOperandB::kElements; ++nk)
|
||||
operand_Bdiff[nk] = B[n*MmaOperandB::kElements + nk].real() - ((kTransformB == ComplexTransform::kConjugate) ?
|
||||
-B[n*MmaOperandB::kElements + nk].imag() : +B[n*MmaOperandB::kElements + nk].imag());
|
||||
|
||||
// accumulator part2
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow) + MmaIterations::kCount;
|
||||
|
||||
mma(*accum, operand_Ar, operand_Bdiff, *accum);
|
||||
}
|
||||
|
||||
// mma(accum.part3(), a.imag(), (b.real() + b.imag()), accum.part3())
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < MmaIterations::kColumn; ++n) {
|
||||
|
||||
// Pack operands together. This may result in actual MOVs
|
||||
MmaOperandA operand_Ai;
|
||||
MmaOperandB operand_Bsum;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mk = 0; mk < MmaOperandA::kElements; ++mk)
|
||||
operand_Ai[mk] = (kTransformA == ComplexTransform::kConjugate) ?
|
||||
-A[m*MmaOperandA::kElements + mk].imag() : +A[m*MmaOperandA::kElements + mk].imag();
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int nk = 0; nk < MmaOperandB::kElements; ++nk)
|
||||
operand_Bsum[nk] = B[n*MmaOperandB::kElements + nk].real() + ((kTransformB == ComplexTransform::kConjugate) ?
|
||||
-B[n*MmaOperandB::kElements + nk].imag() : +B[n*MmaOperandB::kElements + nk].imag());
|
||||
|
||||
// accumulator part3
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow) + 2 * MmaIterations::kCount;
|
||||
|
||||
mma(*accum, operand_Ai, operand_Bsum, *accum);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 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 {
|
||||
dst_A = A;
|
||||
dst_B = B;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
|
||||
@@ -493,7 +493,7 @@ public:
|
||||
|
||||
Index elements_offset = layout_({WmmaShape::kRow, 0});
|
||||
|
||||
byte_offset_ -= (elements_offset + sizeof_bits<Element>::value) / 8;
|
||||
byte_offset_ -= (elements_offset * sizeof_bits<Element>::value) / 8;
|
||||
return *this;
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user