mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
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:
parent
e6ef3653a7
commit
49febba3c2
4 changed files with 71 additions and 55 deletions
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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)>,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@
|
|||
|
||||
namespace onnxruntime {
|
||||
namespace cuda {
|
||||
template <typename T>
|
||||
|
||||
struct Tile final : CudaKernel {
|
||||
Tile(const OpKernelInfo& info) : CudaKernel(info) {
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue