Support int32 and int64 types for Tile CUDA kernel (#4684)

* Support int32 and int64 types for Tile CUDA kernel

* Fix build
This commit is contained in:
Hariharan Seshadri 2020-08-03 17:47:37 -07:00 committed by GitHub
parent e6ef3653a7
commit 49febba3c2
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
4 changed files with 71 additions and 55 deletions

View file

@ -24,7 +24,10 @@ inline bool AreVectorsOverlap(const std::vector<T>& v1, const std::vector<T>& v2
}
} // namespace
//TODO: Tell user why it has conflicts
// TODO: Tell user why it has conflicts
// TODO: Investigate why IsConflict() was not triggered when there were duplicate Tile CUDA
// kernels registered. Removing `InputMemoryType<OrtMemTypeCPUInput>(1)` in the kernel definition
// triggered the conflict.
bool KernelDef::IsConflict(const KernelDef& other) const {
if (op_name_ != other.OpName() || provider_type_ != other.Provider())
return false;

View file

@ -67,7 +67,7 @@ CUDAExecutionProvider::PerThreadContext::PerThreadContext(OrtDevice::DeviceId de
[](OrtDevice::DeviceId id) {
return onnxruntime::make_unique<CUDAAllocator>(id, CUDA);
},
cuda_mem_limit,
cuda_mem_limit,
arena_extend_strategy});
// CUDA malloc/free is expensive so always use an arena
@ -105,19 +105,17 @@ CUDAExecutionProvider::PerThreadContext::~PerThreadContext() {
void CUDAExecutionProvider::UpdateProviderOptionsInfo() {
UnorderedMapStringToString options;
options["device_id"] = std::to_string(device_id_);
options["cuda_mem_limit"] = std::to_string(cuda_mem_limit_);
options["device_id"] = std::to_string(device_id_);
options["cuda_mem_limit"] = std::to_string(cuda_mem_limit_);
std::string strategy;
if (arena_extend_strategy_ == ArenaExtendStrategy::kNextPowerOfTwo) {
strategy = "kNextPowerOfTwo";
}
else if (arena_extend_strategy_ == ArenaExtendStrategy::kSameAsRequested) {
} else if (arena_extend_strategy_ == ArenaExtendStrategy::kSameAsRequested) {
strategy = "kSameAsRequested";
} else {
strategy = "unknown";
}
else {
strategy = "unknown";
}
options["arena_extend_strategy"] = strategy;
options["arena_extend_strategy"] = strategy;
IExecutionProvider::SetProviderOptions(options);
}
@ -347,9 +345,6 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, double, MatMul);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, MLFloat16, MatMul);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, int8_t, MatMulInteger);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, float, Tile);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, double, Tile);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, MLFloat16, Tile);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, float, Elu);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, double, Elu);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, MLFloat16, Elu);
@ -581,9 +576,9 @@ class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kO
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 5, Reshape);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, 4, Reshape_1);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, Shape);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, float, Tile);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, double, Tile);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, MLFloat16, Tile);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, Tile);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, Tile);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, Tile);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, Transpose);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, float, InstanceNormalization);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, double, InstanceNormalization);
@ -834,9 +829,6 @@ static Status RegisterCudaKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 9, MLFloat16, MatMul)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, int8_t, MatMulInteger)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, 10, float, Clip)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, float, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, double, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, MLFloat16, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, float, Elu)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, double, Elu)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, MLFloat16, Elu)>,
@ -1067,9 +1059,9 @@ static Status RegisterCudaKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 5, Reshape)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, 4, Reshape_1)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, Shape)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, float, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, double, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, MLFloat16, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, Transpose)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, float, InstanceNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 6, double, InstanceNormalization)>,

View file

