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:
@@ -29,8 +29,13 @@
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/kernel/default_gemm_planar_complex_universal.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm.h"
|
||||
#include "cutlass/gemm/device/gemm_complex.h"
|
||||
#include "cutlass/gemm/device/gemm_batched.h"
|
||||
#include "cutlass/gemm/device/gemm_array.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
|
||||
#include "cutlass/library/library.h"
|
||||
#include "library_internal.h"
|
||||
@@ -68,8 +73,10 @@ public:
|
||||
GemmOperationBase(char const *name = "unknown_gemm") {
|
||||
|
||||
description_.name = name;
|
||||
description_.provider = Provider::kCUTLASS;
|
||||
description_.kind = OperationKind::kGemm;
|
||||
|
||||
description_.gemm_kind = GemmKind::kGemm;
|
||||
|
||||
description_.tile_description.threadblock_shape = make_Coord(
|
||||
Operator::ThreadblockShape::kM,
|
||||
Operator::ThreadblockShape::kN,
|
||||
@@ -93,22 +100,23 @@ public:
|
||||
description_.tile_description.math_instruction.opcode_class =
|
||||
OpcodeClassMap<typename Operator::OperatorClass>::kId;
|
||||
|
||||
description_.tile_description.math_instruction.math_operation =
|
||||
MathOperationMap<typename Operator::Operator>::kId;
|
||||
|
||||
description_.tile_description.minimum_compute_capability =
|
||||
ArchMap<typename Operator::ArchTag>::kMin;
|
||||
|
||||
description_.tile_description.maximum_compute_capability =
|
||||
ArchMap<typename Operator::ArchTag>::kMax;
|
||||
|
||||
description_.gemm_kind = GemmKind::kGemm;
|
||||
|
||||
description_.A = make_TensorDescription<ElementA, LayoutA>(Operator::kAlignmentA);
|
||||
description_.B = make_TensorDescription<ElementB, LayoutB>(Operator::kAlignmentB);
|
||||
description_.C = make_TensorDescription<ElementC, LayoutC>(Operator::kAlignmentC);
|
||||
description_.element_epilogue = NumericTypeMap<ElementCompute>::kId;
|
||||
|
||||
description_.split_k_mode = Operator::kSplitKSerial ? SplitKMode::kSerial : SplitKMode::kNone;
|
||||
description_.transform_A = ComplexTransform::kNone;
|
||||
description_.transform_B = ComplexTransform::kNone;
|
||||
description_.split_k_mode = SplitKMode::kNone;
|
||||
description_.transform_A = ComplexTransformMap<Operator::kTransformA>::kId;
|
||||
description_.transform_B = ComplexTransformMap<Operator::kTransformB>::kId;
|
||||
}
|
||||
|
||||
/// Returns the description of the GEMM operation
|
||||
@@ -294,8 +302,24 @@ public:
|
||||
|
||||
return op->run(stream);
|
||||
}
|
||||
};
|
||||
|
||||
void print_operator_args(OperatorArguments &operator_args) const {
|
||||
#if 0
|
||||
std::cout << "GemmOperation::OperatorArguments" << std::endl;
|
||||
std::cout << " problem_size: " << operator_args.problem_size.m() << ", "<< operator_args.problem_size.n() << "," << operator_args.problem_size.k() << std::endl;
|
||||
std::cout << " alpha: " << operator_args.epilogue.alpha << std::endl;
|
||||
std::cout << " alpha_ptr: " << operator_args.epilogue.alpha_ptr << std::endl;
|
||||
std::cout << " beta: " << operator_args.epilogue.beta << std::endl;
|
||||
std::cout << " beta_ptr: " << operator_args.epilogue.beta_ptr << std::endl;
|
||||
std::cout << " ref_A.data(): " << operator_args.ref_A.data() << std::endl;
|
||||
std::cout << " ref_A.stride: " << operator_args.ref_A.stride(0) << std::endl;
|
||||
std::cout << " ref_B.data(): " << operator_args.ref_B.data() << std::endl;
|
||||
std::cout << " ref_B.stride: " << operator_args.ref_B.stride(0) << std::endl;
|
||||
std::cout << " ref_C.data(): " << operator_args.ref_C.data() << std::endl;
|
||||
std::cout << " ref_C.stride: " << operator_args.ref_C.stride(0) << std::endl;
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -360,6 +384,7 @@ protected:
|
||||
*static_cast<ElementCompute const *>(arguments->alpha),
|
||||
*static_cast<ElementCompute const *>(arguments->beta)
|
||||
);
|
||||
|
||||
operator_args.epilogue = params;
|
||||
}
|
||||
else if (arguments->pointer_mode == ScalarPointerMode::kDevice){
|
||||
@@ -491,6 +516,593 @@ public:
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Operator_>
|
||||
class GemmArrayOperation : public GemmOperationBase<Operator_> {
|
||||
public:
|
||||
|
||||
using Operator = Operator_;
|
||||
using ElementA = typename Operator::ElementA;
|
||||
using LayoutA = typename Operator::LayoutA;
|
||||
using ElementB = typename Operator::ElementB;
|
||||
using LayoutB = typename Operator::LayoutB;
|
||||
using ElementC = typename Operator::ElementC;
|
||||
using LayoutC = typename Operator::LayoutC;
|
||||
using ElementAccumulator = typename Operator::ElementAccumulator;
|
||||
using ElementCompute = typename Operator::EpilogueOutputOp::ElementCompute;
|
||||
|
||||
using OperatorArguments = typename Operator::Arguments;
|
||||
|
||||
protected:
|
||||
|
||||
///
|
||||
GemmDescription description_;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructor
|
||||
GemmArrayOperation(char const *name = "unknown_gemm"): GemmOperationBase<Operator_>(name) {
|
||||
|
||||
description_.gemm_kind = GemmKind::kArray;
|
||||
}
|
||||
|
||||
protected:
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status construct_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
GemmArrayConfiguration const *configuration) {
|
||||
|
||||
operator_args.problem_size = configuration->problem_size;
|
||||
|
||||
operator_args.batch_count = configuration->batch_count;
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status update_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
GemmArrayArguments const *arguments) {
|
||||
|
||||
if (arguments->pointer_mode == ScalarPointerMode::kHost) {
|
||||
typename Operator::EpilogueOutputOp::Params params(
|
||||
*static_cast<ElementCompute const *>(arguments->alpha),
|
||||
*static_cast<ElementCompute const *>(arguments->beta)
|
||||
);
|
||||
operator_args.epilogue = params;
|
||||
}
|
||||
else if (arguments->pointer_mode == ScalarPointerMode::kDevice){
|
||||
typename Operator::EpilogueOutputOp::Params params(
|
||||
static_cast<ElementCompute const *>(arguments->alpha),
|
||||
static_cast<ElementCompute const *>(arguments->beta)
|
||||
);
|
||||
operator_args.epilogue = params;
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
public:
|
||||
|
||||
/// Returns the description of the GEMM operation
|
||||
virtual OperationDescription const & description() const {
|
||||
return description_;
|
||||
}
|
||||
|
||||
/// Returns success if the operation can proceed
|
||||
virtual Status can_implement(
|
||||
void const *configuration_ptr,
|
||||
void const *arguments_ptr) const {
|
||||
|
||||
GemmArrayConfiguration const *configuration =
|
||||
static_cast<GemmArrayConfiguration const *>(configuration_ptr);
|
||||
|
||||
GemmArrayArguments const *arguments =
|
||||
static_cast<GemmArrayArguments const *>(arguments_ptr);
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = construct_arguments_(args, configuration);
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
status = update_arguments_(args, arguments);
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
return Operator::can_implement(args);
|
||||
}
|
||||
|
||||
/// Gets the host-side workspace
|
||||
virtual uint64_t get_host_workspace_size(
|
||||
void const *configuration) const {
|
||||
|
||||
return sizeof(Operator);
|
||||
}
|
||||
|
||||
/// Gets the device-side workspace
|
||||
virtual uint64_t get_device_workspace_size(
|
||||
void const *configuration_ptr) const {
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = construct_arguments_(
|
||||
args,
|
||||
static_cast<GemmArrayConfiguration const *>(configuration_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
return Operator::get_workspace_size(args);
|
||||
}
|
||||
|
||||
/// Initializes the workspace
|
||||
virtual Status initialize(
|
||||
void const *configuration_ptr,
|
||||
void *host_workspace,
|
||||
void *device_workspace,
|
||||
cudaStream_t stream = nullptr) const {
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = construct_arguments_(
|
||||
args,
|
||||
static_cast<GemmArrayConfiguration const *>(configuration_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
Operator *op = new (host_workspace) Operator;
|
||||
|
||||
return op->initialize(args, device_workspace, stream);
|
||||
}
|
||||
|
||||
/// Runs the kernel
|
||||
virtual Status run(
|
||||
void const *arguments_ptr,
|
||||
void *host_workspace,
|
||||
void *device_workspace = nullptr,
|
||||
cudaStream_t stream = nullptr) const {
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = update_arguments_(
|
||||
args,
|
||||
static_cast<GemmArrayArguments const *>(arguments_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
Operator *op = static_cast<Operator *>(host_workspace);
|
||||
|
||||
status = op->update(args, device_workspace);
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
return op->run(stream);
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Operator_>
|
||||
class GemmPlanarComplexOperation : public GemmOperationBase<Operator_> {
|
||||
public:
|
||||
|
||||
using Operator = Operator_;
|
||||
using ElementA = typename Operator::ElementA;
|
||||
using LayoutA = typename Operator::LayoutA;
|
||||
using ElementB = typename Operator::ElementB;
|
||||
using LayoutB = typename Operator::LayoutB;
|
||||
using ElementC = typename Operator::ElementC;
|
||||
using LayoutC = typename Operator::LayoutC;
|
||||
using ElementAccumulator = typename Operator::ElementAccumulator;
|
||||
using ElementCompute = typename Operator::EpilogueOutputOp::ElementCompute;
|
||||
|
||||
using OperatorArguments = typename Operator::Arguments;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructor
|
||||
GemmPlanarComplexOperation(char const *name = "unknown_gemm"): GemmOperationBase<Operator_>(name) {
|
||||
|
||||
this->description_.gemm_kind = GemmKind::kPlanarComplex;
|
||||
}
|
||||
|
||||
protected:
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status construct_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
GemmPlanarComplexConfiguration const *configuration) {
|
||||
|
||||
operator_args.mode = cutlass::gemm::GemmUniversalMode::kBatched;
|
||||
operator_args.problem_size = configuration->problem_size;
|
||||
operator_args.batch_count = configuration->batch_count;
|
||||
|
||||
operator_args.lda_real = int(configuration->lda_real);
|
||||
operator_args.lda_imag = int(configuration->lda_imag);
|
||||
operator_args.ldb_real = int(configuration->ldb_real);
|
||||
operator_args.ldb_imag = int(configuration->ldb_imag);
|
||||
operator_args.ldc_real = int(configuration->ldc_real);
|
||||
operator_args.ldc_imag = int(configuration->ldc_imag);
|
||||
operator_args.ldd_real = int(configuration->ldd_real);
|
||||
operator_args.ldd_imag = int(configuration->ldd_imag);
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status update_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
GemmPlanarComplexArguments const *arguments) {
|
||||
|
||||
if (arguments->pointer_mode == ScalarPointerMode::kHost) {
|
||||
typename Operator::EpilogueOutputOp::Params params(
|
||||
*static_cast<cutlass::complex<ElementCompute> const *>(arguments->alpha),
|
||||
*static_cast<cutlass::complex<ElementCompute> const *>(arguments->beta)
|
||||
);
|
||||
operator_args.epilogue = params;
|
||||
}
|
||||
else if (arguments->pointer_mode == ScalarPointerMode::kDevice){
|
||||
typename Operator::EpilogueOutputOp::Params params(
|
||||
static_cast<cutlass::complex<ElementCompute> const *>(arguments->alpha),
|
||||
static_cast<cutlass::complex<ElementCompute> const *>(arguments->beta)
|
||||
);
|
||||
operator_args.epilogue = params;
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
// update arguments
|
||||
operator_args.ptr_A_real = arguments->A_real;
|
||||
operator_args.ptr_A_imag = arguments->A_imag;
|
||||
operator_args.ptr_B_real = arguments->B_real;
|
||||
operator_args.ptr_B_imag = arguments->B_imag;
|
||||
operator_args.ptr_C_real = arguments->C_real;
|
||||
operator_args.ptr_C_imag = arguments->C_imag;
|
||||
operator_args.ptr_D_real = arguments->D_real;
|
||||
operator_args.ptr_D_imag = arguments->D_imag;
|
||||
|
||||
operator_args.batch_stride_A = arguments->batch_stride_A_real;
|
||||
operator_args.batch_stride_A_imag = arguments->batch_stride_A_imag;
|
||||
operator_args.batch_stride_B = arguments->batch_stride_B_real;
|
||||
operator_args.batch_stride_B_imag = arguments->batch_stride_B_imag;
|
||||
operator_args.batch_stride_C = arguments->batch_stride_C_real;
|
||||
operator_args.batch_stride_C_imag = arguments->batch_stride_C_imag;
|
||||
operator_args.batch_stride_D = arguments->batch_stride_D_real;
|
||||
operator_args.batch_stride_D_imag = arguments->batch_stride_D_imag;
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
public:
|
||||
|
||||
/// Returns success if the operation can proceed
|
||||
virtual Status can_implement(
|
||||
void const *configuration_ptr,
|
||||
void const *arguments_ptr) const {
|
||||
|
||||
GemmPlanarComplexConfiguration const *configuration =
|
||||
static_cast<GemmPlanarComplexConfiguration const *>(configuration_ptr);
|
||||
|
||||
GemmPlanarComplexArguments const *arguments =
|
||||
static_cast<GemmPlanarComplexArguments const *>(arguments_ptr);
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = construct_arguments_(args, configuration);
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
status = update_arguments_(args, arguments);
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
return Operator::can_implement(args);
|
||||
}
|
||||
|
||||
/// Gets the host-side workspace
|
||||
virtual uint64_t get_host_workspace_size(
|
||||
void const *configuration) const {
|
||||
|
||||
return sizeof(Operator);
|
||||
}
|
||||
|
||||
/// Gets the device-side workspace
|
||||
virtual uint64_t get_device_workspace_size(
|
||||
void const *configuration_ptr) const {
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = construct_arguments_(
|
||||
args,
|
||||
static_cast<GemmPlanarComplexConfiguration const *>(configuration_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
uint64_t size = Operator::get_workspace_size(args);
|
||||
|
||||
return size;
|
||||
}
|
||||
|
||||
/// Initializes the workspace
|
||||
virtual Status initialize(
|
||||
void const *configuration_ptr,
|
||||
void *host_workspace,
|
||||
void *device_workspace,
|
||||
cudaStream_t stream = nullptr) const {
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = construct_arguments_(
|
||||
args,
|
||||
static_cast<GemmPlanarComplexConfiguration const *>(configuration_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
Operator *op = new (host_workspace) Operator;
|
||||
|
||||
status = op->initialize(args, device_workspace, stream);
|
||||
|
||||
return status;
|
||||
}
|
||||
|
||||
/// Runs the kernel
|
||||
virtual Status run(
|
||||
void const *arguments_ptr,
|
||||
void *host_workspace,
|
||||
void *device_workspace = nullptr,
|
||||
cudaStream_t stream = nullptr) const {
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = update_arguments_(
|
||||
args,
|
||||
static_cast<GemmPlanarComplexArguments const *>(arguments_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
Operator *op = static_cast<Operator *>(host_workspace);
|
||||
|
||||
status = op->update(args, device_workspace);
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
status = op->run(stream);
|
||||
|
||||
return status;
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Operator_>
|
||||
class GemmPlanarComplexArrayOperation : public GemmOperationBase<Operator_> {
|
||||
public:
|
||||
|
||||
using Operator = Operator_;
|
||||
using ElementA = typename Operator::ElementA;
|
||||
using LayoutA = typename Operator::LayoutA;
|
||||
using ElementB = typename Operator::ElementB;
|
||||
using LayoutB = typename Operator::LayoutB;
|
||||
using ElementC = typename Operator::ElementC;
|
||||
using LayoutC = typename Operator::LayoutC;
|
||||
using ElementAccumulator = typename Operator::ElementAccumulator;
|
||||
using ElementCompute = typename Operator::EpilogueOutputOp::ElementCompute;
|
||||
|
||||
using OperatorArguments = typename Operator::Arguments;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructor
|
||||
GemmPlanarComplexArrayOperation(char const *name = "unknown_gemm"): GemmOperationBase<Operator_>(name) {
|
||||
|
||||
this->description_.gemm_kind = GemmKind::kPlanarComplexArray;
|
||||
}
|
||||
|
||||
protected:
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status construct_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
GemmPlanarComplexArrayConfiguration const *configuration) {
|
||||
|
||||
operator_args.mode = cutlass::gemm::GemmUniversalMode::kArray;
|
||||
operator_args.problem_size = configuration->problem_size;
|
||||
operator_args.batch_count = configuration->batch_count;
|
||||
|
||||
operator_args.lda_real = int(configuration->lda_real);
|
||||
operator_args.lda_imag = int(configuration->lda_imag);
|
||||
operator_args.ldb_real = int(configuration->ldb_real);
|
||||
operator_args.ldb_imag = int(configuration->ldb_imag);
|
||||
operator_args.ldc_real = int(configuration->ldc_real);
|
||||
operator_args.ldc_imag = int(configuration->ldc_imag);
|
||||
operator_args.ldd_real = int(configuration->ldd_real);
|
||||
operator_args.ldd_imag = int(configuration->ldd_imag);
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status update_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
GemmPlanarComplexArrayArguments const *arguments) {
|
||||
|
||||
if (arguments->pointer_mode == ScalarPointerMode::kHost) {
|
||||
typename Operator::EpilogueOutputOp::Params params(
|
||||
*static_cast<cutlass::complex<ElementCompute> const *>(arguments->alpha),
|
||||
*static_cast<cutlass::complex<ElementCompute> const *>(arguments->beta)
|
||||
);
|
||||
operator_args.epilogue = params;
|
||||
}
|
||||
else if (arguments->pointer_mode == ScalarPointerMode::kDevice){
|
||||
typename Operator::EpilogueOutputOp::Params params(
|
||||
static_cast<cutlass::complex<ElementCompute> const *>(arguments->alpha),
|
||||
static_cast<cutlass::complex<ElementCompute> const *>(arguments->beta)
|
||||
);
|
||||
operator_args.epilogue = params;
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
// update arguments
|
||||
operator_args.ptr_A_real = arguments->A_real;
|
||||
operator_args.ptr_A_imag = arguments->A_imag;
|
||||
operator_args.ptr_B_real = arguments->B_real;
|
||||
operator_args.ptr_B_imag = arguments->B_imag;
|
||||
operator_args.ptr_C_real = arguments->C_real;
|
||||
operator_args.ptr_C_imag = arguments->C_imag;
|
||||
operator_args.ptr_D_real = arguments->D_real;
|
||||
operator_args.ptr_D_imag = arguments->D_imag;
|
||||
|
||||
operator_args.ptr_M = arguments->M;
|
||||
operator_args.ptr_N = arguments->N;
|
||||
operator_args.ptr_K = arguments->K;
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
public:
|
||||
|
||||
/// Returns success if the operation can proceed
|
||||
virtual Status can_implement(
|
||||
void const *configuration_ptr,
|
||||
void const *arguments_ptr) const {
|
||||
|
||||
GemmPlanarComplexArrayConfiguration const *configuration =
|
||||
static_cast<GemmPlanarComplexArrayConfiguration const *>(configuration_ptr);
|
||||
|
||||
GemmPlanarComplexArrayArguments const *arguments =
|
||||
static_cast<GemmPlanarComplexArrayArguments const *>(arguments_ptr);
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = construct_arguments_(args, configuration);
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
status = update_arguments_(args, arguments);
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
return Operator::can_implement(args);
|
||||
}
|
||||
|
||||
/// Gets the host-side workspace
|
||||
virtual uint64_t get_host_workspace_size(
|
||||
void const *configuration) const {
|
||||
|
||||
return sizeof(Operator);
|
||||
}
|
||||
|
||||
/// Gets the device-side workspace
|
||||
virtual uint64_t get_device_workspace_size(
|
||||
void const *configuration_ptr) const {
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = construct_arguments_(
|
||||
args,
|
||||
static_cast<GemmPlanarComplexArrayConfiguration const *>(configuration_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
uint64_t size = Operator::get_workspace_size(args);
|
||||
|
||||
return size;
|
||||
}
|
||||
|
||||
/// Initializes the workspace
|
||||
virtual Status initialize(
|
||||
void const *configuration_ptr,
|
||||
void *host_workspace,
|
||||
void *device_workspace,
|
||||
cudaStream_t stream = nullptr) const {
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = construct_arguments_(
|
||||
args,
|
||||
static_cast<GemmPlanarComplexArrayConfiguration const *>(configuration_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
Operator *op = new (host_workspace) Operator;
|
||||
|
||||
status = op->initialize(args, device_workspace, stream);
|
||||
|
||||
return status;
|
||||
}
|
||||
|
||||
/// Runs the kernel
|
||||
virtual Status run(
|
||||
void const *arguments_ptr,
|
||||
void *host_workspace,
|
||||
void *device_workspace = nullptr,
|
||||
cudaStream_t stream = nullptr) const {
|
||||
|
||||
OperatorArguments args;
|
||||
|
||||
Status status = update_arguments_(
|
||||
args,
|
||||
static_cast<GemmPlanarComplexArrayArguments const *>(arguments_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
Operator *op = static_cast<Operator *>(host_workspace);
|
||||
|
||||
status = op->update(args, device_workspace);
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
status = op->run(stream);
|
||||
|
||||
return status;
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace library
|
||||
|
||||
@@ -0,0 +1,845 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice, this list of
|
||||
* conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright notice, this list of
|
||||
* conditions and the following disclaimer in the documentation and/or other materials
|
||||
* provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
|
||||
* to endorse or promote products derived from this software without specific prior written
|
||||
* permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
|
||||
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
|
||||
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
|
||||
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
|
||||
* STRICT LIABILITY, OR TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
/*! \file
|
||||
\brief CUTLASS Library handle.
|
||||
*/
|
||||
|
||||
#include <stdexcept>
|
||||
#include <cstdint>
|
||||
|
||||
#include "cutlass/library/handle.h"
|
||||
#include "cutlass/library/singleton.h"
|
||||
#include "cutlass/library/util.h"
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Constructor
|
||||
Handle::Handle(
|
||||
cudaStream_t stream,
|
||||
size_t workspace_size
|
||||
):
|
||||
stream_(stream),
|
||||
workspace_(nullptr),
|
||||
workspace_size_(0),
|
||||
scalar_pointer_mode_(ScalarPointerMode::kHost),
|
||||
last_operation_(nullptr) {
|
||||
|
||||
int device_idx = -1;
|
||||
|
||||
cudaError_t error = cudaGetDevice(&device_idx);
|
||||
if (error != cudaSuccess) {
|
||||
throw std::runtime_error("cudaGetDevice() failed");
|
||||
}
|
||||
|
||||
error = cudaGetDeviceProperties(&device_, device_idx);
|
||||
if (error != cudaSuccess) {
|
||||
throw std::runtime_error("cudaGetDeviceProperties() failed");
|
||||
}
|
||||
|
||||
set_workspace_size(workspace_size);
|
||||
|
||||
Singleton::get();
|
||||
}
|
||||
|
||||
/// Destructor
|
||||
Handle::~Handle() {
|
||||
if (workspace_) {
|
||||
|
||||
if (workspace_) {
|
||||
cudaFree(workspace_);
|
||||
}
|
||||
|
||||
workspace_ = nullptr;
|
||||
workspace_size_ = 0;
|
||||
}
|
||||
}
|
||||
|
||||
/// Move constructor
|
||||
Handle::Handle(Handle && handle) {
|
||||
device_ = handle.device_;
|
||||
workspace_size_ = handle.workspace_size_;
|
||||
workspace_ = handle.workspace_;
|
||||
stream_ = handle.stream_;
|
||||
scalar_pointer_mode_ = handle.scalar_pointer_mode_;
|
||||
|
||||
handle.workspace_ = nullptr;
|
||||
handle.workspace_size_ = 0;
|
||||
}
|
||||
|
||||
/// Move assignment operator
|
||||
Handle & Handle::operator=(Handle && handle) {
|
||||
|
||||
device_ = handle.device_;
|
||||
workspace_size_ = handle.workspace_size_;
|
||||
workspace_ = handle.workspace_;
|
||||
stream_ = handle.stream_;
|
||||
scalar_pointer_mode_ = handle.scalar_pointer_mode_;
|
||||
|
||||
handle.workspace_ = nullptr;
|
||||
handle.workspace_size_ = 0;
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
int Handle::compute_capability() const {
|
||||
return device_.major * 10 + device_.minor;
|
||||
}
|
||||
|
||||
/// Sets the current CUDA stream
|
||||
void Handle::set_stream(cudaStream_t stream) {
|
||||
stream_ = stream;
|
||||
}
|
||||
|
||||
/// Gets the current CUDA stream
|
||||
cudaStream_t Handle::get_stream() const {
|
||||
return stream_;
|
||||
}
|
||||
|
||||
/// Gets the device workspace size
|
||||
size_t Handle::get_workspace_size() const {
|
||||
return workspace_size_;
|
||||
}
|
||||
|
||||
/// Gets a pointer to the device workspace allocation in Global Memory
|
||||
void *Handle::get_workspace() const {
|
||||
return workspace_;
|
||||
}
|
||||
|
||||
/// Sets the size of device workspace, invalidating previous calls to get_device_workspace()
|
||||
void Handle::set_workspace_size(size_t bytes) {
|
||||
if (bytes != workspace_size_) {
|
||||
|
||||
if (workspace_) {
|
||||
cudaFree(workspace_);
|
||||
}
|
||||
|
||||
workspace_ = nullptr;
|
||||
workspace_size_ = bytes;
|
||||
|
||||
if (workspace_size_) {
|
||||
|
||||
cudaError_t error = cudaMalloc((void **)&workspace_, workspace_size_);
|
||||
|
||||
if (error != cudaSuccess) {
|
||||
throw std::runtime_error("Failed to allocate workspace");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (workspace_) {
|
||||
cudaError_t error = cudaMemset(workspace_, 0, workspace_size_);
|
||||
|
||||
if (error != cudaSuccess) {
|
||||
throw std::runtime_error("Failed to clear workspace");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Gets the scalar pointer mode
|
||||
ScalarPointerMode Handle::get_scalar_pointer_mode() const {
|
||||
return scalar_pointer_mode_;
|
||||
}
|
||||
|
||||
/// Sets the scalar pointer mode
|
||||
void Handle::set_scalar_pointer_mode(ScalarPointerMode mode) {
|
||||
scalar_pointer_mode_ = mode;
|
||||
}
|
||||
|
||||
/// Gets the last operation
|
||||
Operation const *Handle::get_last_operation() const {
|
||||
return last_operation_;
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Returns the maximum required alignment for each operator
|
||||
static int maximum_alignment_requirement(GemmDescription const &desc) {
|
||||
return std::max(
|
||||
std::max(desc.A.alignment, desc.B.alignment), desc.C.alignment);
|
||||
}
|
||||
|
||||
/// Returns the largest alignment (in units of elements) the problem satisfies, starting from a
|
||||
/// given upper limit.
|
||||
static int gemm_problem_alignment(
|
||||
int M,
|
||||
int N,
|
||||
int K,
|
||||
NumericTypeID element_A,
|
||||
void const *ptr_A,
|
||||
int lda,
|
||||
int64_t batch_stride_A,
|
||||
NumericTypeID element_B,
|
||||
void const *ptr_B,
|
||||
int ldb,
|
||||
int64_t batch_stride_B,
|
||||
NumericTypeID element_C,
|
||||
void const * ptr_C,
|
||||
int ldc,
|
||||
int64_t batch_stride_C,
|
||||
void const * ptr_D,
|
||||
int ldd,
|
||||
int64_t batch_stride_D,
|
||||
int max_alignment_in_bytes = 16
|
||||
) {
|
||||
|
||||
void const *pointers[] = {
|
||||
ptr_A, ptr_B, ptr_C, ptr_D
|
||||
};
|
||||
|
||||
int64_t extents[] = {
|
||||
M, N, K, lda, ldb, ldc, ldd, batch_stride_A, batch_stride_B, batch_stride_C, batch_stride_D
|
||||
};
|
||||
|
||||
NumericTypeID elements[] = {
|
||||
element_A, element_B, element_C
|
||||
};
|
||||
|
||||
for (; max_alignment_in_bytes > 0; max_alignment_in_bytes /= 2) {
|
||||
|
||||
bool satisfied = true;
|
||||
|
||||
// Can pointers satisfy this?
|
||||
for (void const *ptr : pointers) {
|
||||
std::uintptr_t int_ptr = reinterpret_cast<std::uintptr_t>(ptr);
|
||||
|
||||
if (int_ptr % max_alignment_in_bytes) {
|
||||
satisfied = false;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if (!satisfied) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Compute the maximum alignment based on element data types
|
||||
int max_element_alignment = 0;
|
||||
|
||||
for (NumericTypeID type_id : elements) {
|
||||
int element_alignment = max_alignment_in_bytes * 8 / library::sizeof_bits(type_id);
|
||||
max_element_alignment = std::max(max_element_alignment, element_alignment);
|
||||
}
|
||||
|
||||
// Can the problem size and leading dimensions satisfy this?
|
||||
for (int64_t extent : extents) {
|
||||
if (extent % max_element_alignment) {
|
||||
satisfied = false;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if (!satisfied) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Yes
|
||||
return max_element_alignment;
|
||||
}
|
||||
|
||||
// No alignment satisfies this problem
|
||||
return 0;
|
||||
}
|
||||
|
||||
/// Find the best kernel in descending order of preference.
|
||||
static Operation const * find_gemm_operation(
|
||||
GemmOperationFunctionalMap::const_iterator operators_it,
|
||||
GemmPreferenceKey const preference_key) {
|
||||
|
||||
auto cc_it = operators_it->second.upper_bound(preference_key);
|
||||
|
||||
if (cc_it == operators_it->second.begin()) {
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
Operation const *operation = nullptr;
|
||||
|
||||
// Search in descending order of compute capability
|
||||
do {
|
||||
--cc_it;
|
||||
|
||||
// Search tile sizes in order, for now.
|
||||
for (auto const * op : cc_it->second) {
|
||||
|
||||
GemmDescription const &desc = static_cast<GemmDescription const &>(op->description());
|
||||
|
||||
int min_cc = desc.tile_description.minimum_compute_capability;
|
||||
int max_cc = desc.tile_description.maximum_compute_capability;
|
||||
|
||||
int op_alignment = maximum_alignment_requirement(desc);
|
||||
|
||||
if ((min_cc <= preference_key.compute_capability) &&
|
||||
(preference_key.compute_capability <= max_cc) &&
|
||||
(op_alignment <= preference_key.alignment)) {
|
||||
|
||||
operation = op;
|
||||
break;
|
||||
}
|
||||
}
|
||||
} while (!operation && cc_it != operators_it->second.begin());
|
||||
|
||||
return operation;
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Executes a GEMM computation: D <= alpha * A*B + beta * C
|
||||
Status Handle::gemm(
|
||||
|
||||
int M, /// GEMM M dimension
|
||||
int N, /// GEMM N dimension
|
||||
int K, /// GEMM K dimension
|
||||
|
||||
NumericTypeID element_compute, /// Data type of internal accumulation
|
||||
|
||||
NumericTypeID element_scalar, /// Data type of alpha/beta scalars
|
||||
|
||||
void const *alpha, /// Pointer to alpha scalar
|
||||
|
||||
NumericTypeID element_A, /// Data type of A matrix elements
|
||||
LayoutTypeID layout_A, /// Layout of A matrix
|
||||
ComplexTransform transform_A, /// Complex transformation applied to A matrix - ignored for real-valued matrices
|
||||
|
||||
void const * ptr_A, /// Pointer to A matrix in Global Memory
|
||||
int lda, /// Leading dimension of A matrix
|
||||
|
||||
NumericTypeID element_B, /// Data type of B matrix elements
|
||||
LayoutTypeID layout_B, /// Layout of B matrix
|
||||
ComplexTransform transform_B, /// Complex transformation applied to B matrix - ignored for real-valued matrices
|
||||
|
||||
void const * ptr_B, /// Pointer to B matrix in Global Memory
|
||||
int ldb, /// Leading dimension of B matrix
|
||||
|
||||
void const * beta, /// Pointer to beta scalar
|
||||
|
||||
NumericTypeID element_C, /// Data type of C and D matrices
|
||||
|
||||
void const * ptr_C, /// Pointer to C matrix
|
||||
int ldc, /// Leading dimension of C matrix
|
||||
|
||||
void * ptr_D, /// Pointer to D matrix
|
||||
int ldd /// Leading dimension of D matrix
|
||||
) {
|
||||
|
||||
//
|
||||
// Find the operation
|
||||
//
|
||||
|
||||
GemmFunctionalKey key(
|
||||
element_compute,
|
||||
element_scalar,
|
||||
element_A,
|
||||
layout_A,
|
||||
transform_A,
|
||||
element_B,
|
||||
layout_B,
|
||||
transform_B,
|
||||
element_C
|
||||
);
|
||||
|
||||
auto operators_it = Singleton::get().operation_table.gemm_operations.find(key);
|
||||
|
||||
if (operators_it == Singleton::get().operation_table.gemm_operations.end()) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
if (operators_it->second.empty()) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
//
|
||||
// Compute the largest alignment restriction the kernel can satisfy.
|
||||
//
|
||||
|
||||
// Maximum alignment expectation among all kernels (in units of bytes)
|
||||
int const kMaximumAlignmentSize = 16;
|
||||
|
||||
int alignment = gemm_problem_alignment(
|
||||
M, N, K,
|
||||
element_A, ptr_A, lda, 0,
|
||||
element_B, ptr_B, ldb, 0,
|
||||
element_C, ptr_C, ldc, 0,
|
||||
ptr_D, ldd, 0, kMaximumAlignmentSize
|
||||
);
|
||||
|
||||
//
|
||||
// Find the best kernel in descending order of preference.
|
||||
//
|
||||
|
||||
GemmPreferenceKey preference_key(compute_capability(), alignment);
|
||||
|
||||
Operation const *operation = find_gemm_operation(operators_it, preference_key);
|
||||
|
||||
if (!operation) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
last_operation_ = operation;
|
||||
|
||||
//
|
||||
// Configure operation
|
||||
//
|
||||
|
||||
GemmConfiguration configuration{
|
||||
{M, N, K},
|
||||
lda,
|
||||
ldb,
|
||||
ldc,
|
||||
ldd,
|
||||
1
|
||||
};
|
||||
|
||||
// Query host work space size
|
||||
uint64_t host_workspace_size_needed = operation->get_host_workspace_size(&configuration);
|
||||
|
||||
if (uint64_t(kHostWorkspaceSize) < host_workspace_size_needed) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
char host_workspace[kHostWorkspaceSize];
|
||||
|
||||
// Query device workspace size
|
||||
uint64_t device_workspace_size_needed = operation->get_device_workspace_size(&configuration);
|
||||
|
||||
if (uint64_t(workspace_size_) < device_workspace_size_needed) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
// Initialize host and device workspaces
|
||||
Status status = operation->initialize(
|
||||
&configuration,
|
||||
host_workspace,
|
||||
workspace_,
|
||||
stream_);
|
||||
|
||||
if (status != cutlass::Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
// Run the operator
|
||||
GemmArguments arguments{
|
||||
ptr_A,
|
||||
ptr_B,
|
||||
ptr_C,
|
||||
ptr_D,
|
||||
alpha,
|
||||
beta,
|
||||
scalar_pointer_mode_
|
||||
};
|
||||
|
||||
return operation->run(&arguments, host_workspace, workspace_, stream_);
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Planar complex GEMM
|
||||
Status Handle::gemm_planar_complex(
|
||||
|
||||
int M, /// GEMM M dimension
|
||||
int N, /// GEMM N dimension
|
||||
int K, /// GEMM K dimension
|
||||
|
||||
NumericTypeID element_compute, /// Data type of internal accumulation
|
||||
|
||||
NumericTypeID element_scalar, /// Data type of alpha/beta scalars
|
||||
|
||||
void const *alpha, /// Pointer to alpha scalar
|
||||
|
||||
NumericTypeID element_A, /// Data type of A matrix elements
|
||||
LayoutTypeID layout_A, /// Layout of A matrix
|
||||
ComplexTransform transform_A, /// Complex transformation applied to A matrix
|
||||
|
||||
void const * ptr_A_real, /// Pointer to real part of A matrix
|
||||
void const * ptr_A_imag, /// Pointer to imaginary part of A matrix
|
||||
int lda_real, /// Leading dimension of real part of A matrix
|
||||
int lda_imag, /// Leading dimension of imaginary part of A matrix
|
||||
|
||||
NumericTypeID element_B, /// Data type of B matrix elements
|
||||
LayoutTypeID layout_B, /// Layout of B matrix
|
||||
ComplexTransform transform_B, /// Complex transformation applied to B matrix
|
||||
|
||||
void const * ptr_B_real, /// Pointer to real part of B matrix
|
||||
void const * ptr_B_imag, /// Pointer to imaginary part of B matrix
|
||||
int ldb_real, /// Leading dimension of real part of B matrix
|
||||
int ldb_imag, /// Leading dimension of imaginary part of B matrix
|
||||
|
||||
void const * beta, /// Pointer to beta scalar
|
||||
|
||||
NumericTypeID element_C, /// Data type of C and D matrix
|
||||
|
||||
void const * ptr_C_real, /// Pointer to real part of C matrix
|
||||
void const * ptr_C_imag, /// Pointer to imaginary part of C matrix
|
||||
int ldc_real, /// Leading dimension of real part of C matrix
|
||||
int ldc_imag, /// Leading dimension of imaginary part of C matrix
|
||||
|
||||
void * ptr_D_real, /// Pointer to real part of D matrix
|
||||
void * ptr_D_imag, /// Pointer to imaginary part of D matrix
|
||||
int ldd_real, /// Leading dimension of real part of D matrix
|
||||
int ldd_imag, /// Leading dimension of imaginary part of D matrix
|
||||
|
||||
int batch_count, /// Number of batched GEMMs to execute
|
||||
|
||||
int64_t batch_stride_A_real,
|
||||
int64_t batch_stride_A_imag,
|
||||
|
||||
int64_t batch_stride_B_real,
|
||||
int64_t batch_stride_B_imag,
|
||||
|
||||
int64_t batch_stride_C_real,
|
||||
int64_t batch_stride_C_imag,
|
||||
|
||||
int64_t batch_stride_D_real,
|
||||
int64_t batch_stride_D_imag
|
||||
) {
|
||||
|
||||
//
|
||||
// Find the operation
|
||||
//
|
||||
|
||||
GemmFunctionalKey key(
|
||||
element_compute,
|
||||
element_scalar,
|
||||
element_A,
|
||||
layout_A,
|
||||
transform_A,
|
||||
element_B,
|
||||
layout_B,
|
||||
transform_B,
|
||||
element_C
|
||||
);
|
||||
|
||||
auto operators_it = Singleton::get().operation_table.gemm_planar_complex_operations.find(key);
|
||||
|
||||
if (operators_it == Singleton::get().operation_table.gemm_planar_complex_operations.end()) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
if (operators_it->second.empty()) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
//
|
||||
// Compute the largest alignment restriction the kernel can satisfy.
|
||||
//
|
||||
|
||||
// Maximum alignment expectation among all kernels (in units of bytes)
|
||||
int const kMaximumAlignmentSize = 16;
|
||||
|
||||
int alignment = std::max(
|
||||
gemm_problem_alignment(
|
||||
M, N, K,
|
||||
element_A, ptr_A_real, lda_real, batch_stride_A_real,
|
||||
element_B, ptr_B_real, ldb_real, batch_stride_B_real,
|
||||
element_C, ptr_C_real, ldc_real, batch_stride_C_real,
|
||||
ptr_D_real, ldd_real, batch_stride_D_real, kMaximumAlignmentSize
|
||||
),
|
||||
gemm_problem_alignment(
|
||||
M, N, K,
|
||||
element_A, ptr_A_imag, lda_imag, batch_stride_A_imag,
|
||||
element_B, ptr_B_imag, ldb_imag, batch_stride_B_imag,
|
||||
element_C, ptr_C_imag, ldc_imag, batch_stride_C_imag,
|
||||
ptr_D_imag, ldd_imag, batch_stride_D_imag, kMaximumAlignmentSize
|
||||
)
|
||||
);
|
||||
|
||||
//
|
||||
// Find the best kernel in descending order of preference.
|
||||
//
|
||||
|
||||
GemmPreferenceKey preference_key(compute_capability(), alignment);
|
||||
|
||||
Operation const *operation = find_gemm_operation(operators_it, preference_key);
|
||||
|
||||
if (!operation) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
last_operation_ = operation;
|
||||
|
||||
//
|
||||
// Configure operation
|
||||
//
|
||||
|
||||
GemmPlanarComplexConfiguration configuration{
|
||||
GemmUniversalMode::kBatched,
|
||||
{M, N, K},
|
||||
batch_count,
|
||||
lda_real,
|
||||
lda_imag,
|
||||
ldb_real,
|
||||
ldb_imag,
|
||||
ldc_real,
|
||||
ldc_imag,
|
||||
ldd_real,
|
||||
ldd_imag
|
||||
};
|
||||
|
||||
// Query host work space size
|
||||
uint64_t host_workspace_size_needed = operation->get_host_workspace_size(&configuration);
|
||||
|
||||
if (uint64_t(kHostWorkspaceSize) < host_workspace_size_needed) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
char host_workspace[kHostWorkspaceSize];
|
||||
|
||||
// Query device workspace size
|
||||
uint64_t device_workspace_size_needed = operation->get_device_workspace_size(&configuration);
|
||||
|
||||
if (uint64_t(workspace_size_) < device_workspace_size_needed) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
// Initialize host and device workspaces
|
||||
Status status = operation->initialize(
|
||||
&configuration,
|
||||
host_workspace,
|
||||
workspace_,
|
||||
stream_);
|
||||
|
||||
if (status != cutlass::Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
// Run the operator
|
||||
GemmPlanarComplexArguments arguments{
|
||||
ptr_A_real,
|
||||
ptr_A_imag,
|
||||
ptr_B_real,
|
||||
ptr_B_imag,
|
||||
ptr_C_real,
|
||||
ptr_C_imag,
|
||||
ptr_D_real,
|
||||
ptr_D_imag,
|
||||
alpha,
|
||||
beta,
|
||||
scalar_pointer_mode_,
|
||||
batch_stride_A_real,
|
||||
batch_stride_A_imag,
|
||||
batch_stride_B_real,
|
||||
batch_stride_B_imag,
|
||||
batch_stride_C_real,
|
||||
batch_stride_C_imag,
|
||||
batch_stride_D_real,
|
||||
batch_stride_D_imag
|
||||
};
|
||||
|
||||
return operation->run(&arguments, host_workspace, workspace_, stream_);
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Planar complex batched GEMM loading pointers from arrays in global memory
|
||||
Status Handle::gemm_planar_complex_array(
|
||||
|
||||
int expected_M, /// Expected GEMM M dimension (used for sizing CUDA grid)
|
||||
int expected_N, /// Expected GEMM N dimension (used for sizing CUDA grid)
|
||||
int expected_K, /// Expected GEMM K dimension
|
||||
int batch_count, /// Number of independent GEMM computations to execute
|
||||
|
||||
int const *M, /// Array containing the GEMM M dimension for each batch index
|
||||
int const *N, /// Array containing the GEMM N dimension for each batch index
|
||||
int const *K, /// Array containing the GEMM K dimension for each batch index
|
||||
|
||||
NumericTypeID element_compute, /// Data type of internal accumulation
|
||||
|
||||
NumericTypeID element_scalar, /// Data type of alpha/beta scalars
|
||||
|
||||
void const *alpha, /// Pointer to alpha scalar
|
||||
|
||||
NumericTypeID element_A, /// Data type of A matrix elements
|
||||
LayoutTypeID layout_A, /// Layout of A matrix
|
||||
ComplexTransform transform_A, /// Complex transformation applied to A matrix
|
||||
|
||||
void const * const * ptr_A_real, /// Pointer to array containing pointers to real part of A matrices
|
||||
void const * const * ptr_A_imag, /// Pointer to array containing pointers to imaginary part of A matrices
|
||||
|
||||
int lda_real, /// Leading dimension of real part of A matrix
|
||||
int lda_imag, /// Leading dimension of imaginary part of A matrix
|
||||
|
||||
NumericTypeID element_B, /// Data type of B matrix elements
|
||||
LayoutTypeID layout_B, /// Layout of B matrix
|
||||
ComplexTransform transform_B, /// Complex transformation applied to B matrix
|
||||
|
||||
void const * const * ptr_B_real, /// Pointer to array containing pointers to real part of B matrices
|
||||
void const * const * ptr_B_imag, /// Pointer to array containing pointers to imaginary part of B matrices
|
||||
|
||||
int ldb_real, /// Leading dimension of real part of B matrix
|
||||
int ldb_imag, /// Leading dimension of imaginary part of B matrix
|
||||
|
||||
void const * beta, /// Pointer to beta scalar
|
||||
|
||||
NumericTypeID element_C, /// Data type of C and D matrix
|
||||
|
||||
void const * const * ptr_C_real, /// Pointer to array containing pointers to real part of C matrices
|
||||
void const * const * ptr_C_imag, /// Pointer to array containing poitners to imaginary part of C matrices
|
||||
|
||||
int ldc_real, /// Leading dimension of real part of C matrix
|
||||
int ldc_imag, /// Leading dimension of imaginary part of C matrix
|
||||
|
||||
void * const * ptr_D_real, /// Pointer to array containing pointers to real part of D matrices
|
||||
void * const * ptr_D_imag, /// Pointer to array containing poitners to imaginary part of D matrices
|
||||
|
||||
int ldd_real, /// Leading dimension of real part of D matrix
|
||||
int ldd_imag /// Leading dimension of imaginary part of D matrix
|
||||
) {
|
||||
|
||||
//
|
||||
// Find the operation
|
||||
//
|
||||
|
||||
GemmFunctionalKey key(
|
||||
element_compute,
|
||||
element_scalar,
|
||||
element_A,
|
||||
layout_A,
|
||||
transform_A,
|
||||
element_B,
|
||||
layout_B,
|
||||
transform_B,
|
||||
element_C
|
||||
);
|
||||
|
||||
auto operators_it = Singleton::get().operation_table.gemm_planar_complex_array_operations.find(key);
|
||||
|
||||
if (operators_it == Singleton::get().operation_table.gemm_planar_complex_array_operations.end()) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
if (operators_it->second.empty()) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
//
|
||||
// Compute the largest alignment restriction the kernel can satisfy.
|
||||
//
|
||||
|
||||
// Maximum alignment expectation among all kernels (in units of bytes)
|
||||
int const kMaximumAlignmentSize = 16;
|
||||
|
||||
int alignment = std::max(
|
||||
gemm_problem_alignment(
|
||||
expected_M, expected_N, expected_K,
|
||||
element_A, nullptr, lda_real, 0,
|
||||
element_B, nullptr, ldb_real, 0,
|
||||
element_C, nullptr, ldc_real, 0,
|
||||
nullptr, ldd_real, 0, kMaximumAlignmentSize
|
||||
),
|
||||
gemm_problem_alignment(
|
||||
expected_M, expected_N, expected_K,
|
||||
element_A, nullptr, lda_imag, 0,
|
||||
element_B, nullptr, ldb_imag, 0,
|
||||
element_C, nullptr, ldc_imag, 0,
|
||||
nullptr, ldd_imag, 0, kMaximumAlignmentSize
|
||||
)
|
||||
);
|
||||
|
||||
//
|
||||
// Find the best kernel in descending order of preference.
|
||||
//
|
||||
|
||||
GemmPreferenceKey preference_key(compute_capability(), alignment);
|
||||
|
||||
Operation const *operation = find_gemm_operation(operators_it, preference_key);
|
||||
|
||||
if (!operation) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
last_operation_ = operation;
|
||||
|
||||
//
|
||||
// Configure operation
|
||||
//
|
||||
|
||||
GemmPlanarComplexArrayConfiguration configuration{
|
||||
{expected_M, expected_N, expected_K},
|
||||
batch_count,
|
||||
lda_real,
|
||||
lda_imag,
|
||||
ldb_real,
|
||||
ldb_imag,
|
||||
ldc_real,
|
||||
ldc_imag,
|
||||
ldd_real,
|
||||
ldd_imag
|
||||
};
|
||||
|
||||
// Query host work space size
|
||||
uint64_t host_workspace_size_needed = operation->get_host_workspace_size(&configuration);
|
||||
|
||||
if (uint64_t(kHostWorkspaceSize) < host_workspace_size_needed) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
char host_workspace[kHostWorkspaceSize];
|
||||
|
||||
// Query device workspace size
|
||||
uint64_t device_workspace_size_needed = operation->get_device_workspace_size(&configuration);
|
||||
|
||||
if (uint64_t(workspace_size_) < device_workspace_size_needed) {
|
||||
return cutlass::Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
// Initialize host and device workspaces
|
||||
Status status = operation->initialize(
|
||||
&configuration,
|
||||
host_workspace,
|
||||
workspace_,
|
||||
stream_);
|
||||
|
||||
if (status != cutlass::Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
// Run the operator
|
||||
GemmPlanarComplexArrayArguments arguments{
|
||||
M, N, K,
|
||||
ptr_A_real,
|
||||
ptr_A_imag,
|
||||
ptr_B_real,
|
||||
ptr_B_imag,
|
||||
ptr_C_real,
|
||||
ptr_C_imag,
|
||||
ptr_D_real,
|
||||
ptr_D_imag,
|
||||
alpha,
|
||||
beta,
|
||||
scalar_pointer_mode_
|
||||
};
|
||||
|
||||
return operation->run(&arguments, host_workspace, workspace_, stream_);
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace library
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -57,6 +57,10 @@ namespace library {
|
||||
|
||||
template <typename T> struct NumericTypeMap;
|
||||
|
||||
template <> struct NumericTypeMap<cutlass::uint1b_t> {
|
||||
static NumericTypeID const kId = NumericTypeID::kB1;
|
||||
};
|
||||
|
||||
template <> struct NumericTypeMap<cutlass::int4b_t> {
|
||||
static NumericTypeID const kId = NumericTypeID::kS4;
|
||||
};
|
||||
@@ -123,6 +127,28 @@ template <> struct NumericTypeMap<cutlass::complex<double> > {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename T> struct MathOperationMap {
|
||||
static MathOperationID const kId = MathOperationID::kInvalid;
|
||||
};
|
||||
|
||||
template <> struct MathOperationMap<cutlass::arch::OpMultiplyAdd> {
|
||||
static MathOperationID const kId = MathOperationID::kMultiplyAdd;
|
||||
};
|
||||
|
||||
template <> struct MathOperationMap<cutlass::arch::OpMultiplyAddSaturate> {
|
||||
static MathOperationID const kId = MathOperationID::kMultiplyAddSaturate;
|
||||
};
|
||||
|
||||
template <> struct MathOperationMap<cutlass::arch::OpMultiplyAddComplex> {
|
||||
static MathOperationID const kId = MathOperationID::kMultiplyAddComplex;
|
||||
};
|
||||
|
||||
template <> struct MathOperationMap<cutlass::arch::OpXorPopc> {
|
||||
static MathOperationID const kId = MathOperationID::kXorPopc;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename T> struct LayoutMap;
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::ColumnMajor> {
|
||||
@@ -133,6 +159,34 @@ template <> struct LayoutMap<cutlass::layout::RowMajor> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kRowMajor;
|
||||
};
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::ColumnMajorInterleaved<16>> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kColumnMajorInterleavedK16;
|
||||
};
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::RowMajorInterleaved<16>> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kRowMajorInterleavedK16;
|
||||
};
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::ColumnMajorInterleaved<32>> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kColumnMajorInterleavedK32;
|
||||
};
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::RowMajorInterleaved<32>> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kRowMajorInterleavedK32;
|
||||
};
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::ColumnMajorInterleaved<64>> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kColumnMajorInterleavedK64;
|
||||
};
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::RowMajorInterleaved<64>> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kRowMajorInterleavedK64;
|
||||
};
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::TensorNHWC> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kTensorNHWC;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename T> struct OpcodeClassMap;
|
||||
@@ -148,6 +202,19 @@ template <> struct OpcodeClassMap<arch::OpClassTensorOp> {
|
||||
template <> struct OpcodeClassMap<arch::OpClassWmmaTensorOp> {
|
||||
static OpcodeClassID const kId = OpcodeClassID::kWmmaTensorOp;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <cutlass::ComplexTransform Transform> struct ComplexTransformMap;
|
||||
|
||||
template <> struct ComplexTransformMap<cutlass::ComplexTransform::kNone> {
|
||||
static cutlass::library::ComplexTransform const kId = cutlass::library::ComplexTransform::kNone;
|
||||
};
|
||||
|
||||
template <> struct ComplexTransformMap<cutlass::ComplexTransform::kConjugate> {
|
||||
static cutlass::library::ComplexTransform const kId = cutlass::library::ComplexTransform::kConjugate;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename T> struct ArchMap;
|
||||
|
||||
@@ -1,6 +1,4 @@
|
||||
/*!
|
||||
|
||||
*//***************************************************************************************************
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
@@ -37,11 +35,12 @@
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//////////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
void initialize_all(Manifest &manifest);
|
||||
// init and insert all cutlass op in manifest object (procedurally generated using generator.py)
|
||||
void initialize_all(Manifest &manifest);
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Top-level initialization
|
||||
Status Manifest::initialize() {
|
||||
@@ -50,7 +49,13 @@ Status Manifest::initialize() {
|
||||
operations_.clear();
|
||||
}
|
||||
|
||||
initialize_all(*this);
|
||||
switch(provider_) {
|
||||
case Provider::kCUTLASS:
|
||||
initialize_all(*this); break;
|
||||
|
||||
default:
|
||||
break;
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,159 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice, this list of
|
||||
* conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright notice, this list of
|
||||
* conditions and the following disclaimer in the documentation and/or other materials
|
||||
* provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
|
||||
* to endorse or promote products derived from this software without specific prior written
|
||||
* permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
|
||||
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
|
||||
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
|
||||
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
|
||||
* STRICT LIABILITY, OR TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*
|
||||
\file
|
||||
\brief Defines a data structure in which a set of functionally equivalent library::Operation
|
||||
instances may be queried.
|
||||
*/
|
||||
|
||||
#include <fstream>
|
||||
|
||||
#include "cutlass/library/library.h"
|
||||
#include "cutlass/library/operation_table.h"
|
||||
#include "cutlass/library/util.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
std::ostream & operator<<(std::ostream &out, cutlass::library::GemmFunctionalKey const &k) {
|
||||
|
||||
out << "{\n"
|
||||
<< " element_compute: " << to_string(k.element_compute) << "\n"
|
||||
<< " element_scalar: " << to_string(k.element_scalar) << "\n"
|
||||
<< " element_A: " << to_string(k.element_A) << "\n"
|
||||
<< " layout_A: " << to_string(k.layout_A) << "\n"
|
||||
<< " transform_A: " << to_string(k.transform_A) << "\n"
|
||||
<< " element_B: " << to_string(k.element_B) << "\n"
|
||||
<< " layout_B: " << to_string(k.layout_B) << "\n"
|
||||
<< " transform_B: " << to_string(k.transform_B) << "\n"
|
||||
<< " element_C: " << to_string(k.element_C) << "\n"
|
||||
<< "}";
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
void OperationTable::append(Manifest const &manifest) {
|
||||
|
||||
// Insert operations into appropriate data structure
|
||||
for (auto const & operation : manifest) {
|
||||
|
||||
OperationDescription const &desc = operation->description();
|
||||
|
||||
if (desc.kind == OperationKind::kGemm) {
|
||||
GemmDescription const &gemm_desc = static_cast<GemmDescription const &>(desc);
|
||||
|
||||
if (gemm_desc.gemm_kind == GemmKind::kGemm) {
|
||||
|
||||
GemmFunctionalKey functional_key(
|
||||
gemm_desc.tile_description.math_instruction.element_accumulator,
|
||||
gemm_desc.element_epilogue,
|
||||
gemm_desc.A.element,
|
||||
gemm_desc.A.layout,
|
||||
gemm_desc.transform_A,
|
||||
gemm_desc.B.element,
|
||||
gemm_desc.B.layout,
|
||||
gemm_desc.transform_B,
|
||||
gemm_desc.C.element
|
||||
);
|
||||
|
||||
Operation const *op = operation.get();
|
||||
|
||||
int cc = gemm_desc.tile_description.minimum_compute_capability;
|
||||
|
||||
int alignment = std::max(std::max(
|
||||
gemm_desc.A.alignment, gemm_desc.B.alignment), gemm_desc.C.alignment);
|
||||
|
||||
GemmPreferenceKey preference_key(cc, alignment);
|
||||
|
||||
gemm_operations[functional_key][preference_key].push_back(op);
|
||||
}
|
||||
else if (gemm_desc.gemm_kind == GemmKind::kPlanarComplex) {
|
||||
|
||||
GemmFunctionalKey functional_key(
|
||||
gemm_desc.tile_description.math_instruction.element_accumulator,
|
||||
gemm_desc.element_epilogue,
|
||||
gemm_desc.A.element,
|
||||
gemm_desc.A.layout,
|
||||
gemm_desc.transform_A,
|
||||
gemm_desc.B.element,
|
||||
gemm_desc.B.layout,
|
||||
gemm_desc.transform_B,
|
||||
gemm_desc.C.element
|
||||
);
|
||||
|
||||
Operation const *op = operation.get();
|
||||
|
||||
int cc = gemm_desc.tile_description.minimum_compute_capability;
|
||||
|
||||
int alignment = std::max(std::max(
|
||||
gemm_desc.A.alignment, gemm_desc.B.alignment), gemm_desc.C.alignment);
|
||||
|
||||
GemmPreferenceKey preference_key(cc, alignment);
|
||||
|
||||
gemm_planar_complex_operations[functional_key][preference_key].push_back(op);
|
||||
}
|
||||
else if (gemm_desc.gemm_kind == GemmKind::kPlanarComplexArray) {
|
||||
|
||||
GemmFunctionalKey functional_key(
|
||||
gemm_desc.tile_description.math_instruction.element_accumulator,
|
||||
gemm_desc.element_epilogue,
|
||||
gemm_desc.A.element,
|
||||
gemm_desc.A.layout,
|
||||
gemm_desc.transform_A,
|
||||
gemm_desc.B.element,
|
||||
gemm_desc.B.layout,
|
||||
gemm_desc.transform_B,
|
||||
gemm_desc.C.element
|
||||
);
|
||||
|
||||
Operation const *op = operation.get();
|
||||
|
||||
int cc = gemm_desc.tile_description.minimum_compute_capability;
|
||||
|
||||
int alignment = std::max(std::max(
|
||||
gemm_desc.A.alignment, gemm_desc.B.alignment), gemm_desc.C.alignment);
|
||||
|
||||
GemmPreferenceKey preference_key(cc, alignment);
|
||||
|
||||
gemm_planar_complex_array_operations[functional_key][preference_key].push_back(op);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace library
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice, this list of
|
||||
* conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright notice, this list of
|
||||
* conditions and the following disclaimer in the documentation and/or other materials
|
||||
* provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
|
||||
* to endorse or promote products derived from this software without specific prior written
|
||||
* permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
|
||||
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
|
||||
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
|
||||
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
|
||||
* STRICT LIABILITY, OR TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <memory>
|
||||
#include "cutlass/library/library.h"
|
||||
#include "cutlass/library/manifest.h"
|
||||
#include "cutlass/library/operation_table.h"
|
||||
#include "cutlass/library/singleton.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
static std::unique_ptr<Singleton> instance;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
Singleton::Singleton() {
|
||||
|
||||
manifest.initialize();
|
||||
|
||||
operation_table.append(manifest);
|
||||
}
|
||||
|
||||
Singleton const & Singleton::get() {
|
||||
if (!instance.get()) {
|
||||
instance.reset(new Singleton);
|
||||
}
|
||||
return *instance.get();
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace library
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
@@ -25,17 +25,65 @@
|
||||
|
||||
#include <iosfwd>
|
||||
#include <complex>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/complex.h"
|
||||
|
||||
#include "cutlass/library/library.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
|
||||
#include "cutlass/library/library.h"
|
||||
#include "cutlass/library/util.h"
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
static struct {
|
||||
char const *text;
|
||||
char const *pretty;
|
||||
Provider enumerant;
|
||||
}
|
||||
Provider_enumerants[] = {
|
||||
{"cutlass", "CUTLASS", Provider::kCUTLASS},
|
||||
{"host", "reference_host", Provider::kReferenceHost},
|
||||
{"device", "reference_device", Provider::kReferenceDevice},
|
||||
{"cublas", "cuBLAS", Provider::kCUBLAS},
|
||||
};
|
||||
|
||||
/// Converts a Provider enumerant to a string
|
||||
char const *to_string(Provider provider, bool pretty) {
|
||||
|
||||
for (auto const & possible : Provider_enumerants) {
|
||||
if (provider == possible.enumerant) {
|
||||
if (pretty) {
|
||||
return possible.pretty;
|
||||
}
|
||||
else {
|
||||
return possible.text;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return pretty ? "Invalid" : "invalid";
|
||||
}
|
||||
|
||||
/// Parses a Provider enumerant from a string
|
||||
template <>
|
||||
Provider from_string<Provider>(std::string const &str) {
|
||||
|
||||
for (auto const & possible : Provider_enumerants) {
|
||||
if ((str.compare(possible.text) == 0) ||
|
||||
(str.compare(possible.pretty) == 0)) {
|
||||
return possible.enumerant;
|
||||
}
|
||||
}
|
||||
|
||||
return Provider::kInvalid;
|
||||
}
|
||||
|
||||
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
static struct {
|
||||
@@ -44,7 +92,7 @@ static struct {
|
||||
OperationKind enumerant;
|
||||
}
|
||||
OperationKind_enumerants[] = {
|
||||
{"gemm", "Gemm", OperationKind::kGemm},
|
||||
{"gemm", "Gemm", OperationKind::kGemm},
|
||||
};
|
||||
|
||||
/// Converts a Status enumerant to a string
|
||||
@@ -203,6 +251,9 @@ int sizeof_bits(NumericTypeID type) {
|
||||
case NumericTypeID::kF16: return 16;
|
||||
case NumericTypeID::kF32: return 32;
|
||||
case NumericTypeID::kF64: return 64;
|
||||
case NumericTypeID::kCF16: return 32;
|
||||
case NumericTypeID::kCF32: return 64;
|
||||
case NumericTypeID::kCF64: return 128;
|
||||
case NumericTypeID::kS4: return 4;
|
||||
case NumericTypeID::kS8: return 8;
|
||||
case NumericTypeID::kS16: return 16;
|
||||
@@ -291,6 +342,9 @@ bool is_float_type(NumericTypeID type) {
|
||||
case NumericTypeID::kF16: return true;
|
||||
case NumericTypeID::kF32: return true;
|
||||
case NumericTypeID::kF64: return true;
|
||||
case NumericTypeID::kCF16: return true;
|
||||
case NumericTypeID::kCF32: return true;
|
||||
case NumericTypeID::kCF64: return true;
|
||||
default: break;
|
||||
}
|
||||
return false;
|
||||
@@ -309,8 +363,18 @@ layout_aliases[] = {
|
||||
{LayoutTypeID::kColumnMajor, "column"},
|
||||
{LayoutTypeID::kColumnMajor, "col"},
|
||||
{LayoutTypeID::kColumnMajor, "n"},
|
||||
|
||||
{LayoutTypeID::kColumnMajorInterleavedK16, "nk16"},
|
||||
{LayoutTypeID::kRowMajorInterleavedK16, "tk16"},
|
||||
|
||||
{LayoutTypeID::kColumnMajorInterleavedK32, "nk32"},
|
||||
{LayoutTypeID::kRowMajorInterleavedK32, "tk32"},
|
||||
|
||||
{LayoutTypeID::kColumnMajorInterleavedK64, "nk64"},
|
||||
{LayoutTypeID::kRowMajorInterleavedK64, "tk64"},
|
||||
|
||||
{LayoutTypeID::kTensorNCHW, "nchw"},
|
||||
{LayoutTypeID::kTensorNHWC, "packed_nhwc"},
|
||||
{LayoutTypeID::kTensorNHWC, "nhwc"},
|
||||
{LayoutTypeID::kUnknown, "*"},
|
||||
{LayoutTypeID::kInvalid, nullptr}
|
||||
};
|
||||
@@ -344,7 +408,12 @@ int get_layout_stride_rank(LayoutTypeID layout_id) {
|
||||
case LayoutTypeID::kColumnMajorInterleavedK4:
|
||||
case LayoutTypeID::kRowMajorInterleavedK4:
|
||||
case LayoutTypeID::kColumnMajorInterleavedK16:
|
||||
case LayoutTypeID::kRowMajorInterleavedK16: return 1;
|
||||
case LayoutTypeID::kRowMajorInterleavedK16:
|
||||
case LayoutTypeID::kColumnMajorInterleavedK32:
|
||||
case LayoutTypeID::kRowMajorInterleavedK32:
|
||||
case LayoutTypeID::kColumnMajorInterleavedK64:
|
||||
case LayoutTypeID::kRowMajorInterleavedK64:
|
||||
return 1;
|
||||
case LayoutTypeID::kTensorNCHW:
|
||||
case LayoutTypeID::kTensorNHWC: return 3;
|
||||
default : throw std::runtime_error("Unsupported LayoutTypeID in LayoutType::get_stride_rank");
|
||||
@@ -396,8 +465,51 @@ OpcodeClassID from_string<OpcodeClassID>(std::string const &str) {
|
||||
return OpcodeClassID::kInvalid;
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
static struct {
|
||||
char const *text;
|
||||
char const *pretty;
|
||||
ComplexTransform enumerant;
|
||||
}
|
||||
ComplexTransform_enumerants[] = {
|
||||
{"n", "none", ComplexTransform::kNone},
|
||||
{"c", "conj", ComplexTransform::kConjugate}
|
||||
};
|
||||
|
||||
/// Converts a ComplexTransform enumerant to a string
|
||||
char const *to_string(ComplexTransform type, bool pretty) {
|
||||
|
||||
for (auto const & possible : ComplexTransform_enumerants) {
|
||||
if (type == possible.enumerant) {
|
||||
if (pretty) {
|
||||
return possible.pretty;
|
||||
}
|
||||
else {
|
||||
return possible.text;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return pretty ? "Invalid" : "invalid";
|
||||
}
|
||||
|
||||
/// Converts a ComplexTransform enumerant from a string
|
||||
template <>
|
||||
ComplexTransform from_string<ComplexTransform>(std::string const &str) {
|
||||
|
||||
for (auto const & possible : ComplexTransform_enumerants) {
|
||||
if ((str.compare(possible.text) == 0) ||
|
||||
(str.compare(possible.pretty) == 0)) {
|
||||
return possible.enumerant;
|
||||
}
|
||||
}
|
||||
|
||||
return ComplexTransform::kInvalid;
|
||||
}
|
||||
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Lexical cast a string to a byte array. Returns true if cast is successful or false if invalid.
|
||||
bool lexical_cast(std::vector<uint8_t> &bytes, NumericTypeID type, std::string const &str) {
|
||||
int size_bytes = sizeof_bits(type) / 8;
|
||||
@@ -574,25 +686,36 @@ std::string lexical_cast(std::vector<uint8_t> &bytes, NumericTypeID type) {
|
||||
break;
|
||||
case NumericTypeID::kCF16:
|
||||
{
|
||||
std::complex<float> tmp;
|
||||
|
||||
cutlass::complex<half_t> const *x =
|
||||
reinterpret_cast<cutlass::complex<half_t> const *>(bytes.data());
|
||||
|
||||
tmp.real(x->real());
|
||||
tmp.imag(x->imag());
|
||||
ss << float(x->real());
|
||||
|
||||
ss << tmp;
|
||||
if (x->imag() != cutlass::half_t()) {
|
||||
ss << "+i" << float(x->imag());
|
||||
}
|
||||
}
|
||||
break;
|
||||
case NumericTypeID::kCF32:
|
||||
{
|
||||
ss << *reinterpret_cast<std::complex<float>*>(bytes.data());
|
||||
cutlass::complex<float> const * x = reinterpret_cast<cutlass::complex<float> const *>(bytes.data());
|
||||
|
||||
ss << x->real();
|
||||
|
||||
if (x->imag() != float()) {
|
||||
ss << "+i" << x->imag();
|
||||
}
|
||||
}
|
||||
break;
|
||||
case NumericTypeID::kCF64:
|
||||
{
|
||||
ss << *reinterpret_cast<std::complex<double>*>(bytes.data());
|
||||
cutlass::complex<double> const * x = reinterpret_cast<cutlass::complex<double> const *>(bytes.data());
|
||||
|
||||
ss << x->real();
|
||||
|
||||
if (x->imag() != double()) {
|
||||
ss << "+i" << x->imag();
|
||||
}
|
||||
}
|
||||
break;
|
||||
default:
|
||||
Reference in New Issue
Block a user