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:
Manish Gupta
2020-11-19 21:25:25 -08:00
committed by GitHub
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
224 changed files with 43939 additions and 1061 deletions
@@ -335,6 +335,10 @@ public:
using HandlePtr = std::unique_ptr<Handle>;
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Finds conv2d operation instances with Conv2d::ElementC = Reduction::ElementWorkspace
Operation const* find_conv_operation_for_parallel_reduction(Operation const *operation);
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace library
} // namespace cutlass
@@ -53,6 +53,10 @@
#include "cutlass/layout/tensor.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/conv/convolution.h"
#include "cutlass/conv/conv2d_problem_size.h"
#include "cutlass/conv/conv3d_problem_size.h"
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass {
@@ -79,6 +83,10 @@ enum class LayoutTypeID {
kTensorNCDHW,
kTensorNHWC,
kTensorNDHWC,
kTensorNC32HW32,
kTensorC32RSK32,
kTensorNC64HW64,
kTensorC64RSK64,
kInvalid
};
@@ -138,6 +146,7 @@ enum class Provider {
kReferenceHost,
kReferenceDevice,
kCUBLAS,
kCUDNN,
kInvalid
};
@@ -146,6 +155,8 @@ enum class Provider {
/// Enumeration indicating the kind of operation
enum class OperationKind {
kGemm,
kConv2d,
kConv3d,
kEqGemm,
kSparseGemm,
kReduction,
@@ -204,6 +215,30 @@ enum class GemmKind {
/// Mode of Universal GEMM
using GemmUniversalMode = cutlass::gemm::GemmUniversalMode;
/// Enumeration indicating what kind of Conv2d operation to perform
enum class ConvKind {
kUnknown,
kFprop,
kDgrad,
kWgrad,
kInvalid
};
enum class ConvModeID {
kCrossCorrelation,
kConvolution,
kInvalid
};
// Iterator algorithm enum in order of general performance-efficiency
enum class IteratorAlgorithmID {
kNone,
kAnalytic,
kOptimized,
kInvalid
};
enum class EpilogueKind {
kUnknown,
kConversion,
@@ -477,6 +512,66 @@ struct ReductionDescription : public OperationDescription {
};
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Description of all Conv2d operations
struct ConvDescription : public OperationDescription {
/// Describes the convolution dimension support (2D or 3D)
int conv_dim;
/// Describes the kind of convolution
ConvKind conv_kind;
/// Describes the type of iterator algorithm (analytic or precomputed)
IteratorAlgorithmID iterator_algorithm;
/// Describes the A operand
TensorDescription A;
/// Describes the B operand
TensorDescription B;
/// Describes the C operand
TensorDescription C;
/// Describes the data type of the scalars passed to the epilogue
NumericTypeID element_epilogue;
//
// Methods
//
// Returns Activation TensorDescription
TensorDescription activation() const {
switch(conv_kind) {
case library::ConvKind::kFprop : return A;
case library::ConvKind::kDgrad : return C;
case library::ConvKind::kWgrad : return B;
default : throw std::runtime_error("Invalid Conv Operator (fprop, dgrad, wgrad)");
}
}
// Returns Filter TensorDescription
TensorDescription filter() const {
switch(conv_kind) {
case library::ConvKind::kFprop : return B;
case library::ConvKind::kDgrad : return B;
case library::ConvKind::kWgrad : return C;
default : throw std::runtime_error("Invalid Conv Operator (fprop, dgrad, wgrad)");
}
}
// Returns Output TensorDescription
TensorDescription output() const {
switch(conv_kind) {
case library::ConvKind::kFprop : return C;
case library::ConvKind::kDgrad : return A;
case library::ConvKind::kWgrad : return A;
default : throw std::runtime_error("Invalid Conv Operator (fprop, dgrad, wgrad)");
}
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Base class for all operations
@@ -825,6 +920,204 @@ struct SparseGemmArguments {
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Two dimensional convolution
//
// OperationKind: Conv2d
//
struct Conv2dConfiguration {
conv::SplitKMode split_k_mode;
/// Conv2d problem size
// contains strictly conv2d size (N,H,W,C,K,R,S,P,Q,padding,stride,dilation,mode)
// also includes (split_k_slices, groups)
conv::Conv2dProblemSize problem_size;
/// Layout object for activations tensor
layout::TensorNHWC layout_activations;
/// Layout object for filters tensor
layout::TensorNHWC layout_filters;
/// Layout object for source tensor
layout::TensorNHWC layout_source;
/// Layout object for output tensor
layout::TensorNHWC layout_output;
//
// Methods
//
// Mapping functions (A,B,C -> activation,filter,output)
layout::TensorNHWC layout_a(library::ConvKind const &conv_kind) const {
switch (conv_kind) {
case library::ConvKind::kFprop: return layout_activations;
case library::ConvKind::kDgrad: return layout_output;
case library::ConvKind::kWgrad: return layout_output;
default : throw std::runtime_error("Invalid Conv Operator (fprop, dgrad, wgrad)");
}
}
layout::TensorNHWC layout_b(library::ConvKind const &conv_kind) const {
switch (conv_kind) {
case library::ConvKind::kFprop: return layout_filters;
case library::ConvKind::kDgrad: return layout_filters;
case library::ConvKind::kWgrad: return layout_activations;
default : throw std::runtime_error("Invalid Conv Operator (fprop, dgrad, wgrad)");
}
}
layout::TensorNHWC layout_c(library::ConvKind const &conv_kind) const {
switch (conv_kind) {
case library::ConvKind::kFprop: return layout_output;
case library::ConvKind::kDgrad: return layout_activations;
case library::ConvKind::kWgrad: return layout_filters;
default : throw std::runtime_error("Invalid Conv Operator (fprop, dgrad, wgrad)");
}
}
};
/// Three dimensional convolution
//
// OperationKind: Conv3d
//
struct Conv3dConfiguration {
conv::SplitKMode split_k_mode;
/// Conv2d problem size
// contains strictly conv2d size (N,D,H,W,C,K,T,R,S,Z,P,Q,padding,stride,dilation,mode)
// also includes (split_k_slices, groups)
conv::Conv3dProblemSize problem_size;
/// Layout object for activations tensor
layout::TensorNDHWC layout_activations;
/// Layout object for filters tensor
layout::TensorNDHWC layout_filters;
/// Layout object for source tensor
layout::TensorNDHWC layout_source;
/// Layout object for output tensor
layout::TensorNDHWC layout_output;
//
// Methods
//
// Mapping functions (A,B,C -> activation,filter,output)
layout::TensorNDHWC layout_a(library::ConvKind const &conv_kind) const {
switch (conv_kind) {
case library::ConvKind::kFprop: return layout_activations;
case library::ConvKind::kDgrad: return layout_output;
case library::ConvKind::kWgrad: return layout_output;
default : throw std::runtime_error("Invalid Conv Operator (fprop, dgrad, wgrad)");
}
}
layout::TensorNDHWC layout_b(library::ConvKind const &conv_kind) const {
switch (conv_kind) {
case library::ConvKind::kFprop: return layout_filters;
case library::ConvKind::kDgrad: return layout_filters;
case library::ConvKind::kWgrad: return layout_activations;
default : throw std::runtime_error("Invalid Conv Operator (fprop, dgrad, wgrad)");
}
}
layout::TensorNDHWC layout_c(library::ConvKind const &conv_kind) const {
switch (conv_kind) {
case library::ConvKind::kFprop: return layout_output;
case library::ConvKind::kDgrad: return layout_activations;
case library::ConvKind::kWgrad: return layout_filters;
default : throw std::runtime_error("Invalid Conv Operator (fprop, dgrad, wgrad)");
}
}
};
/// Arguments for CONV
struct ConvArguments {
/////////////////////////////////////////////////////////
/// ImplicitGemm matrices A, B, C, D
/////////////////////////////////////////////////////////
/// pointer to implicit gemm matrix A
void const *A;
/// pointer to implicit gemm matrix B
void const *B;
/// pointer to implicit gemm matrix C
void const *C;
/// pointer to implicit gemm desitination matrix D
void *D;
/// Host or device pointer to alpha scalar
void const *alpha;
/// Host or device pointer to beta scalar
void const *beta;
/// Enumerant indicating whether alpha/beta point to host or device memory
ScalarPointerMode pointer_mode;
};
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Configuration for Reduction operations
//
// OperationKind: Reduction
//
struct ReductionConfiguration {
/// Redcution problem size
MatrixCoord problem_size;
/// Number of partitions to reduce
int partitions;
/// Number of lements between each partition
int64_t partition_stride;
/// leading dimension of 'w'orksace operand
int64_t ldw;
/// leading dimension of 's'ource operand
int64_t lds;
/// leading dimension of 'd'estination operand
int64_t ldd;
};
/// Arguments for Reduction
struct ReductionArguments {
/// Pointer to workspace matrix
void const *workspace;
/// Pointer to source matrix
void const *source;
/// Pointer to destination matrix
void *destination;
/// pointer to reference matrix
void *reference;
/// Host or device pointer to alpha scalar
void const *alpha;
/// Host or device pointer to beta scalar
void const *beta;
/// Enumerant indicating whether alpha/beta point to host or device memory
ScalarPointerMode pointer_mode;
};
} // namespace library
} // namespace cutlass
@@ -51,6 +51,9 @@ class Manifest;
// init and insert all cutlass gemm operations in manifest object (procedurally generated using generator.py)
void initialize_all(Manifest &manifest);
// init and insert all reduction op in manifest object (manually instantiated in library/reduction)
void initialize_all_reduction_op(Manifest &manifest);
/////////////////////////////////////////////////////////////////////////////////////////////////////////
/// List of operations
@@ -208,6 +208,262 @@ using GemmOperationFunctionalMap = std::unordered_map<
>;
/////////////////////////////////////////////////////////////////////////////////////////////////
/////////////////////////////////////////////////////////////////////////////////////////////////
// Data Structures for Conv Functional Maps
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Tuple uniquely identifying conv2d functional behavior
struct ConvFunctionalKey {
library::Provider provider;
library::ConvKind conv_kind;
library::NumericTypeID element_A;
library::LayoutTypeID layout_A;
library::NumericTypeID element_B;
library::LayoutTypeID layout_B;
library::NumericTypeID element_C;
library::LayoutTypeID layout_C;
library::NumericTypeID element_accumulator;
library::NumericTypeID element_compute;
//
// Methods
//
inline
ConvFunctionalKey(
library::Provider provider = library::Provider::kInvalid,
library::ConvKind conv_kind = library::ConvKind::kFprop,
library::NumericTypeID element_A = library::NumericTypeID::kF16,
library::LayoutTypeID layout_A = library::LayoutTypeID::kTensorNHWC,
library::NumericTypeID element_B = library::NumericTypeID::kF16,
library::LayoutTypeID layout_B = library::LayoutTypeID::kTensorNHWC,
library::NumericTypeID element_C = library::NumericTypeID::kF16,
library::LayoutTypeID layout_C = library::LayoutTypeID::kTensorNHWC,
library::NumericTypeID element_accumulator = library::NumericTypeID::kF32,
library::NumericTypeID element_compute = library::NumericTypeID::kF32
):
provider(provider),
conv_kind(conv_kind),
element_A(element_A),
layout_A(layout_A),
element_B(element_B),
layout_B(layout_B),
element_C(element_C),
layout_C(layout_C),
element_accumulator(element_accumulator),
element_compute(element_compute)
{ }
inline
bool operator==(ConvFunctionalKey const &rhs) const {
return
(provider == rhs.provider) &&
(conv_kind == rhs.conv_kind) &&
(element_A == rhs.element_A) &&
(layout_A == rhs.layout_A) &&
(element_B == rhs.element_B) &&
(layout_B == rhs.layout_B) &&
(element_C == rhs.element_C) &&
(layout_C == rhs.layout_C) &&
(element_accumulator == rhs.element_accumulator) &&
(element_compute == rhs.element_compute);
}
inline
bool operator!=(ConvFunctionalKey const &rhs) const {
return !(*this == rhs);
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
inline
std::ostream& operator<< (std::ostream& out, const cutlass::library::ConvFunctionalKey& key) {
out << "{\n"
<< "provider: " << to_string(key.provider) << std::endl
<< "conv_kind: " << to_string(key.conv_kind) << std::endl
<< "element_A: " << to_string(key.element_A) << std::endl
<< "layout_A: " << to_string(key.layout_A) << std::endl
<< "element_B: " << to_string(key.element_B) << std::endl
<< "layout_B: " << to_string(key.layout_B) << std::endl
<< "element_C: " << to_string(key.element_C) << std::endl
<< "layout_C: " << to_string(key.layout_C) << std::endl
<< "element_accumulator: " << to_string(key.element_accumulator) << std::endl
<< "element_compute: " << to_string(key.element_compute) << std::endl
<< "}";
return out;
}
/////////////////////////////////////////////////////////////////////////////////////////////////
struct ConvFunctionalKeyHasher {
using IntHash = std::hash<int>;
inline
static size_t rotl(size_t key, int shl) {
return (key << shl) | (key >> (sizeof(key)*8 - shl));
}
inline
size_t operator()(ConvFunctionalKey const &key) const {
IntHash hash;
return
rotl(hash(int(key.provider)), 1) ^
rotl(hash(int(key.conv_kind)), 2) ^
rotl(hash(int(key.element_A)), 3) ^
rotl(hash(int(key.layout_A)), 4) ^
rotl(hash(int(key.element_B)), 5) ^
rotl(hash(int(key.layout_B)), 6) ^
rotl(hash(int(key.element_C)), 7) ^
rotl(hash(int(key.layout_C)), 8) ^
rotl(hash(int(key.element_accumulator)), 9) ^
rotl(hash(int(key.element_compute)), 10);
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Establishes a partial ordering to search for Conv2d operators
struct ConvPreferenceKey {
int compute_capability;
IteratorAlgorithmID iterator_algorithm;
//
// Methods
//
ConvPreferenceKey(): compute_capability(), iterator_algorithm() { }
ConvPreferenceKey(int cc, IteratorAlgorithmID iterator_algorithm):
compute_capability(cc), iterator_algorithm(iterator_algorithm) { }
bool operator<(ConvPreferenceKey const &rhs) const {
return (compute_capability < rhs.compute_capability) ||
((compute_capability == rhs.compute_capability) && (iterator_algorithm < rhs.iterator_algorithm));
}
bool operator==(ConvPreferenceKey const &rhs) const {
return (compute_capability == rhs.compute_capability) &&
(iterator_algorithm == rhs.iterator_algorithm);
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Maps minimum compute capability onto a vector of possible operations
using ConvOperationVectorMap = std::map<
ConvPreferenceKey,
std::vector<Operation const *>
>;
/// Maps a GemmFunctionalKey onto a vector of Operation * objects expected to be of kind kGemm
using ConvOperationFunctionalMap = std::unordered_map<
ConvFunctionalKey,
ConvOperationVectorMap,
ConvFunctionalKeyHasher
>;
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Tuple uniquely identifying conv2d functional behavior
struct ReductionFunctionalKey {
library::Provider provider;
library::NumericTypeID element_workspace;
library::NumericTypeID element_accumulator;
library::NumericTypeID element_output;
library::NumericTypeID element_compute;
library::MathOperationID reduce_math_op;
library::EpilogueKind epilogue_math_op;
//
// Methods
//
inline
ReductionFunctionalKey(
library::Provider provider = library::Provider::kInvalid,
library::NumericTypeID element_workspace = library::NumericTypeID::kF16,
library::NumericTypeID element_accumulator = library::NumericTypeID::kF32,
library::NumericTypeID element_output = library::NumericTypeID::kF16,
library::NumericTypeID element_compute = library::NumericTypeID::kF32,
library::MathOperationID reduce_math_op = library::MathOperationID::kAdd,
library::EpilogueKind epilogue_math_op = library::EpilogueKind::kLinearCombination
):
provider(provider),
element_workspace(element_workspace),
element_accumulator(element_accumulator),
element_output(element_output),
element_compute(element_compute),
reduce_math_op(reduce_math_op),
epilogue_math_op(epilogue_math_op)
{ }
inline
bool operator==(ReductionFunctionalKey const &rhs) const {
return
(provider == rhs.provider) &&
(element_workspace == rhs.element_workspace) &&
(element_accumulator == rhs.element_accumulator) &&
(element_output == rhs.element_output) &&
(element_compute == rhs.element_compute) &&
(reduce_math_op == rhs.reduce_math_op) &&
(epilogue_math_op == rhs.epilogue_math_op);
}
inline
bool operator!=(ReductionFunctionalKey const &rhs) const {
return !(*this == rhs);
}
};
struct ReductionFunctionalKeyHasher {
using IntHash = std::hash<int>;
inline
static size_t rotl(size_t key, int shl) {
return (key << shl) | (key >> (sizeof(key)*8 - shl));
}
inline
size_t operator()(ReductionFunctionalKey const &key) const {
IntHash hash;
return
rotl(hash(int(key.provider)), 1) ^
rotl(hash(int(key.element_workspace)), 2) ^
rotl(hash(int(key.element_accumulator)), 3) ^
rotl(hash(int(key.element_output)), 4) ^
rotl(hash(int(key.element_compute)), 5) ^
rotl(hash(int(key.reduce_math_op)), 6) ^
rotl(hash(int(key.epilogue_math_op)), 7);
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
inline
std::ostream& operator<< (std::ostream& out, const ReductionFunctionalKey& key) {
out << "{\n"
<< "provider: " << library::to_string(key.provider) << std::endl
<< "element_workspace : " << library::to_string(key.element_workspace) << std::endl
<< "element_accumulator : " << library::to_string(key.element_accumulator) << std::endl
<< "element_output : " << library::to_string(key.element_output) << std::endl
<< "element_compute : " << library::to_string(key.element_compute) << std::endl
<< "}";
return out;
}
/////////////////////////////////////////////////////////////////////////////////////////////////
// ReductionOperationFunctionalMap has NO preference key and a single instance per functional key
// i.e. only one tile size configuration per functional key
using ReductionOperationFunctionalMap = std::unordered_map<
ReductionFunctionalKey,
library::Operation const *,
ReductionFunctionalKeyHasher
>;
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Table of cutlass::library::Operation instances
@@ -218,6 +474,18 @@ public:
// provider (kCUTLASS)
GemmOperationFunctionalMap gemm_operations;
/// Map of all operations of type kConv2d
// provider (kCUTLASS, kReferenceHost, kReferenceDevice)
ConvOperationFunctionalMap conv2d_operations;
/// Map of all operations of type kConv3d
// provider (kCUTLASS, kReferenceHost, kReferenceDevice)
ConvOperationFunctionalMap conv3d_operations;
/// Map of all operations of type kConv2d
// provider (kCUTLASS)
ReductionOperationFunctionalMap reduction_operations;
public:
void append(Manifest const &manifest);
@@ -122,6 +122,27 @@ char const *to_string(SplitKMode split_k_mode, bool pretty = false);
template <>
SplitKMode from_string<SplitKMode>(std::string const &str);
/// Converts a ConvModeID enumerant to a string
char const *to_string(ConvModeID type, bool pretty = false);
/// Converts a ConvModeID enumerant from a string
template <>
ConvModeID from_string<ConvModeID>(std::string const &str);
/// Converts a IteratorAlgorithmID enumerant to a string
char const *to_string(IteratorAlgorithmID type, bool pretty = false);
/// Converts a IteratorAlgorithmID enumerant from a string
template <>
IteratorAlgorithmID from_string<IteratorAlgorithmID>(std::string const &str);
/// Converts a ConvKind enumerant to a string
char const *to_string(ConvKind type, bool pretty = false);
/// Converts a ConvKind enumerant from a string
template <>
ConvKind from_string<ConvKind>(std::string const &str);
/// Lexical cast from int64_t to string
std::string lexical_cast(int64_t int_value);