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
@@ -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);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user