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
@@ -0,0 +1,380 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Defines operations for all CONV operation kinds in CUTLASS Library.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/conv/kernel/default_conv2d_fprop.h"
|
||||
#include "cutlass/conv/kernel/default_conv2d_dgrad.h"
|
||||
#include "cutlass/conv/kernel/default_conv2d_wgrad.h"
|
||||
#include "cutlass/conv/device/implicit_gemm_convolution.h"
|
||||
|
||||
#include "cutlass/library/library.h"
|
||||
#include "library_internal.h"
|
||||
#include "cutlass/util/host_tensor.h"
|
||||
|
||||
#include "cutlass/util/reference/host/convolution.h"
|
||||
#include "cutlass/util/reference/host/tensor_compare.h"
|
||||
#include "cutlass/core_io.h"
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Operator_>
|
||||
class Conv2dOperationBase : public Operation {
|
||||
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;
|
||||
static cutlass::conv::IteratorAlgorithm const kIteratorAlgorithm = Operator::kIteratorAlgorithm;
|
||||
static cutlass::conv::Operator const kConvolutionalOperator = Operator::kConvolutionalOperator;
|
||||
|
||||
using OperatorArguments = typename Operator::Arguments;
|
||||
|
||||
protected:
|
||||
|
||||
///
|
||||
ConvDescription description_;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructor
|
||||
Conv2dOperationBase(char const *name = "unknown_conv2d") {
|
||||
|
||||
description_.name = name;
|
||||
description_.provider = Provider::kCUTLASS;
|
||||
description_.kind = OperationKind::kConv2d;
|
||||
description_.conv_dim = Operator::kConvDim;
|
||||
|
||||
description_.iterator_algorithm = IteratorAlgorithmMap<Operator::kIteratorAlgorithm>::kId;
|
||||
|
||||
description_.tile_description.threadblock_shape = make_Coord(
|
||||
Operator::ThreadblockShape::kM,
|
||||
Operator::ThreadblockShape::kN,
|
||||
Operator::ThreadblockShape::kK);
|
||||
|
||||
description_.tile_description.threadblock_stages = Operator::kStages;
|
||||
|
||||
description_.tile_description.warp_count = make_Coord(
|
||||
Operator::ImplicitGemmKernel::WarpCount::kM,
|
||||
Operator::ImplicitGemmKernel::WarpCount::kN,
|
||||
Operator::ImplicitGemmKernel::WarpCount::kK);
|
||||
|
||||
description_.tile_description.math_instruction.instruction_shape = make_Coord(
|
||||
Operator::InstructionShape::kM,
|
||||
Operator::InstructionShape::kN,
|
||||
Operator::InstructionShape::kK);
|
||||
|
||||
description_.tile_description.math_instruction.element_accumulator =
|
||||
NumericTypeMap<ElementAccumulator>::kId;
|
||||
|
||||
description_.tile_description.math_instruction.opcode_class =
|
||||
OpcodeClassMap<typename Operator::OperatorClass>::kId;
|
||||
|
||||
description_.tile_description.math_instruction.math_operation =
|
||||
MathOperationMap<typename Operator::MathOperator>::kId;
|
||||
|
||||
description_.tile_description.minimum_compute_capability =
|
||||
ArchMap<typename Operator::ArchTag, typename Operator::OperatorClass>::kMin;
|
||||
|
||||
description_.tile_description.maximum_compute_capability =
|
||||
ArchMap<typename Operator::ArchTag, typename Operator::OperatorClass>::kMax;
|
||||
|
||||
description_.A = make_TensorDescription<ElementA, LayoutA>();
|
||||
description_.B = make_TensorDescription<ElementB, LayoutB>();
|
||||
description_.C = make_TensorDescription<ElementC, LayoutC>();
|
||||
description_.element_epilogue = NumericTypeMap<ElementCompute>::kId;
|
||||
|
||||
// TODO: Add split k mode Serial and parallel to convolutions
|
||||
// description_.split_k_mode = Operator::kSplitK ? SplitKMode::kSerial : SplitKMode::kNone;
|
||||
|
||||
}
|
||||
|
||||
/// Returns the description of the GEMM operation
|
||||
virtual OperationDescription const & description() const {
|
||||
return description_;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Conv2d library operation class for cutlass profiler
|
||||
//
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
template <typename Operator_>
|
||||
class Conv2dOperation : public Conv2dOperationBase<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;
|
||||
static cutlass::conv::Operator const kConvolutionalOperator = Operator::kConvolutionalOperator;
|
||||
|
||||
using OperatorArguments = typename Operator::Arguments;
|
||||
|
||||
public:
|
||||
/// Constructor
|
||||
Conv2dOperation(char const *name = "unknown_conv2d_fprop") : Conv2dOperationBase<Operator_>(name) {
|
||||
this->description_.conv_kind = ConvKindMap<kConvolutionalOperator>::kId;
|
||||
}
|
||||
|
||||
protected:
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status construct_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
Conv2dConfiguration const *configuration) {
|
||||
|
||||
|
||||
operator_args.problem_size = configuration->problem_size;
|
||||
|
||||
operator_args.ref_A =
|
||||
{
|
||||
nullptr,
|
||||
LayoutA::packed(implicit_gemm_tensor_a_extent(kConvolutionalOperator, configuration->problem_size))
|
||||
};
|
||||
|
||||
operator_args.ref_B =
|
||||
{
|
||||
nullptr,
|
||||
LayoutB::packed(implicit_gemm_tensor_b_extent(kConvolutionalOperator, configuration->problem_size))
|
||||
};
|
||||
|
||||
operator_args.ref_C =
|
||||
{
|
||||
nullptr,
|
||||
LayoutC::packed(implicit_gemm_tensor_c_extent(kConvolutionalOperator, configuration->problem_size))
|
||||
};
|
||||
|
||||
operator_args.ref_D =
|
||||
{
|
||||
nullptr,
|
||||
LayoutC::packed(implicit_gemm_tensor_c_extent(kConvolutionalOperator, configuration->problem_size))
|
||||
};
|
||||
|
||||
operator_args.split_k_mode = configuration->split_k_mode;
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status update_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
ConvArguments 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.output_op = 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.output_op = params;
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
operator_args.ref_A.reset(static_cast<ElementA *>(const_cast<void *>(arguments->A)));
|
||||
operator_args.ref_B.reset(static_cast<ElementB *>(const_cast<void *>(arguments->B)));
|
||||
operator_args.ref_C.reset(static_cast<ElementC *>(const_cast<void *>(arguments->C)));
|
||||
operator_args.ref_D.reset(static_cast<ElementC *>(const_cast<void *>(arguments->D)));
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
public:
|
||||
|
||||
/// Returns success if the operation can proceed
|
||||
virtual Status can_implement(
|
||||
void const *configuration_ptr,
|
||||
void const *arguments_ptr) const {
|
||||
|
||||
Conv2dConfiguration const *configuration =
|
||||
static_cast<Conv2dConfiguration const *>(configuration_ptr);
|
||||
|
||||
ConvArguments const *arguments =
|
||||
static_cast<ConvArguments 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<Conv2dConfiguration 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<Conv2dConfiguration const *>(configuration_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
Operator *op = new (host_workspace) Operator;
|
||||
//std::cout << "initialize library::Conv2dOperation" << std::endl;
|
||||
//print_operator_args(args);
|
||||
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<ConvArguments 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;
|
||||
}
|
||||
//std::cout << "run library::Conv2dOperation" << std::endl;
|
||||
//print_operator_args(args);
|
||||
return op->run(stream);
|
||||
}
|
||||
|
||||
/// Call print_operator_args from the Conv2dOperation::initialize()
|
||||
// to dump arguments passed on to cutlass operator for debugging
|
||||
void print_operator_args(OperatorArguments &operator_args) const {
|
||||
std::cout << "Conv2dOperation::OperatorArguments" << std::endl
|
||||
<< " problem_size:" << std::endl
|
||||
<< operator_args.problem_size << std::endl
|
||||
<< " split_k_mode: "
|
||||
<< (operator_args.split_k_mode == cutlass::conv::SplitKMode::kSerial ? "serial" : "parallel") << std::endl
|
||||
<< " epilouge (alpha, beta): "
|
||||
<< operator_args.output_op.alpha << ", "
|
||||
<< operator_args.output_op.beta << std::endl
|
||||
<< " ref_A (ptr, {stride}): "
|
||||
<< operator_args.ref_A.data() << ", {"
|
||||
<< operator_args.ref_A.stride(0) << ", "
|
||||
<< operator_args.ref_A.stride(1) << ", "
|
||||
<< operator_args.ref_A.stride(2) << "}" << std::endl
|
||||
<< " ref_B (ptr, {stride}): "
|
||||
<< operator_args.ref_B.data() << ", {"
|
||||
<< operator_args.ref_B.stride(0) << ", "
|
||||
<< operator_args.ref_B.stride(1) << ", "
|
||||
<< operator_args.ref_B.stride(2) << "}" << std::endl
|
||||
<< " ref_C (ptr, {stride}): "
|
||||
<< operator_args.ref_C.data() << ", {"
|
||||
<< operator_args.ref_C.stride(0) << ", "
|
||||
<< operator_args.ref_C.stride(1) << ", "
|
||||
<< operator_args.ref_C.stride(2) << "}" << std::endl
|
||||
<< " ref_D (ptr, {stride}): "
|
||||
<< operator_args.ref_D.data() << ", {"
|
||||
<< operator_args.ref_D.stride(0) << ", "
|
||||
<< operator_args.ref_D.stride(1) << ", "
|
||||
<< operator_args.ref_D.stride(2) << "}" << std::endl;
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace library
|
||||
} // namespace cutlass
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,378 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Defines operations for all CONV operation kinds in CUTLASS Library.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/conv/kernel/default_conv3d_fprop.h"
|
||||
#include "cutlass/conv/kernel/default_conv3d_dgrad.h"
|
||||
#include "cutlass/conv/kernel/default_conv3d_wgrad.h"
|
||||
#include "cutlass/conv/device/implicit_gemm_convolution.h"
|
||||
|
||||
#include "cutlass/library/library.h"
|
||||
#include "library_internal.h"
|
||||
#include "cutlass/util/host_tensor.h"
|
||||
|
||||
#include "cutlass/util/reference/host/convolution.h"
|
||||
#include "cutlass/util/reference/host/tensor_compare.h"
|
||||
#include "cutlass/core_io.h"
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Operator_>
|
||||
class Conv3dOperationBase : public Operation {
|
||||
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;
|
||||
static cutlass::conv::IteratorAlgorithm const kIteratorAlgorithm = Operator::kIteratorAlgorithm;
|
||||
static cutlass::conv::Operator const kConvolutionalOperator = Operator::kConvolutionalOperator;
|
||||
|
||||
using OperatorArguments = typename Operator::Arguments;
|
||||
|
||||
protected:
|
||||
|
||||
///
|
||||
ConvDescription description_;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructor
|
||||
Conv3dOperationBase(char const *name = "unknown_conv3d") {
|
||||
|
||||
description_.name = name;
|
||||
description_.provider = Provider::kCUTLASS;
|
||||
description_.kind = OperationKind::kConv3d;
|
||||
description_.conv_dim = Operator::kConvDim;
|
||||
|
||||
description_.iterator_algorithm = IteratorAlgorithmMap<Operator::kIteratorAlgorithm>::kId;
|
||||
|
||||
description_.tile_description.threadblock_shape = make_Coord(
|
||||
Operator::ThreadblockShape::kM,
|
||||
Operator::ThreadblockShape::kN,
|
||||
Operator::ThreadblockShape::kK);
|
||||
|
||||
description_.tile_description.threadblock_stages = Operator::kStages;
|
||||
|
||||
description_.tile_description.warp_count = make_Coord(
|
||||
Operator::ImplicitGemmKernel::WarpCount::kM,
|
||||
Operator::ImplicitGemmKernel::WarpCount::kN,
|
||||
Operator::ImplicitGemmKernel::WarpCount::kK);
|
||||
|
||||
description_.tile_description.math_instruction.instruction_shape = make_Coord(
|
||||
Operator::InstructionShape::kM,
|
||||
Operator::InstructionShape::kN,
|
||||
Operator::InstructionShape::kK);
|
||||
|
||||
description_.tile_description.math_instruction.element_accumulator =
|
||||
NumericTypeMap<ElementAccumulator>::kId;
|
||||
|
||||
description_.tile_description.math_instruction.opcode_class =
|
||||
OpcodeClassMap<typename Operator::OperatorClass>::kId;
|
||||
|
||||
description_.tile_description.minimum_compute_capability =
|
||||
ArchMap<typename Operator::ArchTag, typename Operator::OperatorClass>::kMin;
|
||||
|
||||
description_.tile_description.maximum_compute_capability =
|
||||
ArchMap<typename Operator::ArchTag, typename Operator::OperatorClass>::kMax;
|
||||
|
||||
description_.A = make_TensorDescription<ElementA, LayoutA>();
|
||||
description_.B = make_TensorDescription<ElementB, LayoutB>();
|
||||
description_.C = make_TensorDescription<ElementC, LayoutC>();
|
||||
description_.element_epilogue = NumericTypeMap<ElementCompute>::kId;
|
||||
|
||||
}
|
||||
|
||||
/// Returns the description of the GEMM operation
|
||||
virtual OperationDescription const & description() const {
|
||||
return description_;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Conv2d library operation class for cutlass profiler
|
||||
//
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
template <typename Operator_>
|
||||
class Conv3dOperation : public Conv3dOperationBase<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;
|
||||
static cutlass::conv::Operator const kConvolutionalOperator = Operator::kConvolutionalOperator;
|
||||
|
||||
using OperatorArguments = typename Operator::Arguments;
|
||||
|
||||
public:
|
||||
/// Constructor
|
||||
Conv3dOperation(char const *name = "unknown_conv3d_fprop") : Conv3dOperationBase<Operator_>(name) {
|
||||
this->description_.conv_kind = ConvKindMap<kConvolutionalOperator>::kId;
|
||||
}
|
||||
|
||||
protected:
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status construct_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
Conv3dConfiguration const *configuration) {
|
||||
|
||||
|
||||
operator_args.problem_size = configuration->problem_size;
|
||||
|
||||
operator_args.ref_A =
|
||||
{
|
||||
nullptr,
|
||||
LayoutA::packed(implicit_gemm_tensor_a_extent(kConvolutionalOperator, configuration->problem_size))
|
||||
};
|
||||
|
||||
operator_args.ref_B =
|
||||
{
|
||||
nullptr,
|
||||
LayoutB::packed(implicit_gemm_tensor_b_extent(kConvolutionalOperator, configuration->problem_size))
|
||||
};
|
||||
|
||||
operator_args.ref_C =
|
||||
{
|
||||
nullptr,
|
||||
LayoutC::packed(implicit_gemm_tensor_c_extent(kConvolutionalOperator, configuration->problem_size))
|
||||
};
|
||||
|
||||
operator_args.ref_D =
|
||||
{
|
||||
nullptr,
|
||||
LayoutC::packed(implicit_gemm_tensor_c_extent(kConvolutionalOperator, configuration->problem_size))
|
||||
};
|
||||
|
||||
operator_args.split_k_mode = configuration->split_k_mode;
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status update_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
ConvArguments 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.output_op = 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.output_op = params;
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
operator_args.ref_A.reset(static_cast<ElementA *>(const_cast<void *>(arguments->A)));
|
||||
operator_args.ref_B.reset(static_cast<ElementB *>(const_cast<void *>(arguments->B)));
|
||||
operator_args.ref_C.reset(static_cast<ElementC *>(const_cast<void *>(arguments->C)));
|
||||
operator_args.ref_D.reset(static_cast<ElementC *>(const_cast<void *>(arguments->D)));
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
public:
|
||||
|
||||
/// Returns success if the operation can proceed
|
||||
virtual Status can_implement(
|
||||
void const *configuration_ptr,
|
||||
void const *arguments_ptr) const {
|
||||
|
||||
Conv3dConfiguration const *configuration =
|
||||
static_cast<Conv3dConfiguration const *>(configuration_ptr);
|
||||
|
||||
ConvArguments const *arguments =
|
||||
static_cast<ConvArguments 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<Conv3dConfiguration 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<Conv3dConfiguration const *>(configuration_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
Operator *op = new (host_workspace) Operator;
|
||||
//std::cout << "initialize library::Conv3dOperation" << std::endl;
|
||||
//print_operator_args(args);
|
||||
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<ConvArguments 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;
|
||||
}
|
||||
//std::cout << "run library::Conv3dOperation" << std::endl;
|
||||
//print_operator_args(args);
|
||||
return op->run(stream);
|
||||
}
|
||||
|
||||
/// Call print_operator_args from the Conv3dOperation::initialize()
|
||||
// to dump arguments passed on to cutlass operator for debugging
|
||||
void print_operator_args(OperatorArguments &operator_args) const {
|
||||
std::cout << "Conv3dOperation::OperatorArguments" << std::endl
|
||||
<< " problem_size: "
|
||||
<< operator_args.problem_size << std::endl
|
||||
<< " split_k_mode: "
|
||||
<< (operator_args.split_k_mode == cutlass::conv::SplitKMode::kSerial ? "serial" : "parallel") << std::endl
|
||||
<< " epilouge (alpha, beta): "
|
||||
<< operator_args.output_op.alpha << ", "
|
||||
<< operator_args.output_op.beta << std::endl
|
||||
<< " ref_A (ptr, {stride}): "
|
||||
<< operator_args.ref_A.data() << ", {"
|
||||
<< operator_args.ref_A.stride(0) << ", "
|
||||
<< operator_args.ref_A.stride(1) << ", "
|
||||
<< operator_args.ref_A.stride(2) << ", "
|
||||
<< operator_args.ref_A.stride(3) << "}" << std::endl
|
||||
<< " ref_B (ptr, {stride}): "
|
||||
<< operator_args.ref_B.data() << ", {"
|
||||
<< operator_args.ref_B.stride(0) << ", "
|
||||
<< operator_args.ref_B.stride(1) << ", "
|
||||
<< operator_args.ref_B.stride(2) << ", "
|
||||
<< operator_args.ref_B.stride(3) << "}" << std::endl
|
||||
<< " ref_C (ptr, {stride}): "
|
||||
<< operator_args.ref_C.data() << ", {"
|
||||
<< operator_args.ref_C.stride(0) << ", "
|
||||
<< operator_args.ref_C.stride(1) << ", "
|
||||
<< operator_args.ref_C.stride(2) << ", "
|
||||
<< operator_args.ref_C.stride(3) << "}" << std::endl
|
||||
<< " ref_D (ptr, {stride}): "
|
||||
<< operator_args.ref_D.data() << ", {"
|
||||
<< operator_args.ref_D.stride(0) << ", "
|
||||
<< operator_args.ref_D.stride(1) << ", "
|
||||
<< operator_args.ref_D.stride(2) << ", "
|
||||
<< operator_args.ref_D.stride(3) << "}" << std::endl;
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace library
|
||||
} // namespace cutlass
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -1037,8 +1037,70 @@ Status Handle::gemm_planar_complex_array(
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Finds conv operation instances with Conv::ElementC = Reduction::ElementWorkspace
|
||||
Operation const* find_conv_operation_for_parallel_reduction(Operation const *operation) {
|
||||
|
||||
ConvDescription const &conv_desc =
|
||||
static_cast<ConvDescription const &>(operation->description());
|
||||
|
||||
// if the curren conv operation accumulator and output data type match return operation
|
||||
if(conv_desc.tile_description.math_instruction.element_accumulator == conv_desc.C.element) {
|
||||
return operation;
|
||||
}
|
||||
|
||||
// find conv operation to match conv output and reduction workspace data type
|
||||
ConvFunctionalKey key(
|
||||
library::Provider::kCUTLASS,
|
||||
conv_desc.conv_kind,
|
||||
conv_desc.A.element,
|
||||
conv_desc.A.layout,
|
||||
conv_desc.B.element,
|
||||
conv_desc.B.layout,
|
||||
conv_desc.tile_description.math_instruction.element_accumulator,
|
||||
conv_desc.C.layout,
|
||||
conv_desc.tile_description.math_instruction.element_accumulator,
|
||||
conv_desc.element_epilogue);
|
||||
|
||||
// conv operation table for conv2d or conv3d
|
||||
auto conv_operations = (conv_desc.kind == OperationKind::kConv2d) ?
|
||||
Singleton::get().operation_table.conv2d_operations :
|
||||
Singleton::get().operation_table.conv3d_operations;
|
||||
|
||||
// find ConvFunctionalKey in convolution operation table
|
||||
auto operators_it = conv_operations.find(key);
|
||||
|
||||
if (operators_it == conv_operations.end()) {
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
if (operators_it->second.empty()) {
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
// conv operation for same compute capability and iterator algorithm
|
||||
ConvPreferenceKey preference_key(
|
||||
conv_desc.tile_description.minimum_compute_capability,
|
||||
conv_desc.iterator_algorithm);
|
||||
|
||||
auto it = operators_it->second.find(preference_key);
|
||||
|
||||
if(it == operators_it->second.end()) {
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
// return matching conv opertion (same tile sizes and instruction)
|
||||
for (auto op : it->second) {
|
||||
if (op->description().tile_description == operation->description().tile_description) {
|
||||
return op;
|
||||
}
|
||||
}
|
||||
|
||||
return nullptr;
|
||||
}
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace library
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
@@ -227,6 +227,23 @@ template <> struct LayoutMap<cutlass::layout::TensorNHWC> {
|
||||
template <> struct LayoutMap<cutlass::layout::TensorNDHWC> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kTensorNDHWC;
|
||||
};
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::TensorNCxHWx<32>> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kTensorNC32HW32;
|
||||
};
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::TensorNCxHWx<64>> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kTensorNC64HW64;
|
||||
};
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::TensorCxRSKx<32>> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kTensorC32RSK32;
|
||||
};
|
||||
|
||||
template <> struct LayoutMap<cutlass::layout::TensorCxRSKx<64>> {
|
||||
static LayoutTypeID const kId = LayoutTypeID::kTensorC64RSK64;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename T> struct OpcodeClassMap;
|
||||
@@ -257,6 +274,43 @@ template <> struct ComplexTransformMap<cutlass::ComplexTransform::kConjugate> {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <cutlass::conv::Mode T> struct ConvModeMap;
|
||||
|
||||
template <> struct ConvModeMap<conv::Mode::kCrossCorrelation> {
|
||||
static ConvModeID const kId = ConvModeID::kCrossCorrelation;
|
||||
};
|
||||
|
||||
template <> struct ConvModeMap<conv::Mode::kConvolution> {
|
||||
static ConvModeID const kId = ConvModeID::kConvolution;
|
||||
};
|
||||
|
||||
|
||||
template <cutlass::conv::Operator T> struct ConvKindMap;
|
||||
|
||||
template <> struct ConvKindMap<conv::Operator::kFprop> {
|
||||
static ConvKind const kId = ConvKind::kFprop;
|
||||
};
|
||||
|
||||
template <> struct ConvKindMap<conv::Operator::kDgrad> {
|
||||
static ConvKind const kId = ConvKind::kDgrad;
|
||||
};
|
||||
|
||||
template <> struct ConvKindMap<conv::Operator::kWgrad> {
|
||||
static ConvKind const kId = ConvKind::kWgrad;
|
||||
};
|
||||
|
||||
|
||||
template <cutlass::conv::IteratorAlgorithm T> struct IteratorAlgorithmMap;
|
||||
|
||||
template <> struct IteratorAlgorithmMap<conv::IteratorAlgorithm::kAnalytic> {
|
||||
static IteratorAlgorithmID const kId = IteratorAlgorithmID::kAnalytic;
|
||||
};
|
||||
|
||||
template <> struct IteratorAlgorithmMap<conv::IteratorAlgorithm::kOptimized> {
|
||||
static IteratorAlgorithmID const kId = IteratorAlgorithmID::kOptimized;
|
||||
};
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Element, typename Layout>
|
||||
TensorDescription make_TensorDescription(int alignment = 1) {
|
||||
TensorDescription desc;
|
||||
|
||||
@@ -36,6 +36,11 @@ namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
void initialize_reference_operations(Manifest &manifest);
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Top-level initialization
|
||||
Status Manifest::initialize() {
|
||||
|
||||
@@ -46,6 +51,12 @@ Status Manifest::initialize() {
|
||||
// initialize procedurally generated cutlass op in manifest object
|
||||
initialize_all(*this);
|
||||
|
||||
// initialize manually instanced conv3d reference op in manifest object
|
||||
initialize_reference_operations(*this);
|
||||
|
||||
// initialize manually instanced reduction reference op in manifest object
|
||||
initialize_all_reduction_op(*this);
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
|
||||
@@ -76,6 +76,55 @@ void OperationTable::append(Manifest const &manifest) {
|
||||
}
|
||||
|
||||
|
||||
// insert all conv2d or conv3d operation into operation table
|
||||
if (desc.kind == OperationKind::kConv2d || desc.kind == OperationKind::kConv3d) {
|
||||
auto &conv_desc = static_cast<library::ConvDescription const &>(desc);
|
||||
|
||||
ConvFunctionalKey functional_key(
|
||||
conv_desc.provider,
|
||||
conv_desc.conv_kind,
|
||||
conv_desc.A.element,
|
||||
conv_desc.A.layout,
|
||||
conv_desc.B.element,
|
||||
conv_desc.B.layout,
|
||||
conv_desc.C.element,
|
||||
conv_desc.C.layout,
|
||||
conv_desc.tile_description.math_instruction.element_accumulator,
|
||||
conv_desc.element_epilogue
|
||||
);
|
||||
|
||||
Operation const *op = operation.get();
|
||||
|
||||
int cc = conv_desc.tile_description.minimum_compute_capability;
|
||||
|
||||
ConvPreferenceKey preference_key(cc, conv_desc.iterator_algorithm);
|
||||
|
||||
// insert conv operation to conv2d_operations or conv3d_operations map
|
||||
(desc.kind == OperationKind::kConv2d) ?
|
||||
conv2d_operations[functional_key][preference_key].push_back(op) :
|
||||
conv3d_operations[functional_key][preference_key].push_back(op);
|
||||
}
|
||||
|
||||
// insert all reduction operation into operation table
|
||||
if (desc.kind == OperationKind::kReduction) {
|
||||
auto &reduce_desc = static_cast<library::ReductionDescription const &>(desc);
|
||||
|
||||
ReductionFunctionalKey functional_key(
|
||||
reduce_desc.provider,
|
||||
reduce_desc.element_workspace,
|
||||
reduce_desc.tile_description.math_instruction.element_accumulator,
|
||||
reduce_desc.element_output,
|
||||
reduce_desc.element_epilogue,
|
||||
library::MathOperationID::kAdd,
|
||||
library::EpilogueKind::kLinearCombination
|
||||
);
|
||||
|
||||
Operation const *op = operation.get();
|
||||
|
||||
reduction_operations[functional_key] = op;
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Initialize operations for reduction operation in CUTLASS Library.
|
||||
|
||||
*/
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/library/library.h"
|
||||
#include "cutlass/library/manifest.h"
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// CUTLASS Reduction Instances //
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////
|
||||
void initialize_reduce_add_linear_combination_f32_f32_f16(Manifest &manifest);
|
||||
void initialize_reduce_add_linear_combination_f32_f32_f32(Manifest &manifest);
|
||||
void initialize_reduce_add_linear_combination_cf32_cf32_cf32(Manifest &manifest);
|
||||
|
||||
//
|
||||
// Entry point to construct operations
|
||||
//
|
||||
void initialize_all_reduction_op(Manifest &manifest) {
|
||||
|
||||
initialize_reduce_add_linear_combination_f32_f32_f16(manifest);
|
||||
initialize_reduce_add_linear_combination_f32_f32_f32(manifest);
|
||||
initialize_reduce_add_linear_combination_cf32_cf32_cf32(manifest);
|
||||
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace library
|
||||
} // namespace cutlass
|
||||
@@ -0,0 +1,145 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Defines operations for reduction operation in CUTLASS Library.
|
||||
*/
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/library/library.h"
|
||||
#include "cutlass/library/manifest.h"
|
||||
|
||||
#include "reduction_operation.h"
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
// naming convention initialize_reduce_[ReductionOp]_[EpilogueOp]_[ElementWorkspace]_[ElementAccumulator]_[ElementOutput]
|
||||
|
||||
void initialize_reduce_add_linear_combination_f32_f32_f16(Manifest &manifest) {
|
||||
|
||||
using ElementWorkspace = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementOutput = cutlass::half_t;
|
||||
using ElementCompute = float;
|
||||
|
||||
using EpilogueOutputOp = cutlass::epilogue::thread::LinearCombination<
|
||||
ElementOutput,
|
||||
128 / cutlass::sizeof_bits<ElementOutput>::value,
|
||||
ElementAccumulator,
|
||||
ElementCompute
|
||||
>;
|
||||
|
||||
using ReductionOp = cutlass::reduction::thread::ReduceAdd<
|
||||
ElementAccumulator,
|
||||
typename EpilogueOutputOp::ElementAccumulator,
|
||||
EpilogueOutputOp::kCount
|
||||
>;
|
||||
|
||||
using Operation_reduce_add_linear_combination_f32_f32_f16 = cutlass::reduction::device::ReduceSplitK<
|
||||
cutlass::reduction::kernel::ReduceSplitK<
|
||||
cutlass::MatrixShape<4, 32 * EpilogueOutputOp::kCount>,
|
||||
EpilogueOutputOp,
|
||||
ReductionOp
|
||||
>
|
||||
>;
|
||||
|
||||
manifest.append(new ReductionOperation<
|
||||
Operation_reduce_add_linear_combination_f32_f32_f16>(
|
||||
"reduce_add_linear_combination_f32_f32_f16"
|
||||
));
|
||||
}
|
||||
|
||||
|
||||
void initialize_reduce_add_linear_combination_f32_f32_f32(Manifest &manifest) {
|
||||
|
||||
using ElementWorkspace = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementOutput = float;
|
||||
using ElementCompute = float;
|
||||
|
||||
using EpilogueOutputOp = cutlass::epilogue::thread::LinearCombination<
|
||||
ElementOutput,
|
||||
128 / cutlass::sizeof_bits<ElementOutput>::value,
|
||||
ElementAccumulator,
|
||||
ElementCompute
|
||||
>;
|
||||
|
||||
using ReductionOp = cutlass::reduction::thread::ReduceAdd<
|
||||
ElementAccumulator,
|
||||
typename EpilogueOutputOp::ElementAccumulator,
|
||||
EpilogueOutputOp::kCount
|
||||
>;
|
||||
|
||||
using Operation_reduce_add_linear_combination_f32_f32_f32 = cutlass::reduction::device::ReduceSplitK<
|
||||
cutlass::reduction::kernel::ReduceSplitK<
|
||||
cutlass::MatrixShape<4, 32 * EpilogueOutputOp::kCount>,
|
||||
EpilogueOutputOp,
|
||||
ReductionOp
|
||||
>
|
||||
>;
|
||||
|
||||
manifest.append(new ReductionOperation<
|
||||
Operation_reduce_add_linear_combination_f32_f32_f32>(
|
||||
"reduce_add_linear_combination_f32_f32_f32"
|
||||
));
|
||||
}
|
||||
|
||||
void initialize_reduce_add_linear_combination_cf32_cf32_cf32(Manifest &manifest) {
|
||||
|
||||
using ElementWorkspace = cutlass::complex<float>;
|
||||
using ElementAccumulator = cutlass::complex<float>;
|
||||
using ElementOutput = cutlass::complex<float>;
|
||||
using ElementCompute = cutlass::complex<float>;
|
||||
|
||||
using EpilogueOutputOp = cutlass::epilogue::thread::LinearCombination<
|
||||
ElementOutput,
|
||||
128 / cutlass::sizeof_bits<ElementOutput>::value,
|
||||
ElementAccumulator,
|
||||
ElementCompute
|
||||
>;
|
||||
|
||||
using ReductionOp = cutlass::reduction::thread::ReduceAdd<
|
||||
ElementAccumulator,
|
||||
typename EpilogueOutputOp::ElementAccumulator,
|
||||
EpilogueOutputOp::kCount
|
||||
>;
|
||||
|
||||
using Operation_reduce_add_linear_combination_cf32_cf32_cf32 = cutlass::reduction::device::ReduceSplitK<
|
||||
cutlass::reduction::kernel::ReduceSplitK<
|
||||
cutlass::MatrixShape<4, 32 * EpilogueOutputOp::kCount>,
|
||||
EpilogueOutputOp,
|
||||
ReductionOp
|
||||
>
|
||||
>;
|
||||
|
||||
manifest.append(new ReductionOperation<
|
||||
Operation_reduce_add_linear_combination_cf32_cf32_cf32>(
|
||||
"reduce_add_linear_combination_cf32_cf32_cf32"
|
||||
));
|
||||
}
|
||||
|
||||
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,282 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, 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 operations for reduction operation in CUTLASS Library.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/reduction/thread/reduction_operators.h"
|
||||
#include "cutlass/reduction/device/reduce_split_k.h"
|
||||
|
||||
#include "cutlass/library/library.h"
|
||||
#include "library_internal.h"
|
||||
#include "cutlass/core_io.h"
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Operator_>
|
||||
class ReductionOperation : public Operation {
|
||||
public:
|
||||
using Operator = Operator_;
|
||||
|
||||
using ElementWorkspace = typename Operator::ElementWorkspace;
|
||||
using ElementAccumulator = typename Operator::ElementAccumulator;
|
||||
using ElementOutput = typename Operator::ElementOutput;
|
||||
|
||||
using ElementCompute = typename Operator::OutputOp::ElementCompute;
|
||||
|
||||
using OperatorArguments = typename Operator::Arguments;
|
||||
|
||||
protected:
|
||||
|
||||
///
|
||||
ReductionDescription description_;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructor
|
||||
ReductionOperation(char const *name = "unknown_reduction") {
|
||||
|
||||
description_.name = name;
|
||||
description_.provider = Provider::kCUTLASS;
|
||||
description_.kind = OperationKind::kReduction;
|
||||
|
||||
description_.tile_description.threadblock_shape = make_Coord(Operator::Shape::kRow, Operator::Shape::kColumn, 1);
|
||||
|
||||
description_.tile_description.math_instruction.instruction_shape = make_Coord(1, 1, 1);
|
||||
description_.tile_description.math_instruction.element_accumulator = NumericTypeMap<ElementAccumulator>::kId;
|
||||
description_.tile_description.math_instruction.opcode_class = OpcodeClassID::kSimt;
|
||||
description_.tile_description.math_instruction.math_operation = MathOperationID::kAdd;
|
||||
|
||||
description_.tile_description.minimum_compute_capability = 50;
|
||||
description_.tile_description.maximum_compute_capability = 1024;
|
||||
|
||||
description_.element_workspace = NumericTypeMap<ElementWorkspace>::kId;
|
||||
description_.element_output = NumericTypeMap<ElementOutput>::kId;
|
||||
description_.element_epilogue = NumericTypeMap<ElementCompute>::kId;
|
||||
|
||||
}
|
||||
|
||||
/// Returns the description of the Reduction operation
|
||||
virtual OperationDescription const & description() const {
|
||||
return description_;
|
||||
}
|
||||
|
||||
|
||||
protected:
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status construct_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
ReductionConfiguration const *configuration) {
|
||||
|
||||
operator_args.problem_size = configuration->problem_size;
|
||||
operator_args.partitions = configuration->partitions;
|
||||
operator_args.partition_stride = configuration->partition_stride;
|
||||
|
||||
operator_args.workspace = {nullptr, int(configuration->ldw)};
|
||||
operator_args.source = {nullptr, int(configuration->lds)};
|
||||
operator_args.destination = {nullptr, int(configuration->ldd)};
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Constructs the arguments structure given the configuration and arguments
|
||||
static Status update_arguments_(
|
||||
OperatorArguments &operator_args,
|
||||
ReductionArguments const *arguments) {
|
||||
|
||||
if (arguments->pointer_mode == ScalarPointerMode::kHost) {
|
||||
typename Operator::OutputOp::Params params(
|
||||
*static_cast<ElementCompute const *>(arguments->alpha),
|
||||
*static_cast<ElementCompute const *>(arguments->beta)
|
||||
);
|
||||
operator_args.output = params;
|
||||
}
|
||||
else if (arguments->pointer_mode == ScalarPointerMode::kDevice){
|
||||
typename Operator::OutputOp::Params params(
|
||||
static_cast<ElementCompute const *>(arguments->alpha),
|
||||
static_cast<ElementCompute const *>(arguments->beta)
|
||||
);
|
||||
operator_args.output = params;
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
operator_args.workspace.reset(static_cast<ElementWorkspace *>(const_cast<void *>(arguments->workspace)));
|
||||
operator_args.source.reset(static_cast<ElementOutput *>(const_cast<void *>(arguments->source)));
|
||||
operator_args.destination.reset(static_cast<ElementOutput *>(const_cast<void *>(arguments->destination)));
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
public:
|
||||
|
||||
/// Returns success if the operation can proceed
|
||||
virtual Status can_implement(
|
||||
void const *configuration_ptr,
|
||||
void const *arguments_ptr) const {
|
||||
|
||||
ReductionConfiguration const *configuration =
|
||||
static_cast<ReductionConfiguration const *>(configuration_ptr);
|
||||
|
||||
ReductionArguments const *arguments =
|
||||
static_cast<ReductionArguments 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<ReductionConfiguration 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<ReductionConfiguration const *>(configuration_ptr));
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
Operator *op = new (host_workspace) Operator;
|
||||
//std::cout << "initialize library::Reduction" << std::endl;
|
||||
//print_operator_args(args);
|
||||
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<ReductionArguments 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;
|
||||
}
|
||||
|
||||
//std::cout << "run library::Reduction" << std::endl;
|
||||
//print_operator_args(args);
|
||||
return op->run(stream);
|
||||
}
|
||||
|
||||
/// Call print_operator_args from the Reduction::initialize()
|
||||
// to dump arguments passed on to cutlass operator for debugging
|
||||
void print_operator_args(OperatorArguments &operator_args) const {
|
||||
std::cout << "Reduction::OperatorArguments" << std::endl
|
||||
<< " problem_size: "
|
||||
<< operator_args.problem_size << std::endl
|
||||
<< " partitions: "
|
||||
<< operator_args.partitions << std::endl
|
||||
<< " partition_stride: "
|
||||
<< operator_args.partition_stride << std::endl
|
||||
<< " epilouge (alpha, beta): "
|
||||
<< operator_args.output.alpha << ", "
|
||||
<< operator_args.output.beta << std::endl
|
||||
<< " workspace (ptr, stride): "
|
||||
<< operator_args.workspace.data() << ", "
|
||||
<< operator_args.workspace.stride(0) << std::endl
|
||||
<< " source (ptr, stride): "
|
||||
<< operator_args.source.data() << ", "
|
||||
<< operator_args.source.stride(0) << std::endl
|
||||
<< " destination (ptr, stride): "
|
||||
<< operator_args.destination.data() << ", "
|
||||
<< operator_args.destination.stride(0) << std::endl;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace library
|
||||
} // namespace cutlass
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,223 @@
|
||||
/***************************************************************************************************
|
||||
* 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
|
||||
|
||||
*/
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/library/library.h"
|
||||
#include "cutlass/library/manifest.h"
|
||||
|
||||
#include "conv_reference_operation.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
void initialize_conv2d_reference_operations(Manifest &manifest) {
|
||||
|
||||
make_conv_all<
|
||||
2,
|
||||
cutlass::half_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::half_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::half_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::half_t,
|
||||
cutlass::half_t
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
2,
|
||||
cutlass::half_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::half_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::half_t, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
2,
|
||||
cutlass::half_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::half_t, cutlass::layout::TensorNHWC,
|
||||
float, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
2,
|
||||
cutlass::bfloat16_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::bfloat16_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::bfloat16_t, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
2,
|
||||
cutlass::bfloat16_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::bfloat16_t, cutlass::layout::TensorNHWC,
|
||||
float, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
2,
|
||||
cutlass::tfloat32_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::tfloat32_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::tfloat32_t, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
2,
|
||||
cutlass::tfloat32_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::tfloat32_t, cutlass::layout::TensorNHWC,
|
||||
float, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
2,
|
||||
float, cutlass::layout::TensorNHWC,
|
||||
float, cutlass::layout::TensorNHWC,
|
||||
float, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
2,
|
||||
cutlass::complex<float>, cutlass::layout::TensorNHWC,
|
||||
cutlass::complex<float>, cutlass::layout::TensorNHWC,
|
||||
cutlass::complex<float>, cutlass::layout::TensorNHWC,
|
||||
cutlass::complex<float>,
|
||||
cutlass::complex<float>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
2,
|
||||
int8_t, cutlass::layout::TensorNHWC,
|
||||
int8_t, cutlass::layout::TensorNHWC,
|
||||
int32_t, cutlass::layout::TensorNHWC,
|
||||
int32_t,
|
||||
int32_t,
|
||||
NumericConverterClamp<int32_t, int32_t>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
2,
|
||||
int8_t, cutlass::layout::TensorNHWC,
|
||||
int8_t, cutlass::layout::TensorNHWC,
|
||||
int8_t, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
int32_t,
|
||||
NumericConverterClamp<int8_t, float>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
2,
|
||||
uint8_t, cutlass::layout::TensorNHWC,
|
||||
uint8_t, cutlass::layout::TensorNHWC,
|
||||
uint8_t, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
int32_t,
|
||||
NumericConverterClamp<uint8_t, float>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
2,
|
||||
uint8_t, cutlass::layout::TensorNHWC,
|
||||
uint8_t, cutlass::layout::TensorNHWC,
|
||||
int32_t, cutlass::layout::TensorNHWC,
|
||||
int32_t,
|
||||
int32_t,
|
||||
NumericConverterClamp<int32_t, int32_t>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
2,
|
||||
uint8_t, cutlass::layout::TensorNHWC,
|
||||
uint8_t, cutlass::layout::TensorNHWC,
|
||||
int8_t, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
int32_t,
|
||||
NumericConverterClamp<int8_t, float>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
2,
|
||||
cutlass::int4b_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::int4b_t, cutlass::layout::TensorNHWC,
|
||||
int32_t, cutlass::layout::TensorNHWC,
|
||||
int32_t,
|
||||
int32_t,
|
||||
NumericConverterClamp<int32_t, int32_t>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
2,
|
||||
cutlass::int4b_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::int4b_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::int4b_t, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
int32_t,
|
||||
NumericConverterClamp<cutlass::int4b_t, float>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
2,
|
||||
cutlass::uint4b_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::uint4b_t, cutlass::layout::TensorNHWC,
|
||||
int32_t, cutlass::layout::TensorNHWC,
|
||||
int32_t,
|
||||
int32_t,
|
||||
NumericConverterClamp<int32_t, int32_t>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
2,
|
||||
cutlass::uint4b_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::uint4b_t, cutlass::layout::TensorNHWC,
|
||||
cutlass::uint4b_t, cutlass::layout::TensorNHWC,
|
||||
float,
|
||||
int32_t,
|
||||
NumericConverterClamp<cutlass::uint4b_t, float>
|
||||
>(manifest);
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace library
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -0,0 +1,203 @@
|
||||
/***************************************************************************************************
|
||||
* 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
|
||||
*/
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/library/library.h"
|
||||
#include "cutlass/library/manifest.h"
|
||||
|
||||
#include "conv_reference_operation.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
void initialize_conv3d_reference_operations(Manifest &manifest) {
|
||||
|
||||
make_conv_all<
|
||||
3,
|
||||
cutlass::half_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::half_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::half_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::half_t,
|
||||
cutlass::half_t
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
3,
|
||||
cutlass::half_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::half_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::half_t, cutlass::layout::TensorNDHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
3,
|
||||
cutlass::half_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::half_t, cutlass::layout::TensorNDHWC,
|
||||
float, cutlass::layout::TensorNDHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
3,
|
||||
cutlass::bfloat16_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::bfloat16_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::bfloat16_t, cutlass::layout::TensorNDHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
3,
|
||||
cutlass::bfloat16_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::bfloat16_t, cutlass::layout::TensorNDHWC,
|
||||
float, cutlass::layout::TensorNDHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
3,
|
||||
cutlass::tfloat32_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::tfloat32_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::tfloat32_t, cutlass::layout::TensorNDHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
3,
|
||||
cutlass::tfloat32_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::tfloat32_t, cutlass::layout::TensorNDHWC,
|
||||
float, cutlass::layout::TensorNDHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_all<
|
||||
3,
|
||||
float, cutlass::layout::TensorNDHWC,
|
||||
float, cutlass::layout::TensorNDHWC,
|
||||
float, cutlass::layout::TensorNDHWC,
|
||||
float,
|
||||
float
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
3,
|
||||
int8_t, cutlass::layout::TensorNDHWC,
|
||||
int8_t, cutlass::layout::TensorNDHWC,
|
||||
int32_t, cutlass::layout::TensorNDHWC,
|
||||
int32_t,
|
||||
int32_t,
|
||||
NumericConverterClamp<int32_t, int32_t>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
3,
|
||||
int8_t, cutlass::layout::TensorNDHWC,
|
||||
int8_t, cutlass::layout::TensorNDHWC,
|
||||
int8_t, cutlass::layout::TensorNDHWC,
|
||||
float,
|
||||
int32_t,
|
||||
NumericConverterClamp<int8_t, float>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
3,
|
||||
uint8_t, cutlass::layout::TensorNDHWC,
|
||||
uint8_t, cutlass::layout::TensorNDHWC,
|
||||
int32_t, cutlass::layout::TensorNDHWC,
|
||||
int32_t,
|
||||
int32_t,
|
||||
NumericConverterClamp<int32_t, int32_t>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
3,
|
||||
uint8_t, cutlass::layout::TensorNDHWC,
|
||||
uint8_t, cutlass::layout::TensorNDHWC,
|
||||
int8_t, cutlass::layout::TensorNDHWC,
|
||||
float,
|
||||
int32_t,
|
||||
NumericConverterClamp<int8_t, float>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
3,
|
||||
cutlass::int4b_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::int4b_t, cutlass::layout::TensorNDHWC,
|
||||
int32_t, cutlass::layout::TensorNDHWC,
|
||||
int32_t,
|
||||
int32_t,
|
||||
NumericConverterClamp<int32_t, int32_t>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
3,
|
||||
cutlass::int4b_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::int4b_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::int4b_t, cutlass::layout::TensorNDHWC,
|
||||
float,
|
||||
int32_t,
|
||||
NumericConverterClamp<cutlass::int4b_t, float>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
3,
|
||||
cutlass::uint4b_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::uint4b_t, cutlass::layout::TensorNDHWC,
|
||||
int32_t, cutlass::layout::TensorNDHWC,
|
||||
int32_t,
|
||||
int32_t,
|
||||
NumericConverterClamp<int32_t, int32_t>
|
||||
>(manifest);
|
||||
|
||||
make_conv_fprop<
|
||||
3,
|
||||
cutlass::uint4b_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::uint4b_t, cutlass::layout::TensorNDHWC,
|
||||
cutlass::uint4b_t, cutlass::layout::TensorNDHWC,
|
||||
float,
|
||||
int32_t,
|
||||
NumericConverterClamp<cutlass::uint4b_t, float>
|
||||
>(manifest);
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace library
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -0,0 +1,607 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Defines operations for all CONV operation kinds in CUTLASS Library
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <iostream>
|
||||
#include <sstream>
|
||||
#include <cstring>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/library/library.h"
|
||||
#include "cutlass/library/manifest.h"
|
||||
#include "cutlass/library/util.h"
|
||||
#include "library_internal.h"
|
||||
|
||||
#include "cutlass/util/reference/host/convolution.h"
|
||||
#include "cutlass/util/reference/device/convolution.h"
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
|
||||
template <
|
||||
Provider kProvider,
|
||||
conv::Operator ConvolutionalOperator,
|
||||
int ConvDim,
|
||||
typename ElementA_,
|
||||
typename LayoutA_,
|
||||
typename ElementB_,
|
||||
typename LayoutB_,
|
||||
typename ElementC_,
|
||||
typename LayoutC_,
|
||||
typename ElementCompute_,
|
||||
typename ElementAccumulator_ = ElementCompute_,
|
||||
typename ConvertOp_ = NumericConverter<ElementC_, ElementCompute_>,
|
||||
typename InnerProductOp_ = multiply_add<ElementAccumulator_>
|
||||
>
|
||||
struct ConvReferenceDispatcher;
|
||||
|
||||
/// Dispatcher for Conv2d (partially specialied for kConvDim == 2)
|
||||
template <
|
||||
Provider kProvider,
|
||||
conv::Operator kConvolutionalOperator,
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementCompute,
|
||||
typename ElementAccumulator,
|
||||
typename ConvertOp,
|
||||
typename InnerProductOp
|
||||
>
|
||||
struct ConvReferenceDispatcher<
|
||||
kProvider,
|
||||
kConvolutionalOperator,
|
||||
2,
|
||||
ElementA, LayoutA,
|
||||
ElementB, LayoutB,
|
||||
ElementC, LayoutC,
|
||||
ElementCompute,
|
||||
ElementAccumulator,
|
||||
ConvertOp,
|
||||
InnerProductOp> {
|
||||
|
||||
static Status dispatch(
|
||||
void const *configuration,
|
||||
ElementA *ptr_A,
|
||||
ElementB *ptr_B,
|
||||
ElementC *ptr_C,
|
||||
ElementC *ptr_D,
|
||||
ElementCompute alpha,
|
||||
ElementCompute beta,
|
||||
cudaStream_t stream = nullptr
|
||||
) {
|
||||
|
||||
Conv2dConfiguration const &config =
|
||||
*static_cast<Conv2dConfiguration const *>(configuration);
|
||||
|
||||
ConvKind const conv_kind = ConvKindMap<kConvolutionalOperator>::kId;
|
||||
|
||||
if (kProvider == Provider::kReferenceHost) {
|
||||
|
||||
cutlass::reference::host::Conv2d<
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC ,
|
||||
LayoutC,
|
||||
ElementCompute,
|
||||
ElementAccumulator,
|
||||
ConvertOp,
|
||||
InnerProductOp
|
||||
>(
|
||||
kConvolutionalOperator,
|
||||
config.problem_size,
|
||||
{ptr_A, config.layout_a(conv_kind)},
|
||||
{ptr_B, config.layout_b(conv_kind)},
|
||||
{ptr_C, config.layout_c(conv_kind)},
|
||||
{ptr_D, config.layout_c(conv_kind)},
|
||||
alpha,
|
||||
beta
|
||||
);
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
else if (kProvider == Provider::kReferenceDevice) {
|
||||
return cutlass::reference::device::Conv2d<
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementCompute,
|
||||
ElementAccumulator,
|
||||
ConvertOp,
|
||||
InnerProductOp
|
||||
>(
|
||||
kConvolutionalOperator,
|
||||
config.problem_size,
|
||||
{ptr_A, config.layout_a(conv_kind)},
|
||||
{ptr_B, config.layout_b(conv_kind)},
|
||||
{ptr_C, config.layout_c(conv_kind)},
|
||||
{ptr_D, config.layout_c(conv_kind)},
|
||||
alpha,
|
||||
beta,
|
||||
stream
|
||||
);
|
||||
}
|
||||
return Status::kErrorNotSupported;
|
||||
}
|
||||
};
|
||||
|
||||
/// Dispatcher for Conv3d (partially specialized for kConvDim == 3)
|
||||
template <
|
||||
Provider kProvider,
|
||||
conv::Operator kConvolutionalOperator,
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementCompute,
|
||||
typename ElementAccumulator,
|
||||
typename ConvertOp,
|
||||
typename InnerProductOp
|
||||
>
|
||||
struct ConvReferenceDispatcher<
|
||||
kProvider,
|
||||
kConvolutionalOperator,
|
||||
3,
|
||||
ElementA, LayoutA,
|
||||
ElementB, LayoutB,
|
||||
ElementC, LayoutC,
|
||||
ElementCompute,
|
||||
ElementAccumulator,
|
||||
ConvertOp,
|
||||
InnerProductOp> {
|
||||
|
||||
static Status dispatch(
|
||||
void const *configuration,
|
||||
ElementA *ptr_A,
|
||||
ElementB *ptr_B,
|
||||
ElementC *ptr_C,
|
||||
ElementC *ptr_D,
|
||||
ElementCompute alpha,
|
||||
ElementCompute beta,
|
||||
cudaStream_t stream = nullptr
|
||||
) {
|
||||
|
||||
Conv3dConfiguration const &config =
|
||||
*static_cast<Conv3dConfiguration const *>(configuration);
|
||||
|
||||
ConvKind const conv_kind = ConvKindMap<kConvolutionalOperator>::kId;
|
||||
|
||||
if (kProvider == Provider::kReferenceHost) {
|
||||
cutlass::reference::host::Conv3d<
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC ,
|
||||
LayoutC,
|
||||
ElementCompute,
|
||||
ElementAccumulator,
|
||||
ConvertOp,
|
||||
InnerProductOp
|
||||
>(
|
||||
kConvolutionalOperator,
|
||||
config.problem_size,
|
||||
{ptr_A, config.layout_a(conv_kind)},
|
||||
{ptr_B, config.layout_b(conv_kind)},
|
||||
{ptr_C, config.layout_c(conv_kind)},
|
||||
{ptr_D, config.layout_c(conv_kind)},
|
||||
alpha,
|
||||
beta
|
||||
);
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
else if (kProvider == Provider::kReferenceDevice) {
|
||||
return cutlass::reference::device::Conv3d<
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementCompute,
|
||||
ElementAccumulator,
|
||||
ConvertOp,
|
||||
InnerProductOp
|
||||
>(
|
||||
kConvolutionalOperator,
|
||||
config.problem_size,
|
||||
{ptr_A, config.layout_a(conv_kind)},
|
||||
{ptr_B, config.layout_b(conv_kind)},
|
||||
{ptr_C, config.layout_c(conv_kind)},
|
||||
{ptr_D, config.layout_c(conv_kind)},
|
||||
alpha,
|
||||
beta,
|
||||
stream
|
||||
);
|
||||
}
|
||||
return Status::kErrorNotSupported;
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace detail
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
Provider Provider_,
|
||||
conv::Operator ConvolutionalOperator,
|
||||
int ConvDim,
|
||||
typename ElementA_,
|
||||
typename LayoutA_,
|
||||
typename ElementB_,
|
||||
typename LayoutB_,
|
||||
typename ElementC_,
|
||||
typename LayoutC_,
|
||||
typename ElementCompute_,
|
||||
typename ElementAccumulator_ = ElementCompute_,
|
||||
typename ConvertOp_ = NumericConverter<ElementC_, ElementCompute_>,
|
||||
typename InnerProductOp_ = multiply_add<ElementAccumulator_>
|
||||
>
|
||||
class ConvReferenceOperation : public Operation {
|
||||
public:
|
||||
static Provider const kProvider = Provider_;
|
||||
static conv::Operator const kConvolutionalOperator = ConvolutionalOperator;
|
||||
static int const kConvDim = ConvDim;
|
||||
|
||||
using ElementA = ElementA_;
|
||||
using LayoutA = LayoutA_;
|
||||
using ElementB = ElementB_;
|
||||
using LayoutB = LayoutB_;
|
||||
using ElementC = ElementC_;
|
||||
using LayoutC = LayoutC_;
|
||||
using ElementCompute = ElementCompute_;
|
||||
using ElementAccumulator = ElementAccumulator_;
|
||||
using ConvertOp = ConvertOp_;
|
||||
using InnerProductOp = InnerProductOp_;
|
||||
|
||||
protected:
|
||||
|
||||
/// Storage for the name string
|
||||
std::string name_;
|
||||
|
||||
///
|
||||
ConvDescription description_;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructor
|
||||
ConvReferenceOperation() {
|
||||
|
||||
// Basic information
|
||||
description_.provider = kProvider;
|
||||
description_.kind = (kConvDim == 2 ? OperationKind::kConv2d : OperationKind::kConv3d);
|
||||
description_.conv_kind = ConvKindMap<kConvolutionalOperator>::kId;
|
||||
description_.conv_dim = kConvDim;
|
||||
|
||||
// Tensor description
|
||||
description_.A = make_TensorDescription<ElementA, LayoutA>();
|
||||
description_.B = make_TensorDescription<ElementB, LayoutB>();
|
||||
description_.C = make_TensorDescription<ElementC, LayoutC>();
|
||||
|
||||
// Epilogue compute and accumulator type description
|
||||
description_.element_epilogue = NumericTypeMap<ElementCompute>::kId;
|
||||
|
||||
description_.tile_description.math_instruction.element_accumulator =
|
||||
NumericTypeMap<ElementAccumulator>::kId;
|
||||
|
||||
// Iterator algorithm for convolution reference
|
||||
description_.iterator_algorithm = IteratorAlgorithmID::kNone;
|
||||
|
||||
// Compute capability for convolution reference
|
||||
description_.tile_description.minimum_compute_capability =
|
||||
(kProvider == Provider::kReferenceDevice ? 50 : 0);
|
||||
|
||||
description_.tile_description.maximum_compute_capability = 1024;
|
||||
|
||||
// Procedural name
|
||||
std::stringstream ss;
|
||||
|
||||
ss << "conv" << kConvDim << "d_" << to_string(description_.conv_kind)
|
||||
<< "_reference_" << to_string(description_.provider)
|
||||
<< "_" << to_string(description_.A.element) << to_string(description_.A.layout)
|
||||
<< "_" << to_string(description_.B.element) << to_string(description_.B.layout)
|
||||
<< "_" << to_string(description_.C.element) << to_string(description_.C.layout)
|
||||
<< "_" << to_string(description_.tile_description.math_instruction.element_accumulator);
|
||||
|
||||
name_ = ss.str();
|
||||
|
||||
description_.name = name_.c_str();
|
||||
|
||||
// Epilogue compute and accumulator type description
|
||||
description_.element_epilogue = NumericTypeMap<ElementCompute>::kId;
|
||||
|
||||
description_.tile_description.math_instruction.element_accumulator =
|
||||
NumericTypeMap<ElementAccumulator>::kId;
|
||||
}
|
||||
|
||||
/// Returns the description of the GEMM operation
|
||||
virtual OperationDescription const & description() const {
|
||||
return description_;
|
||||
}
|
||||
|
||||
virtual Status can_implement(
|
||||
void const *configuration,
|
||||
void const *arguments) const {
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
virtual uint64_t get_host_workspace_size(
|
||||
void const *configuration) const {
|
||||
|
||||
switch (kConvDim) {
|
||||
case 2:
|
||||
return sizeof(Conv2dConfiguration);
|
||||
case 3:
|
||||
return sizeof(Conv3dConfiguration);
|
||||
default:
|
||||
break;
|
||||
}
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual uint64_t get_device_workspace_size(
|
||||
void const *configuration) const {
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual Status initialize(
|
||||
void const *configuration,
|
||||
void *host_workspace,
|
||||
void *device_workspace = nullptr,
|
||||
cudaStream_t stream = nullptr) const {
|
||||
|
||||
std::memcpy(host_workspace, configuration, get_host_workspace_size(configuration));
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
virtual Status run(
|
||||
void const *arguments,
|
||||
void *host_workspace,
|
||||
void *device_workspace = nullptr,
|
||||
cudaStream_t stream = nullptr) const {
|
||||
|
||||
ConvArguments const &args = *static_cast<ConvArguments const *>(arguments);
|
||||
|
||||
ElementCompute alpha;
|
||||
ElementCompute beta;
|
||||
|
||||
alpha = *static_cast<ElementCompute const *>(args.alpha);
|
||||
beta = *static_cast<ElementCompute const *>(args.beta);
|
||||
|
||||
// TODO - respect pointer mode
|
||||
|
||||
// Invoke 2D or 3D convolution
|
||||
return detail::ConvReferenceDispatcher<
|
||||
kProvider,
|
||||
kConvolutionalOperator,
|
||||
kConvDim,
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementCompute,
|
||||
ElementAccumulator,
|
||||
ConvertOp,
|
||||
InnerProductOp
|
||||
>::dispatch(
|
||||
host_workspace,
|
||||
static_cast<ElementA *>(const_cast<void *>(args.A)),
|
||||
static_cast<ElementB *>(const_cast<void *>(args.B)),
|
||||
static_cast<ElementC *>(const_cast<void *>(args.C)),
|
||||
static_cast<ElementC *>(args.D),
|
||||
alpha,
|
||||
beta,
|
||||
stream
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Constructs Fprop reference operators.
|
||||
template <
|
||||
int kConvDim,
|
||||
typename ElementA_,
|
||||
typename LayoutA_,
|
||||
typename ElementB_,
|
||||
typename LayoutB_,
|
||||
typename ElementC_,
|
||||
typename LayoutC_,
|
||||
typename ElementCompute_,
|
||||
typename ElementAccumulator_ = ElementCompute_,
|
||||
typename ConvertOp_ = NumericConverter<ElementC_, ElementCompute_>,
|
||||
typename InnerProductOp_ = multiply_add<ElementAccumulator_>
|
||||
>
|
||||
void make_conv_fprop(Manifest &manifest) {
|
||||
|
||||
manifest.append(new ConvReferenceOperation<
|
||||
Provider::kReferenceHost,
|
||||
conv::Operator::kFprop,
|
||||
kConvDim,
|
||||
ElementA_, LayoutA_,
|
||||
ElementB_, LayoutB_,
|
||||
ElementC_, LayoutC_,
|
||||
ElementCompute_,
|
||||
ElementAccumulator_,
|
||||
ConvertOp_,
|
||||
InnerProductOp_
|
||||
>);
|
||||
|
||||
manifest.append(new ConvReferenceOperation<
|
||||
Provider::kReferenceDevice,
|
||||
conv::Operator::kFprop,
|
||||
kConvDim,
|
||||
ElementA_, LayoutA_,
|
||||
ElementB_, LayoutB_,
|
||||
ElementC_, LayoutC_,
|
||||
ElementCompute_,
|
||||
ElementAccumulator_,
|
||||
ConvertOp_,
|
||||
InnerProductOp_
|
||||
>);
|
||||
}
|
||||
|
||||
/// Constructs Dgrad and Wgrad reference operators.
|
||||
template <
|
||||
int kConvDim,
|
||||
typename ElementA_,
|
||||
typename LayoutA_,
|
||||
typename ElementB_,
|
||||
typename LayoutB_,
|
||||
typename ElementC_,
|
||||
typename LayoutC_,
|
||||
typename ElementCompute_,
|
||||
typename ElementAccumulator_ = ElementCompute_,
|
||||
typename ConvertOp_ = NumericConverter<ElementC_, ElementCompute_>,
|
||||
typename InnerProductOp_ = multiply_add<ElementAccumulator_>
|
||||
>
|
||||
void make_conv_backwards(Manifest &manifest) {
|
||||
|
||||
manifest.append(new ConvReferenceOperation<
|
||||
Provider::kReferenceHost,
|
||||
conv::Operator::kDgrad,
|
||||
kConvDim,
|
||||
ElementA_, LayoutA_,
|
||||
ElementB_, LayoutB_,
|
||||
ElementC_, LayoutC_,
|
||||
ElementCompute_,
|
||||
ElementAccumulator_,
|
||||
ConvertOp_,
|
||||
InnerProductOp_
|
||||
>);
|
||||
|
||||
manifest.append(new ConvReferenceOperation<
|
||||
Provider::kReferenceDevice,
|
||||
conv::Operator::kDgrad,
|
||||
kConvDim,
|
||||
ElementA_, LayoutA_,
|
||||
ElementB_, LayoutB_,
|
||||
ElementC_, LayoutC_,
|
||||
ElementCompute_,
|
||||
ElementAccumulator_,
|
||||
ConvertOp_,
|
||||
InnerProductOp_
|
||||
>);
|
||||
|
||||
manifest.append(new ConvReferenceOperation<
|
||||
Provider::kReferenceHost,
|
||||
conv::Operator::kWgrad,
|
||||
kConvDim,
|
||||
ElementA_, LayoutA_,
|
||||
ElementB_, LayoutB_,
|
||||
ElementC_, LayoutC_,
|
||||
ElementCompute_,
|
||||
ElementAccumulator_,
|
||||
ConvertOp_,
|
||||
InnerProductOp_
|
||||
>);
|
||||
|
||||
manifest.append(new ConvReferenceOperation<
|
||||
Provider::kReferenceDevice,
|
||||
conv::Operator::kWgrad,
|
||||
kConvDim,
|
||||
ElementA_, LayoutA_,
|
||||
ElementB_, LayoutB_,
|
||||
ElementC_, LayoutC_,
|
||||
ElementCompute_,
|
||||
ElementAccumulator_,
|
||||
ConvertOp_,
|
||||
InnerProductOp_
|
||||
>);
|
||||
}
|
||||
|
||||
/// Six operators for the price of one.
|
||||
template <
|
||||
int kConvDim,
|
||||
typename ElementA_,
|
||||
typename LayoutA_,
|
||||
typename ElementB_,
|
||||
typename LayoutB_,
|
||||
typename ElementC_,
|
||||
typename LayoutC_,
|
||||
typename ElementCompute_,
|
||||
typename ElementAccumulator_ = ElementCompute_,
|
||||
typename ConvertOp_ = NumericConverter<ElementC_, ElementCompute_>,
|
||||
typename InnerProductOp_ = multiply_add<ElementAccumulator_>
|
||||
>
|
||||
void make_conv_all(Manifest &manifest) {
|
||||
|
||||
make_conv_fprop<
|
||||
kConvDim,
|
||||
ElementA_, LayoutA_,
|
||||
ElementB_, LayoutB_,
|
||||
ElementC_, LayoutC_,
|
||||
ElementCompute_,
|
||||
ElementAccumulator_,
|
||||
ConvertOp_,
|
||||
InnerProductOp_
|
||||
>(manifest);
|
||||
|
||||
make_conv_backwards<
|
||||
kConvDim,
|
||||
ElementA_, LayoutA_,
|
||||
ElementB_, LayoutB_,
|
||||
ElementC_, LayoutC_,
|
||||
ElementCompute_,
|
||||
ElementAccumulator_,
|
||||
ConvertOp_,
|
||||
InnerProductOp_
|
||||
>(manifest);
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace library
|
||||
} // namespace cutlass
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -37,10 +37,14 @@ namespace cutlass {
|
||||
namespace library {
|
||||
|
||||
void initialize_gemm_reference_operations(Manifest &manifest);
|
||||
void initialize_conv2d_reference_operations(Manifest &manifest);
|
||||
void initialize_conv3d_reference_operations(Manifest &manifest);
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
void initialize_reference_operations(Manifest &manifest) {
|
||||
initialize_conv2d_reference_operations(manifest);
|
||||
initialize_conv3d_reference_operations(manifest);
|
||||
initialize_gemm_reference_operations(manifest);
|
||||
}
|
||||
|
||||
|
||||
+168
-8
@@ -50,6 +50,7 @@ Provider_enumerants[] = {
|
||||
{"host", "reference_host", Provider::kReferenceHost},
|
||||
{"device", "reference_device", Provider::kReferenceDevice},
|
||||
{"cublas", "cuBLAS", Provider::kCUBLAS},
|
||||
{"cudnn", "cuDNN", Provider::kCUDNN},
|
||||
};
|
||||
|
||||
/// Converts a Provider enumerant to a string
|
||||
@@ -128,6 +129,9 @@ static struct {
|
||||
OperationKind_enumerants[] = {
|
||||
{"eq_gemm", "EqGemm", OperationKind::kEqGemm},
|
||||
{"gemm", "Gemm", OperationKind::kGemm},
|
||||
{"conv2d", "Conv2d", OperationKind::kConv2d},
|
||||
{"conv3d", "Conv3d", OperationKind::kConv3d},
|
||||
{"spgemm", "SparseGemm", OperationKind::kSparseGemm},
|
||||
};
|
||||
|
||||
/// Converts a Status enumerant to a string
|
||||
@@ -445,6 +449,10 @@ layout_aliases[] = {
|
||||
{LayoutTypeID::kTensorNCDHW, "ncdhw"},
|
||||
{LayoutTypeID::kTensorNHWC, "nhwc"},
|
||||
{LayoutTypeID::kTensorNDHWC, "ndhwc"},
|
||||
{LayoutTypeID::kTensorNC32HW32, "nc32hw32"},
|
||||
{LayoutTypeID::kTensorNC64HW64, "nc64hw64"},
|
||||
{LayoutTypeID::kTensorC32RSK32, "c32rsk32"},
|
||||
{LayoutTypeID::kTensorC64RSK64, "c64rsk64"},
|
||||
|
||||
{LayoutTypeID::kUnknown, "*"},
|
||||
{LayoutTypeID::kInvalid, nullptr}
|
||||
@@ -474,22 +482,46 @@ LayoutTypeID from_string<LayoutTypeID>(std::string const &str) {
|
||||
/// Gets stride rank for the layout_id (static function)
|
||||
int get_layout_stride_rank(LayoutTypeID layout_id) {
|
||||
switch (layout_id) {
|
||||
case LayoutTypeID::kColumnMajor: return cutlass::layout::ColumnMajor::kStrideRank;
|
||||
case LayoutTypeID::kRowMajor: return cutlass::layout::RowMajor::kStrideRank;
|
||||
case LayoutTypeID::kColumnMajor:
|
||||
return cutlass::layout::ColumnMajor::kStrideRank;
|
||||
case LayoutTypeID::kRowMajor:
|
||||
return cutlass::layout::RowMajor::kStrideRank;
|
||||
case LayoutTypeID::kColumnMajorInterleavedK2:
|
||||
return cutlass::layout::ColumnMajorInterleaved<2>::kStrideRank;
|
||||
case LayoutTypeID::kRowMajorInterleavedK2:
|
||||
return cutlass::layout::RowMajorInterleaved<2>::kStrideRank;
|
||||
case LayoutTypeID::kColumnMajorInterleavedK4:
|
||||
return cutlass::layout::ColumnMajorInterleaved<4>::kStrideRank;
|
||||
case LayoutTypeID::kRowMajorInterleavedK4:
|
||||
return cutlass::layout::RowMajorInterleaved<4>::kStrideRank;
|
||||
case LayoutTypeID::kColumnMajorInterleavedK16:
|
||||
return cutlass::layout::ColumnMajorInterleaved<16>::kStrideRank;
|
||||
case LayoutTypeID::kRowMajorInterleavedK16:
|
||||
return cutlass::layout::RowMajorInterleaved<16>::kStrideRank;
|
||||
case LayoutTypeID::kColumnMajorInterleavedK32:
|
||||
return cutlass::layout::ColumnMajorInterleaved<32>::kStrideRank;
|
||||
case LayoutTypeID::kRowMajorInterleavedK32:
|
||||
return cutlass::layout::RowMajorInterleaved<32>::kStrideRank;
|
||||
case LayoutTypeID::kColumnMajorInterleavedK64:
|
||||
case LayoutTypeID::kRowMajorInterleavedK64: return 1;
|
||||
return cutlass::layout::ColumnMajorInterleaved<64>::kStrideRank;
|
||||
case LayoutTypeID::kRowMajorInterleavedK64:
|
||||
return cutlass::layout::RowMajorInterleaved<64>::kStrideRank;
|
||||
case LayoutTypeID::kTensorNCHW:
|
||||
case LayoutTypeID::kTensorNHWC: return 3;
|
||||
case LayoutTypeID::kTensorNDHWC: return 4;
|
||||
default : throw std::runtime_error("Unsupported LayoutTypeID in LayoutType::get_stride_rank");
|
||||
return cutlass::layout::TensorNCHW::kStrideRank;
|
||||
case LayoutTypeID::kTensorNHWC:
|
||||
return cutlass::layout::TensorNHWC::kStrideRank;
|
||||
case LayoutTypeID::kTensorNDHWC:
|
||||
return cutlass::layout::TensorNDHWC::kStrideRank;
|
||||
case LayoutTypeID::kTensorNC32HW32:
|
||||
return cutlass::layout::TensorNCxHWx<32>::kStrideRank;
|
||||
case LayoutTypeID::kTensorNC64HW64:
|
||||
return cutlass::layout::TensorNCxHWx<64>::kStrideRank;
|
||||
case LayoutTypeID::kTensorC32RSK32:
|
||||
return cutlass::layout::TensorCxRSKx<32>::kStrideRank;
|
||||
case LayoutTypeID::kTensorC64RSK64:
|
||||
return cutlass::layout::TensorCxRSKx<64>::kStrideRank;
|
||||
default:
|
||||
throw std::runtime_error("Unsupported LayoutTypeID in LayoutType::get_stride_rank");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -624,6 +656,136 @@ SplitKMode from_string<SplitKMode>(std::string const &str) {
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
static struct {
|
||||
char const *text;
|
||||
char const *pretty;
|
||||
ConvModeID enumerant;
|
||||
}
|
||||
ConvModeID_enumerants[] = {
|
||||
{"cross", "<cross>", ConvModeID::kCrossCorrelation},
|
||||
{"conv", "<conv>", ConvModeID::kConvolution},
|
||||
};
|
||||
|
||||
/// Converts a ConvModeID enumerant to a string
|
||||
char const *to_string(ConvModeID type, bool pretty) {
|
||||
|
||||
for (auto const & possible : ConvModeID_enumerants) {
|
||||
if (type == possible.enumerant) {
|
||||
if (pretty) {
|
||||
return possible.pretty;
|
||||
}
|
||||
else {
|
||||
return possible.text;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return pretty ? "Invalid" : "invalid";
|
||||
}
|
||||
|
||||
/// Converts a ConvModeID enumerant from a string
|
||||
template <>
|
||||
ConvModeID from_string<ConvModeID>(std::string const &str) {
|
||||
|
||||
for (auto const & possible : ConvModeID_enumerants) {
|
||||
if ((str.compare(possible.text) == 0) ||
|
||||
(str.compare(possible.pretty) == 0)) {
|
||||
return possible.enumerant;
|
||||
}
|
||||
}
|
||||
|
||||
return ConvModeID::kInvalid;
|
||||
}
|
||||
|
||||
|
||||
static struct {
|
||||
char const *text;
|
||||
char const *pretty;
|
||||
IteratorAlgorithmID enumerant;
|
||||
}
|
||||
IteratorAlgorithmID_enumerants[] = {
|
||||
{"none", "<none>", IteratorAlgorithmID::kNone},
|
||||
{"analytic", "<analytic>", IteratorAlgorithmID::kAnalytic},
|
||||
{"optimized", "<optimized>", IteratorAlgorithmID::kOptimized},
|
||||
};
|
||||
|
||||
/// Converts a ConvModeID enumerant to a string
|
||||
char const *to_string(IteratorAlgorithmID type, bool pretty) {
|
||||
|
||||
for (auto const & possible : IteratorAlgorithmID_enumerants) {
|
||||
if (type == possible.enumerant) {
|
||||
if (pretty) {
|
||||
return possible.pretty;
|
||||
}
|
||||
else {
|
||||
return possible.text;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return pretty ? "Invalid" : "invalid";
|
||||
}
|
||||
|
||||
/// Converts a ConvModeID enumerant from a string
|
||||
template <>
|
||||
IteratorAlgorithmID from_string<IteratorAlgorithmID>(std::string const &str) {
|
||||
|
||||
for (auto const & possible : IteratorAlgorithmID_enumerants) {
|
||||
if ((str.compare(possible.text) == 0) ||
|
||||
(str.compare(possible.pretty) == 0)) {
|
||||
return possible.enumerant;
|
||||
}
|
||||
}
|
||||
|
||||
return IteratorAlgorithmID::kInvalid;
|
||||
}
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
static struct {
|
||||
char const *text;
|
||||
char const *pretty;
|
||||
ConvKind enumerant;
|
||||
}
|
||||
ConvKind_enumerants[] = {
|
||||
{"unknown", "<unkown>", ConvKind::kUnknown},
|
||||
{"fprop", "<fprop>", ConvKind::kFprop},
|
||||
{"dgrad", "<dgrad>", ConvKind::kDgrad},
|
||||
{"wgrad", "<wgrad>", ConvKind::kWgrad},
|
||||
};
|
||||
|
||||
/// Converts a ConvKind enumerant to a string
|
||||
char const *to_string(ConvKind type, bool pretty) {
|
||||
|
||||
for (auto const & possible : ConvKind_enumerants) {
|
||||
if (type == possible.enumerant) {
|
||||
if (pretty) {
|
||||
return possible.pretty;
|
||||
}
|
||||
else {
|
||||
return possible.text;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return pretty ? "Invalid" : "invalid";
|
||||
}
|
||||
|
||||
|
||||
/// Converts a ConvKind enumerant from a string
|
||||
template <>
|
||||
ConvKind from_string<ConvKind>(std::string const &str) {
|
||||
|
||||
for (auto const & possible : ConvKind_enumerants) {
|
||||
if ((str.compare(possible.text) == 0) ||
|
||||
(str.compare(possible.pretty) == 0)) {
|
||||
return possible.enumerant;
|
||||
}
|
||||
}
|
||||
|
||||
return ConvKind::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;
|
||||
@@ -1224,5 +1386,3 @@ bool cast_from_double(std::vector<uint8_t> &bytes, NumericTypeID type, double sr
|
||||
} // namespace cutlass
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user