@ -8,21 +8,22 @@ using namespace onnxruntime::common;
namespace onnxruntime {
namespace cuda {
#define REGISTER_KERNEL_TYPED(T) \
ONNX_OPERATOR_TYPED_KERNEL_EX( \
Tile, \
kOnnxDomain, \
6, \
T, \
kCudaExecutionProvider, \
KernelDefBuilder() \
.InputMemoryType<OrtMemTypeCPUInput>(1) \
.TypeConstraint("T", DataTypeImpl::GetTensorType<T>()) \
.TypeConstraint("T1", DataTypeImpl::GetTensorType<int64_t>()), \
Tile<T>);
ONNX_OPERATOR_KERNEL_EX(
Tile,
kOnnxDomain,
6,
kCudaExecutionProvider,
KernelDefBuilder()
.InputMemoryType<OrtMemTypeCPUInput>(1)
.TypeConstraint("T", {DataTypeImpl::GetTensorType<float>(),
DataTypeImpl::GetTensorType<double>(),
DataTypeImpl::GetTensorType<int32_t>(),
DataTypeImpl::GetTensorType<int64_t>(),
DataTypeImpl::GetTensorType<MLFloat16>()})
.TypeConstraint("T1", DataTypeImpl::GetTensorType<int64_t>()),
Tile);
template <typename T>
Status Tile<T>::ComputeInternal(OpKernelContext* ctx) const {
Status Tile::ComputeInternal(OpKernelContext* ctx) const {
auto& input_tensor = *ctx->Input<Tensor>(0);
auto& repeats_tensor = *ctx->Input<Tensor>(1);
int32_t rank = static_cast<int32_t>(input_tensor.Shape().NumDimensions());
@ -41,8 +42,8 @@ Status Tile<T>::ComputeInternal(OpKernelContext* ctx) const {
TensorShape outputShape(output_dims);
auto& output_tensor = *ctx->Output(0, outputShape);
T* output_data = output_tensor.template MutableData<T>();
const T* input_data = input_tensor.template Data<T>();
void* output_data = output_tensor.MutableDataRaw();
const void* input_data = input_tensor.DataRaw();
TensorPitches input_pitches(input_shape);
TArray<int64_t> input_strides(input_pitches);
@ -58,27 +59,47 @@ Status Tile<T>::ComputeInternal(OpKernelContext* ctx) const {
fdm_output_strides[i] = fast_divmod(static_cast<int>(output_pitches[i]));
}
static_assert(sizeof(float) == sizeof(int32_t), "Float and Int32 are of different sizes");
static_assert(sizeof(double) == sizeof(int64_t), "Double and Int64 are of different sizes");
if (output_tensor.Shape().Size() > 0) {
TileImpl(
rank,
fdm_input_shape,
input_strides,
reinterpret_cast<const typename ToCudaType<T>::MappedType*>(input_data),
fdm_output_strides,
reinterpret_cast<typename ToCudaType<T>::MappedType*>(output_data),
output_tensor.Shape().Size());
if (input_tensor.IsDataType<float>() ||
input_tensor.IsDataType<int32_t>()) {
TileImpl(
rank,
fdm_input_shape,
input_strides,
reinterpret_cast<const typename ToCudaType<float>::MappedType*>(input_data),
fdm_output_strides,
reinterpret_cast<typename ToCudaType<float>::MappedType*>(output_data),
output_tensor.Shape().Size());
} else if (input_tensor.IsDataType<double>() ||
input_tensor.IsDataType<int64_t>()) {
TileImpl(
rank,
fdm_input_shape,
input_strides,
reinterpret_cast<const typename ToCudaType<double>::MappedType*>(input_data),
fdm_output_strides,
reinterpret_cast<typename ToCudaType<double>::MappedType*>(output_data),
output_tensor.Shape().Size());
} else if (input_tensor.IsDataType<MLFloat16>()) {
TileImpl(
rank,
fdm_input_shape,
input_strides,
reinterpret_cast<const typename ToCudaType<MLFloat16>::MappedType*>(input_data),
fdm_output_strides,
reinterpret_cast<typename ToCudaType<MLFloat16>::MappedType*>(output_data),
output_tensor.Shape().Size());
} else {
// Won't hit this as the kernel doesn't claim support for any type that will trigger this
ORT_THROW("Tile doesn't have an implementation yet for the type: ", input_tensor.DataType());
}
}
return Status::OK();
}
#define SPECIALIZED_COMPUTE(T) \
REGISTER_KERNEL_TYPED(T) \
template Status Tile<T>::ComputeInternal(OpKernelContext* ctx) const;
SPECIALIZED_COMPUTE(float)
SPECIALIZED_COMPUTE(double)
SPECIALIZED_COMPUTE(MLFloat16)
} // namespace cuda
} // namespace onnxruntime

View file

@ -7,7 +7,7 @@
namespace onnxruntime {
namespace cuda {
template <typename T>
struct Tile final : CudaKernel {
Tile(const OpKernelInfo& info) : CudaKernel(info) {
}