CUTLASS 2.1 (#83)
CUTLASS 2.1 contributes: - BLAS-style host-side API added to CUTLASS Library - Planar Complex GEMM kernels targeting Volta and Turing Tensor Cores - Minor enhancements and bug fixes
This commit is contained in:
@@ -33,7 +33,7 @@
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/library/library.h"
|
||||
|
||||
#include "options.h"
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
@@ -86,6 +86,161 @@ public:
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
|
||||
/// Selects one or more cuBLAS algorithms.
|
||||
static void select_cublas_algorithms(
|
||||
std::vector<cublasGemmAlgo_t> &algorithms,
|
||||
Options const &options,
|
||||
library::GemmDescription const &op_desc) {
|
||||
|
||||
library::OpcodeClassID const & opcode_class =
|
||||
op_desc.tile_description.math_instruction.opcode_class;
|
||||
|
||||
switch (options.library.algorithm_mode) {
|
||||
case AlgorithmMode::kMatching:
|
||||
{
|
||||
algorithms.push_back(get_cublas_gemm_algo(
|
||||
op_desc.tile_description.threadblock_shape.m(),
|
||||
op_desc.tile_description.threadblock_shape.n(),
|
||||
op_desc.tile_description.threadblock_shape.k(),
|
||||
opcode_class));
|
||||
break;
|
||||
}
|
||||
|
||||
case AlgorithmMode::kBest:
|
||||
{
|
||||
// Choose first enumerated mode. If none are enumerated, choose based on opcode class
|
||||
// and evaluate all of them.
|
||||
|
||||
if (options.library.algorithms.empty()) {
|
||||
// Enumerate all algorithms
|
||||
if (opcode_class == library::OpcodeClassID::kSimt) {
|
||||
|
||||
for (int algo = CUBLAS_GEMM_DEFAULT;
|
||||
algo <= CUBLAS_GEMM_ALGO23;
|
||||
++algo) {
|
||||
|
||||
algorithms.push_back(cublasGemmAlgo_t(algo));
|
||||
}
|
||||
}
|
||||
else {
|
||||
|
||||
for (int algo = CUBLAS_GEMM_DEFAULT_TENSOR_OP;
|
||||
algo <= CUBLAS_GEMM_ALGO15_TENSOR_OP;
|
||||
++algo) {
|
||||
|
||||
algorithms.push_back(cublasGemmAlgo_t(algo));
|
||||
}
|
||||
}
|
||||
}
|
||||
else {
|
||||
// Use the listed algorithms
|
||||
algorithms.reserve(options.library.algorithms.size());
|
||||
|
||||
for (int algo : options.library.algorithms) {
|
||||
algorithms.push_back(reinterpret_cast<cublasGemmAlgo_t const &>(algo));
|
||||
}
|
||||
}
|
||||
|
||||
break;
|
||||
}
|
||||
|
||||
case AlgorithmMode::kDefault:
|
||||
{
|
||||
|
||||
// Use the library's default algorithm
|
||||
algorithms.push_back((opcode_class == library::OpcodeClassID::kSimt ?
|
||||
CUBLAS_GEMM_DEFAULT : CUBLAS_GEMM_DEFAULT_TENSOR_OP));
|
||||
|
||||
break;
|
||||
}
|
||||
default:
|
||||
{
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Dispatcher to cublasGemmEx()
|
||||
struct cublasGemmExDispatcher {
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
library::GemmConfiguration configuration;
|
||||
library::GemmArguments arguments;
|
||||
|
||||
// cublass-specific data structures to fill cublas API call arguments
|
||||
cublasOperation_t trans_A;
|
||||
cublasOperation_t trans_B;
|
||||
cudaDataType_t data_type_A;
|
||||
cudaDataType_t data_type_B;
|
||||
cudaDataType_t data_type_C;
|
||||
cudaDataType_t compute_type;
|
||||
cublasGemmAlgo_t algo;
|
||||
Status status;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
cublasGemmExDispatcher(
|
||||
library::GemmDescription const &op_desc,
|
||||
library::GemmConfiguration configuration_,
|
||||
library::GemmArguments arguments_,
|
||||
cublasGemmAlgo_t algorithm = CUBLAS_GEMM_DFALT
|
||||
):
|
||||
configuration(configuration_), arguments(arguments_), algo(algorithm), status(Status::kSuccess) {
|
||||
|
||||
trans_A = get_cublas_transpose_operation(op_desc.A.layout);
|
||||
trans_B = get_cublas_transpose_operation(op_desc.B.layout);
|
||||
|
||||
bool good = true;
|
||||
good = (good && get_cublas_datatype(data_type_A, op_desc.A.element));
|
||||
good = (good && get_cublas_datatype(data_type_B, op_desc.B.element));
|
||||
good = (good && get_cublas_datatype(data_type_C, op_desc.C.element));
|
||||
|
||||
good = (good && get_cublas_datatype(
|
||||
compute_type,
|
||||
op_desc.tile_description.math_instruction.element_accumulator));
|
||||
|
||||
if (!good) {
|
||||
status = Status::kErrorNotSupported;
|
||||
}
|
||||
}
|
||||
|
||||
/// Executes GEMM using these arguments
|
||||
cublasStatus_t operator()(cublasHandle_t handle) {
|
||||
|
||||
return cublasGemmEx(
|
||||
handle,
|
||||
trans_A,
|
||||
trans_B,
|
||||
configuration.problem_size.m(),
|
||||
configuration.problem_size.n(),
|
||||
configuration.problem_size.k(),
|
||||
arguments.alpha,
|
||||
arguments.A,
|
||||
data_type_A,
|
||||
int(configuration.lda),
|
||||
arguments.B,
|
||||
data_type_B,
|
||||
int(configuration.ldb),
|
||||
arguments.beta,
|
||||
arguments.D,
|
||||
data_type_C,
|
||||
int(configuration.ldc),
|
||||
compute_type,
|
||||
algo
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace detail
|
||||
|
||||
} // namespace profiler
|
||||
} // namespace cutlass
|
||||
|
||||
|
||||
Reference in New Issue
Block a user