CUTLASS 2.2 (#96)
Adds support for NVIDIA Ampere Architecture features. CUDA 11 Toolkit recommended.
This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -422,6 +422,342 @@ struct DefaultGemmConfiguration<
|
||||
using Operator = arch::OpMultiplyAddSaturate;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementC>
|
||||
struct DefaultGemmConfiguration<
|
||||
arch::OpClassTensorOp,
|
||||
arch::Sm75,
|
||||
uint1b_t,
|
||||
uint1b_t,
|
||||
ElementC,
|
||||
int32_t> {
|
||||
|
||||
static int const kAlignmentA = 128 / sizeof_bits<uint1b_t>::value;
|
||||
static int const kAlignmentB = 128 / sizeof_bits<uint1b_t>::value;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 512>;
|
||||
using WarpShape = GemmShape<64, 64, 512>;
|
||||
using InstructionShape = GemmShape<8, 8, 128>;
|
||||
static int const kStages = 2;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombinationClamp<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, int32_t, float>;
|
||||
|
||||
using Operator = arch::OpXorPopc;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename ElementA, typename ElementB, typename ElementC,
|
||||
typename ElementAccumulator>
|
||||
struct DefaultGemmConfiguration<arch::OpClassTensorOp, arch::Sm80, ElementA,
|
||||
ElementB, ElementC, ElementAccumulator> {
|
||||
|
||||
static int const kAlignmentA = 128 / sizeof_bits<ElementA>::value;
|
||||
static int const kAlignmentB = 128 / sizeof_bits<ElementA>::value;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 64>;
|
||||
using WarpShape = GemmShape<64, 64, 64>;
|
||||
using InstructionShape = GemmShape<16, 8, 16>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombination<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, ElementAccumulator,
|
||||
ElementAccumulator>;
|
||||
|
||||
using Operator = typename platform::conditional<
|
||||
(platform::is_same<ElementA, int8_t>::value ||
|
||||
platform::is_same<ElementA, int4b_t>::value ||
|
||||
platform::is_same<ElementA, uint8_t>::value ||
|
||||
platform::is_same<ElementA, uint4b_t>::value),
|
||||
arch::OpMultiplyAddSaturate, arch::OpMultiplyAdd>::type;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
template <typename ElementC,
|
||||
typename ElementAccumulator>
|
||||
struct DefaultGemmConfiguration<arch::OpClassTensorOp, arch::Sm80, double,
|
||||
double, ElementC, ElementAccumulator> {
|
||||
|
||||
static int const kAlignmentA = 1;
|
||||
static int const kAlignmentB = 1;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 64>;
|
||||
using WarpShape = GemmShape<64, 64, 64>;
|
||||
using InstructionShape = GemmShape<16, 8, 16>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombination<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, ElementAccumulator,
|
||||
ElementAccumulator>;
|
||||
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
};
|
||||
|
||||
|
||||
template <>
|
||||
struct DefaultGemmConfiguration<
|
||||
arch::OpClassTensorOp,
|
||||
arch::Sm80,
|
||||
complex<double>,
|
||||
complex<double>,
|
||||
complex<double>,
|
||||
complex<double>
|
||||
> {
|
||||
|
||||
static int const kAlignmentA = 1;
|
||||
static int const kAlignmentB = 1;
|
||||
|
||||
using ThreadblockShape = GemmShape<64, 64, 16>;
|
||||
using WarpShape = GemmShape<32, 32, 16>;
|
||||
using InstructionShape = GemmShape<8, 8, 4>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombination<
|
||||
complex<double>, 1, complex<double>,
|
||||
complex<double>>;
|
||||
|
||||
using Operator = arch::OpMultiplyAddComplex;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementC>
|
||||
struct DefaultGemmConfiguration<
|
||||
arch::OpClassTensorOp,
|
||||
arch::Sm80,
|
||||
int8_t,
|
||||
int8_t,
|
||||
ElementC,
|
||||
int32_t> {
|
||||
|
||||
static int const kAlignmentA = 128 / sizeof_bits<int8_t>::value;
|
||||
static int const kAlignmentB = 128 / sizeof_bits<int8_t>::value;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 64>;
|
||||
using WarpShape = GemmShape<64, 64, 64>;
|
||||
using InstructionShape = GemmShape<16, 8, 32>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombinationClamp<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, int32_t, float>;
|
||||
|
||||
using Operator = arch::OpMultiplyAddSaturate;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementC>
|
||||
struct DefaultGemmConfiguration<
|
||||
arch::OpClassTensorOp,
|
||||
arch::Sm80,
|
||||
int8_t,
|
||||
uint8_t,
|
||||
ElementC,
|
||||
int32_t> {
|
||||
|
||||
static int const kAlignmentA = 128 / sizeof_bits<int8_t>::value;
|
||||
static int const kAlignmentB = 128 / sizeof_bits<uint8_t>::value;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 64>;
|
||||
using WarpShape = GemmShape<64, 64, 64>;
|
||||
using InstructionShape = GemmShape<16, 8, 32>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombinationClamp<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, int32_t, float>;
|
||||
|
||||
using Operator = arch::OpMultiplyAddSaturate;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementC>
|
||||
struct DefaultGemmConfiguration<
|
||||
arch::OpClassTensorOp,
|
||||
arch::Sm80,
|
||||
uint8_t,
|
||||
int8_t,
|
||||
ElementC,
|
||||
int32_t> {
|
||||
|
||||
static int const kAlignmentA = 128 / sizeof_bits<uint8_t>::value;
|
||||
static int const kAlignmentB = 128 / sizeof_bits<int8_t>::value;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 64>;
|
||||
using WarpShape = GemmShape<64, 64, 64>;
|
||||
using InstructionShape = GemmShape<16, 8, 32>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombinationClamp<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, int32_t, float>;
|
||||
|
||||
using Operator = arch::OpMultiplyAddSaturate;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementC>
|
||||
struct DefaultGemmConfiguration<
|
||||
arch::OpClassTensorOp,
|
||||
arch::Sm80,
|
||||
uint8_t,
|
||||
uint8_t,
|
||||
ElementC,
|
||||
int32_t> {
|
||||
|
||||
static int const kAlignmentA = 128 / sizeof_bits<uint8_t>::value;
|
||||
static int const kAlignmentB = 128 / sizeof_bits<uint8_t>::value;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 64>;
|
||||
using WarpShape = GemmShape<64, 64, 64>;
|
||||
using InstructionShape = GemmShape<16, 8, 32>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombinationClamp<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, int32_t, float>;
|
||||
|
||||
using Operator = arch::OpMultiplyAddSaturate;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementC>
|
||||
struct DefaultGemmConfiguration<
|
||||
arch::OpClassTensorOp,
|
||||
arch::Sm80,
|
||||
int4b_t,
|
||||
int4b_t,
|
||||
ElementC,
|
||||
int32_t> {
|
||||
|
||||
static int const kAlignmentA = 128 / sizeof_bits<int4b_t>::value;
|
||||
static int const kAlignmentB = 128 / sizeof_bits<int4b_t>::value;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 128>;
|
||||
using WarpShape = GemmShape<64, 64, 128>;
|
||||
using InstructionShape = GemmShape<16, 8, 64>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombinationClamp<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, int32_t, float>;
|
||||
|
||||
using Operator = arch::OpMultiplyAddSaturate;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementC>
|
||||
struct DefaultGemmConfiguration<
|
||||
arch::OpClassTensorOp,
|
||||
arch::Sm80,
|
||||
int4b_t,
|
||||
uint4b_t,
|
||||
ElementC,
|
||||
int32_t> {
|
||||
|
||||
static int const kAlignmentA = 128 / sizeof_bits<int4b_t>::value;
|
||||
static int const kAlignmentB = 128 / sizeof_bits<uint4b_t>::value;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 128>;
|
||||
using WarpShape = GemmShape<64, 64, 128>;
|
||||
using InstructionShape = GemmShape<16, 8, 64>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombinationClamp<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, int32_t, float>;
|
||||
|
||||
using Operator = arch::OpMultiplyAddSaturate;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementC>
|
||||
struct DefaultGemmConfiguration<
|
||||
arch::OpClassTensorOp,
|
||||
arch::Sm80,
|
||||
uint4b_t,
|
||||
int4b_t,
|
||||
ElementC,
|
||||
int32_t> {
|
||||
|
||||
static int const kAlignmentA = 128 / sizeof_bits<uint4b_t>::value;
|
||||
static int const kAlignmentB = 128 / sizeof_bits<int4b_t>::value;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 128>;
|
||||
using WarpShape = GemmShape<64, 64, 128>;
|
||||
using InstructionShape = GemmShape<16, 8, 64>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombinationClamp<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, int32_t, float>;
|
||||
|
||||
using Operator = arch::OpMultiplyAddSaturate;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementC>
|
||||
struct DefaultGemmConfiguration<
|
||||
arch::OpClassTensorOp,
|
||||
arch::Sm80,
|
||||
uint4b_t,
|
||||
uint4b_t,
|
||||
ElementC,
|
||||
int32_t> {
|
||||
|
||||
static int const kAlignmentA = 128 / sizeof_bits<uint4b_t>::value;
|
||||
static int const kAlignmentB = 128 / sizeof_bits<uint4b_t>::value;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 128>;
|
||||
using WarpShape = GemmShape<64, 64, 128>;
|
||||
using InstructionShape = GemmShape<16, 8, 64>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombinationClamp<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, int32_t, float>;
|
||||
|
||||
using Operator = arch::OpMultiplyAddSaturate;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementC>
|
||||
struct DefaultGemmConfiguration<
|
||||
arch::OpClassTensorOp,
|
||||
arch::Sm80,
|
||||
uint1b_t,
|
||||
uint1b_t,
|
||||
ElementC,
|
||||
int32_t> {
|
||||
|
||||
static int const kAlignmentA = 128 / sizeof_bits<uint1b_t>::value;
|
||||
static int const kAlignmentB = 128 / sizeof_bits<uint1b_t>::value;
|
||||
|
||||
using ThreadblockShape = GemmShape<128, 256, 512>;
|
||||
using WarpShape = GemmShape<64, 64, 512>;
|
||||
using InstructionShape = GemmShape<16, 8, 256>;
|
||||
static int const kStages = 3;
|
||||
|
||||
using EpilogueOutputOp = epilogue::thread::LinearCombinationClamp<
|
||||
ElementC, 128 / sizeof_bits<ElementC>::value, int32_t, float>;
|
||||
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
} // namespace device
|
||||
} // namespace gemm
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -193,7 +193,7 @@ template <
|
||||
ElementAccumulator_>::EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle_ =
|
||||
typename threadblock::GemmCohortThreadblockSwizzle<LayoutA_, LayoutB_>,
|
||||
typename threadblock::GemmIdentityThreadblockSwizzle<>,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages =
|
||||
DefaultGemmConfiguration<OperatorClass_, ArchTag_, ElementA_, ElementB_,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -192,7 +192,7 @@ template <
|
||||
OperatorClass_, ArchTag_, ElementA_, ElementB_, ElementC_,
|
||||
ElementAccumulator_>::EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle_ = threadblock::GemmIdentityThreadblockSwizzle,
|
||||
typename ThreadblockSwizzle_ = threadblock::GemmIdentityThreadblockSwizzle<>,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages =
|
||||
DefaultGemmConfiguration<OperatorClass_, ArchTag_, ElementA_, ElementB_,
|
||||
@@ -202,6 +202,7 @@ template <
|
||||
/// Complex elementwise transformation on B operand
|
||||
ComplexTransform TransformB = ComplexTransform::kNone,
|
||||
/// Multiply-add operator
|
||||
// (selects complex or gaussian complex)
|
||||
typename Operator_ = arch::OpMultiplyAddComplex,
|
||||
/// If true, kernel supports split-K with serial reduction
|
||||
bool SplitKSerial = false
|
||||
@@ -506,6 +507,7 @@ template <
|
||||
/// Complex elementwise transformation on B operand
|
||||
ComplexTransform TransformB,
|
||||
/// Multiply-add operator
|
||||
// (selects complex or gaussian complex)
|
||||
typename Operator_,
|
||||
/// If true, kernel supports split-K as a serial reduction
|
||||
bool SplitKSerial
|
||||
@@ -571,8 +573,8 @@ public:
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
TransformA,
|
||||
TransformB,
|
||||
TransformA,
|
||||
Operator,
|
||||
SplitKSerial
|
||||
>;
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -89,7 +89,7 @@ template <
|
||||
OperatorClass_, ArchTag_, ElementA_, ElementB_, ElementC_,
|
||||
ElementAccumulator_>::EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle_ = threadblock::GemmIdentityThreadblockSwizzle,
|
||||
typename ThreadblockSwizzle_ = threadblock::GemmIdentityThreadblockSwizzle<>,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages =
|
||||
DefaultGemmConfiguration<OperatorClass_, ArchTag_, ElementA_, ElementB_,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -41,14 +41,78 @@ namespace device {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
|
||||
template <
|
||||
typename ElementA_,
|
||||
typename LayoutA_,
|
||||
ComplexTransform TransformA,
|
||||
int AlignmentA,
|
||||
typename ElementB_,
|
||||
typename LayoutB_,
|
||||
ComplexTransform TransformB,
|
||||
int AlignmentB,
|
||||
typename LayoutC_,
|
||||
bool Transpose
|
||||
>
|
||||
struct MapArguments {
|
||||
using ElementA = ElementA_;
|
||||
using LayoutA = LayoutA_;
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
static int const kAlignmentA = AlignmentA;
|
||||
using ElementB = ElementB_;
|
||||
using LayoutB = LayoutB_;
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
static int const kAlignmentB = AlignmentB;
|
||||
using LayoutC = LayoutC_;
|
||||
};
|
||||
|
||||
template <
|
||||
typename ElementA_,
|
||||
typename LayoutA_,
|
||||
ComplexTransform TransformA,
|
||||
int AlignmentA,
|
||||
typename ElementB_,
|
||||
typename LayoutB_,
|
||||
ComplexTransform TransformB,
|
||||
int AlignmentB,
|
||||
typename LayoutC_
|
||||
>
|
||||
struct MapArguments<
|
||||
ElementA_,
|
||||
LayoutA_,
|
||||
TransformA,
|
||||
AlignmentA,
|
||||
ElementB_,
|
||||
LayoutB_,
|
||||
TransformB,
|
||||
AlignmentB,
|
||||
LayoutC_,
|
||||
true
|
||||
> {
|
||||
using ElementA = ElementB_;
|
||||
using LayoutA = typename layout::LayoutTranspose<LayoutB_>::type;
|
||||
static ComplexTransform const kTransformA = TransformB;
|
||||
static int const kAlignmentA = AlignmentB;
|
||||
using ElementB = ElementA_;
|
||||
using LayoutB = typename layout::LayoutTranspose<LayoutA_>::type;
|
||||
static ComplexTransform const kTransformB = TransformA;
|
||||
static int const kAlignmentB = AlignmentA;
|
||||
using LayoutC = typename layout::LayoutTranspose<LayoutC_>::type;
|
||||
};
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename GemmKernel_>
|
||||
class GemmUniversalAdapter {
|
||||
public:
|
||||
|
||||
using GemmKernel = GemmKernel_;
|
||||
|
||||
static_assert(std::is_same<typename GemmKernel::LayoutC, cutlass::layout::RowMajor>::value,
|
||||
"Universal adapter expects the kernel to be row-major and transposes its arguments.");
|
||||
static bool const kInternalTranspose =
|
||||
std::is_same<typename GemmKernel::LayoutC, cutlass::layout::RowMajor>::value;
|
||||
|
||||
using ThreadblockShape = typename GemmKernel::Mma::Shape;
|
||||
using WarpShape = typename GemmKernel::WarpShape;
|
||||
@@ -56,26 +120,39 @@ public:
|
||||
|
||||
using OperatorClass = typename GemmKernel::OperatorClass;
|
||||
using ArchTag = typename GemmKernel::ArchTag;
|
||||
|
||||
|
||||
// Type, layout, and complex transform deliberately exchanged with B
|
||||
using ElementA = typename GemmKernel::ElementB;
|
||||
using LayoutA = typename layout::LayoutTranspose<typename GemmKernel::LayoutB>::type;
|
||||
using TensorRefA = TensorRef<ElementA const, LayoutA>;
|
||||
static ComplexTransform const kTransformA = GemmKernel::kTransformB;
|
||||
using MapArguments = detail::MapArguments<
|
||||
typename GemmKernel::ElementA,
|
||||
typename GemmKernel::LayoutA,
|
||||
GemmKernel::kTransformA,
|
||||
GemmKernel::kAlignmentA,
|
||||
typename GemmKernel::ElementB,
|
||||
typename GemmKernel::LayoutB,
|
||||
GemmKernel::kTransformB,
|
||||
GemmKernel::kAlignmentB,
|
||||
typename GemmKernel::LayoutC,
|
||||
kInternalTranspose
|
||||
>;
|
||||
|
||||
using ElementA = typename MapArguments::ElementA;
|
||||
using LayoutA = typename MapArguments::LayoutA;
|
||||
static ComplexTransform const kTransformA = MapArguments::kTransformA;
|
||||
static int const kAlignmentA = GemmKernel::kAlignmentA;
|
||||
|
||||
// Type, layout, and complex transform deliberately exchanged with A
|
||||
using ElementB = typename GemmKernel::ElementA;
|
||||
using LayoutB = typename layout::LayoutTranspose<typename GemmKernel::LayoutA>::type;
|
||||
using TensorRefB = TensorRef<ElementB const, LayoutB>;
|
||||
static ComplexTransform const kTransformB = GemmKernel::kTransformA;
|
||||
using ElementB = typename MapArguments::ElementB;
|
||||
using LayoutB = typename MapArguments::LayoutB;
|
||||
static ComplexTransform const kTransformB = MapArguments::kTransformB;
|
||||
static int const kAlignmentB = GemmKernel::kAlignmentB;
|
||||
|
||||
|
||||
using ElementC = typename GemmKernel::ElementC;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = typename MapArguments::LayoutC;
|
||||
static int const kAlignmentC = GemmKernel::kAlignmentC;
|
||||
|
||||
using TensorRefA = TensorRef<ElementA const, LayoutA>;
|
||||
using TensorRefB = TensorRef<ElementB const, LayoutB>;
|
||||
using TensorRefC = TensorRef<ElementC const, LayoutC>;
|
||||
using TensorRefD = TensorRef<ElementC, LayoutC>;
|
||||
static int const kAlignmentC = GemmKernel::kAlignmentC;
|
||||
|
||||
using ElementAccumulator = typename GemmKernel::Mma::Policy::Operator::ElementC;
|
||||
|
||||
@@ -99,7 +176,12 @@ public:
|
||||
|
||||
/// Helper to construct a transposed equivalent for the underying GEMM operator
|
||||
static Arguments to_underlying_arguments(Arguments const &args) {
|
||||
return args.transposed_problem();
|
||||
if (kInternalTranspose) {
|
||||
return args.transposed_problem();
|
||||
}
|
||||
else {
|
||||
return args;
|
||||
}
|
||||
}
|
||||
|
||||
/// Determines whether the GEMM can execute the given problem.
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -400,7 +400,8 @@ enum class GemmUniversalMode {
|
||||
kGemm,
|
||||
kGemmSplitKParallel,
|
||||
kBatched,
|
||||
kArray
|
||||
kArray,
|
||||
kInvalid
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -49,6 +49,7 @@
|
||||
#include "cutlass/gemm/kernel/gemm_pipelined.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_sm75.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_sm70.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_sm80.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_simt.h"
|
||||
#include "cutlass/gemm/threadblock/threadblock_swizzle.h"
|
||||
@@ -116,6 +117,68 @@ template <
|
||||
struct DefaultGemm;
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Ampere Architecture
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for C and D matrix operands
|
||||
typename ElementC,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Epilogue output operator
|
||||
typename EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// If true, kernel is configured to support serial reduction in the
|
||||
/// epilogue
|
||||
bool SplitKSerial,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator>
|
||||
struct DefaultGemm<ElementA, LayoutA, kAlignmentA, ElementB, LayoutB, kAlignmentB, ElementC,
|
||||
layout::RowMajor, ElementAccumulator, arch::OpClassTensorOp,
|
||||
arch::Sm80, ThreadblockShape, WarpShape, InstructionShape,
|
||||
EpilogueOutputOp, ThreadblockSwizzle, Stages, SplitKSerial,
|
||||
Operator> {
|
||||
/// Define the threadblock-scoped matrix multiply-accumulate
|
||||
using Mma = typename cutlass::gemm::threadblock::DefaultMma<
|
||||
ElementA, LayoutA, kAlignmentA, ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator, layout::RowMajor, arch::OpClassTensorOp, arch::Sm80,
|
||||
ThreadblockShape, WarpShape, InstructionShape, Stages,
|
||||
Operator>::ThreadblockMma;
|
||||
|
||||
static const int kPartitionsK = ThreadblockShape::kK / WarpShape::kK;
|
||||
|
||||
/// Define the epilogue
|
||||
using Epilogue =
|
||||
typename cutlass::epilogue::threadblock::DefaultEpilogueTensorOp<
|
||||
ThreadblockShape, typename Mma::Operator, kPartitionsK, EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount>::Epilogue;
|
||||
|
||||
/// Define the kernel-level GEMM operator.
|
||||
using GemmKernel = kernel::Gemm<Mma, Epilogue, ThreadblockSwizzle, SplitKSerial>;
|
||||
};
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Turing Architecture
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
@@ -201,6 +264,75 @@ struct DefaultGemm<
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Ampere Integer Matrix Multiply Interleaved layout
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for C and D matrix operands
|
||||
typename ElementC,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Epilogue output operator
|
||||
typename EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Number of Interleaved k
|
||||
int InterleavedK,
|
||||
/// If true, kernel is configured to support serial reduction in the
|
||||
/// epilogue
|
||||
bool SplitKSerial,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator,
|
||||
/// Is Beta zero or not
|
||||
bool IsBetaZero>
|
||||
struct DefaultGemm<
|
||||
ElementA, layout::ColumnMajorInterleaved<InterleavedK>, kAlignmentA,
|
||||
ElementB, layout::RowMajorInterleaved<InterleavedK>, kAlignmentB, ElementC,
|
||||
layout::ColumnMajorInterleaved<InterleavedK>, int32_t,
|
||||
arch::OpClassTensorOp, arch::Sm80, ThreadblockShape, WarpShape,
|
||||
InstructionShape, EpilogueOutputOp, ThreadblockSwizzle, Stages,
|
||||
SplitKSerial, Operator, IsBetaZero> {
|
||||
using LayoutA = layout::ColumnMajorInterleaved<InterleavedK>;
|
||||
using LayoutB = layout::RowMajorInterleaved<InterleavedK>;
|
||||
using LayoutC = layout::ColumnMajorInterleaved<InterleavedK>;
|
||||
|
||||
using ElementAccumulator = int32_t;
|
||||
|
||||
/// Define the threadblock-scoped matrix multiply-accumulate
|
||||
using Mma = typename cutlass::gemm::threadblock::DefaultMma<
|
||||
ElementA, LayoutA, kAlignmentA, ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator, LayoutC, arch::OpClassTensorOp, arch::Sm80,
|
||||
ThreadblockShape, WarpShape, InstructionShape, Stages, Operator,
|
||||
true>::ThreadblockMma;
|
||||
|
||||
static const int kPartitionsK = ThreadblockShape::kK / WarpShape::kK;
|
||||
|
||||
/// Define the epilogue
|
||||
using Epilogue = typename cutlass::epilogue::threadblock::
|
||||
DefaultInterleavedEpilogueTensorOp<
|
||||
ThreadblockShape, typename Mma::Operator, kPartitionsK, EpilogueOutputOp,
|
||||
64 / sizeof_bits<ElementC>::value, InterleavedK,
|
||||
IsBetaZero>::Epilogue;
|
||||
|
||||
/// Define the kernel-level GEMM operator.
|
||||
using GemmKernel = kernel::Gemm<Mma, Epilogue, ThreadblockSwizzle, SplitKSerial>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Turing Integer Matrix Multiply Interleaved layout
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
@@ -439,6 +571,80 @@ struct DefaultGemm<
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Ampere
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for C and D matrix operands
|
||||
typename ElementC,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Epilogue output operator
|
||||
typename EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages
|
||||
int Stages,
|
||||
/// If true, kernel is configured to support serial reduction in the epilogue
|
||||
bool SplitKSerial,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator>
|
||||
struct DefaultGemm<ElementA,
|
||||
LayoutA,
|
||||
kAlignmentA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
kAlignmentB,
|
||||
ElementC,
|
||||
layout::RowMajor,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
arch::Sm80,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
GemmShape<1, 1, 1>,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
SplitKSerial,
|
||||
Operator> {
|
||||
|
||||
/// Define the threadblock-scoped matrix multiply-accumulate
|
||||
using Mma = typename cutlass::gemm::threadblock::DefaultMma<
|
||||
ElementA, LayoutA, kAlignmentA, ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator, layout::RowMajor, arch::OpClassSimt, arch::Sm80,
|
||||
ThreadblockShape, WarpShape, GemmShape<1, 1, 1>, Stages,
|
||||
Operator>::ThreadblockMma;
|
||||
|
||||
static int const kEpilogueElementsPerAccess = EpilogueOutputOp::kCount;
|
||||
static_assert(kEpilogueElementsPerAccess == 1, "simt epilogue must operate on scalars");
|
||||
|
||||
/// Define the epilogue
|
||||
using Epilogue = typename cutlass::epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
typename Mma::Operator,
|
||||
EpilogueOutputOp,
|
||||
kEpilogueElementsPerAccess
|
||||
>::Epilogue;
|
||||
|
||||
/// Define the kernel-level GEMM operator.
|
||||
using GemmKernel = kernel::Gemm<Mma, Epilogue, ThreadblockSwizzle, SplitKSerial>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
/// Partial specialization for SIMT DP4A
|
||||
|
||||
@@ -516,7 +722,6 @@ struct DefaultGemm<int8_t, LayoutA, kAlignmentA, int8_t, LayoutB, kAlignmentB,
|
||||
using GemmKernel = kernel::Gemm<Mma, Epilogue, ThreadblockSwizzle, SplitKSerial>;
|
||||
};
|
||||
|
||||
|
||||
#if defined(CUTLASS_ARCH_WMMA_ENABLED)
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
/// Partial specialization for Wmma Gemm Kernel
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -49,7 +49,9 @@
|
||||
#include "cutlass/gemm/kernel/gemm_pipelined.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_sm75.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_sm70.h"
|
||||
#include "cutlass/gemm/threadblock/default_multistage_mma_complex_core_sm80.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma.h"
|
||||
#include "cutlass/gemm/threadblock/default_multistage_mma_complex.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_simt.h"
|
||||
#include "cutlass/gemm/threadblock/threadblock_swizzle.h"
|
||||
#include "cutlass/epilogue/threadblock/default_epilogue_complex_tensor_op.h"
|
||||
@@ -101,6 +103,7 @@ template <
|
||||
/// Complex elementwise transformation on B operand
|
||||
ComplexTransform TransformB,
|
||||
/// Multiply-add operator
|
||||
// (arch::OpMultiplyAddComplex, arch::OpMultiplyGaussianComplex)
|
||||
typename Operator,
|
||||
/// If true, kernel is configured to support serial reduction in the epilogue
|
||||
bool SplitKSerial
|
||||
@@ -109,6 +112,64 @@ struct DefaultGemmComplex;
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Ampere Architecture
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Element type for C and D matrix operands
|
||||
typename ElementC,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Epilogue output operator
|
||||
typename EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Complex elementwise transformation on A operand
|
||||
ComplexTransform TransformA,
|
||||
/// Complex elementwise transformation on B operand
|
||||
ComplexTransform TransformB,
|
||||
/// Multiply-add operator
|
||||
// (arch::OpMultiplyAddComplex, arch::OpMultiplyGaussianComplex)
|
||||
typename Operator,
|
||||
/// If true, kernel is configured to support serial reduction in the epilogue
|
||||
bool SplitKSerial
|
||||
>
|
||||
struct DefaultGemmComplex<
|
||||
ElementA, LayoutA, ElementB, LayoutB, ElementC,
|
||||
layout::RowMajor, ElementAccumulator, arch::OpClassTensorOp,
|
||||
arch::Sm80, ThreadblockShape, WarpShape, InstructionShape,
|
||||
EpilogueOutputOp, ThreadblockSwizzle, Stages, TransformA, TransformB, Operator, SplitKSerial> {
|
||||
|
||||
/// Define the threadblock-scoped matrix multiply-accumulate
|
||||
using Mma = typename cutlass::gemm::threadblock::DefaultMultistageMmaComplex<
|
||||
ElementA, LayoutA, ElementB, LayoutB, ElementAccumulator,
|
||||
layout::RowMajor, arch::OpClassTensorOp, arch::Sm80, ThreadblockShape,
|
||||
WarpShape, InstructionShape, Stages, TransformA, TransformB, Operator>::ThreadblockMma;
|
||||
|
||||
/// Define the epilogue
|
||||
using Epilogue =
|
||||
typename cutlass::epilogue::threadblock::DefaultEpilogueComplexTensorOp<
|
||||
ThreadblockShape, typename Mma::Operator, 1, EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount, Operator>::Epilogue;
|
||||
|
||||
/// Define the kernel-level GEMM operator.
|
||||
using GemmKernel = kernel::Gemm<Mma, Epilogue, ThreadblockSwizzle, SplitKSerial>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -49,6 +49,7 @@
|
||||
|
||||
#include "cutlass/epilogue/threadblock/default_epilogue_planar_complex.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_planar_complex_pipelined.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_planar_complex_multistage.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -222,6 +223,122 @@ struct DefaultGemmPlanarComplexUniversal<
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for multiple pipeline stages.
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Complex elementwise transformation on A operand
|
||||
ComplexTransform TransformA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Complex elementwise transformation on B operand
|
||||
ComplexTransform TransformB,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for C and D matrix operands
|
||||
typename ElementC,
|
||||
/// Layout type for C and D matrix operands
|
||||
typename LayoutC,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Operator class tag
|
||||
typename OperatorClass,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Epilogue output operator
|
||||
typename EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator
|
||||
>
|
||||
struct DefaultGemmPlanarComplexUniversal<
|
||||
ElementA,
|
||||
LayoutA,
|
||||
TransformA,
|
||||
kAlignmentA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
TransformB,
|
||||
kAlignmentB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
OperatorClass,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
Operator,
|
||||
typename std::enable_if<(Stages > 2)>::type
|
||||
> {
|
||||
|
||||
/// Define planar complex valued variants instead
|
||||
using Mma = typename gemm::threadblock::DefaultMmaPlanarComplexMultistage<
|
||||
ElementA,
|
||||
LayoutA,
|
||||
kAlignmentA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
kAlignmentB,
|
||||
ElementAccumulator,
|
||||
LayoutC,
|
||||
OperatorClass,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
Stages,
|
||||
TransformA,
|
||||
TransformB,
|
||||
Operator
|
||||
>::ThreadblockMma;
|
||||
|
||||
/// Planar complex epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpiloguePlanarComplex<
|
||||
ThreadblockShape,
|
||||
typename Mma::Policy::Operator,
|
||||
OperatorClass,
|
||||
ArchTag,
|
||||
ThreadblockShape::kK / WarpShape::kK,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
/// Define the kernel in terms of the default kernel
|
||||
using GemmKernel = kernel::GemmPlanarComplex<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle
|
||||
>;
|
||||
|
||||
// Array variant
|
||||
using GemmArrayKernel = kernel::GemmPlanarComplexArray<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -421,6 +421,13 @@ public:
|
||||
|
||||
cutlass::gemm::GemmCoord threadblock_tile_offset = threadblock_swizzle.get_tile_offset();
|
||||
|
||||
// Early exit if CTA is out of range
|
||||
if (params.grid_tiled_shape.m() <= threadblock_tile_offset.m() ||
|
||||
params.grid_tiled_shape.n() <= threadblock_tile_offset.n()) {
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
int offset_k = 0;
|
||||
int problem_size_k = params.problem_size.k();
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -377,6 +377,14 @@ public:
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
cutlass::gemm::GemmCoord threadblock_tile_offset = threadblock_swizzle.get_tile_offset();
|
||||
|
||||
// Early exit if CTA is out of range
|
||||
if (params.grid_tiled_shape.m() <= threadblock_tile_offset.m() ||
|
||||
params.grid_tiled_shape.n() <= threadblock_tile_offset.n()) {
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
int batch_idx = threadblock_tile_offset.k();
|
||||
|
||||
int problem_size_m = params.problem_size.m();
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -71,7 +71,7 @@ public:
|
||||
using OperatorClass = typename Mma::Operator::OperatorClass;
|
||||
using ThreadblockShape = typename Mma::Shape;
|
||||
using WarpShape = typename Mma::Operator::Shape;
|
||||
using InstructionShape = typename Mma::Policy::Operator::Shape;
|
||||
using InstructionShape = typename Mma::Policy::Operator::InstructionShape;
|
||||
using ArchTag = typename Mma::ArchTag;
|
||||
|
||||
static int const kStages = Mma::kStages;
|
||||
@@ -259,9 +259,9 @@ public:
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr) {
|
||||
|
||||
ptr_A = args.ptr_A;
|
||||
ptr_B = args.ptr_B;
|
||||
ptr_C = args.ptr_C;
|
||||
ptr_A = const_cast<void *>(args.ptr_A);
|
||||
ptr_B = const_cast<void *>(args.ptr_B);
|
||||
ptr_C = const_cast<void *>(args.ptr_C);
|
||||
ptr_D = args.ptr_D;
|
||||
|
||||
output_op = args.epilogue;
|
||||
@@ -303,6 +303,10 @@ public:
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
static Status can_implement(Arguments const &args) {
|
||||
return can_implement(args.problem_size);
|
||||
}
|
||||
|
||||
/// Executes one GEMM
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -38,6 +38,8 @@
|
||||
#include "cutlass/transform/threadblock/predicated_tile_iterator_2dthreadtile.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_sm70.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_sm75.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_sm80.h"
|
||||
|
||||
#if defined(CUTLASS_ARCH_WMMA_ENABLED)
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_wmma.h"
|
||||
#endif //CUTLASS_ARCH_WMMA_ENABLED
|
||||
@@ -203,6 +205,58 @@ struct DefaultMma<ElementA, LayoutA, kAlignmentA, ElementB, LayoutB,
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
/// Specialization for row-major output (OperatorClass TensorOp)
|
||||
template <
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Instruction-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator
|
||||
>
|
||||
struct DefaultMma<float, LayoutA, kAlignmentA, float, LayoutB,
|
||||
kAlignmentB, float, layout::RowMajor,
|
||||
arch::OpClassTensorOp, ArchTag, ThreadblockShape, WarpShape,
|
||||
InstructionShape, 2, Operator, false> {
|
||||
// Define the MmaCore components
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, float, LayoutA, float,
|
||||
LayoutB, float, layout::RowMajor, arch::OpClassTensorOp, 2,
|
||||
arch::OpMultiplyAddFastF16>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using IteratorA =
|
||||
cutlass::transform::threadblock::PredicatedTileIterator<
|
||||
cutlass::MatrixShape<MmaCore::Shape::kM, MmaCore::Shape::kK>,
|
||||
float, LayoutA, 1, typename MmaCore::IteratorThreadMapA, kAlignmentA>;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using IteratorB =
|
||||
cutlass::transform::threadblock::PredicatedTileIterator<
|
||||
cutlass::MatrixShape<MmaCore::Shape::kK, MmaCore::Shape::kN>,
|
||||
float, LayoutB, 0, typename MmaCore::IteratorThreadMapB, kAlignmentB>;
|
||||
|
||||
// Define the threadblock-scoped pipelined matrix multiply
|
||||
using ThreadblockMma = cutlass::gemm::threadblock::MmaPipelined<
|
||||
typename MmaCore::Shape, IteratorA, typename MmaCore::SmemIteratorA,
|
||||
IteratorB, typename MmaCore::SmemIteratorB, float,
|
||||
layout::RowMajor, typename MmaCore::MmaPolicy>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization for column-major-interleaved output
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
@@ -271,6 +325,214 @@ struct DefaultMma<ElementA, LayoutA, kAlignmentA, ElementB, LayoutB,
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization for row-major output
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Instruction-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Number of stages used in the multistage mainloop
|
||||
int Stages,
|
||||
/// Operation perfomed by GEMM
|
||||
typename Operator
|
||||
>
|
||||
struct DefaultMma<ElementA, LayoutA, kAlignmentA, ElementB, LayoutB,
|
||||
kAlignmentB, ElementAccumulator, layout::RowMajor,
|
||||
arch::OpClassSimt, ArchTag, ThreadblockShape, WarpShape,
|
||||
InstructionShape, Stages, Operator, false> {
|
||||
// Define the MmaCore components
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, LayoutA,
|
||||
ElementB, LayoutB, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
Stages, Operator>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using AccessTypeA = cutlass::Array<ElementA, kAlignmentA>;
|
||||
using IteratorA =
|
||||
cutlass::transform::threadblock::PredicatedTileAccessIterator<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA, LayoutA, 1, ThreadMapA, AccessTypeA>;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using AccessTypeB = cutlass::Array<ElementB, kAlignmentB>;
|
||||
using IteratorB =
|
||||
cutlass::transform::threadblock::PredicatedTileAccessIterator<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB, LayoutB, 0, ThreadMapB, AccessTypeB>;
|
||||
|
||||
// Define the threadblock-scoped multistage matrix multiply
|
||||
using ThreadblockMma = cutlass::gemm::threadblock::MmaMultistage<
|
||||
typename MmaCore::Shape, IteratorA, typename MmaCore::SmemIteratorA,
|
||||
MmaCore::kCacheOpA, IteratorB, typename MmaCore::SmemIteratorB,
|
||||
MmaCore::kCacheOpB, ElementAccumulator, layout::RowMajor,
|
||||
typename MmaCore::MmaPolicy, Stages>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization for row-major output (OperatorClass TensorOp)
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Instruction-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Number of stages used in the multistage mainloop
|
||||
int Stages,
|
||||
/// Operation perfomed by GEMM
|
||||
typename Operator
|
||||
>
|
||||
struct DefaultMma<ElementA, LayoutA, kAlignmentA, ElementB, LayoutB,
|
||||
kAlignmentB, ElementAccumulator, layout::RowMajor,
|
||||
arch::OpClassTensorOp, ArchTag, ThreadblockShape, WarpShape,
|
||||
InstructionShape, Stages, Operator, false> {
|
||||
static cutlass::arch::CacheOperation::Kind const CacheOpA =
|
||||
((sizeof_bits<ElementA>::value * kAlignmentA) == 128)
|
||||
? cutlass::arch::CacheOperation::Global
|
||||
: cutlass::arch::CacheOperation::Always;
|
||||
|
||||
static cutlass::arch::CacheOperation::Kind const CacheOpB =
|
||||
((sizeof_bits<ElementB>::value * kAlignmentB) == 128)
|
||||
? cutlass::arch::CacheOperation::Global
|
||||
: cutlass::arch::CacheOperation::Always;
|
||||
|
||||
// Define the MmaCore components
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, LayoutA,
|
||||
ElementB, LayoutB, ElementAccumulator, layout::RowMajor, arch::OpClassTensorOp,
|
||||
Stages, Operator, false, CacheOpA, CacheOpB>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using AccessTypeA = cutlass::Array<ElementA, kAlignmentA>;
|
||||
using IteratorA =
|
||||
cutlass::transform::threadblock::PredicatedTileAccessIterator<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA, LayoutA, 1, ThreadMapA, AccessTypeA>;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using AccessTypeB = cutlass::Array<ElementB, kAlignmentB>;
|
||||
using IteratorB =
|
||||
cutlass::transform::threadblock::PredicatedTileAccessIterator<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB, LayoutB, 0, ThreadMapB, AccessTypeB>;
|
||||
|
||||
// Define the threadblock-scoped multistage matrix multiply
|
||||
using ThreadblockMma = cutlass::gemm::threadblock::MmaMultistage<
|
||||
typename MmaCore::Shape, IteratorA, typename MmaCore::SmemIteratorA,
|
||||
MmaCore::kCacheOpA, IteratorB, typename MmaCore::SmemIteratorB,
|
||||
MmaCore::kCacheOpB, ElementAccumulator, layout::RowMajor,
|
||||
typename MmaCore::MmaPolicy, Stages>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization for column-major-interleaved output
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename OperatorClass,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Instruction-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Number of stages used in the multistage mainloop
|
||||
int Stages,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator,
|
||||
/// Number of Interleaved K
|
||||
int InterleavedK>
|
||||
struct DefaultMma<ElementA, LayoutA, kAlignmentA, ElementB, LayoutB,
|
||||
kAlignmentB, ElementAccumulator,
|
||||
layout::ColumnMajorInterleaved<InterleavedK>, OperatorClass,
|
||||
ArchTag, ThreadblockShape, WarpShape, InstructionShape,
|
||||
Stages, Operator, true> {
|
||||
// Define the MmaCore components
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, LayoutA,
|
||||
ElementB, LayoutB, ElementAccumulator,
|
||||
layout::ColumnMajorInterleaved<InterleavedK>, OperatorClass, Stages,
|
||||
Operator, true>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using AccessTypeA = cutlass::Array<ElementA, kAlignmentA>;
|
||||
using IteratorA =
|
||||
cutlass::transform::threadblock::PredicatedTileAccessIterator<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA, LayoutA, 1, ThreadMapA, AccessTypeA>;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using AccessTypeB = cutlass::Array<ElementB, kAlignmentB>;
|
||||
using IteratorB =
|
||||
cutlass::transform::threadblock::PredicatedTileAccessIterator<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB, LayoutB, 0, ThreadMapB, AccessTypeB>;
|
||||
|
||||
// Define the threadblock-scoped multistage matrix multiply
|
||||
using ThreadblockMma = cutlass::gemm::threadblock::MmaMultistage<
|
||||
typename MmaCore::Shape, IteratorA, typename MmaCore::SmemIteratorA,
|
||||
MmaCore::kCacheOpA, IteratorB, typename MmaCore::SmemIteratorB,
|
||||
MmaCore::kCacheOpB, ElementAccumulator, layout::RowMajor,
|
||||
typename MmaCore::MmaPolicy, Stages>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization for SIMT IDP4A Kernels
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -40,6 +40,8 @@
|
||||
#include "cutlass/gemm/warp/mma.h"
|
||||
#include "cutlass/gemm/threadblock/mma_pipelined.h"
|
||||
#include "cutlass/gemm/threadblock/mma_singlestage.h"
|
||||
#include "cutlass/arch/cache_operation.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
@@ -86,6 +88,17 @@ template <
|
||||
/// Store the accumulators in row major or column major. Row major is used
|
||||
/// when output layout is interleaved.
|
||||
bool AccumulatorsInRowMajor = false
|
||||
/// Cache operation of operand A
|
||||
, cutlass::arch::CacheOperation::Kind CacheOpA =
|
||||
cutlass::arch::CacheOperation::Global,
|
||||
/// Cache operation of operand B
|
||||
cutlass::arch::CacheOperation::Kind CacheOpB =
|
||||
cutlass::arch::CacheOperation::Global,
|
||||
/// per-element transformation for elements of A
|
||||
ComplexTransform TransformA = ComplexTransform::kNone,
|
||||
/// per-element transformation for elements of B
|
||||
ComplexTransform TransformB = ComplexTransform::kNone,
|
||||
bool IsComplex = false // (is_complex<ElementA>::value || is_complex<ElementB>::value)
|
||||
>
|
||||
struct DefaultMmaCore;
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -598,6 +598,523 @@ struct DefaultMmaCore<Shape_, WarpShape_, InstructionShape_, ElementA_,
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
/// Below is for arch::OpMultiplyAddFastF16
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization:
|
||||
///
|
||||
/// A: column-major
|
||||
/// B: row-major
|
||||
/// Operator: tensor op class
|
||||
///
|
||||
/// This uses the default warp-level operator given tile sizes
|
||||
template <
|
||||
/// Shape of threadblock-scoped matrix multiply operator (concept:
|
||||
/// GemmShape)
|
||||
typename Shape_,
|
||||
/// Shape of warp-level matrix multiply operator (concept: GemmShape)
|
||||
typename WarpShape_,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC_>
|
||||
struct DefaultMmaCore<Shape_, WarpShape_, InstructionShape_, float,
|
||||
layout::ColumnMajor, float, layout::RowMajor, float,
|
||||
LayoutC_, arch::OpClassTensorOp, 2,
|
||||
arch::OpMultiplyAddFastF16> {
|
||||
using Shape = Shape_;
|
||||
using WarpShape = WarpShape_;
|
||||
using InstructionShape = InstructionShape_;
|
||||
using ElementA = float;
|
||||
using LayoutA = layout::ColumnMajor;
|
||||
using ElementB = float;
|
||||
using LayoutB = layout::RowMajor;
|
||||
using ElementC = float;
|
||||
using LayoutC = LayoutC_;
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Number of warps present
|
||||
using WarpCount = GemmShape<
|
||||
Shape::kM / WarpShape::kM,
|
||||
Shape::kN / WarpShape::kN,
|
||||
Shape::kK / WarpShape::kK
|
||||
>;
|
||||
|
||||
// Divisility requirements
|
||||
static_assert(
|
||||
!(Shape::kM % WarpShape::kM) &&
|
||||
!(Shape::kN % WarpShape::kN),
|
||||
"Threadblock-scoped GEMM should be divisible by warp-scoped GEMM size."
|
||||
);
|
||||
|
||||
/// Number of threads per warp
|
||||
static int const kWarpSize = warp::WarpSize<arch::OpClassTensorOp>::value;
|
||||
|
||||
/// Number of threads total
|
||||
static int const kThreads = WarpCount::kCount * kWarpSize;
|
||||
|
||||
/// Size of a threadblock-scoped access
|
||||
static int const kAccessSizeInBits = 256;
|
||||
|
||||
/// Default Operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
//
|
||||
// Shared memory layouts
|
||||
//
|
||||
|
||||
using SmemLayoutA = layout::ColumnMajorTensorOpMultiplicandCongruous<
|
||||
sizeof_bits<half_t>::value, int(128 / sizeof(half_t))>;
|
||||
|
||||
// Shared memory layout
|
||||
using SmemLayoutB =
|
||||
layout::RowMajorTensorOpMultiplicandCongruous<sizeof_bits<half_t>::value,
|
||||
int(128 / sizeof(half_t))>;
|
||||
|
||||
//
|
||||
// Iterators to write to shared memory
|
||||
//
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using IteratorThreadMapA = transform::PitchLinearWarpRakedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kM, Shape::kK>,
|
||||
kThreads,
|
||||
layout::PitchLinearShape<8, 4>,
|
||||
kAccessSizeInBits / sizeof_bits<ElementA>::value
|
||||
>;
|
||||
|
||||
/// Shared memory iterator to A operand
|
||||
using SmemIteratorA = transform::threadblock::RegularTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kK>,
|
||||
half_t,
|
||||
SmemLayoutA,
|
||||
1,
|
||||
IteratorThreadMapA
|
||||
>;
|
||||
|
||||
/// ThreadMap of iterator B
|
||||
using IteratorThreadMapB = transform::PitchLinearWarpRakedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kN, Shape::kK>,
|
||||
kThreads,
|
||||
layout::PitchLinearShape<8, 4>,
|
||||
kAccessSizeInBits / sizeof_bits<ElementB>::value
|
||||
>;
|
||||
|
||||
/// Shared memory iterator to B operand
|
||||
using SmemIteratorB = transform::threadblock::RegularTileIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>,
|
||||
half_t,
|
||||
SmemLayoutB,
|
||||
0,
|
||||
IteratorThreadMapB
|
||||
>;
|
||||
|
||||
//
|
||||
// Warp-level matrix multiply operator
|
||||
//
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using MmaTensorOp = typename cutlass::gemm::warp::DefaultMmaTensorOp<
|
||||
WarpShape, InstructionShape, half_t, SmemLayoutA, half_t, SmemLayoutB,
|
||||
ElementC, LayoutC, Operator, WarpCount::kK>::Type;
|
||||
|
||||
/// Policy used to define MmaPipelined
|
||||
using MmaPolicy = MmaPolicy<
|
||||
MmaTensorOp,
|
||||
MatrixShape<0, 0>,
|
||||
MatrixShape<0, 0>,
|
||||
WarpCount::kK
|
||||
>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization:
|
||||
///
|
||||
/// A: row-major
|
||||
/// B: column-major
|
||||
/// Operator: tensor op class
|
||||
///
|
||||
/// This uses the default warp-level operator given tile sizes
|
||||
template <
|
||||
/// Shape of threadblock-scoped matrix multiply operator (concept:
|
||||
/// GemmShape)
|
||||
typename Shape_,
|
||||
/// Shape of warp-level matrix multiply operator (concept: GemmShape)
|
||||
typename WarpShape_,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC_>
|
||||
struct DefaultMmaCore<Shape_, WarpShape_, InstructionShape_, float,
|
||||
layout::RowMajor, float, layout::ColumnMajor, float,
|
||||
LayoutC_, arch::OpClassTensorOp, 2,
|
||||
arch::OpMultiplyAddFastF16> {
|
||||
using Shape = Shape_;
|
||||
using WarpShape = WarpShape_;
|
||||
using InstructionShape = InstructionShape_;
|
||||
using ElementA = float;
|
||||
using LayoutA = layout::RowMajor;
|
||||
using ElementB = float;
|
||||
using LayoutB = layout::ColumnMajor;
|
||||
using ElementC = float;
|
||||
using LayoutC = LayoutC_;
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Number of warps present
|
||||
using WarpCount = GemmShape<
|
||||
Shape::kM / WarpShape::kM,
|
||||
Shape::kN / WarpShape::kN,
|
||||
Shape::kK / WarpShape::kK
|
||||
>;
|
||||
|
||||
// Divisility requirements
|
||||
static_assert(
|
||||
!(Shape::kM % WarpShape::kM) &&
|
||||
!(Shape::kN % WarpShape::kN),
|
||||
"Threadblock-scoped GEMM should be divisible by warp-scoped GEMM size."
|
||||
);
|
||||
|
||||
/// Number of threads per warp
|
||||
static int const kWarpSize = warp::WarpSize<arch::OpClassTensorOp>::value;
|
||||
|
||||
/// Number of threads total
|
||||
static int const kThreads = WarpCount::kCount * kWarpSize;
|
||||
|
||||
/// Size of a threadblock-scoped access
|
||||
static int const kAccessSizeInBits = 256;
|
||||
|
||||
/// Default Operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
// Warp thread arrangement
|
||||
static int const kWarpThreadArrangementContiguousA =
|
||||
Shape::kK / (kAccessSizeInBits / sizeof_bits<ElementA>::value);
|
||||
|
||||
static int const kWarpThreadArrangementStridedA =
|
||||
kWarpSize / kWarpThreadArrangementContiguousA;
|
||||
|
||||
static int const kWarpThreadArrangementContiguousB =
|
||||
Shape::kK / (kAccessSizeInBits / sizeof_bits<ElementA>::value);
|
||||
|
||||
static int const kWarpThreadArrangementStridedB =
|
||||
kWarpSize / kWarpThreadArrangementContiguousB;
|
||||
|
||||
//
|
||||
// Shared memory layouts
|
||||
//
|
||||
|
||||
using SmemLayoutA =
|
||||
layout::RowMajorTensorOpMultiplicandCrosswise<sizeof_bits<half_t>::value,
|
||||
Shape::kK>;
|
||||
|
||||
// Shared memory layout
|
||||
using SmemLayoutB = layout::ColumnMajorTensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<half_t>::value, Shape::kK>;
|
||||
|
||||
//
|
||||
// Iterators to write to shared memory
|
||||
//
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using IteratorThreadMapA = transform::PitchLinearWarpRakedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kK, Shape::kM>, kThreads,
|
||||
layout::PitchLinearShape<kWarpThreadArrangementContiguousA,
|
||||
kWarpThreadArrangementStridedA>,
|
||||
kAccessSizeInBits / sizeof_bits<ElementA>::value>;
|
||||
|
||||
/// Shared memory iterator to A operand
|
||||
using SmemIteratorA = transform::threadblock::RegularTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kK>,
|
||||
half_t,
|
||||
SmemLayoutA,
|
||||
0,
|
||||
IteratorThreadMapA
|
||||
>;
|
||||
|
||||
/// ThreadMap of iterator B
|
||||
using IteratorThreadMapB = transform::PitchLinearWarpRakedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kK, Shape::kN>, kThreads,
|
||||
layout::PitchLinearShape<kWarpThreadArrangementContiguousB,
|
||||
kWarpThreadArrangementStridedB>,
|
||||
kAccessSizeInBits / sizeof_bits<ElementB>::value>;
|
||||
|
||||
/// Shared memory iterator to B operand
|
||||
using SmemIteratorB = transform::threadblock::RegularTileIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>,
|
||||
half_t,
|
||||
SmemLayoutB,
|
||||
1,
|
||||
IteratorThreadMapB
|
||||
>;
|
||||
|
||||
//
|
||||
// Warp-level matrix multiply operator
|
||||
//
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using MmaTensorOp = typename cutlass::gemm::warp::DefaultMmaTensorOp<
|
||||
WarpShape, InstructionShape, half_t, SmemLayoutA, half_t, SmemLayoutB,
|
||||
ElementC, LayoutC, Operator, WarpCount::kK>::Type;
|
||||
|
||||
/// Policy used to define MmaPipelined
|
||||
using MmaPolicy = MmaPolicy<
|
||||
MmaTensorOp,
|
||||
MatrixShape<0, 0>,
|
||||
MatrixShape<0, 0>,
|
||||
WarpCount::kK
|
||||
>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization:
|
||||
///
|
||||
/// A: row-major
|
||||
/// B: row-major
|
||||
/// Operator: tensor op class
|
||||
///
|
||||
/// This uses the default warp-level operator given tile sizes
|
||||
template <
|
||||
/// Shape of threadblock-scoped matrix multiply operator (concept:
|
||||
/// GemmShape)
|
||||
typename Shape_,
|
||||
/// Shape of warp-level matrix multiply operator (concept: GemmShape)
|
||||
typename WarpShape_,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC_>
|
||||
struct DefaultMmaCore<Shape_, WarpShape_, InstructionShape_, float,
|
||||
layout::RowMajor, float, layout::RowMajor, float,
|
||||
LayoutC_, arch::OpClassTensorOp, 2,
|
||||
arch::OpMultiplyAddFastF16> {
|
||||
using Shape = Shape_;
|
||||
using WarpShape = WarpShape_;
|
||||
using InstructionShape = InstructionShape_;
|
||||
using ElementA = float;
|
||||
using LayoutA = layout::RowMajor;
|
||||
using ElementB = float;
|
||||
using LayoutB = layout::RowMajor;
|
||||
using ElementC = float;
|
||||
using LayoutC = LayoutC_;
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Number of warps present
|
||||
using WarpCount = GemmShape<
|
||||
Shape::kM / WarpShape::kM,
|
||||
Shape::kN / WarpShape::kN,
|
||||
Shape::kK / WarpShape::kK
|
||||
>;
|
||||
|
||||
// Divisility requirements
|
||||
static_assert(
|
||||
!(Shape::kM % WarpShape::kM) &&
|
||||
!(Shape::kN % WarpShape::kN),
|
||||
"Threadblock-scoped GEMM should be divisible by warp-scoped GEMM size."
|
||||
);
|
||||
|
||||
/// Number of threads per warp
|
||||
static int const kWarpSize = warp::WarpSize<arch::OpClassTensorOp>::value;
|
||||
|
||||
/// Number of threads total
|
||||
static int const kThreads = WarpCount::kCount * kWarpSize;
|
||||
|
||||
/// Size of a threadblock-scoped access
|
||||
static int const kAccessSizeInBits = 256;
|
||||
|
||||
/// Default Operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
// Warp thread arrangement
|
||||
static int const kWarpThreadArrangementContiguousA =
|
||||
Shape::kK / (kAccessSizeInBits / sizeof_bits<ElementA>::value);
|
||||
|
||||
static int const kWarpThreadArrangementStridedA =
|
||||
kWarpSize / kWarpThreadArrangementContiguousA;
|
||||
|
||||
//
|
||||
// Shared memory layouts
|
||||
//
|
||||
|
||||
using SmemLayoutA = layout::RowMajorTensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<half_t>::value, Shape::kK>;
|
||||
|
||||
// Shared memory layout
|
||||
using SmemLayoutB = layout::RowMajorTensorOpMultiplicandCongruous<
|
||||
sizeof_bits<half_t>::value, int(128 / sizeof(half_t))>;
|
||||
|
||||
//
|
||||
// Iterators to write to shared memory
|
||||
//
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using IteratorThreadMapA = transform::PitchLinearWarpRakedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kK, Shape::kM>, kThreads,
|
||||
layout::PitchLinearShape<kWarpThreadArrangementContiguousA,
|
||||
kWarpThreadArrangementStridedA>,
|
||||
kAccessSizeInBits / sizeof_bits<ElementA>::value>;
|
||||
|
||||
/// Shared memory iterator to A operand
|
||||
using SmemIteratorA = transform::threadblock::RegularTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kK>,
|
||||
half_t,
|
||||
SmemLayoutA,
|
||||
0,
|
||||
IteratorThreadMapA
|
||||
>;
|
||||
|
||||
/// ThreadMap of iterator B
|
||||
using IteratorThreadMapB = transform::PitchLinearWarpRakedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kN, Shape::kK>,
|
||||
kThreads,
|
||||
layout::PitchLinearShape<8, 4>,
|
||||
kAccessSizeInBits / sizeof_bits<ElementB>::value
|
||||
>;
|
||||
|
||||
/// Shared memory iterator to B operand
|
||||
using SmemIteratorB = transform::threadblock::RegularTileIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>,
|
||||
half_t,
|
||||
SmemLayoutB,
|
||||
0,
|
||||
IteratorThreadMapB
|
||||
>;
|
||||
|
||||
//
|
||||
// Warp-level matrix multiply operator
|
||||
//
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using MmaTensorOp = typename cutlass::gemm::warp::DefaultMmaTensorOp<
|
||||
WarpShape, InstructionShape, half_t, SmemLayoutA, half_t, SmemLayoutB,
|
||||
ElementC, LayoutC, Operator, WarpCount::kK>::Type;
|
||||
|
||||
/// Policy used to define MmaPipelined
|
||||
using MmaPolicy = MmaPolicy<
|
||||
MmaTensorOp,
|
||||
MatrixShape<0, 0>,
|
||||
MatrixShape<0, 0>,
|
||||
WarpCount::kK
|
||||
>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization:
|
||||
///
|
||||
/// A: column-major
|
||||
/// B: column-major
|
||||
/// Operator: tensor op class
|
||||
///
|
||||
/// This uses the default warp-level operator given tile sizes
|
||||
template <
|
||||
/// Shape of threadblock-scoped matrix multiply operator (concept:
|
||||
/// GemmShape)
|
||||
typename Shape_,
|
||||
/// Shape of warp-level matrix multiply operator (concept: GemmShape)
|
||||
typename WarpShape_,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC_>
|
||||
struct DefaultMmaCore<Shape_, WarpShape_, InstructionShape_, float,
|
||||
layout::ColumnMajor, float, layout::ColumnMajor, float,
|
||||
LayoutC_, arch::OpClassTensorOp, 2,
|
||||
arch::OpMultiplyAddFastF16> {
|
||||
using Shape = Shape_;
|
||||
using WarpShape = WarpShape_;
|
||||
using InstructionShape = InstructionShape_;
|
||||
using ElementA = float;
|
||||
using LayoutA = layout::ColumnMajor;
|
||||
using ElementB = float;
|
||||
using LayoutB = layout::ColumnMajor;
|
||||
using ElementC = float;
|
||||
using LayoutC = LayoutC_;
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Number of warps present
|
||||
using WarpCount = GemmShape<Shape::kM / WarpShape::kM,
|
||||
Shape::kN / WarpShape::kN,
|
||||
Shape::kK / WarpShape::kK>;
|
||||
|
||||
// Divisility requirements
|
||||
static_assert(
|
||||
!(Shape::kM % WarpShape::kM) && !(Shape::kN % WarpShape::kN),
|
||||
"Threadblock-scoped GEMM should be divisible by warp-scoped GEMM size.");
|
||||
|
||||
/// Number of threads per warp
|
||||
static int const kWarpSize = warp::WarpSize<arch::OpClassTensorOp>::value;
|
||||
|
||||
/// Number of threads total
|
||||
static int const kThreads = WarpCount::kCount * kWarpSize;
|
||||
|
||||
/// Size of a threadblock-scoped access
|
||||
static int const kAccessSizeInBits = 256;
|
||||
|
||||
/// Default Operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
// Warp thread arrangement
|
||||
static int const kWarpThreadArrangementContiguousB =
|
||||
Shape::kK / (kAccessSizeInBits / sizeof_bits<ElementA>::value);
|
||||
|
||||
static int const kWarpThreadArrangementStridedB =
|
||||
kWarpSize / kWarpThreadArrangementContiguousB;
|
||||
|
||||
//
|
||||
// Shared memory layouts
|
||||
//
|
||||
|
||||
using SmemLayoutA = layout::ColumnMajorTensorOpMultiplicandCongruous<
|
||||
sizeof_bits<half_t>::value, int(128 / sizeof(half_t))>;
|
||||
|
||||
// Shared memory layout
|
||||
using SmemLayoutB = layout::ColumnMajorTensorOpMultiplicandCrosswise<
|
||||
sizeof_bits<half_t>::value, Shape::kK>;
|
||||
|
||||
//
|
||||
// Iterators to write to shared memory
|
||||
//
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using IteratorThreadMapA = transform::PitchLinearWarpRakedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kM, Shape::kK>, kThreads,
|
||||
layout::PitchLinearShape<8, 4>,
|
||||
kAccessSizeInBits / sizeof_bits<ElementA>::value>;
|
||||
|
||||
/// Shared memory iterator to A operand
|
||||
using SmemIteratorA = transform::threadblock::RegularTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kK>, half_t, SmemLayoutA, 1,
|
||||
IteratorThreadMapA>;
|
||||
|
||||
/// ThreadMap of iterator B
|
||||
using IteratorThreadMapB = transform::PitchLinearWarpRakedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kK, Shape::kN>, kThreads,
|
||||
layout::PitchLinearShape<kWarpThreadArrangementContiguousB,
|
||||
kWarpThreadArrangementStridedB>,
|
||||
kAccessSizeInBits / sizeof_bits<ElementB>::value>;
|
||||
|
||||
/// Shared memory iterator to B operand
|
||||
using SmemIteratorB = transform::threadblock::RegularTileIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>, half_t, SmemLayoutB, 1,
|
||||
IteratorThreadMapB>;
|
||||
|
||||
//
|
||||
// Warp-level matrix multiply operator
|
||||
//
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using MmaTensorOp = typename cutlass::gemm::warp::DefaultMmaTensorOp<
|
||||
WarpShape, InstructionShape, half_t, SmemLayoutA, half_t, SmemLayoutB,
|
||||
ElementC, LayoutC, Operator, WarpCount::kK>::Type;
|
||||
|
||||
/// Policy used to define MmaPipelined
|
||||
using MmaPolicy = MmaPolicy<MmaTensorOp, MatrixShape<0, 0>, MatrixShape<0, 0>,
|
||||
WarpCount::kK>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization:
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -0,0 +1,130 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Template for a multistage GEMM kernel. Does not compute batching or support split-K.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/arch/arch.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_sm80.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma.h"
|
||||
#include "cutlass/gemm/threadblock/mma_planar_complex_multistage.h"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/transform/threadblock/predicated_tile_iterator.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA_,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA_,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB_,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB_,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator_,
|
||||
/// Layout type for C and D matrix operands
|
||||
typename LayoutC_,
|
||||
/// Operator class tag
|
||||
typename OperatorClass_,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag_,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape_,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape_,
|
||||
/// Instruction-level tile size (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Complex transformation on operand A
|
||||
ComplexTransform TransformA = ComplexTransform::kNone,
|
||||
/// Complex transformation on operand B
|
||||
ComplexTransform TransformB = ComplexTransform::kNone,
|
||||
/// Math operator tag (e.g. arch::OpMultiplyAdd)
|
||||
typename Operator = arch::OpMultiplyAdd
|
||||
>
|
||||
struct DefaultMmaPlanarComplexMultistage {
|
||||
|
||||
// Construct a planar complex variant from the real-valued variant
|
||||
using RealMmaMultistage = typename DefaultMma<
|
||||
ElementA_,
|
||||
LayoutA_,
|
||||
kAlignmentA,
|
||||
ElementB_,
|
||||
LayoutB_,
|
||||
kAlignmentB,
|
||||
ElementAccumulator_,
|
||||
LayoutC_,
|
||||
OperatorClass_,
|
||||
ArchTag_,
|
||||
ThreadblockShape_,
|
||||
WarpShape_,
|
||||
InstructionShape_,
|
||||
Stages,
|
||||
Operator
|
||||
>::ThreadblockMma;
|
||||
|
||||
using ThreadblockMma = MmaPlanarComplexMultistage<
|
||||
ThreadblockShape_,
|
||||
typename RealMmaMultistage::IteratorA,
|
||||
typename RealMmaMultistage::SmemIteratorA,
|
||||
cutlass::arch::CacheOperation::Global,
|
||||
typename RealMmaMultistage::IteratorB,
|
||||
typename RealMmaMultistage::SmemIteratorB,
|
||||
cutlass::arch::CacheOperation::Global,
|
||||
ElementAccumulator_,
|
||||
LayoutC_,
|
||||
typename RealMmaMultistage::Policy,
|
||||
Stages,
|
||||
TransformA,
|
||||
TransformB
|
||||
>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,154 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Template for a multistage GEMM kernel. Does not compute batching or support split-K.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/arch/arch.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_sm80.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/transform/threadblock/predicated_tile_iterator.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA_,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA_,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB_,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB_,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator_,
|
||||
/// Layout type for C and D matrix operands
|
||||
typename LayoutC_,
|
||||
/// Operator class tag
|
||||
typename OperatorClass_,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag_,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape_,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape_,
|
||||
/// Instruction-level tile size (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Complex transformation on operand A
|
||||
ComplexTransform TransformA = ComplexTransform::kNone,
|
||||
/// Complex transformation on operand B
|
||||
ComplexTransform TransformB = ComplexTransform::kNone,
|
||||
/// Multiply-add operator (arch::OpMultiplyAddComplex, arch::OpMultiplyGaussianComplex)
|
||||
typename Operator = arch::OpMultiplyAddComplex,
|
||||
/// Store the accumulators in row major or column major. Row major is used
|
||||
/// when output layout is interleaved.
|
||||
bool AccumulatorsInRowMajor = false>
|
||||
struct DefaultMultistageMmaComplex;
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Specialization for row-major output
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename OperatorClass,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Instruction-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Number of stages used in the multistage mainloop
|
||||
int Stages,
|
||||
/// Complex transformation on operand A
|
||||
ComplexTransform TransformA,
|
||||
/// Complex transformation on operand B
|
||||
ComplexTransform TransformB,
|
||||
/// Multiply-add operator (arch::OpMultiplyAddComplex, arch::OpMultiplyGaussianComplex)
|
||||
typename Operator>
|
||||
struct DefaultMultistageMmaComplex<ElementA, LayoutA, ElementB, LayoutB,
|
||||
ElementAccumulator, layout::RowMajor, OperatorClass,
|
||||
ArchTag, ThreadblockShape, WarpShape,
|
||||
InstructionShape, Stages, TransformA, TransformB, Operator> {
|
||||
// Define the MmaCore components
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMultistageMmaComplexCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, LayoutA,
|
||||
ElementB, LayoutB, ElementAccumulator, layout::RowMajor, OperatorClass,
|
||||
Stages, TransformA, TransformB, Operator>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using AccessTypeA = cutlass::Array<ElementA, ThreadMapA::kElementsPerAccess>;
|
||||
using IteratorA =
|
||||
cutlass::transform::threadblock::PredicatedTileAccessIterator<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA, LayoutA, 1, ThreadMapA, AccessTypeA>;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using AccessTypeB = cutlass::Array<ElementB, ThreadMapB::kElementsPerAccess>;
|
||||
using IteratorB =
|
||||
cutlass::transform::threadblock::PredicatedTileAccessIterator<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB, LayoutB, 0, ThreadMapB, AccessTypeB>;
|
||||
|
||||
// Define the threadblock-scoped multistage matrix multiply
|
||||
using ThreadblockMma = cutlass::gemm::threadblock::MmaMultistage<
|
||||
typename MmaCore::Shape, IteratorA, typename MmaCore::SmemIteratorA,
|
||||
MmaCore::kCacheOpA, IteratorB, typename MmaCore::SmemIteratorB,
|
||||
MmaCore::kCacheOpB, ElementAccumulator, layout::RowMajor,
|
||||
typename MmaCore::MmaPolicy, Stages>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,113 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Defines basic properties needed by CTA-level GEMMs assuming
|
||||
expectations about data layout of the global memory fragments, data types,
|
||||
and internal tile sizes.
|
||||
|
||||
Partial specializations for threadblock::Mma operations targeting TensorOp
|
||||
instructions.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/complex.h"
|
||||
|
||||
#include "cutlass/layout/tensor_op_multiplicand_sm75.h"
|
||||
#include "cutlass/layout/tensor_op_multiplicand_sm80.h"
|
||||
|
||||
#include "cutlass/gemm/warp/mma_simt_policy.h"
|
||||
#include "cutlass/gemm/warp/mma_simt.h"
|
||||
#include "cutlass/gemm/warp/default_mma_tensor_op.h"
|
||||
#include "cutlass/gemm/warp/mma_tensor_op_tile_iterator_sm80.h"
|
||||
|
||||
#include "cutlass/gemm/threadblock/default_mma_core.h"
|
||||
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/transform/pitch_linear_thread_map.h"
|
||||
|
||||
#include "cutlass/transform/threadblock/regular_tile_access_iterator_tensor_op.h"
|
||||
#include "cutlass/transform/threadblock/regular_tile_access_iterator_pitch_linear.h"
|
||||
#include "cutlass/transform/threadblock/regular_tile_access_iterator_tensor_op_sm80.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Template defininng default matrix multiply operators inferred from
|
||||
/// threadblock tile size, global memory data layout, and target math
|
||||
/// instruction.
|
||||
template <
|
||||
/// Shape of threadblock-scoped matrix multiply operator
|
||||
typename Shape,
|
||||
/// Shape of warp-level matrix multiply operator
|
||||
typename WarpShape,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Element data type of A operand
|
||||
typename ElementA,
|
||||
/// Layout of operand A
|
||||
typename LayoutA,
|
||||
/// Element data type of B operand
|
||||
typename ElementB,
|
||||
/// Layout of operand B
|
||||
typename LayoutB,
|
||||
/// Data type of accumulator
|
||||
typename ElementC,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC,
|
||||
/// Indicates type of math operator (arch::OpClassSimt or arch::OpClassTensorOp)
|
||||
typename OperatorClass,
|
||||
/// Number of stages
|
||||
int Stages,
|
||||
/// Complex transformation on operand A
|
||||
ComplexTransform TransformA,
|
||||
/// Complex transformation on operand B
|
||||
ComplexTransform TransformB,
|
||||
/// Multiply-add operator (arch::OpMultiplyAddComplex, arch::OpMultiplyGaussianComplex)
|
||||
typename Operator = arch::OpMultiplyAddComplex,
|
||||
/// Cache operation of operand A
|
||||
cutlass::arch::CacheOperation::Kind CacheOpA =
|
||||
cutlass::arch::CacheOperation::Global,
|
||||
/// Cache operation of operand B
|
||||
cutlass::arch::CacheOperation::Kind CacheOpB =
|
||||
cutlass::arch::CacheOperation::Global>
|
||||
struct DefaultMultistageMmaComplexCore;
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,140 +0,0 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 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 Template for a threadblock-scoped GEMV kernel.
|
||||
*/
|
||||
|
||||
#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"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace threadblock {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Structure to compute the matrix-vector product using SIMT math instructions.
|
||||
template <
|
||||
class Core_ //< GemvCore
|
||||
>
|
||||
class Gemv {
|
||||
public:
|
||||
using Shape = typename Core_::Shape;
|
||||
|
||||
/// The MMA operator that computes GEMV
|
||||
using Operator = typename Core_::Operator;
|
||||
|
||||
/// Iterates over A in global memory
|
||||
using IteratorA = typename Core_::IteratorA;
|
||||
|
||||
/// Iterates over B in global memory
|
||||
using IteratorB = typename Core_::IteratorB;
|
||||
|
||||
/// Fragment of operand C loaded from global memory
|
||||
using IteratorC = typename Core_::IteratorC;
|
||||
|
||||
/// Fragment of operand A loaded from global memory
|
||||
using FragmentA = typename IteratorA::Fragment;
|
||||
|
||||
/// Fragment of operand B loaded from global memory
|
||||
using FragmentB = typename IteratorB::Fragment;
|
||||
|
||||
/// Fragment of operand accumulator loaded/stored to global memory
|
||||
using FragmentC = typename Operator::FragmentC;
|
||||
|
||||
/// Shape of the per-thread GEMV operation
|
||||
using ThreadShape = typename Core_::ThreadShape;
|
||||
|
||||
public:
|
||||
CUTLASS_DEVICE
|
||||
Gemv() { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
GemmCoord const &problem_size, ///< problem size of batched GEMV
|
||||
FragmentC &accum, ///< destination accumulator tile
|
||||
IteratorA iterator_A, ///< iterator over A operand in global memory
|
||||
IteratorB iterator_B, ///< iterator over B operand in global memory
|
||||
FragmentC const &src_accum) { ///< source accumualtor tile
|
||||
|
||||
//
|
||||
// Prologue
|
||||
//
|
||||
|
||||
FragmentA frag_A;
|
||||
FragmentB frag_B;
|
||||
frag_A.clear();
|
||||
frag_B.clear();
|
||||
|
||||
iterator_A.load(frag_A);
|
||||
iterator_B.load(frag_B);
|
||||
++iterator_A;
|
||||
++iterator_B;
|
||||
|
||||
//
|
||||
// Mainloop
|
||||
//
|
||||
Operator thread_mma;
|
||||
int gemm_k = problem_size.k();
|
||||
|
||||
if (gemm_k < Shape::kK)
|
||||
{
|
||||
iterator_A.clear_mask();
|
||||
iterator_B.clear_mask();
|
||||
}
|
||||
|
||||
// iterate over K to accumulate result
|
||||
CUTLASS_GEMM_LOOP
|
||||
for (; gemm_k > 0; gemm_k -= Shape::kK) {
|
||||
thread_mma(accum, frag_A, frag_B, accum);
|
||||
|
||||
iterator_A.load(frag_A);
|
||||
iterator_B.load(frag_B);
|
||||
++iterator_A;
|
||||
++iterator_B;
|
||||
|
||||
if (gemm_k < Shape::kK)
|
||||
{
|
||||
iterator_A.clear_mask();
|
||||
iterator_B.clear_mask();
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -0,0 +1,526 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Template for a double-buffered threadblock-scoped GEMM kernel.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/aligned_buffer.h"
|
||||
#include "cutlass/arch/memory.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/threadblock/mma_base.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace threadblock {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Structure to compute the matrix product targeting CUDA cores and SIMT math
|
||||
/// instructions.
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape_,
|
||||
/// Iterates over tiles of A operand in global memory
|
||||
// (concept: ReadableTileIterator | ForwardTileIterator |
|
||||
// MaskedTileIterator)
|
||||
typename IteratorA_,
|
||||
/// Iterates over tiles of A operand in shared memory
|
||||
/// (concept: WriteableTileIterator | RandomAccessTileIterator)
|
||||
typename SmemIteratorA_,
|
||||
/// Cache operation for operand A
|
||||
cutlass::arch::CacheOperation::Kind CacheOpA,
|
||||
/// Iterates over tiles of B operand in global memory
|
||||
// (concept: ReadableTileIterator | ForwardTileIterator |
|
||||
// MaskedTileIterator)
|
||||
typename IteratorB_,
|
||||
/// Iterates over tiles of B operand in shared memory
|
||||
/// (concept: WriteableTileIterator | RandomAccessTileIterator)
|
||||
typename SmemIteratorB_,
|
||||
/// Cache operation for operand B
|
||||
cutlass::arch::CacheOperation::Kind CacheOpB,
|
||||
/// Data type of accumulator matrix
|
||||
typename ElementC_,
|
||||
/// Data type of accumulator matrix
|
||||
typename LayoutC_,
|
||||
/// Policy describing tuning details (concept: MmaPolicy)
|
||||
typename Policy_,
|
||||
/// Number of stages,
|
||||
int Stages,
|
||||
/// Used for partial specialization
|
||||
typename Enable = bool>
|
||||
class MmaMultistage :
|
||||
public MmaBase<Shape_, Policy_, Stages> {
|
||||
public:
|
||||
///< Base class
|
||||
using Base = MmaBase<Shape_, Policy_, Stages>;
|
||||
///< Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
using Shape = Shape_;
|
||||
///< Iterates over tiles of A operand in global memory
|
||||
using IteratorA = IteratorA_;
|
||||
///< Iterates over tiles of B operand in global memory
|
||||
using IteratorB = IteratorB_;
|
||||
///< Data type of accumulator matrix
|
||||
using ElementC = ElementC_;
|
||||
///< Layout of accumulator matrix
|
||||
using LayoutC = LayoutC_;
|
||||
///< Policy describing tuning details
|
||||
using Policy = Policy_;
|
||||
|
||||
using SmemIteratorA = SmemIteratorA_;
|
||||
using SmemIteratorB = SmemIteratorB_;
|
||||
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpA = CacheOpA;
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpB = CacheOpB;
|
||||
|
||||
//
|
||||
// Dependent types
|
||||
//
|
||||
|
||||
/// Fragment of accumulator tile
|
||||
using FragmentC = typename Policy::Operator::FragmentC;
|
||||
|
||||
/// Warp-level Mma
|
||||
using Operator = typename Policy::Operator;
|
||||
|
||||
/// Minimum architecture is Sm80 to support cp.async
|
||||
using ArchTag = arch::Sm80;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = Operator::kTransformA;
|
||||
|
||||
/// Complex transform on B operand
|
||||
static ComplexTransform const kTransformB = Operator::kTransformB;
|
||||
|
||||
/// Internal structure exposed for introspection.
|
||||
struct Detail {
|
||||
|
||||
static_assert(Base::kWarpGemmIterations > 1,
|
||||
"The pipelined structure requires at least two warp-level "
|
||||
"GEMM operations.");
|
||||
|
||||
/// Number of cp.async instructions to load one stage of operand A
|
||||
static int const AsyncCopyIterationsPerStageA =
|
||||
IteratorA::ThreadMap::Iterations::kCount;
|
||||
|
||||
/// Number of cp.async instructions to load one stage of operand B
|
||||
static int const AsyncCopyIterationsPerStageB =
|
||||
IteratorB::ThreadMap::Iterations::kCount;
|
||||
|
||||
/// Number of stages
|
||||
static int const kStages = Stages;
|
||||
|
||||
/// Number of cp.async instructions to load on group of operand A
|
||||
static int const kAccessesPerGroupA =
|
||||
(AsyncCopyIterationsPerStageA + Base::kWarpGemmIterations - 1) / Base::kWarpGemmIterations;
|
||||
|
||||
/// Number of cp.async instructions to load on group of operand B
|
||||
static int const kAccessesPerGroupB =
|
||||
(AsyncCopyIterationsPerStageB + Base::kWarpGemmIterations - 1) / Base::kWarpGemmIterations;
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
using WarpLoadedFragmentA = typename Operator::FragmentA;
|
||||
using WarpLoadedFragmentB = typename Operator::FragmentB;
|
||||
using WarpTransformedFragmentA = typename Operator::TransformedFragmentA;
|
||||
using WarpTransformedFragmentB = typename Operator::TransformedFragmentB;
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Iterator to write threadblock-scoped tile of A operand to shared memory
|
||||
SmemIteratorA smem_iterator_A_;
|
||||
|
||||
/// Iterator to write threadblock-scoped tile of B operand to shared memory
|
||||
SmemIteratorB smem_iterator_B_;
|
||||
|
||||
public:
|
||||
|
||||
/// Construct from tensor references
|
||||
CUTLASS_DEVICE
|
||||
MmaMultistage(
|
||||
///< Shared storage needed for internal use by threadblock-scoped GEMM
|
||||
typename Base::SharedStorage &shared_storage,
|
||||
///< ID within the threadblock
|
||||
int thread_idx,
|
||||
///< ID of warp
|
||||
int warp_idx,
|
||||
///< ID of each thread within a warp
|
||||
int lane_idx
|
||||
):
|
||||
Base(shared_storage, thread_idx, warp_idx, lane_idx),
|
||||
smem_iterator_A_(shared_storage.operand_A_ref(), thread_idx),
|
||||
smem_iterator_B_(shared_storage.operand_B_ref(), thread_idx)
|
||||
{
|
||||
// Compute warp location within threadblock tile by mapping the warp_id to
|
||||
// three coordinates:
|
||||
// _m: the warp's position within the threadblock along the M dimension
|
||||
// _n: the warp's position within the threadblock along the N dimension
|
||||
// _k: the warp's position within the threadblock along the K dimension
|
||||
|
||||
int warp_idx_mn = warp_idx % (Base::WarpCount::kM * Base::WarpCount::kN);
|
||||
int warp_idx_k = warp_idx / (Base::WarpCount::kM * Base::WarpCount::kN);
|
||||
|
||||
int warp_idx_m = warp_idx_mn % Base::WarpCount::kM;
|
||||
int warp_idx_n = warp_idx_mn / Base::WarpCount::kM;
|
||||
|
||||
// Add per-warp offsets in units of warp-level tiles
|
||||
this->warp_tile_iterator_A_.add_tile_offset(
|
||||
{warp_idx_m, Base::kWarpGemmIterations * warp_idx_k});
|
||||
this->warp_tile_iterator_B_.add_tile_offset(
|
||||
{Base::kWarpGemmIterations * warp_idx_k, warp_idx_n});
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void copy_tiles_and_advance(IteratorA &iterator_A, IteratorB &iterator_B,
|
||||
int group_start_A = 0, int group_start_B = 0) {
|
||||
iterator_A.set_iteration_index(group_start_A *
|
||||
IteratorA::kAccessesPerVector);
|
||||
this->smem_iterator_A_.set_iteration_index(group_start_A);
|
||||
|
||||
// Async Copy for operand A
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < Detail::kAccessesPerGroupA; ++j) {
|
||||
if (group_start_A + j < Detail::AsyncCopyIterationsPerStageA) {
|
||||
typename IteratorA::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorA::AccessType *>(
|
||||
this->smem_iterator_A_.get());
|
||||
|
||||
int const kSrcBytes = sizeof_bits<typename IteratorA::Element>::value *
|
||||
IteratorA::ThreadMap::kElementsPerAccess /
|
||||
IteratorA::kAccessesPerVector / 8;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorA::kAccessesPerVector; ++v) {
|
||||
auto gmem_ptr = iterator_A.get();
|
||||
|
||||
cutlass::arch::cp_async<kSrcBytes, kCacheOpA>(
|
||||
dst_ptr + v, gmem_ptr, iterator_A.valid());
|
||||
|
||||
++iterator_A;
|
||||
}
|
||||
|
||||
++this->smem_iterator_A_;
|
||||
}
|
||||
}
|
||||
|
||||
iterator_B.set_iteration_index(group_start_B *
|
||||
IteratorB::kAccessesPerVector);
|
||||
this->smem_iterator_B_.set_iteration_index(group_start_B);
|
||||
|
||||
// Async Copy for operand B
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < Detail::kAccessesPerGroupB; ++j) {
|
||||
if (group_start_B + j < Detail::AsyncCopyIterationsPerStageB) {
|
||||
typename IteratorB::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorB::AccessType *>(
|
||||
this->smem_iterator_B_.get());
|
||||
|
||||
int const kSrcBytes = sizeof_bits<typename IteratorB::Element>::value *
|
||||
IteratorB::ThreadMap::kElementsPerAccess /
|
||||
IteratorB::kAccessesPerVector / 8;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorB::kAccessesPerVector; ++v) {
|
||||
auto gmem_ptr = iterator_B.get();
|
||||
|
||||
cutlass::arch::cp_async<kSrcBytes, kCacheOpB>(
|
||||
dst_ptr + v, gmem_ptr, iterator_B.valid());
|
||||
|
||||
++iterator_B;
|
||||
}
|
||||
++this->smem_iterator_B_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Perform a threadblock-scoped matrix multiply-accumulate
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
///< problem size of GEMM
|
||||
int gemm_k_iterations,
|
||||
///< destination accumulator tile
|
||||
FragmentC &accum,
|
||||
///< iterator over A operand in global memory
|
||||
IteratorA iterator_A,
|
||||
///< iterator over B operand in global memory
|
||||
IteratorB iterator_B,
|
||||
///< initial value of accumulator
|
||||
FragmentC const &src_accum) {
|
||||
|
||||
//
|
||||
// Prologue
|
||||
//
|
||||
|
||||
// Issue several complete stages
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int stage = 0; stage < Base::kStages - 1;
|
||||
++stage, --gemm_k_iterations) {
|
||||
|
||||
if (gemm_k_iterations == 0) {
|
||||
iterator_A.clear_mask();
|
||||
iterator_B.clear_mask();
|
||||
}
|
||||
|
||||
iterator_A.set_iteration_index(0);
|
||||
this->smem_iterator_A_.set_iteration_index(0);
|
||||
|
||||
// Async Copy for operand A
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < Detail::AsyncCopyIterationsPerStageA; ++j) {
|
||||
typename IteratorA::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorA::AccessType *>(
|
||||
this->smem_iterator_A_.get());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorA::kAccessesPerVector; ++v) {
|
||||
int const kSrcBytes =
|
||||
sizeof_bits<typename IteratorA::Element>::value *
|
||||
IteratorA::ThreadMap::kElementsPerAccess /
|
||||
IteratorA::kAccessesPerVector / 8;
|
||||
|
||||
int src_bytes = (iterator_A.valid() ? kSrcBytes : 0);
|
||||
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpA>(
|
||||
dst_ptr + v, iterator_A.get(), iterator_A.valid());
|
||||
|
||||
++iterator_A;
|
||||
}
|
||||
|
||||
++this->smem_iterator_A_;
|
||||
}
|
||||
|
||||
iterator_B.set_iteration_index(0);
|
||||
this->smem_iterator_B_.set_iteration_index(0);
|
||||
|
||||
// Async Copy for operand B
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < Detail::AsyncCopyIterationsPerStageB; ++j) {
|
||||
typename IteratorB::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorB::AccessType *>(
|
||||
this->smem_iterator_B_.get());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorB::kAccessesPerVector; ++v) {
|
||||
int const kSrcBytes =
|
||||
sizeof_bits<typename IteratorB::Element>::value *
|
||||
IteratorB::ThreadMap::kElementsPerAccess /
|
||||
IteratorB::kAccessesPerVector / 8;
|
||||
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpB>(
|
||||
dst_ptr + v, iterator_B.get(), iterator_B.valid());
|
||||
|
||||
++iterator_B;
|
||||
}
|
||||
|
||||
++this->smem_iterator_B_;
|
||||
}
|
||||
|
||||
// Move to the next stage
|
||||
iterator_A.add_tile_offset({0, 1});
|
||||
iterator_B.add_tile_offset({1, 0});
|
||||
|
||||
this->smem_iterator_A_.add_tile_offset({0, 1});
|
||||
this->smem_iterator_B_.add_tile_offset({1, 0});
|
||||
|
||||
// Defines the boundary of a stage of cp.async.
|
||||
cutlass::arch::cp_async_fence();
|
||||
}
|
||||
|
||||
// Perform accumulation in the 'd' output operand
|
||||
accum = src_accum;
|
||||
|
||||
// Waits until kStages-2 stages have committed.
|
||||
cutlass::arch::cp_async_wait<Base::kStages - 2>();
|
||||
__syncthreads();
|
||||
|
||||
// Pair of fragments used to overlap shared memory loads and math
|
||||
// instructions
|
||||
WarpLoadedFragmentA warp_loaded_frag_A[2];
|
||||
WarpLoadedFragmentB warp_loaded_frag_B[2];
|
||||
WarpTransformedFragmentA warp_transformed_frag_A[2];
|
||||
WarpTransformedFragmentB warp_transformed_frag_B[2];
|
||||
|
||||
Operator warp_mma;
|
||||
|
||||
this->warp_tile_iterator_A_.set_kgroup_index(0);
|
||||
this->warp_tile_iterator_B_.set_kgroup_index(0);
|
||||
|
||||
this->warp_tile_iterator_A_.load(warp_loaded_frag_A[0]);
|
||||
this->warp_tile_iterator_B_.load(warp_loaded_frag_B[0]);
|
||||
|
||||
++this->warp_tile_iterator_A_;
|
||||
++this->warp_tile_iterator_B_;
|
||||
|
||||
if (gemm_k_iterations == 0) {
|
||||
iterator_A.clear_mask();
|
||||
iterator_B.clear_mask();
|
||||
}
|
||||
|
||||
int smem_write_stage_idx = Base::kStages - 1;
|
||||
int smem_read_stage_idx = 0;
|
||||
|
||||
warp_mma.transform(warp_transformed_frag_A[0], warp_transformed_frag_B[0],
|
||||
warp_loaded_frag_A[0], warp_loaded_frag_B[0]);
|
||||
|
||||
//
|
||||
// Mainloop
|
||||
//
|
||||
|
||||
CUTLASS_GEMM_LOOP
|
||||
for (; gemm_k_iterations > (-Base::kStages + 1);) {
|
||||
//
|
||||
// Loop over GEMM K dimension
|
||||
//
|
||||
|
||||
// Computes a warp-level GEMM on data held in shared memory
|
||||
// Each "warp_mma_k" refers to a warp-level matrix multiply-accumulate
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int warp_mma_k = 0; warp_mma_k < Base::kWarpGemmIterations;
|
||||
++warp_mma_k) {
|
||||
|
||||
// Load warp-level tiles from shared memory, wrapping to k offset if
|
||||
// this is the last group as the case may be.
|
||||
|
||||
this->warp_tile_iterator_A_.set_kgroup_index((warp_mma_k + 1) % Base::kWarpGemmIterations);
|
||||
this->warp_tile_iterator_B_.set_kgroup_index((warp_mma_k + 1) % Base::kWarpGemmIterations);
|
||||
|
||||
this->warp_tile_iterator_A_.load(warp_loaded_frag_A[(warp_mma_k + 1) % 2]);
|
||||
this->warp_tile_iterator_B_.load(warp_loaded_frag_B[(warp_mma_k + 1) % 2]);
|
||||
|
||||
++this->warp_tile_iterator_A_;
|
||||
++this->warp_tile_iterator_B_;
|
||||
|
||||
if (warp_mma_k > 0)
|
||||
warp_mma.transform(warp_transformed_frag_A[warp_mma_k % 2],
|
||||
warp_transformed_frag_B[warp_mma_k % 2],
|
||||
warp_loaded_frag_A[warp_mma_k % 2],
|
||||
warp_loaded_frag_B[warp_mma_k % 2]);
|
||||
|
||||
warp_mma(
|
||||
accum,
|
||||
warp_transformed_frag_A[warp_mma_k % 2],
|
||||
warp_transformed_frag_B[warp_mma_k % 2],
|
||||
accum
|
||||
);
|
||||
|
||||
// Issue global->shared copies for the this stage
|
||||
if (warp_mma_k < Base::kWarpGemmIterations - 1) {
|
||||
int group_start_iteration_A, group_start_iteration_B;
|
||||
|
||||
group_start_iteration_A = warp_mma_k * Detail::kAccessesPerGroupA;
|
||||
group_start_iteration_B = warp_mma_k * Detail::kAccessesPerGroupB;
|
||||
|
||||
copy_tiles_and_advance(iterator_A, iterator_B, group_start_iteration_A,
|
||||
group_start_iteration_B);
|
||||
}
|
||||
|
||||
if (warp_mma_k + 2 == Base::kWarpGemmIterations) {
|
||||
int group_start_iteration_A, group_start_iteration_B;
|
||||
group_start_iteration_A =
|
||||
(warp_mma_k + 1) * Detail::kAccessesPerGroupA;
|
||||
group_start_iteration_B =
|
||||
(warp_mma_k + 1) * Detail::kAccessesPerGroupB;
|
||||
|
||||
copy_tiles_and_advance(iterator_A, iterator_B, group_start_iteration_A,
|
||||
group_start_iteration_B);
|
||||
|
||||
// Inserts a memory fence between stages of cp.async instructions.
|
||||
cutlass::arch::cp_async_fence();
|
||||
|
||||
// Waits until kStages-2 stages have committed.
|
||||
arch::cp_async_wait<Base::kStages - 2>();
|
||||
__syncthreads();
|
||||
|
||||
// Move to the next stage
|
||||
iterator_A.add_tile_offset({0, 1});
|
||||
iterator_B.add_tile_offset({1, 0});
|
||||
|
||||
this->smem_iterator_A_.add_tile_offset({0, 1});
|
||||
this->smem_iterator_B_.add_tile_offset({1, 0});
|
||||
|
||||
// Add negative offsets to return iterators to the 'start' of the
|
||||
// circular buffer in shared memory
|
||||
if (smem_write_stage_idx == (Base::kStages - 1)) {
|
||||
this->smem_iterator_A_.add_tile_offset({0, -Base::kStages});
|
||||
this->smem_iterator_B_.add_tile_offset({-Base::kStages, 0});
|
||||
smem_write_stage_idx = 0;
|
||||
} else {
|
||||
++smem_write_stage_idx;
|
||||
}
|
||||
|
||||
if (smem_read_stage_idx == (Base::kStages - 1)) {
|
||||
this->warp_tile_iterator_A_.add_tile_offset(
|
||||
{0, -Base::kStages * Policy::kPartitionsK *
|
||||
Base::kWarpGemmIterations});
|
||||
this->warp_tile_iterator_B_.add_tile_offset(
|
||||
{-Base::kStages * Policy::kPartitionsK *
|
||||
Base::kWarpGemmIterations,
|
||||
0});
|
||||
smem_read_stage_idx = 0;
|
||||
} else {
|
||||
++smem_read_stage_idx;
|
||||
}
|
||||
|
||||
--gemm_k_iterations;
|
||||
if (gemm_k_iterations == 0) {
|
||||
iterator_A.clear_mask();
|
||||
iterator_B.clear_mask();
|
||||
}
|
||||
}
|
||||
|
||||
// Do any conversions feeding the first stage at the end of the loop so
|
||||
// we can start right away on mma instructions
|
||||
if (warp_mma_k + 1 == Base::kWarpGemmIterations)
|
||||
warp_mma.transform(warp_transformed_frag_A[(warp_mma_k + 1) % 2],
|
||||
warp_transformed_frag_B[(warp_mma_k + 1) % 2],
|
||||
warp_loaded_frag_A[(warp_mma_k + 1) % 2],
|
||||
warp_loaded_frag_B[(warp_mma_k + 1) % 2]);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -75,7 +75,7 @@ template <
|
||||
typename IteratorA_::Element,
|
||||
IteratorA_::Fragment::kElements>,
|
||||
///
|
||||
/// Transformation applied to A operand
|
||||
/// Transformation applied to B operand
|
||||
typename TransformB_ = NumericArrayConverter<
|
||||
typename SmemIteratorB_::Element,
|
||||
typename IteratorB_::Element,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -0,0 +1,642 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Template for a double-buffered threadblock-scoped GEMM kernel.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/aligned_buffer.h"
|
||||
#include "cutlass/arch/memory.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/array_planar_complex.h"
|
||||
#include "cutlass/functional.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/gemm/threadblock/mma_planar_complex_base.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace threadblock {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Structure to compute the matrix product targeting CUDA cores and SIMT math
|
||||
/// instructions.
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape_,
|
||||
/// Iterates over tiles of A operand in global memory
|
||||
// (concept: ReadableTileIterator | ForwardTileIterator |
|
||||
// MaskedTileIterator)
|
||||
typename IteratorA_,
|
||||
/// Iterates over tiles of A operand in shared memory
|
||||
/// (concept: WriteableTileIterator | RandomAccessTileIterator)
|
||||
typename SmemIteratorA_,
|
||||
/// Cache operation for operand A
|
||||
cutlass::arch::CacheOperation::Kind CacheOpA,
|
||||
/// Iterates over tiles of B operand in global memory
|
||||
// (concept: ReadableTileIterator | ForwardTileIterator |
|
||||
// MaskedTileIterator)
|
||||
typename IteratorB_,
|
||||
/// Iterates over tiles of B operand in shared memory
|
||||
/// (concept: WriteableTileIterator | RandomAccessTileIterator)
|
||||
typename SmemIteratorB_,
|
||||
/// Cache operation for operand B
|
||||
cutlass::arch::CacheOperation::Kind CacheOpB,
|
||||
/// Data type of accumulator matrix
|
||||
typename ElementC_,
|
||||
/// Data type of accumulator matrix
|
||||
typename LayoutC_,
|
||||
/// Policy describing tuning details (concept: MmaPolicy)
|
||||
typename Policy_,
|
||||
/// Number of stages,
|
||||
int Stages,
|
||||
/// Transformation applied to A
|
||||
ComplexTransform TransformA = ComplexTransform::kNone,
|
||||
/// Transformation applied to B
|
||||
ComplexTransform TransformB = ComplexTransform::kNone
|
||||
>
|
||||
class MmaPlanarComplexMultistage :
|
||||
public MmaPlanarComplexBase<Shape_, Policy_, Stages> {
|
||||
public:
|
||||
///< Base class
|
||||
using Base = MmaPlanarComplexBase<Shape_, Policy_, Stages>;
|
||||
|
||||
///< Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
using Shape = Shape_;
|
||||
|
||||
///< Iterates over tiles of A operand in global memory
|
||||
using IteratorA = IteratorA_;
|
||||
|
||||
///< Iterates over tiles of B operand in global memory
|
||||
using IteratorB = IteratorB_;
|
||||
|
||||
///< Data type of accumulator matrix
|
||||
using ElementC = ElementC_;
|
||||
|
||||
///< Layout of accumulator matrix
|
||||
using LayoutC = LayoutC_;
|
||||
|
||||
///< Policy describing tuning details
|
||||
using Policy = Policy_;
|
||||
|
||||
///< Archtecture tag
|
||||
using ArchTag = arch::Sm80;
|
||||
|
||||
using SmemIteratorA = SmemIteratorA_;
|
||||
using SmemIteratorB = SmemIteratorB_;
|
||||
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpA = CacheOpA;
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpB = CacheOpB;
|
||||
|
||||
/// Transformation applied to A
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
|
||||
/// Transformation applied to B
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
|
||||
//
|
||||
// Dependent types
|
||||
//
|
||||
|
||||
/// Fragment of accumulator tile
|
||||
using FragmentC = ArrayPlanarComplex<
|
||||
typename Policy::Operator::FragmentC::Element,
|
||||
Policy::Operator::FragmentC::kElements
|
||||
>;
|
||||
|
||||
/// Warp-level Mma
|
||||
using Operator = typename Policy::Operator;
|
||||
|
||||
/// Internal structure exposed for introspection.
|
||||
struct Detail {
|
||||
|
||||
static_assert(Base::kWarpGemmIterations > 1,
|
||||
"The pipelined structure requires at least two warp-level "
|
||||
"GEMM operations.");
|
||||
|
||||
/// Number of LDGSTS instructions to load one stage of operand A
|
||||
static int const TBLDGSTSIterationsA =
|
||||
IteratorA::ThreadMap::Iterations::kCount;
|
||||
|
||||
/// Number of LDGSTS instructions to load one stage of operand B
|
||||
static int const TBLDGSTSIterationsB =
|
||||
IteratorB::ThreadMap::Iterations::kCount;
|
||||
|
||||
/// Number of stages
|
||||
static int const kStages = Stages;
|
||||
|
||||
/// Number of LDGSTS instructions to load on group of operand A
|
||||
static int const kAccessesPerGroupA =
|
||||
(TBLDGSTSIterationsA + Base::kWarpGemmIterations - 1) / Base::kWarpGemmIterations;
|
||||
|
||||
/// Number of LDGSTS instructions to load on group of operand B
|
||||
static int const kAccessesPerGroupB =
|
||||
(TBLDGSTSIterationsB + Base::kWarpGemmIterations - 1) / Base::kWarpGemmIterations;
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
using WarpFragmentA = typename Operator::FragmentA;
|
||||
using WarpFragmentB = typename Operator::FragmentB;
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Iterator to write threadblock-scoped tile of A operand to shared memory
|
||||
SmemIteratorA smem_iterator_A_;
|
||||
|
||||
/// Iterator to write threadblock-scoped tile of B operand to shared memory
|
||||
SmemIteratorB smem_iterator_B_;
|
||||
|
||||
public:
|
||||
|
||||
/// Construct from tensor references
|
||||
CUTLASS_DEVICE
|
||||
MmaPlanarComplexMultistage(
|
||||
///< Shared storage needed for internal use by threadblock-scoped GEMM
|
||||
typename Base::SharedStorage &shared_storage,
|
||||
///< ID within the threadblock
|
||||
int thread_idx,
|
||||
///< ID of warp
|
||||
int warp_idx,
|
||||
///< ID of each thread within a warp
|
||||
int lane_idx
|
||||
):
|
||||
Base(shared_storage, thread_idx, warp_idx, lane_idx),
|
||||
smem_iterator_A_(shared_storage.operand_A_ref(), thread_idx),
|
||||
smem_iterator_B_(shared_storage.operand_B_ref(), thread_idx)
|
||||
{
|
||||
// Compute warp location within threadblock tile by mapping the warp_id to
|
||||
// three coordinates:
|
||||
// _m: the warp's position within the threadblock along the M dimension
|
||||
// _n: the warp's position within the threadblock along the N dimension
|
||||
// _k: the warp's position within the threadblock along the K dimension
|
||||
|
||||
int warp_idx_mn = warp_idx % (Base::WarpCount::kM * Base::WarpCount::kN);
|
||||
int warp_idx_k = warp_idx / (Base::WarpCount::kM * Base::WarpCount::kN);
|
||||
|
||||
int warp_idx_m = warp_idx_mn % Base::WarpCount::kM;
|
||||
int warp_idx_n = warp_idx_mn / Base::WarpCount::kM;
|
||||
|
||||
// Add per-warp offsets in units of warp-level tiles
|
||||
this->warp_tile_iterator_A_.add_tile_offset({warp_idx_m, Base::kWarpGemmIterations * warp_idx_k});
|
||||
this->warp_tile_iterator_B_.add_tile_offset({Base::kWarpGemmIterations * warp_idx_k, warp_idx_n});
|
||||
}
|
||||
|
||||
private:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void copy_tiles_and_advance(
|
||||
IteratorA &iterator_A_real,
|
||||
IteratorA &iterator_A_imag,
|
||||
|
||||
IteratorB &iterator_B_real,
|
||||
IteratorB &iterator_B_imag,
|
||||
|
||||
int group_start_A = 0,
|
||||
int group_start_B = 0) {
|
||||
|
||||
iterator_A_real.set_iteration_index(group_start_A * IteratorA::kAccessesPerVector);
|
||||
iterator_A_imag.set_iteration_index(group_start_A * IteratorA::kAccessesPerVector);
|
||||
this->smem_iterator_A_.set_iteration_index(group_start_A);
|
||||
|
||||
// LDGSTS for operand A
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < Detail::kAccessesPerGroupA; ++j) {
|
||||
|
||||
typename IteratorA::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorA::AccessType *>(this->smem_iterator_A_.get());
|
||||
|
||||
int const kSrcBytes =
|
||||
sizeof_bits<typename IteratorA::Element>::value *
|
||||
IteratorA::ThreadMap::kElementsPerAccess / IteratorA::kAccessesPerVector / 8;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorA::kAccessesPerVector; ++v) {
|
||||
|
||||
auto gmem_ptr_real = iterator_A_real.get();
|
||||
auto gmem_ptr_imag = iterator_A_imag.get();
|
||||
|
||||
bool pred_guard = iterator_A_real.valid();
|
||||
cutlass::arch::cp_async<kSrcBytes, kCacheOpA>(
|
||||
dst_ptr + v,
|
||||
gmem_ptr_real,
|
||||
pred_guard);
|
||||
cutlass::arch::cp_async<kSrcBytes, kCacheOpA>(
|
||||
dst_ptr + v + (Base::SharedStorage::kImaginaryStrideA / IteratorA::ThreadMap::kElementsPerAccess),
|
||||
reinterpret_cast<char const *>(gmem_ptr_imag),
|
||||
pred_guard);
|
||||
|
||||
++iterator_A_real;
|
||||
++iterator_A_imag;
|
||||
}
|
||||
|
||||
++this->smem_iterator_A_;
|
||||
}
|
||||
|
||||
iterator_B_real.set_iteration_index(group_start_B * IteratorB::kAccessesPerVector);
|
||||
iterator_B_imag.set_iteration_index(group_start_B * IteratorB::kAccessesPerVector);
|
||||
this->smem_iterator_B_.set_iteration_index(group_start_B);
|
||||
|
||||
// LDGSTS for operand B
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < Detail::kAccessesPerGroupB; ++j) {
|
||||
typename IteratorB::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorB::AccessType *>(this->smem_iterator_B_.get());
|
||||
|
||||
int const kSrcBytes =
|
||||
sizeof_bits<typename IteratorB::Element>::value *
|
||||
IteratorB::ThreadMap::kElementsPerAccess / IteratorB::kAccessesPerVector / 8;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorB::kAccessesPerVector; ++v) {
|
||||
auto gmem_ptr_real = iterator_B_real.get();
|
||||
auto gmem_ptr_imag = iterator_B_imag.get();
|
||||
|
||||
bool pred_guard = iterator_B_real.valid();
|
||||
cutlass::arch::cp_async<kSrcBytes, kCacheOpB>(
|
||||
dst_ptr + v,
|
||||
gmem_ptr_real,
|
||||
pred_guard);
|
||||
cutlass::arch::cp_async<kSrcBytes, kCacheOpB>(
|
||||
dst_ptr + v + (Base::SharedStorage::kImaginaryStrideB / IteratorB::ThreadMap::kElementsPerAccess),
|
||||
reinterpret_cast<char const *>(gmem_ptr_imag),
|
||||
pred_guard);
|
||||
|
||||
++iterator_B_real;
|
||||
++iterator_B_imag;
|
||||
}
|
||||
++this->smem_iterator_B_;
|
||||
}
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void warp_mma_planar_complex(
|
||||
Operator & warp_mma,
|
||||
FragmentC &accum,
|
||||
WarpFragmentA const & real_A,
|
||||
WarpFragmentA const & imag_A,
|
||||
WarpFragmentB const & real_B,
|
||||
WarpFragmentB const & imag_B) {
|
||||
|
||||
cutlass::negate<Array<typename WarpFragmentB::Element, WarpFragmentB::kElements>> neg_op_B;
|
||||
|
||||
WarpFragmentB neg_real_B = neg_op_B(real_B);
|
||||
WarpFragmentB neg_imag_B = neg_op_B(imag_B);
|
||||
|
||||
warp_mma(accum.real, real_A, real_B, accum.real);
|
||||
|
||||
if (kTransformB == ComplexTransform::kNone) {
|
||||
warp_mma(accum.imag, real_A, imag_B, accum.imag);
|
||||
}
|
||||
else {
|
||||
warp_mma(accum.imag, real_A, neg_imag_B, accum.imag);
|
||||
}
|
||||
|
||||
if (kTransformA == ComplexTransform::kNone) {
|
||||
warp_mma(accum.imag, imag_A, real_B, accum.imag);
|
||||
}
|
||||
else {
|
||||
warp_mma(accum.imag, imag_A, neg_real_B, accum.imag);
|
||||
}
|
||||
|
||||
if (kTransformA == ComplexTransform::kNone ^ kTransformB == ComplexTransform::kNone) {
|
||||
warp_mma(accum.real, imag_A, imag_B, accum.real);
|
||||
}
|
||||
else {
|
||||
warp_mma(accum.real, imag_A, neg_imag_B, accum.real);
|
||||
}
|
||||
}
|
||||
|
||||
public:
|
||||
|
||||
/// Perform a threadblock-scoped matrix multiply-accumulate
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
///< problem size of GEMM
|
||||
int gemm_k_iterations,
|
||||
///< destination accumulator tile
|
||||
FragmentC &accum,
|
||||
///< iterator over A operand in global memory
|
||||
IteratorA iterator_A_real,
|
||||
///< iterator over A operand in global memory
|
||||
IteratorA iterator_A_imag,
|
||||
///< iterator over B operand in global memory
|
||||
IteratorB iterator_B_real,
|
||||
///< iterator over B operand in global memory
|
||||
IteratorB iterator_B_imag,
|
||||
///< initial value of accumulator
|
||||
FragmentC const &src_accum) {
|
||||
|
||||
//
|
||||
// Prologue
|
||||
//
|
||||
|
||||
// Issue several complete stages
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int stage = 0; stage < Base::kStages - 1;
|
||||
++stage, --gemm_k_iterations) {
|
||||
|
||||
if (gemm_k_iterations == 0) {
|
||||
iterator_A_real.clear_mask();
|
||||
iterator_A_imag.clear_mask();
|
||||
iterator_B_real.clear_mask();
|
||||
iterator_B_imag.clear_mask();
|
||||
}
|
||||
|
||||
iterator_A_real.set_iteration_index(0);
|
||||
iterator_A_imag.set_iteration_index(0);
|
||||
|
||||
this->smem_iterator_A_.set_iteration_index(0);
|
||||
|
||||
// LDGSTS for operand A
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < Detail::TBLDGSTSIterationsA; ++j) {
|
||||
|
||||
typename IteratorA::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorA::AccessType *>(this->smem_iterator_A_.get());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorA::kAccessesPerVector; ++v) {
|
||||
|
||||
int const kSrcBytes =
|
||||
sizeof_bits<typename IteratorA::Element>::value *
|
||||
IteratorA::ThreadMap::kElementsPerAccess / IteratorA::kAccessesPerVector / 8;
|
||||
|
||||
bool pred_guard = iterator_A_real.valid();
|
||||
|
||||
auto src_ptr_real = iterator_A_real.get();
|
||||
auto src_ptr_imag = iterator_A_imag.get();
|
||||
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpA>(
|
||||
dst_ptr + v, src_ptr_real, pred_guard);
|
||||
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpA>(
|
||||
dst_ptr + v +
|
||||
Base::SharedStorage::kImaginaryStrideA /
|
||||
IteratorA::ThreadMap::kElementsPerAccess,
|
||||
reinterpret_cast<char const *>(src_ptr_imag),
|
||||
pred_guard);
|
||||
|
||||
++iterator_A_real;
|
||||
++iterator_A_imag;
|
||||
}
|
||||
|
||||
++this->smem_iterator_A_;
|
||||
}
|
||||
|
||||
iterator_B_real.set_iteration_index(0);
|
||||
iterator_B_imag.set_iteration_index(0);
|
||||
|
||||
this->smem_iterator_B_.set_iteration_index(0);
|
||||
|
||||
// LDGSTS for operand B
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < Detail::TBLDGSTSIterationsB; ++j) {
|
||||
|
||||
typename IteratorB::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorB::AccessType *>(this->smem_iterator_B_.get());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorB::kAccessesPerVector; ++v) {
|
||||
|
||||
int const kSrcBytes =
|
||||
sizeof_bits<typename IteratorB::Element>::value *
|
||||
IteratorB::ThreadMap::kElementsPerAccess / IteratorB::kAccessesPerVector / 8;
|
||||
|
||||
bool pred_guard = iterator_B_real.valid();
|
||||
|
||||
auto src_ptr_real = iterator_B_real.get();
|
||||
auto src_ptr_imag = iterator_B_imag.get();
|
||||
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpB>(
|
||||
dst_ptr + v, src_ptr_real, pred_guard);
|
||||
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpB>(
|
||||
dst_ptr + v +
|
||||
Base::SharedStorage::kImaginaryStrideB /
|
||||
IteratorB::ThreadMap::kElementsPerAccess,
|
||||
reinterpret_cast<char const *>(src_ptr_imag),
|
||||
pred_guard);
|
||||
|
||||
++iterator_B_real;
|
||||
++iterator_B_imag;
|
||||
}
|
||||
|
||||
++this->smem_iterator_B_;
|
||||
}
|
||||
|
||||
// Move to the next stage
|
||||
iterator_A_real.add_tile_offset({0, 1});
|
||||
iterator_A_imag.add_tile_offset({0, 1});
|
||||
|
||||
iterator_B_real.add_tile_offset({1, 0});
|
||||
iterator_B_imag.add_tile_offset({1, 0});
|
||||
|
||||
this->smem_iterator_A_.add_tile_offset({0, 1});
|
||||
this->smem_iterator_B_.add_tile_offset({1, 0});
|
||||
|
||||
// Inserts a memory fence between stages of cp.async instructions
|
||||
cutlass::arch::cp_async_fence();
|
||||
}
|
||||
|
||||
// Perform accumulation in the 'd' output operand
|
||||
accum = src_accum;
|
||||
|
||||
// Blocks until all but kStages-2 cp.async stages have committed.
|
||||
cutlass::arch::cp_async_wait<Base::kStages - 2>();
|
||||
__syncthreads();
|
||||
|
||||
// Pair of fragments used to overlap shared memory loads and math
|
||||
// instructions
|
||||
|
||||
WarpFragmentA warp_frag_real_A[2];
|
||||
WarpFragmentA warp_frag_imag_A[2];
|
||||
|
||||
WarpFragmentB warp_frag_real_B[2];
|
||||
WarpFragmentB warp_frag_imag_B[2];
|
||||
|
||||
this->warp_tile_iterator_A_.set_kgroup_index(0);
|
||||
this->warp_tile_iterator_B_.set_kgroup_index(0);
|
||||
|
||||
this->warp_tile_iterator_A_.load(warp_frag_real_A[0]);
|
||||
this->warp_tile_iterator_A_.load_with_pointer_offset(warp_frag_imag_A[0], Base::SharedStorage::kImaginaryStrideA);
|
||||
|
||||
this->warp_tile_iterator_B_.load(warp_frag_real_B[0]);
|
||||
this->warp_tile_iterator_B_.load_with_pointer_offset(warp_frag_imag_B[0], Base::SharedStorage::kImaginaryStrideB);
|
||||
|
||||
++this->warp_tile_iterator_A_;
|
||||
++this->warp_tile_iterator_B_;
|
||||
|
||||
if (gemm_k_iterations == 0) {
|
||||
iterator_A_real.clear_mask();
|
||||
iterator_A_imag.clear_mask();
|
||||
iterator_B_real.clear_mask();
|
||||
iterator_B_imag.clear_mask();
|
||||
}
|
||||
|
||||
// Start issuing the first group of the next stage outside of the mainloop
|
||||
copy_tiles_and_advance(iterator_A_real, iterator_A_imag, iterator_B_real, iterator_B_imag);
|
||||
|
||||
Operator warp_mma;
|
||||
|
||||
int smem_write_stage_idx = Base::kStages - 1;
|
||||
int smem_read_stage_idx = 0;
|
||||
|
||||
//
|
||||
// Mainloop
|
||||
//
|
||||
|
||||
CUTLASS_GEMM_LOOP
|
||||
for (; gemm_k_iterations > (-Base::kStages + 1);) {
|
||||
//
|
||||
// Loop over GEMM K dimension
|
||||
//
|
||||
|
||||
// Computes a warp-level GEMM on data held in shared memory
|
||||
// Each "warp_mma_k" refers to a warp-level matrix multiply-accumulate
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int warp_mma_k = 0; warp_mma_k < Base::kWarpGemmIterations;
|
||||
++warp_mma_k) {
|
||||
|
||||
// Load warp-level tiles from shared memory, wrapping to k offset if
|
||||
// this is the last group as the case may be.
|
||||
|
||||
this->warp_tile_iterator_A_.set_kgroup_index((warp_mma_k + 1) % Base::kWarpGemmIterations);
|
||||
this->warp_tile_iterator_B_.set_kgroup_index((warp_mma_k + 1) % Base::kWarpGemmIterations);
|
||||
|
||||
this->warp_tile_iterator_A_.load(warp_frag_real_A[(warp_mma_k + 1) % 2]);
|
||||
this->warp_tile_iterator_A_.load_with_pointer_offset(warp_frag_imag_A[(warp_mma_k + 1) % 2], Base::SharedStorage::kImaginaryStrideA);
|
||||
|
||||
this->warp_tile_iterator_B_.load(warp_frag_real_B[(warp_mma_k + 1) % 2]);
|
||||
this->warp_tile_iterator_B_.load_with_pointer_offset(warp_frag_imag_B[(warp_mma_k + 1) % 2], Base::SharedStorage::kImaginaryStrideB);
|
||||
|
||||
++this->warp_tile_iterator_A_;
|
||||
++this->warp_tile_iterator_B_;
|
||||
|
||||
// Issue global->shared copies for the next stage
|
||||
int group_start_iteration_A, group_start_iteration_B;
|
||||
|
||||
if (warp_mma_k + 1 == Base::kWarpGemmIterations) {
|
||||
group_start_iteration_A = 0;
|
||||
group_start_iteration_B = 0;
|
||||
}
|
||||
else {
|
||||
group_start_iteration_A = (warp_mma_k + 1) * Detail::kAccessesPerGroupA;
|
||||
group_start_iteration_B = (warp_mma_k + 1) * Detail::kAccessesPerGroupB;
|
||||
}
|
||||
|
||||
copy_tiles_and_advance(
|
||||
iterator_A_real,
|
||||
iterator_A_imag,
|
||||
iterator_B_real,
|
||||
iterator_B_imag,
|
||||
group_start_iteration_A,
|
||||
group_start_iteration_B);
|
||||
|
||||
if (warp_mma_k + 2 == Base::kWarpGemmIterations) {
|
||||
// Inserts a memory fence between stages of cp.async instructions
|
||||
cutlass::arch::cp_async_fence();
|
||||
|
||||
// Blocks until all but kStages-2 cp.async stages have committed.
|
||||
arch::cp_async_wait<Base::kStages - 2>();
|
||||
__syncthreads();
|
||||
|
||||
// Move to the next stage
|
||||
iterator_A_real.add_tile_offset({0, 1});
|
||||
iterator_A_imag.add_tile_offset({0, 1});
|
||||
|
||||
iterator_B_real.add_tile_offset({1, 0});
|
||||
iterator_B_imag.add_tile_offset({1, 0});
|
||||
|
||||
this->smem_iterator_A_.add_tile_offset({0, 1});
|
||||
this->smem_iterator_B_.add_tile_offset({1, 0});
|
||||
|
||||
// Add negative offsets to return iterators to the 'start' of the
|
||||
// circular buffer in shared memory
|
||||
if (smem_write_stage_idx == (Base::kStages - 1)) {
|
||||
this->smem_iterator_A_.add_tile_offset({0, -Base::kStages});
|
||||
this->smem_iterator_B_.add_tile_offset({-Base::kStages, 0});
|
||||
smem_write_stage_idx = 0;
|
||||
} else {
|
||||
++smem_write_stage_idx;
|
||||
}
|
||||
|
||||
if (smem_read_stage_idx == (Base::kStages - 1)) {
|
||||
|
||||
this->warp_tile_iterator_A_.add_tile_offset(
|
||||
{0, -Base::kStages * Policy::kPartitionsK *
|
||||
Base::kWarpGemmIterations});
|
||||
|
||||
this->warp_tile_iterator_B_.add_tile_offset(
|
||||
{-Base::kStages * Policy::kPartitionsK *
|
||||
Base::kWarpGemmIterations,
|
||||
0});
|
||||
smem_read_stage_idx = 0;
|
||||
} else {
|
||||
++smem_read_stage_idx;
|
||||
}
|
||||
|
||||
--gemm_k_iterations;
|
||||
if (gemm_k_iterations == 0) {
|
||||
iterator_A_real.clear_mask();
|
||||
iterator_A_imag.clear_mask();
|
||||
iterator_B_real.clear_mask();
|
||||
iterator_B_imag.clear_mask();
|
||||
}
|
||||
}
|
||||
|
||||
warp_mma_planar_complex(
|
||||
warp_mma,
|
||||
accum,
|
||||
warp_frag_real_A[warp_mma_k % 2],
|
||||
warp_frag_imag_A[warp_mma_k % 2],
|
||||
warp_frag_real_B[warp_mma_k % 2],
|
||||
warp_frag_imag_B[warp_mma_k % 2]);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -99,61 +99,13 @@ int RematerializeBlockDimZ() {
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Threadblock swizzling function for GEMMs
|
||||
template <int N = 1>
|
||||
struct GemmIdentityThreadblockSwizzle {
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
GemmIdentityThreadblockSwizzle() { }
|
||||
|
||||
int const kTile = 1;
|
||||
|
||||
/// Returns the shape of the problem in units of logical tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
GemmCoord get_tiled_shape(
|
||||
GemmCoord problem_size,
|
||||
GemmCoord tile_size,
|
||||
int split_k_slices) const {
|
||||
|
||||
return GemmCoord(
|
||||
(problem_size.m() + tile_size.m() - 1) / tile_size.m(),
|
||||
(problem_size.n() + tile_size.n() - 1) / tile_size.n(),
|
||||
split_k_slices);
|
||||
}
|
||||
|
||||
/// Computes CUDA grid dimensions given a size in units of logical tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
dim3 get_grid_shape(GemmCoord tiled_shape) const {
|
||||
return dim3(tiled_shape.m() * kTile, (tiled_shape.n() + kTile - 1) / kTile, tiled_shape.k());
|
||||
}
|
||||
|
||||
/// Obtains the threadblock offset (in units of threadblock-scoped tiles)
|
||||
CUTLASS_DEVICE
|
||||
GemmCoord get_tile_offset() const {
|
||||
|
||||
int block_idx_x = RematerializeBlockIdxX();
|
||||
int block_idx_y = RematerializeBlockIdxY();
|
||||
|
||||
return GemmCoord{
|
||||
(block_idx_x / kTile),
|
||||
(block_idx_y * kTile) + (block_idx_x % kTile),
|
||||
RematerializeBlockIdxZ()
|
||||
};
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// A special version of GemmIdentityThreadblockSwizzle. See the choice of kTile below.
|
||||
template <typename LayoutA_, typename LayoutB_>
|
||||
struct GemmCohortThreadblockSwizzle
|
||||
{
|
||||
const int kTile =
|
||||
(platform::is_same<LayoutA_, cutlass::layout::RowMajor>::value ||
|
||||
platform::is_same<LayoutB_, cutlass::layout::ColumnMajor>::value)
|
||||
? 4
|
||||
: 1;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
GemmCohortThreadblockSwizzle() { }
|
||||
int const kTile = N;
|
||||
|
||||
/// Returns the shape of the problem in units of logical tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -271,8 +223,11 @@ struct GemmBatchedIdentityThreadblockSwizzle {
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Threadblock swizzling function for split-K GEMMs
|
||||
template <int N = 1>
|
||||
struct GemmSplitKIdentityThreadblockSwizzle {
|
||||
|
||||
int const kTile = N;
|
||||
|
||||
/// Returns the shape of the problem in units of logical tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
GemmCoord get_tiled_shape(
|
||||
@@ -289,16 +244,20 @@ struct GemmSplitKIdentityThreadblockSwizzle {
|
||||
/// Computes CUDA grid dimensions given a size in units of logical tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
dim3 get_grid_shape(GemmCoord tiled_shape) const {
|
||||
return dim3(tiled_shape.m(), tiled_shape.n(), tiled_shape.k());
|
||||
return dim3(tiled_shape.m() * kTile, (tiled_shape.n() + kTile - 1) / kTile, tiled_shape.k());
|
||||
}
|
||||
|
||||
|
||||
/// Obtains the threadblock offset (in units of threadblock-scoped tiles)
|
||||
CUTLASS_DEVICE
|
||||
GemmCoord get_tile_offset() const {
|
||||
|
||||
int block_idx_x = RematerializeBlockIdxX();
|
||||
int block_idx_y = RematerializeBlockIdxY();
|
||||
|
||||
return GemmCoord{
|
||||
RematerializeBlockIdxX(),
|
||||
RematerializeBlockIdxY(),
|
||||
(block_idx_x / kTile),
|
||||
(block_idx_y * kTile) + (block_idx_x % kTile),
|
||||
RematerializeBlockIdxZ()
|
||||
};
|
||||
}
|
||||
|
||||
@@ -0,0 +1,401 @@
|
||||
/***************************************************************************************************
|
||||
* 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 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_complex_tensor_op.h"
|
||||
#include "cutlass/gemm/warp/mma_gaussian_complex_tensor_op.h"
|
||||
#include "cutlass/layout/tensor_op_multiplicand_sm80.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_,
|
||||
/// Complex transform on A operand
|
||||
ComplexTransform TransformA = ComplexTransform::kNone,
|
||||
/// Complex transform on B operand
|
||||
ComplexTransform TransformB = ComplexTransform::kNone,
|
||||
/// Multiply-add operator (arch::OpMultiplyAddComplex, arch::OpMultiplyGaussianComplex)
|
||||
typename Operator_ = arch::OpMultiplyAddComplex>
|
||||
struct DefaultMmaComplexTensorOp;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for complex<T>*complex<T> case
|
||||
// 4 real-valued mma 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_,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// 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_,
|
||||
InstructionShape_,
|
||||
complex<RealElementA>,
|
||||
LayoutA,
|
||||
complex<RealElementB>,
|
||||
LayoutB,
|
||||
complex<RealElementC>,
|
||||
LayoutC,
|
||||
TransformA,
|
||||
TransformB,
|
||||
arch::OpMultiplyAddComplex> {
|
||||
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
InstructionShape_,
|
||||
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>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for complex<T>*complex<T> case using GaussianComplex operation
|
||||
// 3 real-valued mma 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_,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// 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_,
|
||||
InstructionShape_,
|
||||
complex<RealElementA>,
|
||||
LayoutA,
|
||||
complex<RealElementB>,
|
||||
LayoutB,
|
||||
complex<RealElementC>,
|
||||
LayoutC,
|
||||
TransformA,
|
||||
TransformB,
|
||||
arch::OpMultiplyAddGaussianComplex> {
|
||||
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
InstructionShape_,
|
||||
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>;
|
||||
};
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// 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
|
||||
// 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_,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB,
|
||||
/// 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_,
|
||||
InstructionShape_,
|
||||
complex<float>,
|
||||
LayoutA,
|
||||
complex<float>,
|
||||
LayoutB,
|
||||
complex<float>,
|
||||
LayoutC,
|
||||
TransformA,
|
||||
TransformB,
|
||||
arch::OpMultiplyAddComplex> {
|
||||
|
||||
// Complex floating point tensor operation use MMA.1688.F32.TF32 mma instruction
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
InstructionShape_,
|
||||
32,
|
||||
tfloat32_t,
|
||||
cutlass::layout::RowMajor,
|
||||
tfloat32_t,
|
||||
cutlass::layout::ColumnMajor,
|
||||
float,
|
||||
cutlass::layout::RowMajor,
|
||||
arch::OpMultiplyAdd>,
|
||||
cutlass::MatrixShape<1, 1>
|
||||
>;
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using Type = cutlass::gemm::warp::MmaComplexTensorOp<
|
||||
WarpShape_,
|
||||
complex<float>,
|
||||
LayoutA,
|
||||
complex<float>,
|
||||
LayoutB,
|
||||
complex<float>,
|
||||
LayoutC,
|
||||
Policy,
|
||||
TransformA,
|
||||
TransformB>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// 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
|
||||
// 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_,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB,
|
||||
/// 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_,
|
||||
InstructionShape_,
|
||||
complex<float>,
|
||||
LayoutA,
|
||||
complex<float>,
|
||||
LayoutB,
|
||||
complex<float>,
|
||||
LayoutC,
|
||||
TransformA,
|
||||
TransformB,
|
||||
arch::OpMultiplyAddFastBF16> {
|
||||
|
||||
// Complex floating point tensor operation use MMA.1688.F32.BF16 mma instruction
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
InstructionShape_,
|
||||
32,
|
||||
bfloat16_t,
|
||||
cutlass::layout::RowMajor,
|
||||
bfloat16_t,
|
||||
cutlass::layout::ColumnMajor,
|
||||
float,
|
||||
cutlass::layout::RowMajor,
|
||||
arch::OpMultiplyAdd>,
|
||||
cutlass::MatrixShape<1, 1>
|
||||
>;
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using Type = cutlass::gemm::warp::MmaComplexTensorOp<
|
||||
WarpShape_,
|
||||
complex<float>,
|
||||
LayoutA,
|
||||
complex<float>,
|
||||
LayoutB,
|
||||
complex<float>,
|
||||
LayoutC,
|
||||
Policy,
|
||||
TransformA,
|
||||
TransformB>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// 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
|
||||
// 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_,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB,
|
||||
/// 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_,
|
||||
InstructionShape_,
|
||||
complex<float>,
|
||||
LayoutA,
|
||||
complex<float>,
|
||||
LayoutB,
|
||||
complex<float>,
|
||||
LayoutC,
|
||||
TransformA,
|
||||
TransformB,
|
||||
arch::OpMultiplyAddFastF16> {
|
||||
|
||||
// Complex floating point tensor operation use MMA.1688.F32.F16 mma instruction
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
InstructionShape_,
|
||||
32,
|
||||
half_t,
|
||||
cutlass::layout::RowMajor,
|
||||
half_t,
|
||||
cutlass::layout::ColumnMajor,
|
||||
float,
|
||||
cutlass::layout::RowMajor,
|
||||
arch::OpMultiplyAdd>,
|
||||
cutlass::MatrixShape<1, 1>
|
||||
>;
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using Type = cutlass::gemm::warp::MmaComplexTensorOp<
|
||||
WarpShape_,
|
||||
complex<float>,
|
||||
LayoutA,
|
||||
complex<float>,
|
||||
LayoutB,
|
||||
complex<float>,
|
||||
LayoutC,
|
||||
Policy,
|
||||
TransformA,
|
||||
TransformB>;
|
||||
};
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -60,10 +60,7 @@ template <
|
||||
int PartitionsK = 1,
|
||||
/// Store the accumulators in row major or column major. Row major is used
|
||||
/// when output layout is interleaved.
|
||||
bool AccumulatorsInRowMajor = false,
|
||||
/// Number of partitions along N dimension per warp
|
||||
int PartitionsN = 1
|
||||
>
|
||||
bool AccumulatorsInRowMajor = false>
|
||||
struct DefaultMmaTensorOp;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -92,9 +89,7 @@ template <
|
||||
int PartitionsK,
|
||||
/// Store the accumulators in row major or column major. Row major is used
|
||||
/// when output layout is interleaved.
|
||||
bool AccumulatorsInRowMajor,
|
||||
/// Number of partitions along N dimension per warp
|
||||
int PartitionsN>
|
||||
bool AccumulatorsInRowMajor>
|
||||
struct DefaultMmaTensorOp {
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<InstructionShape_, 32, ElementA,
|
||||
@@ -106,7 +101,7 @@ struct DefaultMmaTensorOp {
|
||||
// Define the warp-level tensor op
|
||||
using Type = cutlass::gemm::warp::MmaTensorOp<
|
||||
WarpShape_, ElementA, LayoutA, ElementB, LayoutB, ElementC, LayoutC,
|
||||
Policy, PartitionsK, AccumulatorsInRowMajor, PartitionsN>;
|
||||
Policy, PartitionsK, AccumulatorsInRowMajor>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -117,3 +112,6 @@ struct DefaultMmaTensorOp {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#include "default_mma_tensor_op_sm80.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -0,0 +1,186 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Default warp-level GEMM operators selected by data type, size, and layouts of operands.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/arch/mma.h"
|
||||
#include "cutlass/gemm/warp/mma_tensor_op.h"
|
||||
#include "cutlass/gemm/warp/default_mma_tensor_op.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace warp {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial Specialization - inputs and output types are float - uses BF16 internally
|
||||
template <
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename WarpShape_,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB,
|
||||
/// Layout of C matrix (concept: MatrixLayout)
|
||||
typename LayoutC,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK,
|
||||
/// Store the accumulators in row major or column major. Row major is used
|
||||
/// when output layout is interleaved.
|
||||
bool AccumulatorsInRowMajor>
|
||||
struct DefaultMmaTensorOp<
|
||||
WarpShape_,
|
||||
GemmShape<16, 8, 8>,
|
||||
float, LayoutA,
|
||||
float, LayoutB,
|
||||
float, LayoutC,
|
||||
arch::OpMultiplyAddFastBF16,
|
||||
PartitionsK, AccumulatorsInRowMajor> {
|
||||
|
||||
// Uses BF16 internally
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
GemmShape<16, 8, 8>,
|
||||
32,
|
||||
bfloat16_t, cutlass::layout::RowMajor,
|
||||
bfloat16_t, cutlass::layout::ColumnMajor,
|
||||
float, cutlass::layout::RowMajor,
|
||||
arch::OpMultiplyAdd
|
||||
>,
|
||||
cutlass::MatrixShape<1, 1> >;
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using Type = cutlass::gemm::warp::MmaTensorOp<
|
||||
WarpShape_, float, LayoutA, float, LayoutB, float, LayoutC,
|
||||
Policy, PartitionsK, AccumulatorsInRowMajor>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial Specialization - inputs and output types are float - uses F16 internally
|
||||
template <
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename WarpShape_,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB,
|
||||
/// Layout of C matrix (concept: MatrixLayout)
|
||||
typename LayoutC,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK,
|
||||
/// Store the accumulators in row major or column major. Row major is used
|
||||
/// when output layout is interleaved.
|
||||
bool AccumulatorsInRowMajor>
|
||||
struct DefaultMmaTensorOp<
|
||||
WarpShape_,
|
||||
GemmShape<16, 8, 8>,
|
||||
float, LayoutA,
|
||||
float, LayoutB,
|
||||
float, LayoutC,
|
||||
arch::OpMultiplyAddFastF16,
|
||||
PartitionsK, AccumulatorsInRowMajor> {
|
||||
|
||||
// Uses F16 internally
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
GemmShape<16, 8, 8>,
|
||||
32,
|
||||
half_t, cutlass::layout::RowMajor,
|
||||
half_t, cutlass::layout::ColumnMajor,
|
||||
float, cutlass::layout::RowMajor,
|
||||
arch::OpMultiplyAdd
|
||||
>,
|
||||
cutlass::MatrixShape<1, 1> >;
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using Type = cutlass::gemm::warp::MmaTensorOp<
|
||||
WarpShape_, float, LayoutA, float, LayoutB, float, LayoutC,
|
||||
Policy, PartitionsK, AccumulatorsInRowMajor>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial Specialization - inputs and output types are float - uses TF32 internally
|
||||
template <
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename WarpShape_,
|
||||
/// Shape of target matrix multiply instruction (concept: GemmShape)
|
||||
typename InstructionShape_,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB,
|
||||
/// Layout of C matrix (concept: MatrixLayout)
|
||||
typename LayoutC,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK,
|
||||
/// Store the accumulators in row major or column major. Row major is used
|
||||
/// when output layout is interleaved.
|
||||
bool AccumulatorsInRowMajor>
|
||||
struct DefaultMmaTensorOp<
|
||||
WarpShape_,
|
||||
InstructionShape_,
|
||||
float, LayoutA,
|
||||
float, LayoutB,
|
||||
float, LayoutC,
|
||||
arch::OpMultiplyAdd, PartitionsK, AccumulatorsInRowMajor> {
|
||||
|
||||
// Uses TF32 internally
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Mma<
|
||||
InstructionShape_,
|
||||
32,
|
||||
tfloat32_t, cutlass::layout::RowMajor,
|
||||
tfloat32_t, cutlass::layout::ColumnMajor,
|
||||
float, cutlass::layout::RowMajor,
|
||||
arch::OpMultiplyAdd
|
||||
>,
|
||||
cutlass::MatrixShape<1, 1> >;
|
||||
|
||||
// Define the warp-level tensor op
|
||||
using Type = cutlass::gemm::warp::MmaTensorOp<
|
||||
WarpShape_, float, LayoutA, float, LayoutB, float, LayoutC,
|
||||
Policy, PartitionsK, AccumulatorsInRowMajor>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#include "cutlass/gemm/warp/mma_complex_tensor_op_tile_iterator_sm80.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -61,9 +61,7 @@ template <
|
||||
/// Operator describing the tensor operation
|
||||
typename Operator_ = arch::OpMultiplyAdd,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK = 1,
|
||||
/// Number of partitions along N dimension per warp
|
||||
int PartitionsN = 1
|
||||
int PartitionsK = 1
|
||||
>
|
||||
struct DefaultMmaTensorOpWmma;
|
||||
|
||||
@@ -90,9 +88,7 @@ template <
|
||||
/// Operator describing the tensor operation
|
||||
typename Operator_,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK,
|
||||
/// Number of partitions along N dimension per warp
|
||||
int PartitionsN>
|
||||
int PartitionsK>
|
||||
struct DefaultMmaTensorOpWmma {
|
||||
using Policy = cutlass::gemm::warp::MmaTensorOpPolicy<
|
||||
cutlass::arch::Wmma<
|
||||
@@ -116,8 +112,7 @@ struct DefaultMmaTensorOpWmma {
|
||||
ElementC,
|
||||
LayoutC,
|
||||
Policy,
|
||||
PartitionsK,
|
||||
PartitionsN>;
|
||||
PartitionsK>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -127,4 +122,3 @@ struct DefaultMmaTensorOpWmma {
|
||||
} // namespace cutlass
|
||||
|
||||
#endif
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -0,0 +1,843 @@
|
||||
/***************************************************************************************************
|
||||
* 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 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/functional.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"
|
||||
#include "cutlass/gemm/warp/mma_complex_tensor_op_tile_iterator_sm80.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace warp {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
|
||||
template <
|
||||
/// Data type of real & imag members of complex numbers in the SourceFragment
|
||||
typename RealElement,
|
||||
/// Destination fragment required by the mma operation
|
||||
typename DestinationFragment,
|
||||
/// Source fragment holding complex<RealElement> elements
|
||||
typename SourceFragment,
|
||||
/// Number of mma operations performed
|
||||
typename MmaIterations,
|
||||
/// Shape of operand elements
|
||||
typename MmaOperandShape,
|
||||
/// Complex transform on A operand
|
||||
ComplexTransform Transform_,
|
||||
/// Operand A or Operand B
|
||||
Operand Operand_,
|
||||
/// Floating-point rounding style
|
||||
FloatRoundStyle Round_>
|
||||
struct UnpackComplexConvertAndPackForMma;
|
||||
|
||||
// Partial specialization for OperandA and Congruous smem layout
|
||||
template <
|
||||
typename RealElement,
|
||||
typename DestinationFragment,
|
||||
typename SourceFragment,
|
||||
typename MmaIterations,
|
||||
typename MmaOperandShape,
|
||||
ComplexTransform Transform_,
|
||||
FloatRoundStyle Round_>
|
||||
struct UnpackComplexConvertAndPackForMma <
|
||||
RealElement,
|
||||
DestinationFragment,
|
||||
SourceFragment,
|
||||
MmaIterations,
|
||||
MmaOperandShape,
|
||||
Transform_,
|
||||
Operand::kA,
|
||||
Round_> {
|
||||
|
||||
//
|
||||
// Type definitions
|
||||
//
|
||||
static Operand const kOperand = Operand::kA;
|
||||
static ComplexTransform const kTransform = Transform_;
|
||||
static FloatRoundStyle const kRound = Round_;
|
||||
|
||||
// Data type of elements in the destination fragment
|
||||
using MmaElement = typename DestinationFragment::Element;
|
||||
|
||||
// Numeric convertor MmaElement <= RealElement
|
||||
using Converter = NumericConverter<MmaElement, RealElement, kRound>;
|
||||
|
||||
// Operand layout parameters
|
||||
using SourceFragmentLayout = layout::ColumnMajor;
|
||||
static int const kLdm = MmaIterations::kRow * MmaOperandShape::kRow;
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_DEVICE
|
||||
UnpackComplexConvertAndPackForMma() {}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void operator()(DestinationFragment *dest, SourceFragment const &source) {
|
||||
|
||||
Converter convert_op;
|
||||
SourceFragmentLayout layout(kLdm);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int i=0; i<MmaIterations::kRow; i++) {
|
||||
int pos = 0;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int c=0; c<MmaOperandShape::kColumn; c++) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int r=0; r<MmaOperandShape::kRow; r++) {
|
||||
// Logical position of element in source fragment
|
||||
int row = r + i * MmaOperandShape::kRow;
|
||||
int col = c;
|
||||
|
||||
// Access complex<RealElement> and apply rounding on real and imag parts
|
||||
MmaElement a = convert_op(source[layout(MatrixCoord{row,col})].real());
|
||||
MmaElement b = convert_op(source[layout(MatrixCoord{row,col})].imag());
|
||||
|
||||
// Unpack rounded complex<MmaElement> and pack into DestinationFragment for mma operation
|
||||
dest[i][pos] = a;
|
||||
dest[i+MmaIterations::kRow][pos++] = (kTransform == ComplexTransform::kConjugate ? -b : b);
|
||||
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// Partial specialization for OperandB and Congruous smem layout
|
||||
template <
|
||||
typename RealElement,
|
||||
typename DestinationFragment,
|
||||
typename SourceFragment,
|
||||
typename MmaIterations,
|
||||
typename MmaOperandShape,
|
||||
ComplexTransform Transform_,
|
||||
FloatRoundStyle Round_>
|
||||
struct UnpackComplexConvertAndPackForMma <
|
||||
RealElement,
|
||||
DestinationFragment,
|
||||
SourceFragment,
|
||||
MmaIterations,
|
||||
MmaOperandShape,
|
||||
Transform_,
|
||||
Operand::kB,
|
||||
Round_> {
|
||||
|
||||
//
|
||||
// Type definitions
|
||||
//
|
||||
static Operand const kOperand = Operand::kB;
|
||||
static ComplexTransform const kTransform = Transform_;
|
||||
static FloatRoundStyle const kRound = Round_;
|
||||
|
||||
// Data type of elements in the destination fragment
|
||||
using MmaElement = typename DestinationFragment::Element;
|
||||
|
||||
// Numeric convertor MmaElement <= RealElement
|
||||
using Converter = NumericConverter<MmaElement, RealElement, kRound>;
|
||||
|
||||
// Operand layout parameters
|
||||
using SourceFragmentLayout = layout::RowMajor;
|
||||
static int const kLdm = MmaIterations::kColumn * MmaOperandShape::kColumn;
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_DEVICE
|
||||
UnpackComplexConvertAndPackForMma() {}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void operator()(DestinationFragment *dest, SourceFragment const &source) {
|
||||
|
||||
Converter convert_op;
|
||||
SourceFragmentLayout layout(kLdm);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int i=0; i<MmaIterations::kColumn; i++) {
|
||||
int pos = 0;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int c=0; c<MmaOperandShape::kColumn; c++) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int r=0; r<MmaOperandShape::kRow; r++) {
|
||||
// Logical position of element in source fragment
|
||||
int row = r;
|
||||
int col = c + i * MmaOperandShape::kColumn;
|
||||
|
||||
// Access complex<RealElement> apply rounding on real and imag parts
|
||||
MmaElement a = convert_op(source[layout(MatrixCoord{row,col})].real());
|
||||
MmaElement b = convert_op(source[layout(MatrixCoord{row,col})].imag());
|
||||
|
||||
// Unpack rounded complex<MmaElement> and pack into DestinationFragment for mma operation
|
||||
dest[i][pos] = a;
|
||||
dest[i+MmaIterations::kColumn][pos++] = (kTransform == ComplexTransform::kConjugate ? -b : b);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
} // namespace detail
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
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: MmaLanePolicyTensorOp)
|
||||
using Policy = Policy_;
|
||||
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename Policy::Operator::Shape;
|
||||
|
||||
/// 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;
|
||||
|
||||
/// 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<Policy::Operator::Shape::kK, Policy::Operator::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 % 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 B 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());
|
||||
|
||||
// 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;
|
||||
|
||||
operand_A[0] = (kTransformA == ComplexTransform::kConjugate ? -A[m].imag() : A[m].imag());
|
||||
operand_B[0] = B[n].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 {
|
||||
//TODO: Implement this
|
||||
dst_A = A;
|
||||
dst_B = B;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// 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
|
||||
// Output data type: complex<float>
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
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,
|
||||
/// Used for partial specialization
|
||||
typename Enable
|
||||
>
|
||||
class MmaComplexTensorOp<
|
||||
Shape_,
|
||||
complex<float>,
|
||||
LayoutA_,
|
||||
complex<float>,
|
||||
LayoutB_,
|
||||
complex<float>,
|
||||
LayoutC_,
|
||||
Policy_,
|
||||
TransformA,
|
||||
TransformB,
|
||||
Enable> {
|
||||
public:
|
||||
/// Shape of warp-level matrix operation (concept: GemmShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Data type of members of complex multiplicand A
|
||||
using RealElementA = float;
|
||||
|
||||
/// 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 = float;
|
||||
|
||||
/// 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 = float;
|
||||
|
||||
/// 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_;
|
||||
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename Policy::Operator::Shape;
|
||||
|
||||
/// Underlying arch tag
|
||||
using ArchTag = typename Policy::Operator::ArchTag;
|
||||
|
||||
/// 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;
|
||||
|
||||
/// Storage for transformed A tile
|
||||
using TransformedFragmentA =
|
||||
Array<typename Policy::Operator::ElementA, FragmentA::kElements * 2>;
|
||||
|
||||
/// 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;
|
||||
|
||||
/// Storage for transformed B tile
|
||||
using TransformedFragmentB =
|
||||
Array<typename Policy::Operator::ElementB, FragmentB::kElements * 2>;
|
||||
|
||||
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 complex products operations performed (one complex product needs four mma instructions)
|
||||
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;
|
||||
|
||||
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,
|
||||
TransformedFragmentA const &A,
|
||||
TransformedFragmentB const &B,
|
||||
FragmentC const &C
|
||||
) const {
|
||||
|
||||
// Alias types for underlying real-valued matrix multiply operator
|
||||
using InstMmaOperandA = typename Policy::Operator::FragmentA;
|
||||
using InstMmaOperandB = typename Policy::Operator::FragmentB;
|
||||
using MmaOperandC = typename Policy::Operator::FragmentC;
|
||||
|
||||
static_assert(platform::is_same<cutlass::gemm::GemmShape<16, 8, 8>, typename Policy::Operator::Shape>::value,
|
||||
"This implementation only supports MMA.1688 math instructions.");
|
||||
|
||||
static_assert(InstMmaOperandA::kElements == 4,
|
||||
"This implementation only supports math instructions in which exactly four element is needed for the A operand."
|
||||
"We can geneneralize later.");
|
||||
|
||||
static_assert(InstMmaOperandB::kElements == 2,
|
||||
"This implementation only supports math instructions in which exactly two element is needed for the B operand."
|
||||
"We can geneneralize later.");
|
||||
|
||||
// Instruction Operands A & B holding real part followed by imaginary part for mma operations
|
||||
InstMmaOperandA const *operand_A = reinterpret_cast<InstMmaOperandA const *>(&A);
|
||||
InstMmaOperandB const *operand_B = reinterpret_cast<InstMmaOperandB const *>(&B);
|
||||
|
||||
//
|
||||
// Accumulate in place
|
||||
//
|
||||
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) {
|
||||
|
||||
// Real-valued accumulator part
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow);
|
||||
|
||||
mma(*accum, operand_A[m], operand_B[n], *accum);
|
||||
}
|
||||
|
||||
// mma(accum.imag(), a.real(), b.imag(), accum.imag());
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = MmaIterations::kColumn - 1; n >= 0; --n) {
|
||||
|
||||
// Complex-valued accumulator part
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow) + MmaIterations::kCount;
|
||||
|
||||
mma(*accum, operand_A[m], operand_B[n+MmaIterations::kColumn], *accum);
|
||||
}
|
||||
|
||||
// mma(accum.real(), a.imag(), -b.imag(), accum.real())
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < MmaIterations::kColumn; ++n) {
|
||||
|
||||
// negate OperandB to accumulate -(a.imag()*b.imag())
|
||||
// negating OperandB emits less instrucitons than negating OperandA as OperandB has less elements
|
||||
negate<InstMmaOperandB> negate_op;
|
||||
|
||||
// Real-valued accumulator part
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow);
|
||||
|
||||
mma(*accum, operand_A[m+MmaIterations::kRow], negate_op(operand_B[n+MmaIterations::kColumn]), *accum);
|
||||
}
|
||||
|
||||
// mma(accum.imag(), a.imag(), b.real(), accum.imag())
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = MmaIterations::kColumn - 1; n >= 0; --n) {
|
||||
|
||||
// Complex-valued accumulator part
|
||||
MmaOperandC *accum = reinterpret_cast<MmaOperandC *>(&D) +
|
||||
(m + n * MmaIterations::kRow) + MmaIterations::kCount;
|
||||
|
||||
mma(*accum, operand_A[m+MmaIterations::kRow], operand_B[n], *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 {
|
||||
// Alias types for underlying real-valued matrix multiply operator
|
||||
using InstMmaOperandA = typename Policy::Operator::FragmentA;
|
||||
using InstMmaOperandB = typename Policy::Operator::FragmentB;
|
||||
|
||||
//
|
||||
// Define conversions from source type to instruction operands' type
|
||||
//
|
||||
|
||||
FloatRoundStyle const kRoundA = FloatRoundStyle::round_half_ulp_trunc_dntz;
|
||||
FloatRoundStyle const kRoundB = FloatRoundStyle::round_half_ulp_trunc_dntz;
|
||||
|
||||
detail::UnpackComplexConvertAndPackForMma <
|
||||
RealElementA,
|
||||
InstMmaOperandA,
|
||||
FragmentA,
|
||||
MmaIterations,
|
||||
MatrixShape<2, 2>,
|
||||
kTransformA,
|
||||
Operand::kA,
|
||||
kRoundA> convert_A;
|
||||
|
||||
detail::UnpackComplexConvertAndPackForMma <
|
||||
RealElementB,
|
||||
InstMmaOperandB,
|
||||
FragmentB,
|
||||
MmaIterations,
|
||||
MatrixShape<2, 1>,
|
||||
kTransformB,
|
||||
Operand::kB,
|
||||
kRoundB> convert_B;
|
||||
|
||||
// Convert Fragment[A|B] holding complex<RealElement[A|B]> to InstMmaOperand[A|B] holding InstMmaOperand[A|B]::Element
|
||||
convert_A(reinterpret_cast<InstMmaOperandA *>(&dst_A), A);
|
||||
convert_B(reinterpret_cast<InstMmaOperandB *>(&dst_B), B);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// TODO - partial specializations of real*complex and complex*real
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,357 @@
|
||||
/***************************************************************************************************
|
||||
* 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 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/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_gaussian_complex_tensor_op_tile_iterator_sm80.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 MmaGaussianComplexTensorOp;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// 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 MmaGaussianComplexTensorOp<
|
||||
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_;
|
||||
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename Policy::Operator::Shape;
|
||||
|
||||
/// Underlying architecture tag
|
||||
using ArchTag = typename Policy::Operator::ArchTag;
|
||||
|
||||
/// 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;
|
||||
|
||||
/// 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<Policy::Operator::Shape::kK, Policy::Operator::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 % 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 = MmaTensorOpGaussianComplexAccumulatorTileIterator<
|
||||
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 '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 * Policy::Operator::FragmentC::kElements,
|
||||
"Unexpected gaussian complex fragment length.");
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Underlying real-valued matrix multiply operator (concept: arch::Mma)
|
||||
typename Policy::Operator 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 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 B operand."
|
||||
"We can geneneralize later.");
|
||||
|
||||
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;
|
||||
|
||||
operand_Asum[0] = A[m].real() + ((kTransformA == ComplexTransform::kConjugate) ? -A[m].imag() : +A[m].imag());
|
||||
operand_Br[0] = B[n].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;
|
||||
|
||||
operand_Ar[0] = -A[m].real();
|
||||
operand_Bdiff[0] = B[n].real() - ((kTransformB == ComplexTransform::kConjugate) ? -B[n].imag() : +B[n].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;
|
||||
|
||||
operand_Ai[0] = (kTransformA == ComplexTransform::kConjugate) ? -A[m].imag() : +A[m].imag();
|
||||
operand_Bsum[0] = B[n].real() + ((kTransformB == ComplexTransform::kConjugate) ? -B[n].imag() : +B[n].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 {
|
||||
//TODO: Implement this
|
||||
dst_A = A;
|
||||
dst_B = B;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// TODO - partial specializations of real*complex and complex*real
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,384 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Defines iterators used by warp-level matrix multiply operations targeting Tensor Cores.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
|
||||
#include "cutlass/arch/memory_sm75.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/layout/tensor.h"
|
||||
#include "cutlass/layout/pitch_linear.h"
|
||||
#include "cutlass/layout/tensor_op_multiplicand_sm80.h"
|
||||
#include "cutlass/gemm/warp/mma_complex_tensor_op_tile_iterator_sm80.h"
|
||||
|
||||
#include "cutlass/platform/platform.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace warp {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Element type
|
||||
typename Element_,
|
||||
/// Layout of operand in memory
|
||||
typename Layout_,
|
||||
/// 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 MmaTensorOpGaussianComplexAccumulatorTileIterator;
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
///
|
||||
/// Partial specialization for complex<T>
|
||||
///
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Data type of underlying field of reals.
|
||||
typename RealElement,
|
||||
/// 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 MmaTensorOpGaussianComplexAccumulatorTileIterator<
|
||||
Shape_, complex<RealElement>, cutlass::layout::RowMajor, 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 = complex<RealElement>;
|
||||
|
||||
/// 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_assert(
|
||||
!(Shape::kRow % InstructionShape::kM) &&
|
||||
!(Shape::kColumn % InstructionShape::kN),
|
||||
"Shape of warp-level Mma must be divisible by operator shape.");
|
||||
|
||||
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,
|
||||
Shape::kColumn / 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. It is assumed that the accumulators
|
||||
/// are stored in a gaussian complex arrangement with parts 1, 2, and 3 as entirely contiguous
|
||||
/// arranged as [part1, part2, part3]
|
||||
using Fragment = Array<RealElement, (Shape::kCount / kThreads) * 3>;
|
||||
|
||||
static int const kPart1Index = (Shape::kCount / kThreads) * 0;
|
||||
static int const kPart2Index = (Shape::kCount / kThreads) * 1;
|
||||
static int const kPart3Index = (Shape::kCount / kThreads) * 2;
|
||||
|
||||
private:
|
||||
|
||||
/// Reference to output tensor
|
||||
TensorRef ref_;
|
||||
|
||||
public:
|
||||
|
||||
/// Default ctor constructs null iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpGaussianComplexAccumulatorTileIterator() { }
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpGaussianComplexAccumulatorTileIterator(
|
||||
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
|
||||
MmaTensorOpGaussianComplexAccumulatorTileIterator &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
|
||||
MmaTensorOpGaussianComplexAccumulatorTileIterator &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
|
||||
MmaTensorOpGaussianComplexAccumulatorTileIterator & operator++() {
|
||||
// deliberate no-op
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpGaussianComplexAccumulatorTileIterator & operator--() {
|
||||
// deliberate no-op
|
||||
return *this;
|
||||
}
|
||||
|
||||
///< advances in units of whole tiles along the logical coordinate space of the tensor
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpGaussianComplexAccumulatorTileIterator & 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
|
||||
MmaTensorOpGaussianComplexAccumulatorTileIterator & 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;
|
||||
|
||||
Element z = offset_ref.at({accum_m, accum_n});
|
||||
|
||||
frag[mma_accum_start + row * kElementsPerAccess + col + kPart1Index] = z.real() + z.imag();
|
||||
frag[mma_accum_start + row * kElementsPerAccess + col + kPart2Index] = -z.real();
|
||||
frag[mma_accum_start + row * kElementsPerAccess + col + kPart3Index] = z.imag();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 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;
|
||||
|
||||
Element z(frag[kPart1Index + idx] - frag[kPart3Index + idx],
|
||||
frag[kPart1Index + idx] + frag[kPart2Index + idx]);
|
||||
|
||||
offset_ref.at({accum_m, accum_n}) = z;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 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);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -147,6 +147,9 @@ public:
|
||||
dp4a_type
|
||||
>;
|
||||
|
||||
/// Shape of the underlying instruction
|
||||
using InstructionShape = GemmShape<1,1,use_dp4a ? 4 : 1>;
|
||||
|
||||
public:
|
||||
|
||||
/// Iterates over the A operand in memory
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -39,12 +39,16 @@
|
||||
|
||||
#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_tile_iterator.h"
|
||||
#include "cutlass/gemm/warp/mma_tensor_op_tile_iterator_sm80.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
@@ -77,6 +81,27 @@ struct ConvertAndPack<T, T, N, Round> {
|
||||
}
|
||||
};
|
||||
|
||||
template <int N, FloatRoundStyle Round>
|
||||
struct ConvertAndPack<bfloat16_t, float, N, Round> {
|
||||
|
||||
using Converter = NumericArrayConverter<bfloat16_t, float, N, Round>;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<bfloat16_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);
|
||||
}
|
||||
};
|
||||
|
||||
template <int N, FloatRoundStyle Round>
|
||||
struct ConvertAndPack<half_t, float, N, Round> {
|
||||
|
||||
@@ -130,8 +155,6 @@ template <
|
||||
/// Store the accumulators in row major or column major. Row major is used
|
||||
/// when output layout is interleaved.
|
||||
bool AccumulatorsInRowMajor = false,
|
||||
/// PartitionsN indicating how many PartitionsN for multiplicand B
|
||||
int PartitionsN_ = 1,
|
||||
/// Used for partial specialization
|
||||
typename Enable = bool
|
||||
>
|
||||
@@ -167,6 +190,9 @@ public:
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename Policy::Operator::Shape;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = ComplexTransform::kNone;
|
||||
|
||||
@@ -179,9 +205,6 @@ public:
|
||||
/// Number of partitions along K dimension
|
||||
static int const kPartitionsK = PartitionsK_;
|
||||
|
||||
/// PartitionsN indicating how many PartitionsN for multiplicand B
|
||||
static int const kPartitionsN = PartitionsN_;
|
||||
|
||||
public:
|
||||
|
||||
/// Iterates over the A operand in memory
|
||||
@@ -228,9 +251,7 @@ private:
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = MatrixShape<
|
||||
Shape::kM / Policy::Operator::Shape::kM,
|
||||
(Shape::kN / Policy::Operator::Shape::kN / kPartitionsN > 0) ?
|
||||
Shape::kN / Policy::Operator::Shape::kN / kPartitionsN :
|
||||
1
|
||||
Shape::kN / Policy::Operator::Shape::kN
|
||||
>;
|
||||
|
||||
public:
|
||||
@@ -254,8 +275,8 @@ public:
|
||||
FragmentC &D,
|
||||
TransformedFragmentA const &A,
|
||||
TransformedFragmentB const &B,
|
||||
FragmentC const &C,
|
||||
int const &partitionN_idx = 0) const {
|
||||
FragmentC const &C
|
||||
) const {
|
||||
|
||||
using MmaOperandA = typename Policy::Operator::FragmentA;
|
||||
using MmaOperandB = typename Policy::Operator::FragmentB;
|
||||
@@ -267,8 +288,7 @@ public:
|
||||
MmaOperandB const *ptr_B = reinterpret_cast<MmaOperandB const *>(&B);
|
||||
MmaOperandC *ptr_D = reinterpret_cast<MmaOperandC *>(&D);
|
||||
|
||||
// The offset of multilicand B for current partition
|
||||
const int n_off = partitionN_idx * FragmentB::kElements / MmaOperandB::kElements / kPartitionsN;
|
||||
#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) {
|
||||
@@ -286,24 +306,46 @@ public:
|
||||
ptr_D[n + m_serpentine * MmaIterations::kColumn]);
|
||||
} else {
|
||||
mma(
|
||||
ptr_D[m_serpentine + (n + n_off) * MmaIterations::kRow],
|
||||
ptr_D[m_serpentine + n * MmaIterations::kRow],
|
||||
ptr_A[m_serpentine],
|
||||
ptr_B[n + n_off],
|
||||
ptr_D[m_serpentine + (n + n_off) * MmaIterations::kRow]);
|
||||
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]);
|
||||
}
|
||||
}
|
||||
}
|
||||
#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 {
|
||||
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
|
||||
@@ -314,6 +356,7 @@ public:
|
||||
FloatRoundStyle const kRoundB =
|
||||
PreferredRoundingMode<typename Policy::Operator::ElementB,
|
||||
ElementB>::kRound;
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ < 800)
|
||||
detail::ConvertAndPack<typename Policy::Operator::ElementA, ElementA,
|
||||
FragmentA::kElements, kRoundA>
|
||||
convert_A;
|
||||
@@ -331,6 +374,26 @@ public:
|
||||
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 Policy::Operator::ElementA, ElementA,
|
||||
FragmentA::kElements / 2, kRoundA>
|
||||
convert_A;
|
||||
NumericArrayConverter<typename Policy::Operator::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 Policy::Operator::ElementA, FragmentA::kElements / 2> *
|
||||
ptr_dst_A = reinterpret_cast<Array<typename Policy::Operator::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
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -0,0 +1,428 @@
|
||||
/*! \file
|
||||
\brief This defines a "fragment" iterator for visiting the fragments of a warp tile
|
||||
that participate in one warp-level mma operation.
|
||||
|
||||
Typically, this is used to access the accumulator tile/fragement of a warp-level mma operation.
|
||||
The accumulator tile is then partitioned into smaller tiles/fragments that can be fed into
|
||||
next warp-level mma operation.
|
||||
|
||||
This iterator is necessary to accomplish warp-level mma fusion where the accumulator tile is
|
||||
reused as multiplicand tile for the next mma.
|
||||
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/layout/tensor.h"
|
||||
#include "cutlass/numeric_conversion.h"
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace warp {
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Size of the accumulation tile shape (concept: MatrixShape)
|
||||
typename AccumulatorShape_,
|
||||
/// KBlocks columns to compute residual
|
||||
int KBlocksColumn_,
|
||||
/// Accumulator Element type
|
||||
typename ElementAccumulator_,
|
||||
/// Element type
|
||||
typename Element_,
|
||||
/// Layout of operand in memory
|
||||
typename Layout_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Output operation on the fragment
|
||||
typename OutputOp_,
|
||||
/// Whether beta is zero
|
||||
bool IsBetaZero_ >
|
||||
class MmaTensorOpFragmentIterator;
|
||||
|
||||
|
||||
// Partial specialization for col-major accumulator tile
|
||||
// And Element type is the same as Accumulator Element type
|
||||
|
||||
template <
|
||||
/// Shape of warp tile to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Shape of the warp accumulation tile (concept: MatrixShape)
|
||||
typename AccumulatorShape_,
|
||||
/// KBlocks columns to compute residual
|
||||
int KBlocksColumn_,
|
||||
/// Element type
|
||||
typename Element_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Output operation on fragment
|
||||
typename OutputOp_>
|
||||
class MmaTensorOpFragmentIterator<Shape_, AccumulatorShape_, KBlocksColumn_, Element_, Element_,
|
||||
cutlass::layout::ColumnMajor,
|
||||
InstructionShape_, OutputOp_, true> {
|
||||
public:
|
||||
|
||||
/// Shape of warp tile to load (concept: MatrixShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Shape of the warp accumulation tile (concept: MatrixShape)
|
||||
using AccumulatorShape = AccumulatorShape_;
|
||||
|
||||
/// KBlocks columns to compute residual
|
||||
static int const kKBlockColumn = KBlocksColumn_;
|
||||
|
||||
/// Element type
|
||||
using Element = Element_;
|
||||
|
||||
/// Layout of source tile
|
||||
using Layout = cutlass::layout::ColumnMajor;
|
||||
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
using InstructionShape = InstructionShape_;
|
||||
|
||||
/// 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;
|
||||
|
||||
/// Internal structure of iterator - made public to enable introspection
|
||||
struct Policy {
|
||||
static_assert(
|
||||
!(Shape::kRow % InstructionShape::kM) &&
|
||||
!(Shape::kColumn % InstructionShape::kN),
|
||||
"Shape of warp-level Mma must be divisible by operator shape.");
|
||||
static_assert(
|
||||
!(AccumulatorShape::kRow % Shape::kRow) &&
|
||||
!(AccumulatorShape::kColumn % Shape::kColumn),
|
||||
"Shape of Warp Accumulator must be divisible by warp shape.");
|
||||
static_assert(
|
||||
!(kKBlockColumn % Shape::kColumn),
|
||||
"KBlock size must be divisible by warp shape.");
|
||||
|
||||
/// Number of times this iterator can be incremented
|
||||
static int const kIterations = AccumulatorShape::kCount / Shape::kCount;
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
static int const kElementsPerAccess = InstructionShape::kM * InstructionShape::kN / kThreads;
|
||||
|
||||
/// Number of mma operations performed by a warp
|
||||
using MmaIterations = MatrixShape<Shape::kRow / InstructionShape::kM,
|
||||
Shape::kColumn / InstructionShape::kN>;
|
||||
/// Number of mma operations performed by the entire accumulator
|
||||
using AccumulatorIterations = MatrixShape<AccumulatorShape::kRow / InstructionShape::kM,
|
||||
AccumulatorShape::kColumn / InstructionShape::kN>;
|
||||
|
||||
/// Number of K iterations
|
||||
static int const kKBlockIterations = (AccumulatorShape::kColumn + kKBlockColumn - 1) / kKBlockColumn;
|
||||
static int const kResidualColumn = AccumulatorShape::kColumn - (kKBlockIterations - 1) * kKBlockColumn;
|
||||
static int const kKBlockColumnIterations = kKBlockColumn / Shape::kColumn
|
||||
* (AccumulatorShape::kRow / Shape::kRow);
|
||||
static int const kResidualIndex = kResidualColumn / Shape::kColumn
|
||||
* (AccumulatorShape::kRow / Shape::kRow);
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Derived quantities
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
/// This is the fragment size produced by one access of the iterator.
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads>;
|
||||
|
||||
/// Accumulator Fragment object
|
||||
using AccumulatorFragment = Array<Element, AccumulatorShape::kCount / kThreads>;
|
||||
|
||||
|
||||
private:
|
||||
|
||||
/// Internal access type
|
||||
using AccessType = Array<Element, kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Accumulator tile
|
||||
AccessType const *accumulators_;
|
||||
|
||||
/// Internal index
|
||||
int index_;
|
||||
|
||||
/// Used to access residual tile first
|
||||
bool is_residual_tile_;
|
||||
|
||||
public:
|
||||
/// Constructs an iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpFragmentIterator(AccumulatorFragment const &accum)
|
||||
: accumulators_(reinterpret_cast<AccessType const *>(&accum)),
|
||||
index_(0), is_residual_tile_(true) {}
|
||||
|
||||
/// Add offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_offset(int index_offset) {
|
||||
index_ += index_offset;
|
||||
if(is_residual_tile_ && index_ >= kKBlockColumnIterations) {
|
||||
index_ = index_ - kKBlockColumnIterations + kResidualIndex;
|
||||
is_residual_tile_ = false;
|
||||
}
|
||||
}
|
||||
|
||||
/// Increments
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpFragmentIterator &operator++() {
|
||||
add_offset(1);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Decrements
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpFragmentIterator &operator--() {
|
||||
add_offset(-1);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Loads a fragment from the referenced part of the accumulator tile
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag, OutputOp output_op) const {
|
||||
|
||||
if (output_op.is_source_needed()) //beta must be zero
|
||||
assert(0);
|
||||
|
||||
AccessType src_fragment;
|
||||
src_fragment.clear();
|
||||
|
||||
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
|
||||
int index_m = (index_ * MmaIterations::kRow) % AccumulatorIterations::kRow;
|
||||
int index_n = (index_ * MmaIterations::kRow) / AccumulatorIterations::kRow
|
||||
* MmaIterations::kColumn;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < MmaIterations::kColumn; n++) {
|
||||
for (int m = 0; m < MmaIterations::kRow; m++) {
|
||||
int accumulator_access_offset =
|
||||
(n + index_n) * AccumulatorIterations::kRow + m + index_m;
|
||||
|
||||
frag_ptr[n * MmaIterations::kRow + m].clear();
|
||||
if(!(is_residual_tile_ && index_ >= kResidualIndex))
|
||||
//frag_ptr[n * MmaIterations::kRow + m] = accumulators_[accumulator_access_offset];
|
||||
frag_ptr[n * MmaIterations::kRow + m] = output_op(accumulators_[accumulator_access_offset], src_fragment);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
// Partial specialization for row-major accumulator tile
|
||||
|
||||
template <
|
||||
/// Shape of warp tile to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Shape of the warp accumulation tile (concept: MatrixShape)
|
||||
typename AccumulatorShape_,
|
||||
/// KBlocks columns to compute residual
|
||||
int KBlocksColumn_,
|
||||
/// Accumulator Element type
|
||||
typename ElementAccumulator_,
|
||||
/// Element type
|
||||
typename Element_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Output operation on fragment
|
||||
typename OutputOp_>
|
||||
class MmaTensorOpFragmentIterator<Shape_, AccumulatorShape_, KBlocksColumn_, ElementAccumulator_, Element_,
|
||||
cutlass::layout::RowMajor,
|
||||
InstructionShape_, OutputOp_, true> {
|
||||
public:
|
||||
|
||||
/// Shape of warp tile to load (concept: MatrixShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Shape of the warp accumulation tile (concept: MatrixShape)
|
||||
using AccumulatorShape = AccumulatorShape_;
|
||||
|
||||
/// KBlocks columns to compute residual
|
||||
static int const kKBlockColumn = KBlocksColumn_;
|
||||
|
||||
/// Accumulator Element type
|
||||
using ElementAccumulator = ElementAccumulator_;
|
||||
|
||||
/// 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_;
|
||||
|
||||
/// 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;
|
||||
|
||||
/// Internal structure of iterator - made public to enable introspection
|
||||
struct Policy {
|
||||
static_assert(
|
||||
!(Shape::kRow % InstructionShape::kM) &&
|
||||
!(Shape::kColumn % InstructionShape::kN),
|
||||
"Shape of warp-level Mma must be divisible by operator shape.");
|
||||
static_assert(
|
||||
!(AccumulatorShape::kRow % Shape::kRow) &&
|
||||
!(AccumulatorShape::kColumn % Shape::kColumn),
|
||||
"Shape of Warp Accumulator must be divisible by warp shape.");
|
||||
static_assert(
|
||||
!(kKBlockColumn % Shape::kColumn),
|
||||
"KBlock size must be divisible by warp shape.");
|
||||
|
||||
/// Number of times this iterator can be incremented
|
||||
static int const kIterations = AccumulatorShape::kCount / Shape::kCount;
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
static int const kElementsPerAccess = InstructionShape::kM * InstructionShape::kN / kThreads;
|
||||
|
||||
/// Number of mma operations performed by a warp
|
||||
using MmaIterations = MatrixShape<Shape::kRow / InstructionShape::kM,
|
||||
Shape::kColumn / InstructionShape::kN>;
|
||||
/// Number of mma operations performed by the entire accumulator
|
||||
using AccumulatorIterations = MatrixShape<AccumulatorShape::kRow / InstructionShape::kM,
|
||||
AccumulatorShape::kColumn / InstructionShape::kN>;
|
||||
|
||||
/// Number of K iterations
|
||||
static int const kKBlockIterations = (AccumulatorShape::kColumn + kKBlockColumn - 1) / kKBlockColumn;
|
||||
static int const kResidualColumn = AccumulatorShape::kColumn - (kKBlockIterations - 1) * kKBlockColumn;
|
||||
static int const kKBlockColumnIterations = kKBlockColumn / Shape::kColumn
|
||||
* (AccumulatorShape::kRow / Shape::kRow);
|
||||
static int const kResidualIndex = kResidualColumn / Shape::kColumn
|
||||
* (AccumulatorShape::kRow / Shape::kRow);
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Derived quantities
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
/// This is the fragment size produced by one access of the iterator.
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads>;
|
||||
|
||||
/// Accumulator Fragment object
|
||||
using AccumulatorFragment = Array<ElementAccumulator, AccumulatorShape::kCount / kThreads>;
|
||||
|
||||
|
||||
private:
|
||||
|
||||
/// Internal access type
|
||||
using AccessType = Array<ElementAccumulator, kElementsPerAccess>;
|
||||
using FragmentAccessType = Array<Element, kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Accumulator tile
|
||||
AccessType const *accumulators_;
|
||||
|
||||
/// Internal index
|
||||
int index_;
|
||||
|
||||
/// Used to access residual tile first
|
||||
bool is_residual_tile_;
|
||||
|
||||
public:
|
||||
/// Constructs an iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpFragmentIterator(AccumulatorFragment const &accum)
|
||||
: accumulators_(reinterpret_cast<AccessType const *>(&accum)),
|
||||
index_(0), is_residual_tile_(true) {}
|
||||
|
||||
/// Add offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_offset(int index_offset) {
|
||||
index_ += index_offset;
|
||||
if(is_residual_tile_ && index_ >= kKBlockColumnIterations) {
|
||||
index_ = index_ - kKBlockColumnIterations + kResidualIndex;
|
||||
is_residual_tile_ = false;
|
||||
}
|
||||
}
|
||||
|
||||
/// Increments
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpFragmentIterator &operator++() {
|
||||
add_offset(1);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Decrements
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpFragmentIterator &operator--() {
|
||||
add_offset(-1);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Loads a fragment from the referenced part of the accumulator tile
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag, OutputOp output_op) const {
|
||||
|
||||
if (output_op.is_source_needed()) //beta must be zero
|
||||
assert(0);
|
||||
|
||||
FragmentAccessType src_fragment;
|
||||
src_fragment.clear();
|
||||
|
||||
FragmentAccessType *frag_ptr = reinterpret_cast<FragmentAccessType *>(&frag);
|
||||
// NumericArrayConverter<Element, ElementAccumulator, kElementsPerAccess, FloatRoundStyle::round_indeterminate> fragmentConverter;
|
||||
|
||||
int index_m = (index_ * MmaIterations::kRow) % AccumulatorIterations::kRow;
|
||||
int index_n = (index_ * MmaIterations::kRow) / AccumulatorIterations::kRow
|
||||
* MmaIterations::kColumn;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int m = 0; m < MmaIterations::kRow; m++) {
|
||||
for (int n = 0; n < MmaIterations::kColumn; n++) {
|
||||
int accumulator_access_offset =
|
||||
(m + index_m) * AccumulatorIterations::kColumn + n + index_n;
|
||||
|
||||
frag_ptr[m * MmaIterations::kColumn + n].clear();
|
||||
if(!(is_residual_tile_ && index_ >= kResidualIndex))
|
||||
// frag_ptr[m * MmaIterations::kColumn + n] = fragmentConverter(accumulators_[accumulator_access_offset]);
|
||||
frag_ptr[m * MmaIterations::kColumn + n] = output_op(accumulators_[accumulator_access_offset], src_fragment);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -106,6 +106,9 @@ public:
|
||||
/// Architecture tag
|
||||
using ArchTag = arch::Sm70;
|
||||
|
||||
/// Underlying instruction shape
|
||||
using InstructionShape = typename Policy::Operator::Shape;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = ComplexTransform::kNone;
|
||||
|
||||
@@ -210,8 +213,7 @@ public:
|
||||
FragmentC &D,
|
||||
FragmentA const &A,
|
||||
FragmentB const &B,
|
||||
FragmentC const &C,
|
||||
int const &partitionN_idx = 0) {
|
||||
FragmentC const &C) {
|
||||
|
||||
using MmaOperandA = typename Policy::Operator::FragmentA;
|
||||
using MmaOperandB = typename Policy::Operator::FragmentB;
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -229,8 +229,11 @@ public:
|
||||
k_group_idx_(0) {
|
||||
|
||||
int quad_pair = (lane_id >> 3);
|
||||
int quad_quad = (lane_id >> 4);
|
||||
int lane_in_quad = (lane_id & 3);
|
||||
int lane_in_quad_pair = (lane_id & 7);
|
||||
int lane_in_quad_quad = (lane_id & 15);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kPointerCount; ++i) {
|
||||
int partition_contiguous_idx = -1;
|
||||
@@ -242,6 +245,24 @@ public:
|
||||
access_contiguous_idx = (quad_pair ^ lane_in_quad);
|
||||
access_strided_idx = lane_in_quad_pair;
|
||||
}
|
||||
else if (Policy::LdsmShape::kContiguous == 2 &&
|
||||
kOperand == Operand::kA) {
|
||||
// Matrix multiply 16816 A
|
||||
// Q0 Q2
|
||||
// Q1 Q3
|
||||
partition_contiguous_idx = ((lane_in_quad_pair >> 2) ^ (i >> 1));
|
||||
access_contiguous_idx =
|
||||
(((quad_pair & 1) + ((i & 1) << 1)) ^ lane_in_quad);
|
||||
access_strided_idx = lane_in_quad_pair + (lane_id >> 4 << 3);
|
||||
} else if (Policy::LdsmShape::kContiguous == 2 &&
|
||||
kOperand == Operand::kB) {
|
||||
// Matrix multiply 16816 B
|
||||
// Q0 Q1
|
||||
// Q2 Q3
|
||||
partition_contiguous_idx = ((lane_in_quad_pair >> 2) ^ (i >> 1));
|
||||
access_contiguous_idx = ((quad_quad + ((i & 1) << 1)) ^ lane_in_quad);
|
||||
access_strided_idx = lane_in_quad_quad;
|
||||
}
|
||||
int access_contiguous =
|
||||
partition_contiguous_idx * Layout::PartitionShape::kContiguous +
|
||||
access_contiguous_idx;
|
||||
@@ -436,6 +457,364 @@ public:
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// This tile iterator is specialized for 32-thread MMA.TF32 NT TensorOps. It
|
||||
/// uses LDS.32 to load from shared memory and therefore must be initialized
|
||||
/// with a TensorRef to shared memory.
|
||||
///
|
||||
/// Satisfies:
|
||||
/// ReadableRandomAccessContiguousTileIteratorConcept
|
||||
///
|
||||
template <
|
||||
/// Size of the matrix to load (concept: PitchLinearShape)
|
||||
typename Shape_,
|
||||
/// Identifies A or B multiplicand
|
||||
Operand Operand_,
|
||||
/// Data type of elements
|
||||
typename Element_,
|
||||
/// Shape of one matrix product operation (concept: PitchLinearShape)
|
||||
typename InstructionShape_,
|
||||
/// Interval between adjacent *MMA instructions (in units of MMA
|
||||
/// instructions)
|
||||
int OpDelta_,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK_>
|
||||
class MmaTensorOpMultiplicandTileIterator<
|
||||
Shape_, Operand_, Element_,
|
||||
cutlass::layout::TensorOpMultiplicandCongruous<32, 32>, InstructionShape_,
|
||||
OpDelta_, 32, PartitionsK_> {
|
||||
public:
|
||||
/// Shape of tile to load (concept: PitchLinearShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Operand tag
|
||||
static Operand const kOperand = Operand_;
|
||||
|
||||
static_assert(kOperand == Operand::kA || kOperand == Operand::kB,
|
||||
"MmaTensorOpMultiplicandIterator may only be instantiated for "
|
||||
"A or B operands to warp-level Mma.");
|
||||
|
||||
/// Element type
|
||||
using Element = Element_;
|
||||
|
||||
/// Layout of source tile
|
||||
using Layout = cutlass::layout::TensorOpMultiplicandCongruous<32, 32>;
|
||||
|
||||
/// Shape of one matrix product operation (concept: GemmShape)
|
||||
using InstructionShape = InstructionShape_;
|
||||
|
||||
/// Delta between *MMA operations (in units of *MMA operations, concept:
|
||||
/// MatrixShape)
|
||||
static int const kOpDelta = OpDelta_;
|
||||
|
||||
/// Number of participating threads
|
||||
static int const kThreads = 32;
|
||||
|
||||
/// Number of partitions along K dimension
|
||||
static int const kPartitionsK = PartitionsK_;
|
||||
|
||||
/// 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_assert(
|
||||
!(Shape::kContiguous % InstructionShape::kContiguous),
|
||||
"Shape of warp-level Mma must be divisible by operator shape.");
|
||||
|
||||
// Determine number of elements along outer dimension per individual LDS.32
|
||||
// op. Every one warp of LDS.32 loads 8x4 elements
|
||||
static int const kLdsOpInner = Layout::TileShape::kStrided;
|
||||
static int const kLdsOpOuter = kThreads / kLdsOpInner;
|
||||
|
||||
static_assert(!(Shape::kContiguous % kLdsOpOuter),
|
||||
"Shape of warp-level mma must be divisible by LDS.32's "
|
||||
"fundamental tile size.");
|
||||
|
||||
static_assert(!(Shape::kStrided % kLdsOpInner),
|
||||
"Shape of warp-level mma must be divisible by LDS.32's "
|
||||
"fundamental tile size.");
|
||||
|
||||
/// Number of LDS.32 instructions needed by one MMA instruction
|
||||
/// 1684 A 2x1
|
||||
/// 1684 B 1x1
|
||||
/// 1688 A 2x2
|
||||
/// 1688 B 1x2
|
||||
static int const LdsShapeContiguous =
|
||||
InstructionShape::kContiguous / kLdsOpOuter;
|
||||
static int const LdsShapeStrided = InstructionShape::kStrided / kLdsOpInner;
|
||||
using LdsShape =
|
||||
layout::PitchLinearShape<LdsShapeContiguous, LdsShapeStrided>;
|
||||
|
||||
/// Number and arrangement of LDS instructions
|
||||
using LdsIterations = layout::PitchLinearShape<
|
||||
Shape::kContiguous / LdsShapeContiguous / kLdsOpOuter, 1>;
|
||||
|
||||
/// Number of groups for each tile
|
||||
static int const kGroupsPerTile =
|
||||
Shape::kStrided / InstructionShape::kStrided;
|
||||
};
|
||||
|
||||
private:
|
||||
/// Not working on this feature at the moment.
|
||||
static_assert(kOpDelta == 1,
|
||||
"Alternative arrangements not supported at present.");
|
||||
|
||||
/// Number of internal pointers needed to reference shared memory
|
||||
static int const kPointerCount = Layout::TileShape::kContiguous *
|
||||
Layout::kElementsPerAccess /
|
||||
Policy::kLdsOpOuter;
|
||||
|
||||
/// Vectorized access is not used
|
||||
static int const kElementsPerAccess = 1;
|
||||
|
||||
/// Pointer type used for accesses
|
||||
using AccessType = Element;
|
||||
|
||||
/// Internal counter used to jump to next K partition
|
||||
int k_group_idx_;
|
||||
|
||||
public:
|
||||
//
|
||||
// Derived quantities
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment =
|
||||
Array<Element, Shape::kContiguous * InstructionShape::kStrided / kThreads>;
|
||||
|
||||
private:
|
||||
/// Layout object storing stride values
|
||||
Index stride_;
|
||||
|
||||
/// Shared memory base pointers - not advanced
|
||||
AccessType const *pointer_[kPointerCount];
|
||||
|
||||
/// Byte offset incremented as iterator advances
|
||||
Index byte_offset_;
|
||||
|
||||
public:
|
||||
/// Default ctor constructs null iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator() : stride_(0), byte_offset_(0) {}
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator(TensorRef const &ref, int lane_id)
|
||||
: stride_(ref.stride(0)), byte_offset_(0), k_group_idx_(0) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kPointerCount; ++i) {
|
||||
int access_strided = lane_id % Policy::kLdsOpInner;
|
||||
int access_contiguous = (lane_id / Policy::kLdsOpInner) +
|
||||
(access_strided ^ i) * Policy::kLdsOpOuter;
|
||||
|
||||
pointer_[i] = reinterpret_cast<AccessType const *>(ref.data()) +
|
||||
access_contiguous + access_strided * stride_;
|
||||
}
|
||||
}
|
||||
|
||||
/// Adds a pointer offset to internal pointer(s) to advance through memory
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator &add_pointer_offset(LongIndex offset) {
|
||||
byte_offset_ += offset * sizeof(Element);
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole
|
||||
/// tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator &add_tile_offset(
|
||||
TensorCoord const &tile_offset) {
|
||||
int contiguous_offset = tile_offset.contiguous();
|
||||
if (Shape::kContiguous ==
|
||||
Layout::TileShape::kContiguous * Layout::kElementsPerAccess / 2) {
|
||||
if (tile_offset.contiguous() % 2) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kPointerCount / 2; ++i) {
|
||||
AccessType const *tmp_pointer = pointer_[i];
|
||||
pointer_[i] = pointer_[i + kPointerCount / 2];
|
||||
pointer_[i + kPointerCount / 2] = tmp_pointer;
|
||||
}
|
||||
}
|
||||
contiguous_offset = (tile_offset.contiguous() >> 1) << 1;
|
||||
}
|
||||
|
||||
int offset = (tile_offset.strided() * InstructionShape::kStrided) * stride_ +
|
||||
contiguous_offset * Shape::kContiguous;
|
||||
|
||||
add_pointer_offset(offset);
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator &operator++() {
|
||||
add_tile_offset({0, 1});
|
||||
|
||||
if (kPartitionsK > 1) {
|
||||
++k_group_idx_;
|
||||
// Jump to next stage
|
||||
if (k_group_idx_ == Policy::kGroupsPerTile) {
|
||||
k_group_idx_ = 0;
|
||||
add_tile_offset(
|
||||
{0, ((kPartitionsK - 1) * Policy::kGroupsPerTile)});
|
||||
}
|
||||
}
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the opposite of the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator &operator--() {
|
||||
byte_offset_ -= stride_ * InstructionShape::kStrided * sizeof(Element) *
|
||||
kElementsPerAccess;
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
///< advances in units of whole tiles along the logical coordinate space of
|
||||
///< the tensor
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator &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
|
||||
MmaTensorOpMultiplicandTileIterator &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_byte_offset(frag, 0); }
|
||||
|
||||
/// 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 {
|
||||
Element *fetch_ptr = reinterpret_cast<Element *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < Policy::LdsIterations::kStrided; ++s) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int c = 0; c < Policy::LdsIterations::kContiguous; ++c) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ss = 0; ss < Policy::LdsShape::kStrided; ++ss) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int cc = 0; cc < Policy::LdsShape::kContiguous; ++cc) {
|
||||
int access_idx =
|
||||
cc + (ss + (c + s * Policy::LdsIterations::kContiguous) *
|
||||
Policy::LdsShape::kStrided) *
|
||||
Policy::LdsShape::kContiguous;
|
||||
int access_idx_contiguous = cc + c * Policy::LdsShape::kContiguous;
|
||||
int access_idx_strided =
|
||||
(ss + s * Policy::LdsShape::kStrided) * Policy::kLdsOpInner;
|
||||
|
||||
AccessType const *source_ptr =
|
||||
pointer_[access_idx_contiguous % kPointerCount] +
|
||||
Layout::TileShape::kContiguous * Layout::kElementsPerAccess *
|
||||
(access_idx_contiguous / kPointerCount) +
|
||||
access_idx_strided * stride_;
|
||||
|
||||
char const *source_byte_ptr =
|
||||
reinterpret_cast<char const *>(source_ptr) + byte_offset +
|
||||
byte_offset_;
|
||||
|
||||
fetch_ptr[access_idx] =
|
||||
*reinterpret_cast<Element const *>(source_byte_ptr);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 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 {
|
||||
load_with_byte_offset(frag, pointer_offset * sizeof(Element));
|
||||
}
|
||||
|
||||
/// 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 {
|
||||
load_with_byte_offset(frag, tile_offset, 0);
|
||||
}
|
||||
|
||||
/// 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 {
|
||||
load_with_byte_offset(frag, tile_offset, pointer_offset * sizeof(Element));
|
||||
}
|
||||
|
||||
/// 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 {
|
||||
Index pointer_offset =
|
||||
tile_offset.contiguous() * Shape::kContiguous /
|
||||
Layout::kElementsPerAccess +
|
||||
tile_offset.strided() * InstructionShape::kStrided * stride_;
|
||||
|
||||
byte_offset += sizeof(AccessType) * pointer_offset;
|
||||
|
||||
load_with_byte_offset(frag, byte_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) {
|
||||
// no op
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// This tile iterator is specialized for 32-thread TensorOps. It uses LDSM to load from shared
|
||||
/// memory and therefore must be initialized with a TensorRef to shared memory.
|
||||
///
|
||||
@@ -1069,7 +1448,6 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
k_group_idx_(0) {
|
||||
// Warp level iterator at most use double buffer to hide latency. If there
|
||||
// are more than 2 sections, every stage should have more than 1 section.
|
||||
// TODO: refactor code after every case is implemented
|
||||
|
||||
// Turing silicon requires all 32 threads in a warp provide valid addresses
|
||||
// even for LDSM.1 and LDSM.2
|
||||
@@ -1077,6 +1455,8 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
lane_id = lane_id % (Policy::LdsmShape::kCount * Policy::kLdsmOpInner);
|
||||
#endif
|
||||
|
||||
int quad_quad = (lane_id >> 4);
|
||||
int quad_pair = (lane_id >> 3);
|
||||
int lane_in_pair = (lane_id & 1);
|
||||
int lane_in_quad = (lane_id & 3);
|
||||
int lane_in_quad_pair = (lane_id & 7);
|
||||
@@ -1100,6 +1480,26 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
(lane_in_quad_quad / Layout::kFactor));
|
||||
access_strided_idx = lane_id / Layout::kFactor;
|
||||
}
|
||||
else if (Policy::LdsmShape::kStrided ==
|
||||
(Policy::LdsmShape::kCount / 2) &&
|
||||
kOperand == Operand::kA) {
|
||||
// Integer matrix multiply 16832 A
|
||||
partition_contiguous_idx = lane_in_quad / factor_in_partition;
|
||||
access_strided_idx = lane_in_quad_quad / Layout::kFactor;
|
||||
access_contiguous_idx =
|
||||
((lane_in_pair * factor_in_partition + quad_quad) ^
|
||||
access_strided_idx);
|
||||
}
|
||||
else if (Policy::LdsmShape::kStrided ==
|
||||
(Policy::LdsmShape::kCount / 2) &&
|
||||
kOperand == Operand::kB) {
|
||||
// Integer matrix multiply 16832 B
|
||||
partition_contiguous_idx = lane_in_quad / factor_in_partition;
|
||||
access_strided_idx = lane_in_quad_pair / Layout::kFactor + quad_quad * 2;
|
||||
access_contiguous_idx =
|
||||
((lane_in_pair * factor_in_partition + ((lane_id & 8) >> 3)) ^
|
||||
access_strided_idx);
|
||||
}
|
||||
} else if (Layout::kFactor == 2) {
|
||||
// Super Matrix multiply kBlock = 32
|
||||
if (Policy::LdsmShape::kStrided == Policy::LdsmShape::kCount) {
|
||||
@@ -1113,6 +1513,28 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
access_contiguous_idx = (lane_in_quad_pair / Layout::kFactor);
|
||||
access_strided_idx = lane_id / Layout::kFactor;
|
||||
}
|
||||
else if (Policy::LdsmShape::kStrided ==
|
||||
(Policy::LdsmShape::kCount / 2) &&
|
||||
kOperand == Operand::kA) {
|
||||
// Matrix multiply 16816|1688.TF32 A
|
||||
// Q0 Q2
|
||||
// Q1 Q3
|
||||
partition_contiguous_idx = (lane_id % Layout::kFactor);
|
||||
access_contiguous_idx =
|
||||
(quad_quad ^ (lane_in_quad_pair / Layout::kFactor));
|
||||
access_strided_idx = (lane_in_quad_quad / Layout::kFactor);
|
||||
} else if (Policy::LdsmShape::kStrided ==
|
||||
(Policy::LdsmShape::kCount / 2) &&
|
||||
kOperand == Operand::kB) {
|
||||
// Matrix multiply 16816|1688.TF32 B
|
||||
// Q0 Q1
|
||||
// Q2 Q3
|
||||
partition_contiguous_idx = (lane_id % Layout::kFactor);
|
||||
access_contiguous_idx =
|
||||
((quad_pair & 1) ^ (lane_in_quad_pair / Layout::kFactor));
|
||||
access_strided_idx =
|
||||
(lane_in_quad_pair + (lane_id >> 4 << 3)) / Layout::kFactor;
|
||||
}
|
||||
} else if (Layout::kFactor == 1) {
|
||||
// Super Matrix multiply kBlock = 64
|
||||
if (Policy::LdsmShape::kStrided == Policy::LdsmShape::kCount) {
|
||||
@@ -1124,6 +1546,25 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
access_contiguous_idx = lane_in_quad;
|
||||
access_strided_idx = lane_id;
|
||||
}
|
||||
else if (Policy::LdsmShape::kStrided ==
|
||||
(Policy::LdsmShape::kCount / 2) &&
|
||||
kOperand == Operand::kA) {
|
||||
// Matrix multiply 16816|1688.TF32 A
|
||||
// Q0 Q2
|
||||
// Q1 Q3
|
||||
partition_contiguous_idx = (lane_in_quad_pair >> 2);
|
||||
access_contiguous_idx = (quad_quad ^ lane_in_quad);
|
||||
access_strided_idx = lane_in_quad_quad;
|
||||
} else if (Policy::LdsmShape::kStrided ==
|
||||
(Policy::LdsmShape::kCount / 2) &&
|
||||
kOperand == Operand::kB) {
|
||||
// Matrix multiply 16816|1688.TF32 B
|
||||
// Q0 Q1
|
||||
// Q2 Q3
|
||||
partition_contiguous_idx = (lane_in_quad_pair >> 2);
|
||||
access_contiguous_idx = ((quad_pair & 1) ^ lane_in_quad);
|
||||
access_strided_idx = lane_in_quad_pair + (lane_id >> 4 << 3);
|
||||
}
|
||||
}
|
||||
|
||||
int access_contiguous =
|
||||
@@ -1161,16 +1602,68 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
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) {
|
||||
|
||||
int whole_tiles = tile_offset.contiguous() / Policy::kGroupsPerTile;
|
||||
int k_groups_delta = tile_offset.contiguous() % Policy::kGroupsPerTile;
|
||||
if (k_groups_delta < 0) {
|
||||
whole_tiles -= 1;
|
||||
k_groups_delta += Policy::kGroupsPerTile;
|
||||
}
|
||||
|
||||
if ((Policy::kGroupsPerTile / kPartitionsK) >= 2) {
|
||||
byte_offset_ ^= (k_groups_delta & 1) * Policy::LdsmShape::kContiguous *
|
||||
sizeof_bits<Element>::value *
|
||||
Layout::kElementsPerAccess / 8;
|
||||
}
|
||||
if ((Policy::kGroupsPerTile / kPartitionsK) >= 4) {
|
||||
byte_offset_ ^= ((k_groups_delta + (k_group_idx_ & 1)) & 2) *
|
||||
Policy::LdsmShape::kContiguous *
|
||||
sizeof_bits<Element>::value *
|
||||
Layout::kElementsPerAccess / 8;
|
||||
}
|
||||
if ((Policy::kGroupsPerTile / kPartitionsK) == 8) {
|
||||
byte_offset_ ^= ((k_groups_delta + (k_group_idx_ & 3)) & 4) *
|
||||
Policy::LdsmShape::kContiguous *
|
||||
sizeof_bits<Element>::value *
|
||||
Layout::kElementsPerAccess / 8;
|
||||
}
|
||||
|
||||
k_group_idx_ += k_groups_delta;
|
||||
whole_tiles += k_group_idx_ / (Policy::kGroupsPerTile / kPartitionsK);
|
||||
k_group_idx_ = k_group_idx_ % (Policy::kGroupsPerTile / kPartitionsK);
|
||||
|
||||
pointer_ +=
|
||||
tile_offset.strided() * stride_ * Shape::kStrided / Layout::kFactor +
|
||||
whole_tiles * stride_ / sections_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpMultiplicandTileIterator &operator++() {
|
||||
|
||||
// Integer matrix multiply 16832 Interleaved-32
|
||||
// NONE
|
||||
// Integer matrix multiply 16816 Interleaved-32 || Integer matrix multiply 16816 kblock=32
|
||||
|
||||
// Integer matrix multiply 8816 Interleaved-32
|
||||
// ^1 ^1
|
||||
// Matrix multiply 1684.TF32 kblock=16 || Integer matrix multiply 16816 kblock=64
|
||||
// Matrix multiply 1688 kblock=32 || Integer matrix multiply 8816 kblock=64
|
||||
// ^1 ^3 ^1 ^3
|
||||
// Matrix multiply 1688 kblock=64
|
||||
// ^1 ^3 ^1 ^7 ^1 ^3 ^1 ^7
|
||||
|
||||
// Matrix multiply 16816 kblock=32 | 1688.TF32 kblock=16 || Integer matrix multiply 16832 kblock=64
|
||||
// ^2 ^2
|
||||
// Matrix multiply 16816 kblock=64 | 1688.TF32 kblock=32 || Integer matrix multiply 16832 kblock=128
|
||||
// ^2 ^6 ^2 ^6
|
||||
|
||||
if ((Policy::kGroupsPerTile / kPartitionsK) > 1) {
|
||||
int mask = ((Policy::kGroupsPerTile / kPartitionsK) == 8)
|
||||
? 3
|
||||
@@ -1443,6 +1936,16 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
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) {
|
||||
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++() {
|
||||
@@ -1673,6 +2176,16 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
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) {
|
||||
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++() {
|
||||
@@ -1782,6 +2295,7 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
@@ -2682,6 +3196,7 @@ public:
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* 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:
|
||||
@@ -40,6 +40,8 @@
|
||||
|
||||
#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"
|
||||
|
||||
@@ -75,8 +77,6 @@ template <
|
||||
typename Policy_,
|
||||
///< Number of partitions along K dimension
|
||||
int PartitionsK_ = 1,
|
||||
///< Number of partitions along N dimension
|
||||
int PartitionsN_ = 1,
|
||||
///< Used for partial specialization
|
||||
typename Enable = bool
|
||||
>
|
||||
@@ -106,6 +106,9 @@ public:
|
||||
/// Shape of the warp in units of thread (concept: MmaTensorOpPolicy)
|
||||
using Policy = Policy_;
|
||||
|
||||
/// Underlying instruction shape
|
||||
using InstructionShape = typename Policy::Operator::Shape;
|
||||
|
||||
/// Underlying architecture tag
|
||||
using ArchTag = typename Policy::Operator::ArchTag;
|
||||
|
||||
@@ -116,7 +119,7 @@ public:
|
||||
static ComplexTransform const kTransformB = ComplexTransform::kNone;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
using OperatorClass = arch::OpClassWmmaTensorOp;
|
||||
|
||||
/// Number of threads participating in warp-level matrix product
|
||||
static int const kThreadCount = 32;
|
||||
@@ -124,9 +127,6 @@ public:
|
||||
/// Number of partitions along K dimension
|
||||
static int const kPartitionsK = PartitionsK_;
|
||||
|
||||
/// PartitionsN indicating how many PartitionsN for multiplicand B
|
||||
static int const kPartitionsN = PartitionsN_;
|
||||
|
||||
public:
|
||||
|
||||
/// Iterates over the A operand in memory
|
||||
@@ -163,9 +163,7 @@ private:
|
||||
/// Number of wmma operations performed
|
||||
using WmmaIterations = MatrixShape<
|
||||
Shape::kM / Policy::Operator::Shape::kM,
|
||||
(Shape::kN / Policy::Operator::Shape::kN / kPartitionsN > 0) ?
|
||||
Shape::kN / Policy::Operator::Shape::kN / kPartitionsN :
|
||||
1
|
||||
Shape::kN / Policy::Operator::Shape::kN
|
||||
>;
|
||||
|
||||
public:
|
||||
@@ -189,8 +187,7 @@ public:
|
||||
FragmentC &D,
|
||||
FragmentA const &A,
|
||||
FragmentB const &B,
|
||||
FragmentC const &C,
|
||||
int const &partitionN_idx = 0) const {
|
||||
FragmentC const &C) const {
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < WmmaIterations::kColumn; ++n) {
|
||||
|
||||
Reference in New Issue
Block a user