CUTLASS 2.4 (Implicit GEMM convolution) (#147)
CUTLASS 2.4 (Implicit GEMM Convolution) Co-authored-by: Manish Gupta <manigupta@nvidia.com>, Haicheng Wu <haichengw@nvidia.com>, Dustyn Blasig <dblasig@nvidia.com>, Andrew Kerr <akerr@nvidia.com>
This commit is contained in:
co-authored by
Manish Gupta <manigupta@nvidia.com>, Haicheng Wu <haichengw@nvidia.com>, Dustyn Blasig <dblasig@nvidia.com>, Andrew Kerr <akerr@nvidia.com>
parent
c2b80ad4e4
commit
6615010cd0
@@ -429,6 +429,25 @@ public:
|
||||
args.epilogue,
|
||||
static_cast<int *>(workspace)
|
||||
};
|
||||
|
||||
int smem_size = int(sizeof(typename GemmKernel::SharedStorage));
|
||||
if (smem_size >= (48 << 10)) {
|
||||
cudaError_t result = cudaFuncSetAttribute(Kernel<GemmKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
|
||||
result = cudaFuncSetAttribute(
|
||||
Kernel<GemmKernel>,
|
||||
cudaFuncAttributePreferredSharedMemoryCarveout, 100);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
@@ -461,30 +480,11 @@ public:
|
||||
dim3 grid = threadblock_swizzle.get_grid_shape(params_.grid_tiled_shape);
|
||||
dim3 block(GemmKernel::kThreadCount, 1, 1);
|
||||
|
||||
cudaError_t result;
|
||||
|
||||
int smem_size = int(sizeof(typename GemmKernel::SharedStorage));
|
||||
if (smem_size >= (48 << 10)) {
|
||||
result = cudaFuncSetAttribute(Kernel<GemmKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
|
||||
result = cudaFuncSetAttribute(
|
||||
Kernel<GemmKernel>,
|
||||
cudaFuncAttributePreferredSharedMemoryCarveout, 100);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
|
||||
cutlass::Kernel<GemmKernel><<<grid, block, smem_size, stream>>>(params_);
|
||||
|
||||
result = cudaGetLastError();
|
||||
cudaError_t result = cudaGetLastError();
|
||||
|
||||
return result == cudaSuccess ? Status::kSuccess : Status::kErrorInternal;
|
||||
}
|
||||
|
||||
@@ -117,9 +117,16 @@ public:
|
||||
using ThreadblockShape = typename GemmKernel::Mma::Shape;
|
||||
using WarpShape = typename GemmKernel::WarpShape;
|
||||
using InstructionShape = typename GemmKernel::InstructionShape;
|
||||
|
||||
using OperatorClass = typename GemmKernel::OperatorClass;
|
||||
using ArchTag = typename GemmKernel::ArchTag;
|
||||
|
||||
// warp-level, arch-level (instruction), math operator
|
||||
using WarpMmaOperator = typename GemmKernel::Mma::Policy::Operator;
|
||||
using ArchMmaOperator = typename WarpMmaOperator::ArchMmaOperator;
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
// Operator class and arch tag extract bottom-up
|
||||
// set it for top-level gemm device-level template
|
||||
using OperatorClass = typename WarpMmaOperator::OperatorClass;
|
||||
using ArchTag = typename WarpMmaOperator::ArchTag;
|
||||
|
||||
// Type, layout, and complex transform deliberately exchanged with B
|
||||
using MapArguments = detail::MapArguments<
|
||||
|
||||
@@ -311,6 +311,27 @@ public:
|
||||
gemm_k_size,
|
||||
static_cast<int *>(workspace)
|
||||
);
|
||||
|
||||
// Specify shared memory capacity for kernel.
|
||||
int smem_size = int(sizeof(typename GemmKernel::SharedStorage));
|
||||
|
||||
if (smem_size >= (48 << 10)) {
|
||||
cudaError_t result = cudaFuncSetAttribute(Kernel<GemmKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
|
||||
result = cudaFuncSetAttribute(
|
||||
Kernel<GemmKernel>,
|
||||
cudaFuncAttributePreferredSharedMemoryCarveout, 100);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
@@ -335,38 +356,31 @@ public:
|
||||
Status run(cudaStream_t stream = nullptr) {
|
||||
CUTLASS_TRACE_HOST("GemmUniversalBase::run()");
|
||||
|
||||
//
|
||||
// Configure grid and block dimensions
|
||||
//
|
||||
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
dim3 grid = threadblock_swizzle.get_grid_shape(params_.grid_tiled_shape);
|
||||
dim3 block(GemmKernel::kThreadCount, 1, 1);
|
||||
|
||||
cudaError_t result;
|
||||
|
||||
int smem_size = int(sizeof(typename GemmKernel::SharedStorage));
|
||||
if (smem_size >= (48 << 10)) {
|
||||
result = cudaFuncSetAttribute(Kernel<GemmKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
|
||||
result = cudaFuncSetAttribute(
|
||||
Kernel<GemmKernel>,
|
||||
cudaFuncAttributePreferredSharedMemoryCarveout, 100);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
//
|
||||
// Launch kernel
|
||||
//
|
||||
|
||||
CUTLASS_TRACE_HOST(" grid: (" << grid << "), block: (" << block
|
||||
<< "), SMEM: " << smem_size << " bytes");
|
||||
|
||||
// Launch
|
||||
cutlass::Kernel<GemmKernel><<<grid, block, smem_size, stream>>>(params_);
|
||||
|
||||
result = cudaGetLastError();
|
||||
//
|
||||
// Query for errors
|
||||
//
|
||||
cudaError_t result = cudaGetLastError();
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
CUTLASS_TRACE_HOST(" grid launch failed with error " << cudaGetErrorString(result));
|
||||
|
||||
@@ -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_simt.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"
|
||||
@@ -112,6 +113,101 @@ 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::OpClassSimt,
|
||||
arch::Sm50, ThreadblockShape, WarpShape, InstructionShape,
|
||||
EpilogueOutputOp, ThreadblockSwizzle, Stages, TransformA, TransformB, Operator, SplitKSerial> {
|
||||
|
||||
/// Define the threadblock-scoped matrix multiply-accumulate
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
ElementA, LayoutA,
|
||||
ElementB, LayoutB,
|
||||
ElementAccumulator, layout::RowMajor,
|
||||
arch::OpClassSimt,
|
||||
Stages,
|
||||
Operator,
|
||||
false,
|
||||
cutlass::arch::CacheOperation::Global,
|
||||
cutlass::arch::CacheOperation::Global,
|
||||
TransformA,
|
||||
TransformB
|
||||
>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using IteratorA =
|
||||
cutlass::transform::threadblock::PredicatedTileIterator<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA, LayoutA, 1,
|
||||
typename MmaCore::IteratorThreadMapA>;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using IteratorB =
|
||||
cutlass::transform::threadblock::PredicatedTileIterator<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB, LayoutB, 0,
|
||||
typename MmaCore::IteratorThreadMapB>;
|
||||
|
||||
// Define the threadblock-scoped pipelined matrix multiply
|
||||
using Mma = cutlass::gemm::threadblock::MmaPipelined<
|
||||
typename MmaCore::Shape, IteratorA, typename MmaCore::SmemIteratorA,
|
||||
IteratorB, typename MmaCore::SmemIteratorB, ElementAccumulator,
|
||||
layout::RowMajor, typename MmaCore::MmaPolicy>;
|
||||
|
||||
/// Define the epilogue
|
||||
using Epilogue =
|
||||
typename cutlass::epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
typename Mma::Operator,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
/// Define the kernel-level GEMM operator.
|
||||
using GemmKernel = kernel::Gemm<Mma, Epilogue, ThreadblockSwizzle, SplitKSerial>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for Ampere Architecture
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
@@ -170,6 +266,70 @@ struct DefaultGemmComplex<
|
||||
using GemmKernel = kernel::Gemm<Mma, Epilogue, ThreadblockSwizzle, SplitKSerial>;
|
||||
};
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// 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::OpClassSimt,
|
||||
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::OpClassSimt, arch::Sm80, ThreadblockShape,
|
||||
WarpShape, InstructionShape, Stages, TransformA, TransformB, Operator>::ThreadblockMma;
|
||||
|
||||
/// Define the epilogue
|
||||
using Epilogue =
|
||||
typename cutlass::epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
typename Mma::Operator,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
/// Define the kernel-level GEMM operator.
|
||||
using GemmKernel = kernel::Gemm<Mma, Epilogue, ThreadblockSwizzle, SplitKSerial>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
|
||||
@@ -138,8 +138,20 @@ struct Gemm {
|
||||
typename Epilogue::OutputTileIterator::TensorRef ref_C,
|
||||
typename Epilogue::OutputTileIterator::TensorRef ref_D) {
|
||||
|
||||
static int const kAlignmentA = Mma::IteratorA::AccessType::kElements;
|
||||
static int const kAlignmentB = Mma::IteratorB::AccessType::kElements;
|
||||
static int const kAlignmentA = (platform::is_same<typename Mma::IteratorA::Layout,
|
||||
layout::ColumnMajorInterleaved<32>>::value)
|
||||
? 32
|
||||
: (platform::is_same<typename Mma::IteratorA::Layout,
|
||||
layout::ColumnMajorInterleaved<64>>::value)
|
||||
? 64
|
||||
: Mma::IteratorA::AccessType::kElements;
|
||||
static int const kAlignmentB = (platform::is_same<typename Mma::IteratorB::Layout,
|
||||
layout::RowMajorInterleaved<32>>::value)
|
||||
? 32
|
||||
: (platform::is_same<typename Mma::IteratorB::Layout,
|
||||
layout::RowMajorInterleaved<64>>::value)
|
||||
? 64
|
||||
: Mma::IteratorB::AccessType::kElements;
|
||||
static int const kAlignmentC = Epilogue::OutputTileIterator::kElementsPerAccess;
|
||||
|
||||
if (!TensorRef_aligned(ref_A, kAlignmentA)) {
|
||||
@@ -274,7 +286,7 @@ struct Gemm {
|
||||
semaphore.fetch();
|
||||
|
||||
// Indicate which position in a serial reduction the output operator is currently updating
|
||||
output_op.set_k_partition(threadblock_tile_offset.k());
|
||||
output_op.set_k_partition(threadblock_tile_offset.k(), params.grid_tiled_shape.k());
|
||||
}
|
||||
|
||||
// Tile iterator loading from source tensor.
|
||||
|
||||
@@ -582,7 +582,7 @@ public:
|
||||
semaphore.fetch();
|
||||
|
||||
// Indicate which position in a serial reduction the output operator is currently updating
|
||||
output_op.set_k_partition(threadblock_tile_offset.k());
|
||||
output_op.set_k_partition(threadblock_tile_offset.k(), params.grid_tiled_shape.k());
|
||||
}
|
||||
}
|
||||
else if (params.mode == GemmUniversalMode::kGemmSplitKParallel) {
|
||||
|
||||
@@ -302,8 +302,20 @@ public:
|
||||
|
||||
CUTLASS_TRACE_HOST("GemmUniversal::can_implement()");
|
||||
|
||||
static int const kAlignmentA = Mma::IteratorA::AccessType::kElements;
|
||||
static int const kAlignmentB = Mma::IteratorB::AccessType::kElements;
|
||||
static int const kAlignmentA = (platform::is_same<typename Mma::IteratorA::Layout,
|
||||
layout::ColumnMajorInterleaved<32>>::value)
|
||||
? 32
|
||||
: (platform::is_same<typename Mma::IteratorA::Layout,
|
||||
layout::ColumnMajorInterleaved<64>>::value)
|
||||
? 64
|
||||
: Mma::IteratorA::AccessType::kElements;
|
||||
static int const kAlignmentB = (platform::is_same<typename Mma::IteratorB::Layout,
|
||||
layout::RowMajorInterleaved<32>>::value)
|
||||
? 32
|
||||
: (platform::is_same<typename Mma::IteratorB::Layout,
|
||||
layout::RowMajorInterleaved<64>>::value)
|
||||
? 64
|
||||
: Mma::IteratorB::AccessType::kElements;
|
||||
static int const kAlignmentC = Epilogue::OutputTileIterator::kElementsPerAccess;
|
||||
|
||||
if ((problem_size.m() % kAlignmentA) || (problem_size.k() % kAlignmentA) ||
|
||||
@@ -468,7 +480,7 @@ public:
|
||||
semaphore.fetch();
|
||||
|
||||
// Indicate which position in a serial reduction the output operator is currently updating
|
||||
output_op.set_k_partition(threadblock_tile_offset.k());
|
||||
output_op.set_k_partition(threadblock_tile_offset.k(), params.grid_tiled_shape.k());
|
||||
}
|
||||
}
|
||||
else if (params.mode == GemmUniversalMode::kGemmSplitKParallel) {
|
||||
|
||||
@@ -319,7 +319,7 @@ struct SparseGemm {
|
||||
semaphore.fetch();
|
||||
|
||||
// Indicate which position in a serial reduction the output operator is currently updating
|
||||
output_op.set_k_partition(threadblock_tile_offset.k());
|
||||
output_op.set_k_partition(threadblock_tile_offset.k(), params.grid_tiled_shape.k());
|
||||
}
|
||||
|
||||
// Tile iterator loading from source tensor.
|
||||
|
||||
@@ -93,6 +93,9 @@ struct Mma_HFMA2 <
|
||||
/// C operand storage
|
||||
using FragmentC = Array<half_t, Shape::kMN>;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -179,6 +182,9 @@ struct Mma_HFMA2<
|
||||
/// C operand storage
|
||||
using FragmentC = Array<half_t, Shape::kMN>;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -270,6 +276,9 @@ struct Mma_HFMA2 <
|
||||
/// C operand storage
|
||||
using FragmentC = Array<half_t, Shape::kMN>;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -356,6 +365,8 @@ struct Mma_HFMA2<
|
||||
/// C operand storage
|
||||
using FragmentC = Array<half_t, Shape::kMN>;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -443,6 +454,9 @@ struct Mma_HFMA2 <
|
||||
/// C operand storage
|
||||
using FragmentC = Array<half_t, Shape::kMN>;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -533,6 +547,9 @@ struct Mma_HFMA2 <
|
||||
/// C operand storage
|
||||
using FragmentC = Array<half_t, Shape::kMN>;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -623,6 +640,9 @@ struct Mma_HFMA2 <
|
||||
/// C operand storage
|
||||
using FragmentC = Array<half_t, Shape::kMN>;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -714,6 +734,9 @@ struct Mma_HFMA2<
|
||||
/// C operand storage
|
||||
using FragmentC = Array<half_t, Shape::kMN>;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -800,6 +823,9 @@ struct Mma_HFMA2<
|
||||
/// C operand storage
|
||||
using FragmentC = Array<half_t, Shape::kMN>;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -879,6 +905,9 @@ struct Mma_HFMA2<
|
||||
/// C operand storage
|
||||
using FragmentC = Array<half_t, Shape::kMN>;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
@@ -389,7 +389,7 @@ struct DefaultMmaCore<Shape_, WarpShape_, GemmShape<1, 1, 1>, ElementA_,
|
||||
/// Policy used to define MmaPipelined
|
||||
using MmaPolicy = MmaPolicy<
|
||||
MmaWarpSimt,
|
||||
MatrixShape<kPaddingN, 0>, // skew for A matrix to avoid SMEM bank conflicts
|
||||
MatrixShape<kPaddingM, 0>, // skew for A matrix to avoid SMEM bank conflicts
|
||||
MatrixShape<0, kPaddingN>, // skew for B matrix to avoid SMEM bank conflicts
|
||||
WarpCount::kK
|
||||
>;
|
||||
|
||||
@@ -34,6 +34,7 @@
|
||||
#include "cutlass/gemm/threadblock/default_mma_core_sm80.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/transform/threadblock/predicated_tile_iterator.h"
|
||||
#include "cutlass/gemm/threadblock/default_multistage_mma_complex_core_sm80.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
@@ -1105,6 +1105,676 @@ struct DefaultMultistageMmaComplexCore<
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for complex double-precision
|
||||
///
|
||||
/// A: column-major
|
||||
/// B: row-major
|
||||
/// Operator: arch::OpMultiplyAddComplex or arch::OpMultiplyGaussianComplex
|
||||
///
|
||||
/// 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_,
|
||||
typename RealA,
|
||||
typename RealB,
|
||||
typename RealC,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC_,
|
||||
/// 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_,
|
||||
/// Cache operation of operand A
|
||||
cutlass::arch::CacheOperation::Kind CacheOpA,
|
||||
/// Cache operation of operand B
|
||||
cutlass::arch::CacheOperation::Kind CacheOpB>
|
||||
struct DefaultMultistageMmaComplexCore<
|
||||
Shape_, WarpShape_, GemmShape<1, 1, 1>,
|
||||
complex<RealA>, layout::ColumnMajor,
|
||||
complex<RealB>, layout::ColumnMajor,
|
||||
complex<RealC>, LayoutC_,
|
||||
arch::OpClassSimt,
|
||||
Stages,
|
||||
TransformA, TransformB,
|
||||
Operator_,
|
||||
CacheOpA, CacheOpB> {
|
||||
|
||||
using Shape = Shape_;
|
||||
using WarpShape = WarpShape_;
|
||||
using InstructionShape = GemmShape<1, 1, 1>;
|
||||
using ElementA = complex<RealA>;
|
||||
using LayoutA = layout::ColumnMajor;
|
||||
using ElementB = complex<RealB>;
|
||||
using LayoutB = layout::ColumnMajor;
|
||||
using ElementC = complex<RealC>;
|
||||
using LayoutC = LayoutC_;
|
||||
static int const kStages = Stages;
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
using Operator = Operator_;
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpA = cutlass::arch::CacheOperation::Always;
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpB = cutlass::arch::CacheOperation::Always;
|
||||
|
||||
/// 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.");
|
||||
|
||||
static_assert(WarpCount::kCount > 1,
|
||||
"This specialization requires at least two warps.");
|
||||
|
||||
/// 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 access
|
||||
static int const kAccessSizeInBits = sizeof_bits<ElementA>::value;
|
||||
|
||||
/// No vectorized accesses
|
||||
static int const kElementsPerAccess = 1;
|
||||
|
||||
//
|
||||
// Shared memory layouts
|
||||
//
|
||||
|
||||
using SmemLayoutA = layout::ColumnMajor;
|
||||
|
||||
using SmemLayoutB = layout::RowMajor;
|
||||
|
||||
//
|
||||
// Iterators to write to shared memory
|
||||
//
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using IteratorThreadMapA = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kM, Shape::kK>,
|
||||
kThreads,
|
||||
kElementsPerAccess
|
||||
>;
|
||||
|
||||
/// Shared memory iterator to A operand
|
||||
using SmemIteratorA = transform::threadblock::RegularTileAccessIterator<
|
||||
MatrixShape<Shape::kM, Shape::kK>, ElementA, SmemLayoutA, 0,
|
||||
IteratorThreadMapA>;
|
||||
|
||||
/// Policy of iterator B
|
||||
using IteratorThreadMapB = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kK, Shape::kN>,
|
||||
kThreads,
|
||||
kElementsPerAccess
|
||||
>;
|
||||
|
||||
/// Transpose the ThreadMap of iterator B
|
||||
using SmemThreadMapB = transform::TransposePitchLinearThreadMapSimt<IteratorThreadMapB>;
|
||||
|
||||
/// Shared memory iterator to B operand
|
||||
using SmemIteratorB = transform::threadblock::RegularTileAccessIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>, ElementB, SmemLayoutB, 1,
|
||||
SmemThreadMapB>;
|
||||
|
||||
//
|
||||
// Warp-level matrix multiply operator
|
||||
//
|
||||
|
||||
// Define the warp-level op
|
||||
static const int WarpNumThreadsM = 4; // TODO need to extract these from template data
|
||||
static const int WarpNumThreadsN = 8;
|
||||
static_assert(!(WarpShape::kM % WarpNumThreadsM) && !(WarpShape::kN % WarpNumThreadsN),
|
||||
"WarpShape must be divisible by ThreadTile shape.");
|
||||
static const int ThreadTileM = WarpShape::kM / WarpNumThreadsM;
|
||||
static const int ThreadTileN = WarpShape::kN / WarpNumThreadsN;
|
||||
static const int LaneLayout = ThreadTileM > 4 && ThreadTileN > 4 ? 2 : 1;
|
||||
static const int numElementsA = 128 / sizeof_bits<ElementA>::value;
|
||||
static const int numElementsB = 128 / sizeof_bits<ElementB>::value;
|
||||
static const int LaneM = cutlass::const_min(numElementsA, ThreadTileM);
|
||||
static const int LaneN = cutlass::const_min(numElementsB, ThreadTileN);
|
||||
// these should have max of thread tile also
|
||||
using LaneMmaShape = cutlass::gemm::GemmShape<
|
||||
LaneM,
|
||||
LaneN,
|
||||
1>;
|
||||
using Policy = cutlass::gemm::warp::MmaSimtPolicy<
|
||||
cutlass::MatrixShape<WarpNumThreadsM, WarpNumThreadsN>, // WarpShape
|
||||
cutlass::layout::RowMajorInterleaved<LaneLayout>, // LaneLayout
|
||||
LaneMmaShape
|
||||
>;
|
||||
|
||||
using MmaWarpSimt = cutlass::gemm::warp::MmaSimt<
|
||||
WarpShape, /// Size of the Gemm problem - concept: gemm::GemmShape<> 128, 128, 8
|
||||
ElementA, /// Data type of A elements
|
||||
SmemLayoutA, /// Layout of A matrix (concept: MatrixLayout)
|
||||
ElementB, /// Data type of B elements
|
||||
SmemLayoutB, /// Layout of B matrix (concept: MatrixLayout)
|
||||
ElementC, /// Element type of C matrix
|
||||
LayoutC, /// Layout of C matrix (concept: MatrixLayout)
|
||||
Policy /// Policy describing warp-level MmaTensorOp (concept: MmaTensorOp policy)
|
||||
>; /// Used for partial specialization
|
||||
|
||||
/// Policy used to define MmaPipelined
|
||||
using MmaPolicy = MmaPolicy<
|
||||
MmaWarpSimt,
|
||||
MatrixShape<0, 0>,
|
||||
MatrixShape<0, Shape::kK / 32>,
|
||||
WarpCount::kK>;
|
||||
};
|
||||
|
||||
/// Partial specialization for complex double-precision
|
||||
///
|
||||
/// A: column-major
|
||||
/// B: row-major
|
||||
/// Operator: arch::OpMultiplyAddComplex or arch::OpMultiplyGaussianComplex
|
||||
///
|
||||
/// 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_,
|
||||
typename RealA,
|
||||
typename RealB,
|
||||
typename RealC,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC_,
|
||||
/// 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_,
|
||||
/// Cache operation of operand A
|
||||
cutlass::arch::CacheOperation::Kind CacheOpA,
|
||||
/// Cache operation of operand B
|
||||
cutlass::arch::CacheOperation::Kind CacheOpB>
|
||||
struct DefaultMultistageMmaComplexCore<
|
||||
Shape_, WarpShape_, GemmShape<1, 1, 1>,
|
||||
complex<RealA>, layout::ColumnMajor,
|
||||
complex<RealB>, layout::RowMajor,
|
||||
complex<RealC>, LayoutC_,
|
||||
arch::OpClassSimt,
|
||||
Stages,
|
||||
TransformA, TransformB,
|
||||
Operator_,
|
||||
CacheOpA, CacheOpB> {
|
||||
|
||||
using Shape = Shape_;
|
||||
using WarpShape = WarpShape_;
|
||||
using InstructionShape = GemmShape<1, 1, 1>;
|
||||
using ElementA = complex<RealA>;
|
||||
using LayoutA = layout::ColumnMajor;
|
||||
using ElementB = complex<RealB>;
|
||||
using LayoutB = layout::RowMajor;
|
||||
using ElementC = complex<RealC>;
|
||||
using LayoutC = LayoutC_;
|
||||
static int const kStages = Stages;
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
using Operator = Operator_;
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpA = cutlass::arch::CacheOperation::Always;
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpB = cutlass::arch::CacheOperation::Always;
|
||||
|
||||
/// 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.");
|
||||
|
||||
static_assert(WarpCount::kCount > 1,
|
||||
"This specialization requires at least two warps.");
|
||||
|
||||
/// 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 access
|
||||
static int const kAccessSizeInBits = sizeof_bits<ElementA>::value;
|
||||
|
||||
/// No vectorized accesses
|
||||
static int const kElementsPerAccess = 1;
|
||||
|
||||
//
|
||||
// Shared memory layouts
|
||||
//
|
||||
|
||||
using SmemLayoutA = layout::ColumnMajor;
|
||||
|
||||
using SmemLayoutB = layout::RowMajor;
|
||||
|
||||
//
|
||||
// Iterators to write to shared memory
|
||||
//
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using IteratorThreadMapA = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kM, Shape::kK>,
|
||||
kThreads,
|
||||
kElementsPerAccess
|
||||
>;
|
||||
|
||||
/// Shared memory iterator to A operand
|
||||
using SmemIteratorA = transform::threadblock::RegularTileAccessIterator<
|
||||
MatrixShape<Shape::kM, Shape::kK>, ElementA, SmemLayoutA, 0,
|
||||
IteratorThreadMapA>;
|
||||
|
||||
/// Policy of iterator B
|
||||
using IteratorThreadMapB = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kN, Shape::kK>,
|
||||
kThreads,
|
||||
kElementsPerAccess
|
||||
>;
|
||||
|
||||
/// Shared memory iterator to B operand
|
||||
using SmemIteratorB = transform::threadblock::RegularTileAccessIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>, ElementB, SmemLayoutB, 1,
|
||||
IteratorThreadMapB>;
|
||||
|
||||
//
|
||||
// Warp-level matrix multiply operator
|
||||
//
|
||||
|
||||
// Define the warp-level op
|
||||
static const int WarpNumThreadsM = 4; // TODO need to extract these from template data
|
||||
static const int WarpNumThreadsN = 8;
|
||||
static_assert(!(WarpShape::kM % WarpNumThreadsM) && !(WarpShape::kN % WarpNumThreadsN),
|
||||
"WarpShape must be divisible by ThreadTile shape.");
|
||||
static const int ThreadTileM = WarpShape::kM / WarpNumThreadsM;
|
||||
static const int ThreadTileN = WarpShape::kN / WarpNumThreadsN;
|
||||
static const int LaneLayout = ThreadTileM > 4 && ThreadTileN > 4 ? 2 : 1;
|
||||
static const int numElementsA = 128 / sizeof_bits<ElementA>::value;
|
||||
static const int numElementsB = 128 / sizeof_bits<ElementB>::value;
|
||||
static const int LaneM = cutlass::const_min(numElementsA, ThreadTileM);
|
||||
static const int LaneN = cutlass::const_min(numElementsB, ThreadTileN);
|
||||
// these should have max of thread tile also
|
||||
using LaneMmaShape = cutlass::gemm::GemmShape<
|
||||
LaneM,
|
||||
LaneN,
|
||||
1>;
|
||||
using Policy = cutlass::gemm::warp::MmaSimtPolicy<
|
||||
cutlass::MatrixShape<WarpNumThreadsM, WarpNumThreadsN>, // WarpShape
|
||||
cutlass::layout::RowMajorInterleaved<LaneLayout>, // LaneLayout
|
||||
LaneMmaShape
|
||||
>;
|
||||
|
||||
using MmaWarpSimt = cutlass::gemm::warp::MmaSimt<
|
||||
WarpShape, /// Size of the Gemm problem - concept: gemm::GemmShape<> 128, 128, 8
|
||||
ElementA, /// Data type of A elements
|
||||
SmemLayoutA, /// Layout of A matrix (concept: MatrixLayout)
|
||||
ElementB, /// Data type of B elements
|
||||
SmemLayoutB, /// Layout of B matrix (concept: MatrixLayout)
|
||||
ElementC, /// Element type of C matrix
|
||||
LayoutC, /// Layout of C matrix (concept: MatrixLayout)
|
||||
Policy /// Policy describing warp-level MmaTensorOp (concept: MmaTensorOp policy)
|
||||
>; /// Used for partial specialization
|
||||
|
||||
/// Policy used to define MmaPipelined
|
||||
using MmaPolicy = MmaPolicy<
|
||||
MmaWarpSimt,
|
||||
MatrixShape<0, 0>,
|
||||
MatrixShape<0, 0>, // or Shape::kK / 32
|
||||
WarpCount::kK>;
|
||||
};
|
||||
|
||||
/// Partial specialization for complex double-precision
|
||||
///
|
||||
/// A: column-major
|
||||
/// B: row-major
|
||||
/// Operator: arch::OpMultiplyAddComplex or arch::OpMultiplyGaussianComplex
|
||||
///
|
||||
/// 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_,
|
||||
typename RealA,
|
||||
typename RealB,
|
||||
typename RealC,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC_,
|
||||
/// 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_,
|
||||
/// Cache operation of operand A
|
||||
cutlass::arch::CacheOperation::Kind CacheOpA,
|
||||
/// Cache operation of operand B
|
||||
cutlass::arch::CacheOperation::Kind CacheOpB>
|
||||
struct DefaultMultistageMmaComplexCore<
|
||||
Shape_, WarpShape_, GemmShape<1, 1, 1>,
|
||||
complex<RealA>, layout::RowMajor,
|
||||
complex<RealB>, layout::ColumnMajor,
|
||||
complex<RealC>, LayoutC_,
|
||||
arch::OpClassSimt,
|
||||
Stages,
|
||||
TransformA, TransformB,
|
||||
Operator_,
|
||||
CacheOpA, CacheOpB> {
|
||||
|
||||
using Shape = Shape_;
|
||||
using WarpShape = WarpShape_;
|
||||
using InstructionShape = GemmShape<1, 1, 1>;
|
||||
using ElementA = complex<RealA>;
|
||||
using LayoutA = layout::RowMajor;
|
||||
using ElementB = complex<RealB>;
|
||||
using LayoutB = layout::ColumnMajor;
|
||||
using ElementC = complex<RealC>;
|
||||
using LayoutC = LayoutC_;
|
||||
static int const kStages = Stages;
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
using Operator = Operator_;
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpA = cutlass::arch::CacheOperation::Always;
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpB = cutlass::arch::CacheOperation::Always;
|
||||
|
||||
/// 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.");
|
||||
|
||||
static_assert(WarpCount::kCount > 1,
|
||||
"This specialization requires at least two warps.");
|
||||
|
||||
/// 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 access
|
||||
static int const kAccessSizeInBits = sizeof_bits<ElementA>::value;
|
||||
|
||||
/// No vectorized accesses
|
||||
static int const kElementsPerAccess = 1;
|
||||
|
||||
//
|
||||
// Shared memory layouts
|
||||
//
|
||||
|
||||
using SmemLayoutA = layout::ColumnMajor;
|
||||
|
||||
using SmemLayoutB = layout::RowMajor;
|
||||
|
||||
//
|
||||
// Iterators to write to shared memory
|
||||
//
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using IteratorThreadMapA = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kK, Shape::kM>,
|
||||
kThreads,
|
||||
kElementsPerAccess
|
||||
>;
|
||||
|
||||
/// Transpose the ThreadMap of iterator A
|
||||
using SmemThreadMapA = transform::TransposePitchLinearThreadMapSimt<IteratorThreadMapA>;
|
||||
|
||||
/// Shared memory iterator to A operand
|
||||
using SmemIteratorA = transform::threadblock::RegularTileAccessIterator<
|
||||
MatrixShape<Shape::kM, Shape::kK>, ElementA, SmemLayoutA, 0,
|
||||
SmemThreadMapA>;
|
||||
|
||||
/// Policy of iterator B
|
||||
using IteratorThreadMapB = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kK, Shape::kN>,
|
||||
kThreads,
|
||||
kElementsPerAccess
|
||||
>;
|
||||
|
||||
/// Transpose the ThreadMap of iterator B
|
||||
using SmemThreadMapB = transform::TransposePitchLinearThreadMapSimt<IteratorThreadMapB>;
|
||||
|
||||
/// Shared memory iterator to B operand
|
||||
using SmemIteratorB = transform::threadblock::RegularTileAccessIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>, ElementB, SmemLayoutB, 1,
|
||||
SmemThreadMapB>;
|
||||
|
||||
//
|
||||
// Warp-level matrix multiply operator
|
||||
//
|
||||
|
||||
// Define the warp-level op
|
||||
static const int WarpNumThreadsM = 4; // TODO need to extract these from template data
|
||||
static const int WarpNumThreadsN = 8;
|
||||
static_assert(!(WarpShape::kM % WarpNumThreadsM) && !(WarpShape::kN % WarpNumThreadsN),
|
||||
"WarpShape must be divisible by ThreadTile shape.");
|
||||
static const int ThreadTileM = WarpShape::kM / WarpNumThreadsM;
|
||||
static const int ThreadTileN = WarpShape::kN / WarpNumThreadsN;
|
||||
static const int LaneLayout = ThreadTileM > 4 && ThreadTileN > 4 ? 2 : 1;
|
||||
static const int numElementsA = 128 / sizeof_bits<ElementA>::value;
|
||||
static const int numElementsB = 128 / sizeof_bits<ElementB>::value;
|
||||
static const int LaneM = cutlass::const_min(numElementsA, ThreadTileM);
|
||||
static const int LaneN = cutlass::const_min(numElementsB, ThreadTileN);
|
||||
// these should have max of thread tile also
|
||||
using LaneMmaShape = cutlass::gemm::GemmShape<
|
||||
LaneM,
|
||||
LaneN,
|
||||
1>;
|
||||
using Policy = cutlass::gemm::warp::MmaSimtPolicy<
|
||||
cutlass::MatrixShape<WarpNumThreadsM, WarpNumThreadsN>, // WarpShape
|
||||
cutlass::layout::RowMajorInterleaved<LaneLayout>, // LaneLayout
|
||||
LaneMmaShape
|
||||
>;
|
||||
|
||||
using MmaWarpSimt = cutlass::gemm::warp::MmaSimt<
|
||||
WarpShape, /// Size of the Gemm problem - concept: gemm::GemmShape<> 128, 128, 8
|
||||
ElementA, /// Data type of A elements
|
||||
SmemLayoutA, /// Layout of A matrix (concept: MatrixLayout)
|
||||
ElementB, /// Data type of B elements
|
||||
SmemLayoutB, /// Layout of B matrix (concept: MatrixLayout)
|
||||
ElementC, /// Element type of C matrix
|
||||
LayoutC, /// Layout of C matrix (concept: MatrixLayout)
|
||||
Policy /// Policy describing warp-level MmaTensorOp (concept: MmaTensorOp policy)
|
||||
>; /// Used for partial specialization
|
||||
|
||||
/// Policy used to define MmaPipelined
|
||||
using MmaPolicy = MmaPolicy<
|
||||
MmaWarpSimt,
|
||||
MatrixShape<Shape::kK / 32, 0>,
|
||||
MatrixShape<0, Shape::kK / 32>,
|
||||
WarpCount::kK>;
|
||||
};
|
||||
|
||||
/// Partial specialization for complex double-precision
|
||||
///
|
||||
/// A: column-major
|
||||
/// B: row-major
|
||||
/// Operator: arch::OpMultiplyAddComplex or arch::OpMultiplyGaussianComplex
|
||||
///
|
||||
/// 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_,
|
||||
typename RealA,
|
||||
typename RealB,
|
||||
typename RealC,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC_,
|
||||
/// 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_,
|
||||
/// Cache operation of operand A
|
||||
cutlass::arch::CacheOperation::Kind CacheOpA,
|
||||
/// Cache operation of operand B
|
||||
cutlass::arch::CacheOperation::Kind CacheOpB>
|
||||
struct DefaultMultistageMmaComplexCore<
|
||||
Shape_, WarpShape_, GemmShape<1, 1, 1>,
|
||||
complex<RealA>, layout::RowMajor,
|
||||
complex<RealB>, layout::RowMajor,
|
||||
complex<RealC>, LayoutC_,
|
||||
arch::OpClassSimt,
|
||||
Stages,
|
||||
TransformA, TransformB,
|
||||
Operator_,
|
||||
CacheOpA, CacheOpB> {
|
||||
|
||||
using Shape = Shape_;
|
||||
using WarpShape = WarpShape_;
|
||||
using InstructionShape = GemmShape<1, 1, 1>;
|
||||
using ElementA = complex<RealA>;
|
||||
using LayoutA = layout::RowMajor;
|
||||
using ElementB = complex<RealB>;
|
||||
using LayoutB = layout::RowMajor;
|
||||
using ElementC = complex<RealC>;
|
||||
using LayoutC = LayoutC_;
|
||||
static int const kStages = Stages;
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
using Operator = Operator_;
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpA = cutlass::arch::CacheOperation::Always;
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpB = cutlass::arch::CacheOperation::Always;
|
||||
|
||||
/// 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.");
|
||||
|
||||
static_assert(WarpCount::kCount > 1,
|
||||
"This specialization requires at least two warps.");
|
||||
|
||||
/// 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 access
|
||||
static int const kAccessSizeInBits = sizeof_bits<ElementA>::value;
|
||||
|
||||
/// No vectorized accesses
|
||||
static int const kElementsPerAccess = 1;
|
||||
|
||||
//
|
||||
// Shared memory layouts
|
||||
//
|
||||
|
||||
using SmemLayoutA = layout::ColumnMajor;
|
||||
|
||||
using SmemLayoutB = layout::RowMajor;
|
||||
|
||||
//
|
||||
// Iterators to write to shared memory
|
||||
//
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using IteratorThreadMapA = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kK, Shape::kM>,
|
||||
kThreads,
|
||||
kElementsPerAccess
|
||||
>;
|
||||
|
||||
/// Transpose the ThreadMap of iterator A
|
||||
using SmemThreadMapA = transform::TransposePitchLinearThreadMapSimt<IteratorThreadMapA>;
|
||||
|
||||
/// Shared memory iterator to A operand
|
||||
using SmemIteratorA = transform::threadblock::RegularTileAccessIterator<
|
||||
MatrixShape<Shape::kM, Shape::kK>, ElementA, SmemLayoutA, 0,
|
||||
SmemThreadMapA>;
|
||||
|
||||
/// Policy of iterator B
|
||||
using IteratorThreadMapB = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kN, Shape::kK>,
|
||||
kThreads,
|
||||
kElementsPerAccess
|
||||
>;
|
||||
|
||||
/// Shared memory iterator to B operand
|
||||
using SmemIteratorB = transform::threadblock::RegularTileAccessIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>, ElementB, SmemLayoutB, 1,
|
||||
IteratorThreadMapB>;
|
||||
|
||||
//
|
||||
// Warp-level matrix multiply operator
|
||||
//
|
||||
|
||||
// Define the warp-level op
|
||||
static const int WarpNumThreadsM = 4; // TODO need to extract these from template data
|
||||
static const int WarpNumThreadsN = 8;
|
||||
static_assert(!(WarpShape::kM % WarpNumThreadsM) && !(WarpShape::kN % WarpNumThreadsN),
|
||||
"WarpShape must be divisible by ThreadTile shape.");
|
||||
static const int ThreadTileM = WarpShape::kM / WarpNumThreadsM;
|
||||
static const int ThreadTileN = WarpShape::kN / WarpNumThreadsN;
|
||||
static const int LaneLayout = ThreadTileM > 4 && ThreadTileN > 4 ? 2 : 1;
|
||||
static const int numElementsA = 128 / sizeof_bits<ElementA>::value;
|
||||
static const int numElementsB = 128 / sizeof_bits<ElementB>::value;
|
||||
static const int LaneM = cutlass::const_min(numElementsA, ThreadTileM);
|
||||
static const int LaneN = cutlass::const_min(numElementsB, ThreadTileN);
|
||||
// these should have max of thread tile also
|
||||
using LaneMmaShape = cutlass::gemm::GemmShape<
|
||||
LaneM,
|
||||
LaneN,
|
||||
1>;
|
||||
using Policy = cutlass::gemm::warp::MmaSimtPolicy<
|
||||
cutlass::MatrixShape<WarpNumThreadsM, WarpNumThreadsN>, // WarpShape
|
||||
cutlass::layout::RowMajorInterleaved<LaneLayout>, // LaneLayout
|
||||
LaneMmaShape
|
||||
>;
|
||||
|
||||
using MmaWarpSimt = cutlass::gemm::warp::MmaSimt<
|
||||
WarpShape, /// Size of the Gemm problem - concept: gemm::GemmShape<> 128, 128, 8
|
||||
ElementA, /// Data type of A elements
|
||||
SmemLayoutA, /// Layout of A matrix (concept: MatrixLayout)
|
||||
ElementB, /// Data type of B elements
|
||||
SmemLayoutB, /// Layout of B matrix (concept: MatrixLayout)
|
||||
ElementC, /// Element type of C matrix
|
||||
LayoutC, /// Layout of C matrix (concept: MatrixLayout)
|
||||
Policy /// Policy describing warp-level MmaTensorOp (concept: MmaTensorOp policy)
|
||||
>; /// Used for partial specialization
|
||||
|
||||
/// Policy used to define MmaPipelined
|
||||
using MmaPolicy = MmaPolicy<
|
||||
MmaWarpSimt,
|
||||
MatrixShape<Shape::kK / 32, 0>,
|
||||
MatrixShape<0, 0>, // or Shape::kK / 32
|
||||
WarpCount::kK>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace gemm
|
||||
|
||||
@@ -228,7 +228,7 @@ public:
|
||||
for (int v = 0; v < IteratorA::kAccessesPerVector; ++v) {
|
||||
auto gmem_ptr = iterator_A.get();
|
||||
|
||||
cutlass::arch::cp_async<kSrcBytes, kCacheOpA>(
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpA>(
|
||||
dst_ptr + v, gmem_ptr, iterator_A.valid());
|
||||
|
||||
++iterator_A;
|
||||
@@ -258,7 +258,7 @@ public:
|
||||
for (int v = 0; v < IteratorB::kAccessesPerVector; ++v) {
|
||||
auto gmem_ptr = iterator_B.get();
|
||||
|
||||
cutlass::arch::cp_async<kSrcBytes, kCacheOpB>(
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpB>(
|
||||
dst_ptr + v, gmem_ptr, iterator_B.valid());
|
||||
|
||||
++iterator_B;
|
||||
@@ -513,6 +513,11 @@ public:
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
// commit and drain all pending and predicated LDGSTS pnz from the GEMM mainloop
|
||||
cutlass::arch::cp_async_fence();
|
||||
cutlass::arch::cp_async_wait<0>();
|
||||
__syncthreads();
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
@@ -105,6 +105,14 @@ public:
|
||||
/// Warp-level Mma
|
||||
using Operator = typename Policy::Operator;
|
||||
|
||||
using ArchTag = arch::Sm70;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = Operator::kTransformA;
|
||||
|
||||
/// Complex transform on B operand
|
||||
static ComplexTransform const kTransformB = Operator::kTransformB;
|
||||
|
||||
// staticaly assert kStages for MmaSingleStage is 1 (single stage mma pipeline)
|
||||
static_assert((Base::kStages==1), "MmaSingleStage requires kStages set to value 1");
|
||||
private:
|
||||
|
||||
@@ -314,8 +314,17 @@ public:
|
||||
/// Shape of the warp in units of thread (concept: MmaLanePolicyTensorOp)
|
||||
using Policy = Policy_;
|
||||
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename Policy::Operator;
|
||||
|
||||
/// Architecture tag from underlying instruction
|
||||
using ArchTag = typename ArchMmaOperator::ArchTag;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename Policy::Operator::Shape;
|
||||
using InstructionShape = typename ArchMmaOperator::Shape;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
@@ -323,9 +332,6 @@ public:
|
||||
/// 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;
|
||||
|
||||
@@ -337,7 +343,7 @@ public:
|
||||
Operand::kA,
|
||||
ElementA,
|
||||
LayoutA,
|
||||
MatrixShape<Policy::Operator::Shape::kM, Policy::Operator::Shape::kK>,
|
||||
MatrixShape<ArchMmaOperator::Shape::kM, ArchMmaOperator::Shape::kK>,
|
||||
Policy::OpDelta::kRow,
|
||||
32,
|
||||
1
|
||||
@@ -355,7 +361,7 @@ public:
|
||||
Operand::kB,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
MatrixShape<Policy::Operator::Shape::kK, Policy::Operator::Shape::kN>,
|
||||
MatrixShape<ArchMmaOperator::Shape::kK, ArchMmaOperator::Shape::kN>,
|
||||
Policy::OpDelta::kColumn,
|
||||
32,
|
||||
1
|
||||
@@ -368,14 +374,14 @@ public:
|
||||
using TransformedFragmentB = FragmentB;
|
||||
|
||||
static_assert(
|
||||
!(Shape::kM % Policy::Operator::Shape::kM) &&
|
||||
!(Shape::kN % Policy::Operator::Shape::kN),
|
||||
!(Shape::kM % ArchMmaOperator::Shape::kM) &&
|
||||
!(Shape::kN % ArchMmaOperator::Shape::kN),
|
||||
"Shape of warp-level Mma must be divisible by operator shape.");
|
||||
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = MatrixShape<
|
||||
Shape::kM / Policy::Operator::Shape::kM,
|
||||
Shape::kN / Policy::Operator::Shape::kN
|
||||
Shape::kM / ArchMmaOperator::Shape::kM,
|
||||
Shape::kN / ArchMmaOperator::Shape::kN
|
||||
>;
|
||||
|
||||
/// Iterates over the C operand in memory
|
||||
@@ -383,7 +389,7 @@ public:
|
||||
MatrixShape<Shape::kM, Shape::kN>,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
typename Policy::Operator::Shape,
|
||||
typename ArchMmaOperator::Shape,
|
||||
typename Policy::OpDelta>;
|
||||
|
||||
/// Storage for C tile, the accumulator. Note, regardless of multiplicand type, this
|
||||
@@ -393,7 +399,7 @@ public:
|
||||
using FragmentC = typename IteratorC::Fragment;
|
||||
|
||||
static_assert(
|
||||
FragmentC::kElements == 2 * MmaIterations::kCount * Policy::Operator::FragmentC::kElements,
|
||||
FragmentC::kElements == 2 * MmaIterations::kCount * ArchMmaOperator::FragmentC::kElements,
|
||||
"Unexpected planar complex fragment length.");
|
||||
|
||||
private:
|
||||
@@ -403,7 +409,7 @@ private:
|
||||
//
|
||||
|
||||
/// Underlying real-valued matrix multiply operator (concept: arch::Mma)
|
||||
typename Policy::Operator mma;
|
||||
ArchMmaOperator mma;
|
||||
|
||||
public:
|
||||
|
||||
@@ -425,9 +431,9 @@ public:
|
||||
) 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;
|
||||
using MmaOperandA = typename ArchMmaOperator::FragmentA;
|
||||
using MmaOperandB = typename ArchMmaOperator::FragmentB;
|
||||
using MmaOperandC = typename ArchMmaOperator::FragmentC;
|
||||
|
||||
static_assert(MmaOperandA::kElements == 1,
|
||||
"This implementation only supports math instructions in which exactly one element is needed for the A operand."
|
||||
@@ -599,12 +605,18 @@ public:
|
||||
|
||||
/// Shape of the warp in units of thread (concept: MmaLanePolicySimt)
|
||||
using Policy = Policy_;
|
||||
|
||||
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename Policy::Operator;
|
||||
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename Policy::Operator::Shape;
|
||||
using InstructionShape = typename ArchMmaOperator::Shape;
|
||||
|
||||
/// Underlying arch tag
|
||||
using ArchTag = typename Policy::Operator::ArchTag;
|
||||
using ArchTag = typename ArchMmaOperator::ArchTag;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
@@ -612,9 +624,6 @@ public:
|
||||
/// 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;
|
||||
|
||||
@@ -626,7 +635,7 @@ public:
|
||||
Operand::kA,
|
||||
ElementA,
|
||||
LayoutA,
|
||||
MatrixShape<Policy::Operator::Shape::kM, Policy::Operator::Shape::kK>,
|
||||
MatrixShape<ArchMmaOperator::Shape::kM, ArchMmaOperator::Shape::kK>,
|
||||
Policy::OpDelta::kRow,
|
||||
32,
|
||||
1
|
||||
@@ -637,7 +646,7 @@ public:
|
||||
|
||||
/// Storage for transformed A tile
|
||||
using TransformedFragmentA =
|
||||
Array<typename Policy::Operator::ElementA, FragmentA::kElements * 2>;
|
||||
Array<typename ArchMmaOperator::ElementA, FragmentA::kElements * 2>;
|
||||
|
||||
/// Iterates over the B operand in memory
|
||||
using IteratorB = MmaTensorOpMultiplicandTileIterator<
|
||||
@@ -645,7 +654,7 @@ public:
|
||||
Operand::kB,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
MatrixShape<Policy::Operator::Shape::kK, Policy::Operator::Shape::kN>,
|
||||
MatrixShape<ArchMmaOperator::Shape::kK, ArchMmaOperator::Shape::kN>,
|
||||
Policy::OpDelta::kColumn,
|
||||
32,
|
||||
1
|
||||
@@ -656,17 +665,17 @@ public:
|
||||
|
||||
/// Storage for transformed B tile
|
||||
using TransformedFragmentB =
|
||||
Array<typename Policy::Operator::ElementB, FragmentB::kElements * 2>;
|
||||
Array<typename ArchMmaOperator::ElementB, FragmentB::kElements * 2>;
|
||||
|
||||
static_assert(
|
||||
!(Shape::kM % Policy::Operator::Shape::kM) &&
|
||||
!(Shape::kN % Policy::Operator::Shape::kN),
|
||||
!(Shape::kM % ArchMmaOperator::Shape::kM) &&
|
||||
!(Shape::kN % ArchMmaOperator::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
|
||||
Shape::kM / ArchMmaOperator::Shape::kM,
|
||||
Shape::kN / ArchMmaOperator::Shape::kN
|
||||
>;
|
||||
|
||||
/// Iterates over the C operand in memory
|
||||
@@ -674,7 +683,7 @@ public:
|
||||
MatrixShape<Shape::kM, Shape::kN>,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
typename Policy::Operator::Shape,
|
||||
typename ArchMmaOperator::Shape,
|
||||
typename Policy::OpDelta>;
|
||||
|
||||
/// Storage for C tile, the accumulator. Note, regardless of multiplicand type, this
|
||||
@@ -690,7 +699,7 @@ private:
|
||||
//
|
||||
|
||||
/// Underlying real-valued matrix multiply operator (concept: arch::Mma)
|
||||
typename Policy::Operator mma;
|
||||
ArchMmaOperator mma;
|
||||
|
||||
public:
|
||||
|
||||
@@ -712,11 +721,11 @@ public:
|
||||
) 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;
|
||||
using InstMmaOperandA = typename ArchMmaOperator::FragmentA;
|
||||
using InstMmaOperandB = typename ArchMmaOperator::FragmentB;
|
||||
using MmaOperandC = typename ArchMmaOperator::FragmentC;
|
||||
|
||||
static_assert(platform::is_same<cutlass::gemm::GemmShape<16, 8, 8>, typename Policy::Operator::Shape>::value,
|
||||
static_assert(platform::is_same<cutlass::gemm::GemmShape<16, 8, 8>, typename ArchMmaOperator::Shape>::value,
|
||||
"This implementation only supports MMA.1688 math instructions.");
|
||||
|
||||
static_assert(InstMmaOperandA::kElements == 4,
|
||||
@@ -794,8 +803,8 @@ public:
|
||||
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;
|
||||
using InstMmaOperandA = typename ArchMmaOperator::FragmentA;
|
||||
using InstMmaOperandB = typename ArchMmaOperator::FragmentB;
|
||||
|
||||
//
|
||||
// Define conversions from source type to instruction operands' type
|
||||
|
||||
@@ -147,11 +147,17 @@ public:
|
||||
/// Shape of the warp in units of thread (concept: MmaLanePolicySimt)
|
||||
using Policy = Policy_;
|
||||
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename Policy::Operator::Shape;
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename Policy::Operator;
|
||||
|
||||
/// Underlying architecture tag
|
||||
using ArchTag = typename Policy::Operator::ArchTag;
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename ArchMmaOperator::Shape;
|
||||
|
||||
/// Underlying arch tag
|
||||
using ArchTag = typename ArchMmaOperator::ArchTag;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
@@ -159,8 +165,6 @@ public:
|
||||
/// 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;
|
||||
@@ -173,7 +177,7 @@ public:
|
||||
Operand::kA,
|
||||
ElementA,
|
||||
LayoutA,
|
||||
MatrixShape<Policy::Operator::Shape::kM, Policy::Operator::Shape::kK>,
|
||||
MatrixShape<ArchMmaOperator::Shape::kM, ArchMmaOperator::Shape::kK>,
|
||||
Policy::OpDelta::kRow,
|
||||
32,
|
||||
1
|
||||
@@ -191,7 +195,7 @@ public:
|
||||
Operand::kB,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
MatrixShape<Policy::Operator::Shape::kK, Policy::Operator::Shape::kN>,
|
||||
MatrixShape<ArchMmaOperator::Shape::kK, ArchMmaOperator::Shape::kN>,
|
||||
Policy::OpDelta::kColumn,
|
||||
32,
|
||||
1
|
||||
@@ -204,14 +208,14 @@ public:
|
||||
using TransformedFragmentB = FragmentB;
|
||||
|
||||
static_assert(
|
||||
!(Shape::kM % Policy::Operator::Shape::kM) &&
|
||||
!(Shape::kN % Policy::Operator::Shape::kN),
|
||||
!(Shape::kM % ArchMmaOperator::Shape::kM) &&
|
||||
!(Shape::kN % ArchMmaOperator::Shape::kN),
|
||||
"Shape of warp-level Mma must be divisible by operator shape.");
|
||||
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = MatrixShape<
|
||||
Shape::kM / Policy::Operator::Shape::kM,
|
||||
Shape::kN / Policy::Operator::Shape::kN
|
||||
Shape::kM / ArchMmaOperator::Shape::kM,
|
||||
Shape::kN / ArchMmaOperator::Shape::kN
|
||||
>;
|
||||
|
||||
/// Iterates over the C operand in memory
|
||||
@@ -219,7 +223,7 @@ public:
|
||||
MatrixShape<Shape::kM, Shape::kN>,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
typename Policy::Operator::Shape,
|
||||
typename ArchMmaOperator::Shape,
|
||||
typename Policy::OpDelta>;
|
||||
|
||||
/// Storage for C tile, the accumulator. Note, regardless of multiplicand type, this
|
||||
@@ -229,7 +233,7 @@ public:
|
||||
using FragmentC = typename IteratorC::Fragment;
|
||||
|
||||
static_assert(
|
||||
FragmentC::kElements == 3 * MmaIterations::kCount * Policy::Operator::FragmentC::kElements,
|
||||
FragmentC::kElements == 3 * MmaIterations::kCount * ArchMmaOperator::FragmentC::kElements,
|
||||
"Unexpected gaussian complex fragment length.");
|
||||
|
||||
private:
|
||||
@@ -239,7 +243,7 @@ private:
|
||||
//
|
||||
|
||||
/// Underlying real-valued matrix multiply operator (concept: arch::Mma)
|
||||
typename Policy::Operator mma;
|
||||
ArchMmaOperator mma;
|
||||
|
||||
public:
|
||||
|
||||
@@ -261,9 +265,9 @@ public:
|
||||
) 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;
|
||||
using MmaOperandA = typename ArchMmaOperator::FragmentA;
|
||||
using MmaOperandB = typename ArchMmaOperator::FragmentB;
|
||||
using MmaOperandC = typename ArchMmaOperator::FragmentC;
|
||||
|
||||
static_assert(MmaOperandA::kElements == 1,
|
||||
"This implementation only supports math instructions in which exactly one element is needed for the A operand."
|
||||
@@ -346,8 +350,6 @@ public:
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// TODO - partial specializations of real*complex and complex*real
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
|
||||
@@ -68,6 +68,10 @@ template <
|
||||
typename Policy_,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK = 1,
|
||||
/// Complex transformation on operand A
|
||||
ComplexTransform TransformA = ComplexTransform::kNone,
|
||||
/// Complex transformation on operand B
|
||||
ComplexTransform TransformB = ComplexTransform::kNone,
|
||||
/// Used for partial specialization
|
||||
typename Enable = bool
|
||||
>
|
||||
@@ -104,10 +108,10 @@ public:
|
||||
using ArchTag = arch::Sm50;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = ComplexTransform::kNone;
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
|
||||
/// Complex transform on B operand
|
||||
static ComplexTransform const kTransformB = ComplexTransform::kNone;
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
|
||||
/// Layout of threads
|
||||
using ThreadLayoutA = typename platform::conditional< platform::is_same< layout::ColumnMajorInterleaved<4>, LayoutA >::value,
|
||||
@@ -215,12 +219,22 @@ public:
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
FragmentC &d,
|
||||
FragmentA const &a,
|
||||
FragmentB const &b,
|
||||
FragmentA a,
|
||||
FragmentB b,
|
||||
FragmentC const &c, int group_idx = 0) const {
|
||||
|
||||
ThreadMma mma;
|
||||
|
||||
if (kTransformA == ComplexTransform::kConjugate) {
|
||||
conjugate<FragmentA> conj_a;
|
||||
a = conj_a(a);
|
||||
}
|
||||
|
||||
if (kTransformB == ComplexTransform::kConjugate) {
|
||||
conjugate<FragmentB> conj_b;
|
||||
b = conj_b(b);
|
||||
}
|
||||
|
||||
mma(d, a, b, c);
|
||||
}
|
||||
|
||||
|
||||
@@ -111,17 +111,28 @@ public:
|
||||
/// Shape of the warp in units of thread (concept: MmaLanePolicySimt)
|
||||
using Policy = Policy_;
|
||||
|
||||
/// Equivalant base dense mma
|
||||
using Base = MmaTensorOp<Shape, ElementA, LayoutA, ElementB, LayoutB,
|
||||
ElementC, LayoutC, Policy, PartitionsK_,
|
||||
AccumulatorsInRowMajor, Enable>;
|
||||
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename Base::ArchMmaOperator;
|
||||
|
||||
/// Architecture tag from underlying instruction
|
||||
using ArchTag = typename Policy::Operator::ArchTag;
|
||||
using ArchTag = typename Base::ArchTag;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassTensorOp;
|
||||
using OperatorClass = typename Base::OperatorClass;
|
||||
|
||||
/// Shape of underlying instruction
|
||||
using InstructionShape = typename Base::InstructionShape;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = ComplexTransform::kNone;
|
||||
static ComplexTransform const kTransformA = Base::kTransformA;
|
||||
|
||||
/// Complex transform on B operand
|
||||
static ComplexTransform const kTransformB = ComplexTransform::kNone;
|
||||
static ComplexTransform const kTransformB = Base::kTransformB;
|
||||
|
||||
/// Number of threads participating in warp-level matrix product
|
||||
static int const kThreadCount = 32;
|
||||
@@ -171,25 +182,19 @@ public:
|
||||
Array<typename Policy::Operator::ElementA, FragmentA::kElements>;
|
||||
|
||||
/// Iterates over the B operand in memory
|
||||
using IteratorB = MmaTensorOpMultiplicandTileIterator<
|
||||
MatrixShape<Shape::kK, Shape::kN>, Operand::kB, ElementB, LayoutB,
|
||||
MatrixShape<Policy::Operator::Shape::kK, Policy::Operator::Shape::kN>,
|
||||
Policy::OpDelta::kRow, kThreadCount, kPartitionsK>;
|
||||
using IteratorB = typename Base::IteratorB;
|
||||
|
||||
/// Storage for B tile
|
||||
using FragmentB = typename IteratorB::Fragment;
|
||||
using FragmentB = typename Base::FragmentB;
|
||||
|
||||
/// Storage for transformed B tile
|
||||
using TransformedFragmentB =
|
||||
Array<typename Policy::Operator::ElementB, FragmentB::kElements>;
|
||||
using TransformedFragmentB = typename Base::TransformedFragmentB;
|
||||
|
||||
/// Iterates over the C operand in memory
|
||||
using IteratorC = MmaTensorOpAccumulatorTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kN>, ElementC, LayoutC,
|
||||
typename Policy::Operator::Shape, typename Policy::OpDelta>;
|
||||
using IteratorC = typename Base::IteratorC;
|
||||
|
||||
/// Storage for C tile
|
||||
using FragmentC = typename IteratorC::Fragment;
|
||||
using FragmentC = typename Base::FragmentC;
|
||||
|
||||
/// Iterates over the E operand in memory
|
||||
using IteratorE = SparseMmaTensorOpMetaTileIterator<
|
||||
@@ -204,23 +209,13 @@ public:
|
||||
/// Storage for E tile
|
||||
using FragmentE = typename IteratorE::Fragment;
|
||||
|
||||
private:
|
||||
|
||||
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
|
||||
>;
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = typename Base::MmaIterations;
|
||||
|
||||
public:
|
||||
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
typename Policy::Operator mma;
|
||||
ArchMmaOperator mma;
|
||||
|
||||
public:
|
||||
|
||||
@@ -299,21 +294,21 @@ public:
|
||||
// Define conversions from source type to instruction type
|
||||
//
|
||||
FloatRoundStyle const kRoundA =
|
||||
PreferredRoundingMode<typename Policy::Operator::ElementA,
|
||||
PreferredRoundingMode<typename ArchMmaOperator::ElementA,
|
||||
ElementA>::kRound;
|
||||
FloatRoundStyle const kRoundB =
|
||||
PreferredRoundingMode<typename Policy::Operator::ElementB,
|
||||
PreferredRoundingMode<typename ArchMmaOperator::ElementB,
|
||||
ElementB>::kRound;
|
||||
detail::ConvertAndPack<typename Policy::Operator::ElementA, ElementA,
|
||||
detail::ConvertAndPack<typename ArchMmaOperator::ElementA, ElementA,
|
||||
FragmentA::kElements / 2, kRoundA>
|
||||
convert_A;
|
||||
NumericArrayConverter<typename Policy::Operator::ElementB, ElementB,
|
||||
NumericArrayConverter<typename ArchMmaOperator::ElementB, ElementB,
|
||||
FragmentB::kElements, kRoundB>
|
||||
convert_B;
|
||||
Array<ElementA, FragmentA::kElements / 2> const *ptr_A =
|
||||
reinterpret_cast<Array<ElementA, FragmentA::kElements / 2> const *>(&A);
|
||||
Array<typename Policy::Operator::ElementA, FragmentA::kElements / 2> *
|
||||
ptr_dst_A = reinterpret_cast<Array<typename Policy::Operator::ElementA,
|
||||
Array<typename ArchMmaOperator::ElementA, FragmentA::kElements / 2> *
|
||||
ptr_dst_A = reinterpret_cast<Array<typename ArchMmaOperator::ElementA,
|
||||
FragmentA::kElements / 2> *>(&dst_A);
|
||||
|
||||
dst_B = convert_B(B);
|
||||
|
||||
@@ -244,8 +244,6 @@ public:
|
||||
/// Storage for C tile
|
||||
using FragmentC = typename IteratorC::Fragment;
|
||||
|
||||
private:
|
||||
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = MatrixShape<
|
||||
(Shape::kM + ArchMmaOperator::Shape::kM - 1) / ArchMmaOperator::Shape::kM,
|
||||
|
||||
@@ -1518,6 +1518,7 @@ class MmaTensorOpMultiplicandTileIterator<
|
||||
} else if (Layout::kFactor == 2) {
|
||||
// Super Matrix multiply kBlock = 32
|
||||
if (Policy::LdsmShape::kStrided == Policy::LdsmShape::kCount) {
|
||||
// Matrix multiply 1688 A/B
|
||||
// (Q stands for 1 8x128bit block).
|
||||
// Q0
|
||||
// Q1
|
||||
@@ -3191,10 +3192,430 @@ public:
|
||||
|
||||
int idx = mma_m + mma_n * Policy::MmaIterations::kRow;
|
||||
|
||||
AccessType* access_ptr = reinterpret_cast<AccessType *>(offset_ref.data() +
|
||||
offset_ref.offset(TensorCoord(accum_m, accum_n)));
|
||||
AccessType* access_ptr = reinterpret_cast<AccessType *>(offset_ref.data() +
|
||||
offset_ref.offset(TensorCoord(accum_m, accum_n)));
|
||||
|
||||
access_ptr[0] = frag_ptr[idx];
|
||||
access_ptr[0] = frag_ptr[idx];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Stores a fragment to memory with additional pointer offset
|
||||
CUTLASS_DEVICE
|
||||
void store_with_byte_offset(
|
||||
Fragment const &frag, ///< fragment to store from the tensor
|
||||
Index byte_offset) const { ///< store a tile with a linear offset
|
||||
|
||||
store_with_pointer_offset(byte_offset / sizeof(Element));
|
||||
}
|
||||
|
||||
/// Stores a fragment to memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void store(
|
||||
Fragment &frag, ///< fragment to store to the tensor
|
||||
TensorCoord const &tile_offset) const { ///< stores a tile with a logical offset in units of whole tiles
|
||||
|
||||
store(frag, tile_offset, 0);
|
||||
}
|
||||
|
||||
/// Stores a fragment from memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void store(
|
||||
/// fragment to store to the tensor
|
||||
Fragment const &frag,
|
||||
/// stores a tile with a logical offset in units of whole tiles
|
||||
TensorCoord const &tile_offset,
|
||||
/// stores a tile with a logical offset AND a pointer offset
|
||||
Index pointer_offset) const {
|
||||
store_with_pointer_offset(frag, ref_.offset(tile_offset) + pointer_offset);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// This tile iterator is specialized for 32-thread TensorOps. It is used to load or store
|
||||
/// accumulators from memory and is agnostic to layout. It could be faster if it assumed row-major
|
||||
/// accumulator layout.
|
||||
///
|
||||
/// Satisfies:
|
||||
/// ReadableRandomAccessContiguousTileIteratorConcept |
|
||||
/// WriteableRandomAccessContiguousTileIteratorConcept
|
||||
///
|
||||
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Element typ
|
||||
typename Element_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Interval between adjacent *MMA instructions (in units of MMA
|
||||
/// instructions, concept: MatrixShape)
|
||||
typename OpDelta_,
|
||||
/// Interleaved N
|
||||
int InterleavedN>
|
||||
class MmaTensorOpAccumulatorTileIterator<
|
||||
Shape_, Element_, cutlass::layout::TensorNCxHWx<InterleavedN>,
|
||||
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 = int8_t;
|
||||
|
||||
/// Layout of source tile
|
||||
using Layout = cutlass::layout::TensorNCxHWx<InterleavedN>;
|
||||
|
||||
/// 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.");
|
||||
|
||||
/// Number of elements in strided dimension that each STG writes
|
||||
static int const kStridedPerSTG = 8;
|
||||
|
||||
/// Factor to calculate reorder index to pack accumulator.
|
||||
static int const kPackedFactor = Shape::kColumn / 32;
|
||||
|
||||
/// Number of mma operations performed
|
||||
using MmaIterations = MatrixShape<Shape::kRow / kStridedPerSTG,
|
||||
Shape::kColumn / InterleavedN>;
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
static int const kElementsPerAccess = InterleavedN / 4;
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Derived quantities
|
||||
//
|
||||
|
||||
struct alignas((kElementsPerAccess * sizeof_bits<Element>::value / 8)) AccessType {
|
||||
Array<Element, kElementsPerAccess> storage;
|
||||
};
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<int32_t, Shape::kCount / kThreads>;
|
||||
|
||||
private:
|
||||
|
||||
/// Reference to output tensor
|
||||
TensorRef ref_;
|
||||
|
||||
/// Row offset index globally
|
||||
LongIndex global_offset_row_;
|
||||
|
||||
/// Column offset index globally
|
||||
LongIndex global_offset_col_;
|
||||
|
||||
/// Output tensor size
|
||||
TensorCoord extent_;
|
||||
|
||||
/// Alpha
|
||||
float alpha_;
|
||||
|
||||
/// Beta
|
||||
float beta_;
|
||||
|
||||
public:
|
||||
|
||||
/// Default ctor constructs null iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator() { }
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator(
|
||||
TensorRef const &ref,
|
||||
int const lane_id,
|
||||
TensorCoord extent,
|
||||
float alpha = 1.0f,
|
||||
float beta = 0.0f
|
||||
):
|
||||
ref_(ref),
|
||||
extent_(extent),
|
||||
alpha_(alpha),
|
||||
beta_(beta) {
|
||||
|
||||
int quad = (lane_id >> 2);
|
||||
int lane_in_quad = (lane_id & 3);
|
||||
|
||||
global_offset_row_ = quad;
|
||||
|
||||
global_offset_col_ = lane_in_quad * kElementsPerAccess;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset to internal pointer(s) to advance through memory
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator &add_pointer_offset(LongIndex offset) {
|
||||
ref_.add_pointer_offset(offset);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator &add_tile_offset(MatrixCoord const &tile_offset) {
|
||||
|
||||
global_offset_row_ += tile_offset.row() * Shape::kRow;
|
||||
|
||||
global_offset_col_ += tile_offset.column() * Shape::kColumn;
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator & operator++() {
|
||||
// deliberate no-op
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator & operator--() {
|
||||
// deliberate no-op
|
||||
return *this;
|
||||
}
|
||||
|
||||
///< advances in units of whole tiles along the logical coordinate space of the tensor
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator & operator+=(TensorCoord const &tile_offset) {
|
||||
add_tile_offset(tile_offset);
|
||||
return *this;
|
||||
}
|
||||
|
||||
///< advances in units of whole tiles along the logical coordinate space of the tensor
|
||||
CUTLASS_DEVICE
|
||||
MmaTensorOpAccumulatorTileIterator & operator-=(TensorCoord const &tile_offset) {
|
||||
add_tile_offset(-tile_offset);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory at the location pointed to by the iterator.
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) const {
|
||||
load_with_pointer_offset(frag);
|
||||
}
|
||||
|
||||
/// 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);
|
||||
|
||||
AccessType* frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_n = 0; mma_n < Policy::MmaIterations::kN; ++mma_n) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_m = 0; mma_m < Policy::MmaIterations::kM; ++mma_m) {
|
||||
int accum_m = mma_m * InstructionShape::kM;
|
||||
int accum_n = mma_n * InstructionShape::kN;
|
||||
|
||||
int idx = mma_m + mma_n * Policy::MmaIterations::kM;
|
||||
|
||||
AccessType* access_ptr = reinterpret_cast<AccessType *>(offset_ref.data() +
|
||||
accum_m * offset_ref.stride(0) + accum_n);
|
||||
|
||||
frag_ptr[idx] = access_ptr[0];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 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);
|
||||
|
||||
Array<float, Shape::kCount / kThreads> output_frag_f;
|
||||
Array<Element, Shape::kCount / kThreads> output_frag;
|
||||
|
||||
LongIndex pq = extent_.h() * extent_.w();
|
||||
|
||||
LongIndex extent_row = extent_.n() * pq;
|
||||
LongIndex extent_col = extent_.c();
|
||||
|
||||
LongIndex k_major = (global_offset_col_ / InterleavedN) * pq;
|
||||
Index k_minor = global_offset_col_ % InterleavedN;
|
||||
LongIndex k_offset = k_major * InterleavedN + k_minor;
|
||||
LongIndex k_offset_delta = pq * InterleavedN;
|
||||
|
||||
LongIndex stride_n = pq * extent_.c();
|
||||
|
||||
Index n;
|
||||
LongIndex pq_rem;
|
||||
|
||||
unsigned int pq_mul, pq_shr;
|
||||
find_divisor(pq_mul, pq_shr, pq);
|
||||
|
||||
if(beta_ == 0.0f) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int i = 0; i < frag.size(); ++i) {
|
||||
output_frag_f[i] = frag[i];
|
||||
}
|
||||
|
||||
if(InstructionShape::kM == Policy::kStridedPerSTG) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int i = 0; i < frag.size(); ++i) {
|
||||
output_frag[i] = (Element)(output_frag_f[i] * alpha_);
|
||||
}
|
||||
} else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int i = 0; i < frag.size(); ++i) {
|
||||
int map_i = (i / (16 * Policy::kPackedFactor)) * (16 * Policy::kPackedFactor)
|
||||
+ (i % (8 * Policy::kPackedFactor)) / 2 * 4
|
||||
+ (i % (8 * Policy::kPackedFactor)) % 2
|
||||
+ (i / (8 * Policy::kPackedFactor)) % 2 * 2;
|
||||
output_frag[i] = (Element)(output_frag_f[map_i] * alpha_);
|
||||
}
|
||||
}
|
||||
|
||||
AccessType const *frag_ptr = reinterpret_cast<AccessType const*>(&output_frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_m = 0; mma_m < Policy::MmaIterations::kRow; ++mma_m) {
|
||||
int accum_m = mma_m * Policy::kStridedPerSTG;
|
||||
|
||||
fast_divmod(n, pq_rem, global_offset_row_ + accum_m, pq, pq_mul, pq_shr);
|
||||
LongIndex offset_m = n * stride_n + k_offset + pq_rem * InterleavedN;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_n = 0; mma_n < Policy::MmaIterations::kColumn; ++mma_n) {
|
||||
|
||||
int accum_n = mma_n * InterleavedN;
|
||||
|
||||
int idx = mma_n + mma_m * Policy::MmaIterations::kColumn;
|
||||
|
||||
if((global_offset_row_ + accum_m < extent_row) && (global_offset_col_ + accum_n < extent_col)) {
|
||||
AccessType* access_ptr = reinterpret_cast<AccessType *>(offset_ref.data() +
|
||||
offset_m + mma_n * k_offset_delta);
|
||||
|
||||
access_ptr[0] = frag_ptr[idx];
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
if(InstructionShape::kM == Policy::kStridedPerSTG) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int i = 0; i < frag.size(); ++i) {
|
||||
output_frag_f[i] = frag[i];
|
||||
}
|
||||
} else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int i = 0; i < frag.size(); ++i) {
|
||||
int map_i = (i / (16 * Policy::kPackedFactor)) * (16 * Policy::kPackedFactor)
|
||||
+ (i % (8 * Policy::kPackedFactor)) / 2 * 4
|
||||
+ (i % (8 * Policy::kPackedFactor)) % 2
|
||||
+ (i / (8 * Policy::kPackedFactor)) % 2 * 2;
|
||||
output_frag_f[i] = frag[map_i];
|
||||
}
|
||||
}
|
||||
|
||||
AccessType const *frag_ptr = reinterpret_cast<AccessType const*>(&output_frag);
|
||||
|
||||
Array<Element, kElementsPerAccess> ref_frag;
|
||||
AccessType *ref_frag_ptr = reinterpret_cast<AccessType *>(&ref_frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_m = 0; mma_m < Policy::MmaIterations::kRow; ++mma_m) {
|
||||
int accum_m = mma_m * Policy::kStridedPerSTG;
|
||||
|
||||
fast_divmod(n, pq_rem, global_offset_row_ + accum_m, pq, pq_mul, pq_shr);
|
||||
LongIndex offset_m = n * stride_n + k_offset + pq_rem * InterleavedN;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_n = 0; mma_n < Policy::MmaIterations::kColumn; ++mma_n) {
|
||||
|
||||
int accum_n = mma_n * InterleavedN;
|
||||
|
||||
int idx = mma_n + mma_m * Policy::MmaIterations::kColumn;
|
||||
|
||||
if((global_offset_row_ + accum_m < extent_row) && (global_offset_col_ + accum_n < extent_col)) {
|
||||
AccessType* access_ptr = reinterpret_cast<AccessType *>(offset_ref.data() +
|
||||
offset_m + mma_n * k_offset_delta);
|
||||
|
||||
ref_frag_ptr[0] = access_ptr[0];
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int i = 0; i < kElementsPerAccess; ++i) {
|
||||
output_frag[idx * kElementsPerAccess + i] = Element(alpha_ * output_frag_f[idx * kElementsPerAccess + i]
|
||||
+ beta_ * ref_frag[i]);
|
||||
}
|
||||
|
||||
access_ptr[0] = frag_ptr[idx];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2243,6 +2243,847 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tile iterator specialized for 'TN' arrangement
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Operand identity
|
||||
Operand Operand_,
|
||||
/// Data type of A elements
|
||||
typename Element_,
|
||||
/// Layout of matrix operand
|
||||
typename Layout_,
|
||||
/// Shape of one matrix production operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Delta between *MMA operations (in units of *MMA operations, concept:
|
||||
/// MatrixShape)
|
||||
int OpDelta_,
|
||||
/// Number of threads participating in one matrix operation
|
||||
int Threads = 32,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK_ = 1>
|
||||
class MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner {
|
||||
public:
|
||||
|
||||
/// Shape of tile to load (concept: MatrixShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Operand tag
|
||||
static Operand const kOperand = Operand_;
|
||||
|
||||
/// Basic check
|
||||
static_assert(kOperand == Operand::kA || kOperand== Operand::kB,
|
||||
"MmaVoltaTensorOpMultiplicandIterator may only be instantiated for A or B operands to warp-level Mma.");
|
||||
|
||||
/// Element type
|
||||
using Element = Element_;
|
||||
|
||||
/// Layout of source tile
|
||||
using Layout = Layout_;
|
||||
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
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;
|
||||
|
||||
/// 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;
|
||||
|
||||
/// Number of elements accessed per Shared Memory load
|
||||
static int const kElementsPerAccess = 4;
|
||||
|
||||
private:
|
||||
|
||||
static int const kInterleavedTileRows = 32;
|
||||
static int const kInterleavedTileColumns = 32;
|
||||
static int const kInstructionsPerTile = 2;
|
||||
|
||||
/// Rounded up instruction counts
|
||||
using TileCount = MatrixShape<
|
||||
Shape::kRow / kInterleavedTileRows,
|
||||
Shape::kColumn / kInterleavedTileColumns
|
||||
>;
|
||||
|
||||
using FragmentCount = MatrixShape<
|
||||
TileCount::kRow * kInstructionsPerTile,
|
||||
TileCount::kColumn * kInstructionsPerTile
|
||||
>;
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Derived quantities
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<
|
||||
Element,
|
||||
(kOperand == Operand::kA ? FragmentCount::kRow : FragmentCount::kColumn) * kElementsPerAccess
|
||||
>;
|
||||
|
||||
/// Memory access type
|
||||
using AccessType = AlignedArray<Element, kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
|
||||
/// Underlying tensor reference
|
||||
TensorRef ref_;
|
||||
|
||||
/// Extent of tensor
|
||||
MatrixCoord extent_;
|
||||
|
||||
/// Origin
|
||||
MatrixCoord origin_;
|
||||
|
||||
/// Used to conditionally enable extents checking
|
||||
bool divisible_;
|
||||
|
||||
public:
|
||||
|
||||
/// Default ctor constructs null iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner(): divisible_(true) { }
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner(
|
||||
TensorRef const &ref,
|
||||
int lane_id
|
||||
):
|
||||
ref_(ref), extent_(Shape::kRow, Shape::kColumn), divisible_(true) {
|
||||
|
||||
int quad_id = lane_id / 4;
|
||||
int lane_in_quad = (lane_id % 4);
|
||||
|
||||
if (kOperand == Operand::kA) {
|
||||
|
||||
int row_idx = ((quad_id & 1) + ((quad_id & 4) / 2)) * 4 * kInstructionsPerTile + lane_in_quad;
|
||||
int col_idx = 0;
|
||||
|
||||
origin_ = MatrixCoord(row_idx, col_idx);
|
||||
}
|
||||
else {
|
||||
|
||||
int row_idx = 0;
|
||||
int col_idx = (quad_id / 2) * 4 * kInstructionsPerTile + lane_in_quad;
|
||||
|
||||
origin_ = MatrixCoord(row_idx, col_idx);
|
||||
}
|
||||
|
||||
ref_.add_coord_offset(origin_);
|
||||
}
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner(
|
||||
TensorRef const &ref,
|
||||
TensorCoord extent,
|
||||
int lane_id
|
||||
): ref_(ref), extent_(extent), divisible_(false) {
|
||||
|
||||
int quad_id = lane_id / 4;
|
||||
int lane_in_quad = (lane_id % 4);
|
||||
|
||||
if (kOperand == Operand::kA) {
|
||||
|
||||
int row_idx = ((quad_id & 1) + ((quad_id & 4) / 2)) * 4 * kInstructionsPerTile + lane_in_quad;
|
||||
int col_idx = 0;
|
||||
|
||||
origin_ = MatrixCoord(row_idx, col_idx);
|
||||
}
|
||||
else {
|
||||
|
||||
int row_idx = 0;
|
||||
int col_idx = (quad_id / 2) * 4 * kInstructionsPerTile + lane_in_quad;
|
||||
|
||||
origin_ = MatrixCoord(row_idx, col_idx);
|
||||
}
|
||||
|
||||
#if defined(__CUDA_ARCH__)
|
||||
__syncthreads();
|
||||
#endif
|
||||
|
||||
ref_.add_coord_offset(origin_);
|
||||
}
|
||||
|
||||
/// Adds a pointer offset to internal pointer(s) to advance through memory
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner &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
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner &add_tile_offset(TensorCoord const &tile_offset) {
|
||||
|
||||
TensorCoord coord_offset(tile_offset.row() * Shape::kRow, tile_offset.column() * Shape::kColumn);
|
||||
origin_ += coord_offset;
|
||||
|
||||
ref_.add_coord_offset(coord_offset);
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner & operator++() {
|
||||
|
||||
if (kOperand == Operand::kA) {
|
||||
add_tile_offset({0, 1});
|
||||
}
|
||||
else {
|
||||
add_tile_offset({1, 0});
|
||||
}
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner & operator--() {
|
||||
|
||||
if (kOperand == Operand::kA) {
|
||||
add_tile_offset({0, -1});
|
||||
}
|
||||
else {
|
||||
add_tile_offset({-1, 0});
|
||||
}
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
///< advances in units of whole tiles along the logical coordinate space of the tensor
|
||||
CUTLASS_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner & 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
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner & 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 to load from the tensor
|
||||
Fragment &frag,
|
||||
/// loads a tile with a linear offset
|
||||
Index pointer_offset) const {
|
||||
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
AccessType const *access_ptr = reinterpret_cast<AccessType const *>(ref_.data());
|
||||
int ldm = ref_.stride()[0];
|
||||
|
||||
if (kOperand == Operand::kA) {
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int idx = 0; idx < FragmentCount::kRow; ++idx) {
|
||||
|
||||
int tile_idx = idx / 2;
|
||||
int quad_idx = idx % 2;
|
||||
|
||||
int row_offset = tile_idx * kInterleavedTileRows + quad_idx * 4;
|
||||
frag_ptr[idx] = access_ptr[row_offset * ldm / kElementsPerAccess];
|
||||
}
|
||||
}
|
||||
else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int idx = 0; idx < FragmentCount::kColumn; ++idx) {
|
||||
|
||||
int tile_idx = idx / 2;
|
||||
int quad_idx = idx % 2;
|
||||
|
||||
int col_offset = tile_idx * kInterleavedTileColumns + quad_idx * 4;
|
||||
frag_ptr[idx] = access_ptr[col_offset * ldm / kElementsPerAccess];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 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
|
||||
Index byte_offset) const {
|
||||
|
||||
load_with_pointer_offset(frag, byte_offset * 8 / sizeof_bits<Element>::value);
|
||||
}
|
||||
|
||||
/// 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 {
|
||||
|
||||
TensorCoord coord_offset(tile_offset.row() * Shape::kRow, tile_offset.column() * Shape::kColumn);
|
||||
|
||||
load_with_pointer_offset(frag, ref_.offset(coord_offset));
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void load(
|
||||
/// fragment to load from the tensor
|
||||
Fragment &frag,
|
||||
/// loads a tile with a logical offset in units of whole tiles
|
||||
TensorCoord const &tile_offset,
|
||||
/// loads a tile with a logical offset AND a pointer offset
|
||||
Index pointer_offset) const {
|
||||
|
||||
TensorCoord coord_offset(tile_offset.row() * Shape::kRow, tile_offset.column() * Shape::kColumn);
|
||||
|
||||
load_with_pointer_offset(frag, ref_.offset(coord_offset) + pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void load_with_byte_offset(
|
||||
/// fragment to load from the tensor
|
||||
Fragment &frag,
|
||||
/// loads a tile with a logical offset in units of whole tiles
|
||||
TensorCoord const &tile_offset,
|
||||
/// loads a tile with a logical offset AND a pointer offset
|
||||
Index byte_offset) const {
|
||||
|
||||
TensorCoord coord_offset(tile_offset.row() * Shape::kRow, tile_offset.column() * Shape::kColumn);
|
||||
|
||||
load_with_pointer_offset(frag, ref_.offset(coord_offset) + byte_offset * 8 / sizeof_bits<Element>::value);
|
||||
}
|
||||
|
||||
/// 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 operation
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
/// Tile iterator specialized for 'NT' arrangement
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Operand identity
|
||||
Operand Operand_,
|
||||
/// Data type of A elements
|
||||
typename Element_,
|
||||
/// Layout of matrix operand
|
||||
typename Layout_,
|
||||
/// Shape of one matrix production operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Delta between *MMA operations (in units of *MMA operations, concept:
|
||||
/// MatrixShape)
|
||||
int OpDelta_,
|
||||
/// Number of threads participating in one matrix operation
|
||||
int Threads = 32,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK_ = 1>
|
||||
class MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter {
|
||||
public:
|
||||
|
||||
/// Shape of tile to load (concept: MatrixShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Operand tag
|
||||
static Operand const kOperand = Operand_;
|
||||
|
||||
/// Basic check
|
||||
static_assert(kOperand == Operand::kA || kOperand== Operand::kB,
|
||||
"MmaVoltaTensorOpMultiplicandIterator may only be instantiated for A or B operands to warp-level Mma.");
|
||||
|
||||
/// Element type
|
||||
using Element = Element_;
|
||||
|
||||
/// Layout of source tile
|
||||
using Layout = Layout_;
|
||||
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
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;
|
||||
|
||||
/// 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;
|
||||
|
||||
/// Number of elements accessed per Shared Memory load
|
||||
static int const kElementsPerAccess = 4;
|
||||
|
||||
private:
|
||||
|
||||
static int const kInterleavedTileRows = 32;
|
||||
static int const kInterleavedTileColumns = 32;
|
||||
static int const kInstructionsPerTile = 2;
|
||||
|
||||
/// Rounded up instruction counts
|
||||
using TileCount = MatrixShape<
|
||||
Shape::kRow / kInterleavedTileRows,
|
||||
Shape::kColumn / kInterleavedTileColumns
|
||||
>;
|
||||
|
||||
using FragmentCount = MatrixShape<
|
||||
TileCount::kRow * kInstructionsPerTile,
|
||||
TileCount::kColumn * kInstructionsPerTile
|
||||
>;
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Derived quantities
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<
|
||||
Element,
|
||||
(kOperand == Operand::kA ? FragmentCount::kRow : FragmentCount::kColumn) * kElementsPerAccess
|
||||
>;
|
||||
|
||||
/// Memory access type
|
||||
using AccessType = AlignedArray<Element, kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
|
||||
/// Underlying tensor reference
|
||||
TensorRef ref_;
|
||||
|
||||
/// Extent of tensor
|
||||
MatrixCoord extent_;
|
||||
|
||||
/// Origin
|
||||
MatrixCoord origin_;
|
||||
|
||||
/// Used to conditionally enable extents checking
|
||||
bool divisible_;
|
||||
|
||||
public:
|
||||
|
||||
/// Default ctor constructs null iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter(): divisible_(true) { }
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter(
|
||||
TensorRef const &ref,
|
||||
int lane_id
|
||||
):
|
||||
ref_(ref), extent_(Shape::kRow, Shape::kColumn), divisible_(true) {
|
||||
|
||||
int quad_id = lane_id / 4;
|
||||
int lane_in_quad = (lane_id % 4);
|
||||
|
||||
if (kOperand == Operand::kA) {
|
||||
|
||||
int row_idx = ((quad_id & 1) + ((quad_id & 4) / 2)) * 4 * kInstructionsPerTile;
|
||||
int col_idx = lane_in_quad;
|
||||
|
||||
origin_ = MatrixCoord(row_idx, col_idx);
|
||||
}
|
||||
else {
|
||||
|
||||
int row_idx = lane_in_quad;
|
||||
int col_idx = (quad_id / 2) * 4 * kInstructionsPerTile;
|
||||
|
||||
origin_ = MatrixCoord(row_idx, col_idx);
|
||||
}
|
||||
|
||||
ref_.add_coord_offset(origin_);
|
||||
}
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter(
|
||||
TensorRef const &ref,
|
||||
TensorCoord extent,
|
||||
int lane_id
|
||||
): ref_(ref), extent_(extent), divisible_(false) {
|
||||
|
||||
int quad_id = lane_id / 4;
|
||||
int lane_in_quad = (lane_id % 4);
|
||||
|
||||
if (kOperand == Operand::kA) {
|
||||
|
||||
int row_idx = ((quad_id & 1) + ((quad_id & 4) / 2)) * 4 * kInstructionsPerTile;
|
||||
int col_idx = lane_in_quad;
|
||||
|
||||
origin_ = MatrixCoord(row_idx, col_idx);
|
||||
}
|
||||
else {
|
||||
|
||||
int row_idx = lane_in_quad;
|
||||
int col_idx = (quad_id / 2) * 4 * kInstructionsPerTile;
|
||||
|
||||
origin_ = MatrixCoord(row_idx, col_idx);
|
||||
}
|
||||
|
||||
#if defined(__CUDA_ARCH__)
|
||||
__syncthreads();
|
||||
#endif
|
||||
|
||||
ref_.add_coord_offset(origin_);
|
||||
}
|
||||
|
||||
/// Adds a pointer offset to internal pointer(s) to advance through memory
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter &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
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter &add_tile_offset(TensorCoord const &tile_offset) {
|
||||
|
||||
TensorCoord coord_offset(tile_offset.row() * Shape::kRow, tile_offset.column() * Shape::kColumn);
|
||||
origin_ += coord_offset;
|
||||
|
||||
ref_.add_coord_offset(coord_offset);
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter & operator++() {
|
||||
|
||||
if (kOperand == Operand::kA) {
|
||||
add_tile_offset({0, 1});
|
||||
}
|
||||
else {
|
||||
add_tile_offset({1, 0});
|
||||
}
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter & operator--() {
|
||||
|
||||
if (kOperand == Operand::kA) {
|
||||
add_tile_offset({0, -1});
|
||||
}
|
||||
else {
|
||||
add_tile_offset({-1, 0});
|
||||
}
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
///< advances in units of whole tiles along the logical coordinate space of the tensor
|
||||
CUTLASS_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter & 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
|
||||
MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter & 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 to load from the tensor
|
||||
Fragment &frag,
|
||||
/// loads a tile with a linear offset
|
||||
Index pointer_offset) const {
|
||||
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
AccessType const *access_ptr = reinterpret_cast<AccessType const *>(ref_.data());
|
||||
int ldm = ref_.stride()[0];
|
||||
|
||||
if (kOperand == Operand::kA) {
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int idx = 0; idx < FragmentCount::kRow; ++idx) {
|
||||
|
||||
int tile_idx = idx / 2;
|
||||
int quad_idx = idx % 2;
|
||||
|
||||
int row_offset = tile_idx * kInterleavedTileRows;
|
||||
frag_ptr[idx] = access_ptr[row_offset / kElementsPerAccess + quad_idx];
|
||||
}
|
||||
}
|
||||
else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int idx = 0; idx < FragmentCount::kColumn; ++idx) {
|
||||
|
||||
int tile_idx = idx / 2;
|
||||
int quad_idx = idx % 2;
|
||||
|
||||
int col_offset = tile_idx * kInterleavedTileColumns;
|
||||
frag_ptr[idx] = access_ptr[col_offset / kElementsPerAccess + quad_idx];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 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
|
||||
Index byte_offset) const {
|
||||
|
||||
load_with_pointer_offset(frag, byte_offset * 8 / sizeof_bits<Element>::value);
|
||||
}
|
||||
|
||||
/// 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 {
|
||||
|
||||
TensorCoord coord_offset(tile_offset.row() * Shape::kRow, tile_offset.column() * Shape::kColumn);
|
||||
|
||||
load_with_pointer_offset(frag, ref_.offset(coord_offset));
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void load(
|
||||
/// fragment to load from the tensor
|
||||
Fragment &frag,
|
||||
/// loads a tile with a logical offset in units of whole tiles
|
||||
TensorCoord const &tile_offset,
|
||||
/// loads a tile with a logical offset AND a pointer offset
|
||||
Index pointer_offset) const {
|
||||
|
||||
TensorCoord coord_offset(tile_offset.row() * Shape::kRow, tile_offset.column() * Shape::kColumn);
|
||||
|
||||
load_with_pointer_offset(frag, ref_.offset(coord_offset) + pointer_offset);
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory with logical offset in units of whole tiles.
|
||||
CUTLASS_DEVICE
|
||||
void load_with_byte_offset(
|
||||
/// fragment to load from the tensor
|
||||
Fragment &frag,
|
||||
/// loads a tile with a logical offset in units of whole tiles
|
||||
TensorCoord const &tile_offset,
|
||||
/// loads a tile with a logical offset AND a pointer offset
|
||||
Index byte_offset) const {
|
||||
|
||||
TensorCoord coord_offset(tile_offset.row() * Shape::kRow, tile_offset.column() * Shape::kColumn);
|
||||
|
||||
load_with_pointer_offset(frag, ref_.offset(coord_offset) + byte_offset * 8 / sizeof_bits<Element>::value);
|
||||
}
|
||||
|
||||
/// 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 operation
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Data type of elements
|
||||
typename Element_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Interval between adjacent *MMA instructions (in units of MMA
|
||||
/// instructions)
|
||||
int OpDelta_>
|
||||
class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
Shape_,
|
||||
Operand::kA,
|
||||
Element_,
|
||||
cutlass::layout::RowMajor,
|
||||
InstructionShape_,
|
||||
OpDelta_,
|
||||
32
|
||||
> : public MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner<
|
||||
Shape_, Operand::kA, Element_, cutlass::layout::RowMajor, InstructionShape_, OpDelta_> {
|
||||
|
||||
public:
|
||||
using Base = MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner<
|
||||
Shape_, Operand::kA, Element_, cutlass::layout::RowMajor, InstructionShape_, OpDelta_> ;
|
||||
|
||||
using TensorRef = typename Base::TensorRef;
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIterator(
|
||||
TensorRef const &ref,
|
||||
int lane_id
|
||||
): Base(ref, lane_id) { }
|
||||
|
||||
};
|
||||
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Data type of elements
|
||||
typename Element_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Interval between adjacent *MMA instructions (in units of MMA
|
||||
/// instructions)
|
||||
int OpDelta_>
|
||||
class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
Shape_,
|
||||
Operand::kA,
|
||||
Element_,
|
||||
cutlass::layout::ColumnMajor,
|
||||
InstructionShape_,
|
||||
OpDelta_,
|
||||
32
|
||||
> : public MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter<
|
||||
Shape_, Operand::kA, Element_, cutlass::layout::ColumnMajor, InstructionShape_, OpDelta_> {
|
||||
|
||||
public:
|
||||
using Base = MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter<
|
||||
Shape_, Operand::kA, Element_, cutlass::layout::ColumnMajor, InstructionShape_, OpDelta_> ;
|
||||
|
||||
using TensorRef = typename Base::TensorRef;
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIterator(
|
||||
TensorRef const &ref,
|
||||
int lane_id
|
||||
): Base(ref, lane_id) { }
|
||||
|
||||
};
|
||||
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Data type of elements
|
||||
typename Element_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Interval between adjacent *MMA instructions (in units of MMA
|
||||
/// instructions)
|
||||
int OpDelta_>
|
||||
class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
Shape_, Operand::kB, Element_,
|
||||
cutlass::layout::ColumnMajor,
|
||||
InstructionShape_, OpDelta_, 32
|
||||
> : public MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner<
|
||||
Shape_, Operand::kB, Element_, cutlass::layout::ColumnMajor, InstructionShape_, OpDelta_> {
|
||||
|
||||
public:
|
||||
using Base = MmaVoltaTensorOpMultiplicandTileIteratorCanonicalInner<
|
||||
Shape_, Operand::kB, Element_, cutlass::layout::ColumnMajor, InstructionShape_, OpDelta_>;
|
||||
|
||||
using TensorRef = typename Base::TensorRef;
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIterator(
|
||||
TensorRef const &ref,
|
||||
int lane_id
|
||||
): Base(ref, lane_id) { }
|
||||
};
|
||||
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Data type of elements
|
||||
typename Element_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Interval between adjacent *MMA instructions (in units of MMA
|
||||
/// instructions)
|
||||
int OpDelta_>
|
||||
class MmaVoltaTensorOpMultiplicandTileIterator<
|
||||
Shape_, Operand::kB, Element_,
|
||||
cutlass::layout::RowMajor,
|
||||
InstructionShape_, OpDelta_, 32
|
||||
> : public MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter<
|
||||
Shape_, Operand::kB, Element_, cutlass::layout::RowMajor, InstructionShape_, OpDelta_> {
|
||||
|
||||
public:
|
||||
using Base = MmaVoltaTensorOpMultiplicandTileIteratorCanonicalOuter<
|
||||
Shape_, Operand::kB, Element_, cutlass::layout::RowMajor, InstructionShape_, OpDelta_>;
|
||||
|
||||
using TensorRef = typename Base::TensorRef;
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaVoltaTensorOpMultiplicandTileIterator(
|
||||
TensorRef const &ref,
|
||||
int lane_id
|
||||
): Base(ref, lane_id) { }
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
Reference in New Issue
Block a user