From 2fa10fb8033561a2021f161a7f9bc30b15d4ad9e Mon Sep 17 00:00:00 2001 From: Chen Fu <1316708+chenfucn@users.noreply.github.com> Date: Tue, 25 Apr 2023 14:01:47 -0700 Subject: [PATCH] Fp16 onnx pool operators, relu, leakyrelu (#15498) ### Description Adding the fp16 onnx operator implementations: maxpool, averagepool, global average pool, relu, leaky relu ### Motivation and Context Continue with support for fp16. Standard onnx operator implementations are needed as a basis for the graph optimizers to work. --- .../contrib_ops/cpu/cpu_contrib_kernels.cc | 4 +- .../providers/cpu/activation/activations.cc | 11 +- .../providers/cpu/cpu_execution_provider.cc | 57 +- .../providers/cpu/fp16/fp16_activations.h | 83 +++ .../providers/cpu/fp16/fp16_pool.cc} | 176 ++++-- .../test/contrib_ops/nhwc_pool_in_op_test.cc | 2 +- .../cpu/activation/activation_op_test.cc | 35 ++ .../cpu/activation/activation_op_test.h | 12 + .../providers/cpu/nn/pool_fp16_op_test.cc | 570 ++++++++++++++++++ tools/ci_build/op_registration_validator.py | 2 +- 10 files changed, 903 insertions(+), 49 deletions(-) create mode 100644 onnxruntime/core/providers/cpu/fp16/fp16_activations.h rename onnxruntime/{contrib_ops/cpu/fp16/nhwc_pool_fp16.cc => core/providers/cpu/fp16/fp16_pool.cc} (57%) create mode 100644 onnxruntime/test/providers/cpu/nn/pool_fp16_op_test.cc diff --git a/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc b/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc index f6a0501bf7..31b1b12d38 100644 --- a/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc @@ -80,7 +80,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, #ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, MLFloat16, NhwcFusedConv); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSInternalNHWCDomain, 11, MLFloat16, MaxPool); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSInternalNHWCDomain, 12, MLFloat16, MaxPool); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSInternalNHWCDomain, 11, MLFloat16, AveragePool); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSInternalNHWCDomain, 1, MLFloat16, GlobalAveragePool); #endif @@ -160,7 +160,7 @@ Status RegisterFp16Kernels(KernelRegistry& kernel_registry) { static const BuildKernelCreateInfoFn function_table[] = { BuildKernelCreateInfo, // default entry to avoid the list become empty after ops-reducing BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, }; diff --git a/onnxruntime/core/providers/cpu/activation/activations.cc b/onnxruntime/core/providers/cpu/activation/activations.cc index e16bec0880..049fee4b95 100644 --- a/onnxruntime/core/providers/cpu/activation/activations.cc +++ b/onnxruntime/core/providers/cpu/activation/activations.cc @@ -1,12 +1,13 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +#include "core/mlas/inc/mlas.h" #include "core/providers/cpu/activation/activations.h" +#include "core/providers/cpu/fp16/fp16_activations.h" #include "core/providers/cpu/math/element_wise_ops.h" #ifndef DISABLE_CONTRIB_OPS #include "contrib_ops/cpu/activations.h" #endif -#include "core/mlas/inc/mlas.h" using namespace onnxruntime::common; @@ -43,6 +44,14 @@ REGISTER_UNARY_ELEMENTWISE_TYPED_KERNEL(Relu, 14, float); REGISTER_UNARY_ELEMENTWISE_TYPED_KERNEL(Relu, 14, double); REGISTER_UNARY_ELEMENTWISE_TYPED_KERNEL(Relu, 14, int8_t); REGISTER_UNARY_ELEMENTWISE_TYPED_KERNEL(Relu, 14, int32_t); +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED +REGISTER_VERSIONED_UNARY_ELEMENTWISE_TYPED_KERNEL(Relu, 6, 12, MLFloat16); +REGISTER_VERSIONED_UNARY_ELEMENTWISE_TYPED_KERNEL(Relu, 13, 13, MLFloat16); +REGISTER_UNARY_ELEMENTWISE_TYPED_KERNEL(Relu, 14, MLFloat16); +REGISTER_VERSIONED_UNARY_ELEMENTWISE_TYPED_KERNEL(LeakyRelu, 6, 15, MLFloat16); +REGISTER_UNARY_ELEMENTWISE_TYPED_KERNEL(LeakyRelu, 16, MLFloat16); +#endif // MLAS_F16VEC_INTRINSICS_SUPPORTED + REGISTER_UNARY_ELEMENTWISE_KERNEL(Selu, 6); REGISTER_VERSIONED_UNARY_ELEMENTWISE_TYPED_KERNEL(Sigmoid, 6, 12, float); REGISTER_VERSIONED_UNARY_ELEMENTWISE_TYPED_KERNEL(Sigmoid, 6, 12, double); diff --git a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc index 7f99f069f2..8e1b412dd2 100644 --- a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc +++ b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc @@ -69,6 +69,10 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, Har class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 15, LeakyRelu); class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12, float, Relu); class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12, double, Relu); +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12, MLFloat16, Relu); +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 15, MLFloat16, LeakyRelu); +#endif class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, Selu); class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12, float, Sigmoid); class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12, double, Sigmoid); @@ -179,9 +183,15 @@ class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDoma class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, AveragePool); class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 7, MaxPool); class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 8, 11, MaxPool); +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 8, 11, MLFloat16, MaxPool); +#endif class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, 10, LpPool); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, GlobalLpPool); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, GlobalAveragePool); +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, MLFloat16, GlobalAveragePool); +#endif class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, GlobalMaxPool); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, MaxRoiPool); class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, float, ReduceL1); @@ -456,6 +466,7 @@ class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDoma class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Conv); #ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MLFloat16, Conv); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MLFloat16, AveragePool); #endif class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, ConvTranspose); class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 12, If); @@ -504,6 +515,9 @@ class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDoma class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, 12, Min); class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, 12, Max); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, MaxPool); +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, MLFloat16, MaxPool); +#endif class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, 12, Pow); class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, 12, float, ReduceMax); class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, 12, double, ReduceMax); @@ -650,6 +664,9 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, double, Sqrt); class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, 13, float, Relu); class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, 13, double, Relu); +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED +class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, 13, MLFloat16, Relu); +#endif class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, float, Sigmoid); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, double, Sigmoid); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, float, Tanh); @@ -742,6 +759,9 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 14, double, Relu); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 14, int8_t, Relu); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 14, int32_t, Relu); +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 14, MLFloat16, Relu); +#endif class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 14, Trilu); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 14, float, Add); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 14, double, Add); @@ -795,6 +815,9 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 16, int64_t, Where); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 16, uint8_t, Where); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 16, LeakyRelu); +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 16, MLFloat16, LeakyRelu); +#endif class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 16, PRelu); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 16, Scan); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 16, float, GreaterOrEqual); @@ -1510,9 +1533,6 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, -#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED - BuildKernelCreateInfo, -#endif BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -2288,6 +2308,32 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { return Status::OK(); } +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED +Status RegisterFp16Kernels(KernelRegistry& kernel_registry) { + static const BuildKernelCreateInfoFn function_table[] = { + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + }; + + for (auto& function_table_entry : function_table) { + KernelCreateInfo info = function_table_entry(); + if (info.kernel_def != nullptr) { // filter disabled entries where type is void + ORT_RETURN_IF_ERROR(kernel_registry.Register(std::move(info))); + } + } + + return Status::OK(); +} +#endif + // Forward declarations of ml op kernels #ifndef DISABLE_ML_OPS namespace ml { @@ -2464,6 +2510,11 @@ Status RegisterOnnxMLOperatorKernels(KernelRegistry& kernel_registry) { Status RegisterCPUKernels(KernelRegistry& kernel_registry) { ORT_RETURN_IF_ERROR(RegisterOnnxOperatorKernels(kernel_registry)); +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED + if (MlasFp16AccelerationSupported()) { + ORT_RETURN_IF_ERROR(RegisterFp16Kernels(kernel_registry)); + } +#endif #ifndef DISABLE_ML_OPS ORT_RETURN_IF_ERROR(::onnxruntime::ml::RegisterOnnxMLOperatorKernels(kernel_registry)); #endif diff --git a/onnxruntime/core/providers/cpu/fp16/fp16_activations.h b/onnxruntime/core/providers/cpu/fp16/fp16_activations.h new file mode 100644 index 0000000000..16bcd171e3 --- /dev/null +++ b/onnxruntime/core/providers/cpu/fp16/fp16_activations.h @@ -0,0 +1,83 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/mlas/inc/mlas.h" +#include "core/framework/float16.h" +#include "core/providers/cpu/activation/activations.h" + +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED + +namespace onnxruntime { +namespace functors { + +template <> +struct Relu : public ElementWiseRangedTransform { + MLAS_ACTIVATION Activation; + Status Init(const onnxruntime::NodeAttributes&) { + Activation.ActivationKind = MlasReluActivation; + return Status::OK(); + } + GSL_SUPPRESS(r .11) + ElementWiseRangedTransform* Copy() const final { + using T1 = typename std::remove_pointer::type; + using T2 = typename std::remove_const::type; // redundant? + return new T2(*this); + } + float Cost() const final { + return 1.0f; + } + void operator()(std::ptrdiff_t first, std::ptrdiff_t last) const final { + ptrdiff_t len = last - first; + MLFloat16* output_ptr = this->output + first; + const MLFloat16* input_ptr = this->input + first; + + // Linux compilation pipeline complained memcpy_s does not exists?! + // memcpy_s(output_ptr, len * sizeof(MLFloat16), input_ptr, len * sizeof(MLFloat16)); + memcpy(output_ptr, input_ptr, len * sizeof(MLFloat16)); + + MlasFp16Activation(&Activation, output_ptr, 1, len, len); + } +}; + +template <> +struct LeakyRelu : public ElementWiseRangedTransform { + MLAS_ACTIVATION Activation; + Status Init(const onnxruntime::NodeAttributes& attributes) { + Activation.ActivationKind = MlasLeakyReluActivation; + return (GetFloatParam("alpha", attributes, Activation.Parameters.LeakyRelu.alpha)); + } + GSL_SUPPRESS(r .11) + ElementWiseRangedTransform* Copy() const final { + using T1 = typename std::remove_pointer::type; + using T2 = typename std::remove_const::type; + return new T2(*this); + }; + + float Cost() const final { + return 2.0f; + } + + void operator()(std::ptrdiff_t first, std::ptrdiff_t last) const final { + ptrdiff_t len = last - first; + MLFloat16* output_ptr = this->output + first; + const MLFloat16* input_ptr = this->input + first; + // Linux compilation pipeline complained memcpy_s does not exists?! + // memcpy_s(output_ptr, len * sizeof(MLFloat16), input_ptr, len * sizeof(MLFloat16)); + memcpy(output_ptr, input_ptr, len * sizeof(MLFloat16)); + + MlasFp16Activation(&Activation, output_ptr, 1, len, len); + } +}; + +// TODO Add the following activations: +// MlasTanhActivation, +// MlasLogisticActivation, +// MlasClipActivation, +// MlasHardSigmoidActivation, + +} // namespace functors +} // namespace onnxruntime + +#endif // MLAS_F16VEC_INTRINSICS_SUPPORTED diff --git a/onnxruntime/contrib_ops/cpu/fp16/nhwc_pool_fp16.cc b/onnxruntime/core/providers/cpu/fp16/fp16_pool.cc similarity index 57% rename from onnxruntime/contrib_ops/cpu/fp16/nhwc_pool_fp16.cc rename to onnxruntime/core/providers/cpu/fp16/fp16_pool.cc index b8d11a5aa6..68762a4c8e 100644 --- a/onnxruntime/contrib_ops/cpu/fp16/nhwc_pool_fp16.cc +++ b/onnxruntime/core/providers/cpu/fp16/fp16_pool.cc @@ -1,20 +1,20 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -#include "core/common/common.h" -#include "core/framework/op_kernel.h" -#include "core/providers/cpu/nn/pool_attributes.h" -#include "core/common/safeint.h" -#include "core/util/math.h" #include "core/mlas/inc/mlas.h" -namespace onnxruntime { -namespace contrib { - #ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED +#include "core/framework/op_kernel.h" +#include "core/util/math.h" +#include "core/providers/cpu/nn/pool_attributes.h" +#include "core/platform/threadpool.h" +#include "core/common/safeint.h" + +namespace onnxruntime { + /** - * @brief Pooling operator for type FP16, layout NHWC. + * @brief Pooling operator for type FP16, * Only max pool and average pool supported. * * Single threadded operation for now. @@ -22,21 +22,23 @@ namespace contrib { * TODO!! implemente thread partition similar with * fp16 conv operator */ -class NhwcPoolFp16 : public OpKernel { +class PoolFp16 : public OpKernel { public: - explicit NhwcPoolFp16(const OpKernelInfo& info) + explicit PoolFp16(const OpKernelInfo& info) : OpKernel(info), pool_attrs_(info, info.GetKernelDef().OpName(), info.node().SinceVersion()), - is_max_pool_(info.GetKernelDef().OpName() == "MaxPool") {} + is_max_pool_(info.GetKernelDef().OpName() == "MaxPool"), + channels_last_(info.GetKernelDef().Domain() == kMSInternalNHWCDomain) {} Status Compute(OpKernelContext* context) const override; protected: PoolAttributes pool_attrs_; bool is_max_pool_; // either max pool or average pool + bool channels_last_; }; -Status NhwcPoolFp16::Compute(OpKernelContext* context) const { +Status PoolFp16::Compute(OpKernelContext* context) const { const auto* X = context->Input(0); const TensorShape& input_shape = X->Shape(); @@ -44,31 +46,55 @@ Status NhwcPoolFp16::Compute(OpKernelContext* context) const { ORT_RETURN_IF_NOT(input_rank >= 3, "Input dimension cannot be less than 3."); const int64_t N = input_shape[0]; - const int64_t C = input_shape[input_rank - 1]; + const int64_t C = channels_last_ ? input_shape[input_rank - 1] : input_shape[1]; ORT_ENFORCE(input_shape.Size() > 0 || N == 0, "Invalid input shape. Only N can be zero. Got:", input_shape); const size_t spatial_dims = input_rank - 2; + const size_t spatial_dim_start = channels_last_ ? 1 : 2; // Compute the output size and effective padding for this pooling operation. TensorShapeVector output_dims({N}); + if (!channels_last_) { + output_dims.push_back(C); + } TensorShapeVector pads = pool_attrs_.pads; TensorShapeVector kernel_shape = pool_attrs_.kernel_shape; TensorShapeVector strides = pool_attrs_.strides; TensorShapeVector dilations = pool_attrs_.dilations; if (pool_attrs_.global_pooling) { const auto& input_dims = input_shape.GetDims(); - kernel_shape.assign(input_dims.begin() + 1, input_dims.end() - 1); + if (channels_last_) { + kernel_shape.assign(input_dims.begin() + 1, input_dims.end() - 1); + } else { + kernel_shape.assign(input_dims.begin() + 2, input_dims.end()); + } pads.resize(kernel_shape.size() * 2, 0); strides.resize(kernel_shape.size(), 1); dilations.resize(kernel_shape.size(), 1); } + if (kernel_shape.size() != spatial_dims) { + std::ostringstream ss; + ss << "Invalid kernel shape. Input shape "; + ss << (channels_last_ ? "(NHWC):[" : "(NCHW):["); + for (int64_t i = 0; i < input_shape.Size(); i++) { + ss << input_shape[i] << ", "; + } + ss << "] Kernel shape:["; + for (size_t i = 0; i < kernel_shape.size(); i++) { + ss << kernel_shape[i] << ", "; + } + ss << "]"; + + ORT_THROW(ss.str()); + } + int64_t kernel_size = 1; int64_t input_image_size = 1; int64_t output_image_size = 1; for (size_t dim = 0; dim < spatial_dims; ++dim) { int64_t kernel = kernel_shape[dim]; - int64_t input_dim = input_shape[dim + 1]; + int64_t input_dim = input_shape[dim + spatial_dim_start]; kernel_size *= kernel; input_image_size *= input_dim; @@ -85,19 +111,9 @@ Status NhwcPoolFp16::Compute(OpKernelContext* context) const { output_image_size *= output_dim; } - output_dims.push_back(C); - - Tensor* Y = context->Output(0, output_dims); - - constexpr int64_t output_batch_count = 512; - - // Allocate indirection buffer pointers and prepare a padding vector for the - // im2col transform. - AllocatorPtr alloc; - ORT_RETURN_IF_ERROR(context->GetTempSpaceAllocator(&alloc)); - int64_t col_buffer_batch_count = std::min(output_image_size, output_batch_count); - auto* col_data = alloc->Alloc(SafeInt(sizeof(const MLFloat16*)) * kernel_size * col_buffer_batch_count); - BufferUniquePtr col_buffer(col_data, BufferDeleter(std::move(alloc))); + if (channels_last_) { + output_dims.push_back(C); + } const bool need_padding = !is_max_pool_ && pool_attrs_.count_include_pad; std::vector padding_data; @@ -106,16 +122,51 @@ Status NhwcPoolFp16::Compute(OpKernelContext* context) const { } const auto* Xdata = X->Data(); + auto* Y = context->Output(0, output_dims); auto* Ydata = Y->MutableData(); + // Allocate temporary buffers for transposing to channels last format. + AllocatorPtr alloc; + ORT_RETURN_IF_ERROR(context->GetTempSpaceAllocator(&alloc)); + BufferUniquePtr transpose_input_buffer; + BufferUniquePtr transpose_output_buffer; + if (!channels_last_) { + auto* transpose_input = alloc->Alloc(SafeInt(sizeof(MLFloat16)) * C * input_image_size + MLAS_SYMM_QGEMM_BUF_OVERRUN); + transpose_input_buffer = BufferUniquePtr(transpose_input, BufferDeleter(alloc)); + auto* transpose_output = alloc->Alloc(SafeInt(sizeof(MLFloat16)) * C * output_image_size); + transpose_output_buffer = BufferUniquePtr(transpose_output, BufferDeleter(alloc)); + } + + // Allocate indirection buffer pointers and prepare a padding vector for the + // im2col transform. + constexpr int64_t output_batch_count = 512; + int64_t col_buffer_batch_count = std::min(output_image_size, output_batch_count); + auto* col_data = alloc->Alloc(SafeInt(sizeof(const MLFloat16*)) * kernel_size * col_buffer_batch_count); + BufferUniquePtr col_buffer(col_data, BufferDeleter(std::move(alloc))); + for (int64_t image_id = 0; image_id < N; ++image_id) { + const auto* input_data = Xdata; + auto* output_data = Ydata; + + if (!channels_last_) { + // Transpose the input from channels first (CHW) to channels last (HWC). + MlasTranspose( + Xdata, + static_cast(transpose_input_buffer.get()), + static_cast(C), + static_cast(input_image_size)); + input_data = static_cast(transpose_input_buffer.get()); + output_data = static_cast(transpose_output_buffer.get()); + } + + auto* outputptr = output_data; for (int64_t output_start = 0; output_start < output_image_size;) { int64_t output_count = std::min(output_image_size - output_start, output_batch_count); math::Im2col()( - Xdata, + input_data, C, - input_shape.GetDims().data() + 1, - output_dims.data() + 1, + input_shape.GetDims().data() + spatial_dim_start, + output_dims.data() + spatial_dim_start, kernel_shape.data(), strides.data(), dilations.data(), @@ -128,38 +179,80 @@ Status NhwcPoolFp16::Compute(OpKernelContext* context) const { if (is_max_pool_) { MlasNhwcMaxPool( static_cast(col_buffer.get()), - Ydata, + outputptr, static_cast(C), static_cast(output_count), static_cast(kernel_size)); } else { MlasNhwcAvgPool( static_cast(col_buffer.get()), - Ydata, + outputptr, static_cast(C), static_cast(output_count), static_cast(kernel_size)); } - Ydata += output_count * C; + outputptr += output_count * C; output_start += output_count; } + if (!channels_last_) { + // Transpose the output from channels last (NHWC) to channels first (NCHW). + MlasTranspose( + output_data, + Ydata, + static_cast(output_image_size), + static_cast(C)); + } Xdata += input_image_size * C; + Ydata += output_image_size * C; } return Status::OK(); } +// +// Operator definitions +// +ONNX_CPU_OPERATOR_VERSIONED_TYPED_KERNEL( + MaxPool, 8, 11, + MLFloat16, + KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), + PoolFp16); + +ONNX_CPU_OPERATOR_TYPED_KERNEL( + MaxPool, + 12, + MLFloat16, + KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), + PoolFp16); + +ONNX_CPU_OPERATOR_TYPED_KERNEL( + AveragePool, + 11, + MLFloat16, + KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), + PoolFp16); + +ONNX_CPU_OPERATOR_TYPED_KERNEL( + GlobalAveragePool, + 1, + MLFloat16, + KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), + PoolFp16); + +#ifndef DISABLE_CONTRIB_OPS +namespace contrib { + ONNX_OPERATOR_TYPED_KERNEL_EX( MaxPool, kMSInternalNHWCDomain, - 11, + 12, MLFloat16, kCpuExecutionProvider, KernelDefBuilder() .TypeConstraint("T", DataTypeImpl::GetTensorType()), - NhwcPoolFp16); + PoolFp16); ONNX_OPERATOR_TYPED_KERNEL_EX( AveragePool, @@ -169,7 +262,7 @@ ONNX_OPERATOR_TYPED_KERNEL_EX( kCpuExecutionProvider, KernelDefBuilder() .TypeConstraint("T", DataTypeImpl::GetTensorType()), - NhwcPoolFp16); + PoolFp16); ONNX_OPERATOR_TYPED_KERNEL_EX( GlobalAveragePool, @@ -179,9 +272,10 @@ ONNX_OPERATOR_TYPED_KERNEL_EX( kCpuExecutionProvider, KernelDefBuilder() .TypeConstraint("T", DataTypeImpl::GetTensorType()), - NhwcPoolFp16); - -#endif + PoolFp16); } // namespace contrib +#endif // DISABLE_CONTRIB_OPS + } // namespace onnxruntime +#endif // MLAS_F16VEC_INTRINSICS_SUPPORTED diff --git a/onnxruntime/test/contrib_ops/nhwc_pool_in_op_test.cc b/onnxruntime/test/contrib_ops/nhwc_pool_in_op_test.cc index 3868bf7d4a..e40e635c79 100644 --- a/onnxruntime/test/contrib_ops/nhwc_pool_in_op_test.cc +++ b/onnxruntime/test/contrib_ops/nhwc_pool_in_op_test.cc @@ -168,7 +168,7 @@ class NhwcFp16PoolOpTester { std::vector Y_shape; ComputeExpectedOutput(Y_data, Y_shape); - OpTester test(is_max_pool_ ? "MaxPool" : "AveragePool", 11, onnxruntime::kMSInternalNHWCDomain); + OpTester test(is_max_pool_ ? "MaxPool" : "AveragePool", is_max_pool_ ? 12 : 11, onnxruntime::kMSInternalNHWCDomain); test.AddInput("x", X_shape_, X_data_); test.AddOutput("y", Y_shape, Y_data); test.AddAttribute("kernel_shape", kernel_shape_); diff --git a/onnxruntime/test/providers/cpu/activation/activation_op_test.cc b/onnxruntime/test/providers/cpu/activation/activation_op_test.cc index 0a6b7ed3ab..145d1241e0 100644 --- a/onnxruntime/test/providers/cpu/activation/activation_op_test.cc +++ b/onnxruntime/test/providers/cpu/activation/activation_op_test.cc @@ -121,6 +121,18 @@ TEST_F(ActivationOpTest, Relu) { {}, /*is_tensorrt_supported=*/false, /*opset_version= */ 14); +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED + TestActivationOp( + "Relu", + input_values_fp16, + [](MLFloat16 x) { + if (x.ToFloat() > 0.0f) return x; + return MLFloat16(); + }, + {}, + /*is_tensorrt_supported=*/false, + /*opset_version= */ 11); +#endif // MLAS_F16VEC_INTRINSICS_SUPPORTED } #if defined(USE_CUDA) || defined(USE_ROCM) @@ -396,6 +408,29 @@ TEST_F(ActivationOpTest, LeakyRelu) { {{"alpha", alpha}}); } +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED +TEST_F(ActivationOpTest, LeakyRelu_fp16) { + OpTester test("LeakyRelu", 11); + float alpha = 0.01f; // oneDNN set alpha equal to 0.01 + auto formula = [alpha](float x) { return (x >= 0) ? x : alpha * x; }; + + std::vector X = input_values.front(); + std::vector Y; + for (unsigned i = 0; i < X.size(); i++) + Y.push_back(formula(X[i])); + std::vector dims{(int64_t)X.size()}; + + std::vector bf_X(X.size()); + ConvertFloatToMLFloat16(X.data(), bf_X.data(), (int)X.size()); + std::vector bf_Y(Y.size()); + ConvertFloatToMLFloat16(Y.data(), bf_Y.data(), (int)Y.size()); + + test.AddInput("X", dims, bf_X); + test.AddOutput("Y", dims, bf_Y); + test.Run(); +} +#endif // MLAS_F16VEC_INTRINSICS_SUPPORTED + TEST_F(ActivationOpTest, ThresholdedRelu) { float alpha = 0.1f; TestActivationOp( diff --git a/onnxruntime/test/providers/cpu/activation/activation_op_test.h b/onnxruntime/test/providers/cpu/activation/activation_op_test.h index 8b2cf01e8e..c7991c05f1 100644 --- a/onnxruntime/test/providers/cpu/activation/activation_op_test.h +++ b/onnxruntime/test/providers/cpu/activation/activation_op_test.h @@ -7,6 +7,7 @@ #include #include #include +#include "core/mlas/inc/mlas.h" #include "core/graph/constants.h" #include "test/providers/provider_test_utils.h" @@ -96,6 +97,17 @@ class ActivationOpTest : public ::testing::Test { DBL_MAX, -DBL_MAX, std::numeric_limits::infinity()}}; // max, -max, inf std::vector> input_values_int8{{-1, -5, 0, 1, 5, 100, -100, // normal input values for activation std::numeric_limits::min(), std::numeric_limits::max()}}; // min, max +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED + std::vector> input_values_fp16{{MLFloat16(-1.0f), + MLFloat16(-5.f), + MLFloat16(), + MLFloat16(1.f), + MLFloat16(5.f), + MLFloat16(100.f), + MLFloat16(-100.f), + MLFloat16(65504.f), + MLFloat16(-65504.f)}}; +#endif // MLAS_F16VEC_INTRINSICS_SUPPORTED void SetUp() override { float low = -1.0f, high = 1.0f; diff --git a/onnxruntime/test/providers/cpu/nn/pool_fp16_op_test.cc b/onnxruntime/test/providers/cpu/nn/pool_fp16_op_test.cc new file mode 100644 index 0000000000..b033ddbca2 --- /dev/null +++ b/onnxruntime/test/providers/cpu/nn/pool_fp16_op_test.cc @@ -0,0 +1,570 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/mlas/inc/mlas.h" + +#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED + +#include "core/providers/cpu/nn/pool.h" +#include "gtest/gtest.h" +#include "test/providers/provider_test_utils.h" +#include "test/common/cuda_op_test_utils.h" + +namespace onnxruntime { +namespace test { + +// Disable TensorRT on some of the tests because "pads" attribute is not supported + +TEST(PoolFp16Test, MaxPool) { + // TODO: Unskip when fixed #41968513 + if (DefaultDmlExecutionProvider().get() != nullptr) { + GTEST_SKIP() << "Skipping because of the following error: MLOperatorAuthorImpl.cpp(2100): The parameter is incorrect."; + } + + OpTester test("MaxPool", 12); + + test.AddAttribute("auto_pad", ""); + test.AddAttribute("strides", std::vector{1, 1}); + test.AddAttribute("pads", std::vector{0, 0, 0, 0}); + test.AddAttribute("kernel_shape", std::vector{8, 8}); + + std::vector x_vals = { + MLFloat16(0x1.884p-3f), MLFloat16(0x1.3e8p-1f), MLFloat16(0x1.c04p-2f), MLFloat16(0x1.92p-1f), + MLFloat16(0x1.8f4p-1f), MLFloat16(0x1.174p-2f), MLFloat16(0x1.1bp-2f), MLFloat16(0x1.9a8p-1f), + MLFloat16(0x1.ea8p-1f), MLFloat16(0x1.c08p-1f), MLFloat16(0x1.6e8p-2f), MLFloat16(0x1.008p-1f), + MLFloat16(0x1.5ep-1f), MLFloat16(0x1.6dp-1f), MLFloat16(0x1.7b4p-2f), MLFloat16(0x1.1f4p-1f), + MLFloat16(0x1.018p-1f), MLFloat16(0x1.c34p-7f), MLFloat16(0x1.8bcp-1f), MLFloat16(0x1.c4p-1f), + MLFloat16(0x1.75cp-2f), MLFloat16(0x1.3bp-1f), MLFloat16(0x1.34cp-4f), MLFloat16(0x1.79cp-2f), + MLFloat16(0x1.ddcp-1f), MLFloat16(0x1.4d8p-1f), MLFloat16(0x1.96cp-2f), MLFloat16(0x1.93cp-1f), + MLFloat16(0x1.448p-2f), MLFloat16(0x1.22cp-1f), MLFloat16(0x1.bdp-1f), MLFloat16(0x1.becp-2f), + MLFloat16(0x1.9acp-1f), MLFloat16(0x1.268p-3f), MLFloat16(0x1.688p-1f), MLFloat16(0x1.68cp-1f), + MLFloat16(0x1.cp-3f), MLFloat16(0x1.d98p-1f), MLFloat16(0x1.c4cp-2f), MLFloat16(0x1.d18p-1f), + MLFloat16(0x1.eap-5f), MLFloat16(0x1.798p-3f), MLFloat16(0x1.84p-5f), MLFloat16(0x1.598p-1f), + MLFloat16(0x1.308p-1f), MLFloat16(0x1.11p-1f), MLFloat16(0x1.63p-5f), MLFloat16(0x1.1f8p-1f), + MLFloat16(0x1.518p-2f), MLFloat16(0x1.018p-1f), MLFloat16(0x1.ca4p-4f), MLFloat16(0x1.37p-1f), + MLFloat16(0x1.21cp-1f), MLFloat16(0x1.bb4p-8f), MLFloat16(0x1.3c4p-1f), MLFloat16(0x1.d3p-1f), + MLFloat16(0x1.94cp-1f), MLFloat16(0x1.fcp-1f), MLFloat16(0x1.ebp-1f), MLFloat16(0x1.958p-1f), + MLFloat16(0x1.24p-2f), MLFloat16(0x1.4p-1f), MLFloat16(0x1.e98p-2f), MLFloat16(0x1.90cp-3f), + + MLFloat16(0x1.878p-2f), MLFloat16(0x1.b94p-5f), MLFloat16(0x1.ce8p-2f), MLFloat16(0x1.f6cp-1f), + MLFloat16(0x1.fbcp-4f), MLFloat16(0x1.e9p-4f), MLFloat16(0x1.7ap-1f), MLFloat16(0x1.2ccp-1f), + MLFloat16(0x1.e3p-2f), MLFloat16(0x1.b6cp-4f), MLFloat16(0x1.d58p-3f), MLFloat16(0x1.cccp-1f), + MLFloat16(0x1.aacp-2f), MLFloat16(0x1.124p-1f), MLFloat16(0x1.97p-8f), MLFloat16(0x1.33cp-2f), + MLFloat16(0x1.bf8p-2f), MLFloat16(0x1.398p-1f), MLFloat16(0x1.d6p-1f), MLFloat16(0x1.408p-1f), + MLFloat16(0x1.698p-1f), MLFloat16(0x1.32cp-3f), MLFloat16(0x1.7ep-1f), MLFloat16(0x1.a98p-1f), + MLFloat16(0x1.448p-1f), MLFloat16(0x1.c0cp-2f), MLFloat16(0x1.388p-3f), MLFloat16(0x1.23p-1f), + MLFloat16(0x1.0e8p-1f), MLFloat16(0x1.e74p-1f), MLFloat16(0x1.ecp-2f), MLFloat16(0x1.014p-1f), + MLFloat16(0x1.13p-1f), MLFloat16(0x1.a38p-1f), MLFloat16(0x1.d4p-5f), MLFloat16(0x1.56cp-1f), + MLFloat16(0x1.88cp-1f), MLFloat16(0x1.6a8p-1f), MLFloat16(0x1.98p-1f), MLFloat16(0x1.1d8p-1f), + MLFloat16(0x1.ee8p-1f), MLFloat16(0x1.2d8p-3f), MLFloat16(0x1.e5cp-6f), MLFloat16(0x1.3p-1f), + MLFloat16(0x1.d34p-4f), MLFloat16(0x1.e6cp-1f), MLFloat16(0x1.4d8p-2f), MLFloat16(0x1.8c8p-3f), + MLFloat16(0x1.d4cp-2f), MLFloat16(0x1.d74p-1f), MLFloat16(0x1.c2p-1f), MLFloat16(0x1.02cp-2f), + MLFloat16(0x1.644p-2f), MLFloat16(0x1.76p-3f), MLFloat16(0x1.cdcp-1f), MLFloat16(0x1.69cp-1f), + MLFloat16(0x1.74p-1f), MLFloat16(0x1.cccp-1f), MLFloat16(0x1.8fp-1f), MLFloat16(0x1.32cp-1f), + MLFloat16(0x1.2ap-2f), MLFloat16(0x1.36p-3f), MLFloat16(0x1.574p-2f), MLFloat16(0x1.50cp-1f), + + MLFloat16(0x1.2c8p-4f), MLFloat16(0x1.c28p-5f), MLFloat16(0x1.4bp-2f), MLFloat16(0x1.2e4p-1f), + MLFloat16(0x1.b54p-1f), MLFloat16(0x1.26p-2f), MLFloat16(0x1.628p-3f), MLFloat16(0x1.128p-3f), + MLFloat16(0x1.fd4p-1f), MLFloat16(0x1.6f8p-3f), MLFloat16(0x1.454p-2f), MLFloat16(0x1.23p-1f), + MLFloat16(0x1.324p-7f), MLFloat16(0x1.cd4p-1f), MLFloat16(0x1.f44p-1f), MLFloat16(0x1.1d4p-1f), + MLFloat16(0x1.5b4p-4f), MLFloat16(0x1.55p-2f), MLFloat16(0x1.75p-1f), MLFloat16(0x1.23cp-3f), + MLFloat16(0x1.1acp-1f), MLFloat16(0x1.178p-2f), MLFloat16(0x1.f3p-1f), MLFloat16(0x1.56p-1f), + MLFloat16(0x1.05cp-2f), MLFloat16(0x1.bbcp-4f), MLFloat16(0x1.8d8p-1f), MLFloat16(0x1.90cp-1f), + MLFloat16(0x1.86p-1f), MLFloat16(0x1.d44p-1f), MLFloat16(0x1.514p-1f), MLFloat16(0x1.23p-1f), + MLFloat16(0x1.9d4p-3f), MLFloat16(0x1.658p-1f), MLFloat16(0x1.e78p-1f), MLFloat16(0x1.c7cp-1f), + MLFloat16(0x1.fccp-1f), MLFloat16(0x1.a34p-1f), MLFloat16(0x1.17p-1f), MLFloat16(0x1.cep-2f), + MLFloat16(0x1.c8p-1f), MLFloat16(0x1.f24p-1f), MLFloat16(0x1.2fcp-1f), MLFloat16(0x1.76cp-2f), + MLFloat16(0x1.4acp-2f), MLFloat16(0x1.be4p-1f), MLFloat16(0x1.b98p-3f), MLFloat16(0x1.784p-1f), + MLFloat16(0x1.768p-2f), MLFloat16(0x1.9a8p-1f), MLFloat16(0x1.90cp-1f), MLFloat16(0x1.67p-1f), + MLFloat16(0x1.3ecp-1f), MLFloat16(0x1.f98p-2f), MLFloat16(0x1.ae4p-1f), MLFloat16(0x1.6c8p-1f), + MLFloat16(0x1.c68p-2f), MLFloat16(0x1.fc8p-6f), MLFloat16(0x1.74p-2f), MLFloat16(0x1.764p-1f), + MLFloat16(0x1.e7p-2f), MLFloat16(0x1.60cp-2f), MLFloat16(0x1.484p-1f), MLFloat16(0x1.028p-3f)}; + std::vector x_dims = {1, 3, 8, 8}; + std::vector expected_dims = {1, 3, 1, 1}; + std::vector expected_vals = {MLFloat16(0x1.fcp-1f), MLFloat16(0x1.f6cp-1f), MLFloat16(0x1.fd4p-1f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); // TensorRT: result differs +} + +TEST(PoolFp16Test, MaxPool_10_Dilation_1d) { + // TODO: Unskip when fixed #41968513 + if (DefaultDmlExecutionProvider().get() != nullptr) { + GTEST_SKIP() << "Skipping because of the following error: MLOperatorAuthorImpl.cpp(2100): The parameter is incorrect."; + } + + OpTester test("MaxPool", 10); + + test.AddAttribute("auto_pad", ""); + test.AddAttribute("strides", std::vector{1}); + test.AddAttribute("pads", std::vector{0, 0}); + test.AddAttribute("kernel_shape", std::vector{3}); + test.AddAttribute("dilations", std::vector{3}); + + std::vector x_vals = { + MLFloat16(1.f), MLFloat16(3.f), MLFloat16(2.f), MLFloat16(4.f), + MLFloat16(-1.f), MLFloat16(-3.f), MLFloat16(-2.f), MLFloat16(-4.f), + MLFloat16(-6.f), MLFloat16(-5.f), MLFloat16(-4.f), MLFloat16(-2.f)}; + std::vector x_dims = {1, 1, 12}; + std::vector expected_dims = {1, 1, 6}; + std::vector expected_vals = { + MLFloat16(4.f), MLFloat16(3.f), MLFloat16(2.f), + MLFloat16(4.f), MLFloat16(-1.f), MLFloat16(-2.f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); +} + +TEST(PoolFp16Test, MaxPool_DefaultDilations) { + OpTester test("MaxPool", 12); + + test.AddAttribute("kernel_shape", std::vector{2}); + + std::vector x_dims = {1, 3, 3}; + std::vector x_vals = { + MLFloat16(0.f), MLFloat16(1.f), MLFloat16(2.f), + MLFloat16(3.f), MLFloat16(4.f), MLFloat16(5.f), + MLFloat16(6.f), MLFloat16(7.f), MLFloat16(8.f)}; + + std::vector expected_dims = {1, 3, 2}; + std::vector expected_vals = { + MLFloat16(1.f), MLFloat16(2.f), + MLFloat16(4.f), MLFloat16(5.f), + MLFloat16(7.f), MLFloat16(8.f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); +} + +TEST(PoolFp16Test, MaxPool_DilationPadding_1d) { + // TODO: Unskip when fixed #41968513 + if (DefaultDmlExecutionProvider().get() != nullptr) { + GTEST_SKIP() << "Skipping because of the following error: MLOperatorAuthorImpl.cpp(2100): The parameter is incorrect."; + } + + OpTester test("MaxPool", 12); + + test.AddAttribute("auto_pad", ""); + test.AddAttribute("strides", std::vector{1}); + test.AddAttribute("pads", std::vector{1, 1}); + test.AddAttribute("kernel_shape", std::vector{3}); + test.AddAttribute("dilations", std::vector{3}); + + std::vector x_vals = { + MLFloat16(1.f), MLFloat16(3.f), MLFloat16(2.f), MLFloat16(4.f), + MLFloat16(-1.f), MLFloat16(-3.f), MLFloat16(-2.f), MLFloat16(-4.f), + MLFloat16(-6.f), MLFloat16(-5.f), MLFloat16(-4.f), MLFloat16(-2.f)}; + std::vector x_dims = {1, 1, 12}; + std::vector expected_dims = {1, 1, 8}; + std::vector expected_vals = { + MLFloat16(2.f), MLFloat16(4.f), MLFloat16(3.f), MLFloat16(2.f), + MLFloat16(4.f), MLFloat16(-1.f), MLFloat16(-2.f), MLFloat16(-2.f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", + {kCudaExecutionProvider, kTensorrtExecutionProvider, kRocmExecutionProvider}); +} + +TEST(PoolFp16Test, MaxPool_Dilation_2d) { + // TODO: Unskip when fixed #41968513 + if (DefaultDmlExecutionProvider().get() != nullptr) { + GTEST_SKIP() << "Skipping because of the following error: MLOperatorAuthorImpl.cpp(2100): The parameter is incorrect."; + } + + OpTester test("MaxPool", 12); + + test.AddAttribute("auto_pad", ""); + test.AddAttribute("strides", std::vector{1, 1}); + test.AddAttribute("pads", std::vector{0, 0, 0, 0}); + test.AddAttribute("kernel_shape", std::vector{2, 2}); + test.AddAttribute("dilations", std::vector{2, 2}); + + std::vector x_vals = { + MLFloat16(1.f), MLFloat16(3.f), MLFloat16(2.f), MLFloat16(4.f), MLFloat16(-1.f), + MLFloat16(5.f), MLFloat16(7.f), MLFloat16(6.f), MLFloat16(8.f), MLFloat16(-2.f), + MLFloat16(9.f), MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(-3.f), + MLFloat16(13.f), MLFloat16(15.f), MLFloat16(14.f), MLFloat16(16.f), MLFloat16(-4.f)}; + std::vector x_dims = {1, 1, 4, 5}; + std::vector expected_dims = {1, 1, 2, 3}; + std::vector expected_vals = { + MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f), MLFloat16(14.f), MLFloat16(16.f), MLFloat16(14.f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); +} + +TEST(PoolFp16Test, MaxPool_DilationPadding_2d) { + OpTester test("MaxPool", 12); + + test.AddAttribute("auto_pad", ""); + test.AddAttribute("strides", std::vector{1, 1}); + test.AddAttribute("pads", std::vector{1, 1, 1, 1}); + test.AddAttribute("kernel_shape", std::vector{2, 2}); + test.AddAttribute("dilations", std::vector{2, 2}); + + std::vector x_vals = { + MLFloat16(1.f), MLFloat16(3.f), MLFloat16(2.f), MLFloat16(4.f), MLFloat16(-1.f), + MLFloat16(5.f), MLFloat16(7.f), MLFloat16(6.f), MLFloat16(8.f), MLFloat16(-2.f), + MLFloat16(9.f), MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(-3.f), + MLFloat16(13.f), MLFloat16(15.f), MLFloat16(14.f), MLFloat16(16.f), MLFloat16(-4.f)}; + std::vector x_dims = {1, 1, 4, 5}; + std::vector expected_dims = {1, 1, 4, 5}; + std::vector expected_vals = { + MLFloat16(7.f), MLFloat16(6.f), MLFloat16(8.f), MLFloat16(6.f), MLFloat16(8.f), + MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f), MLFloat16(12.f), + MLFloat16(15.f), MLFloat16(14.f), MLFloat16(16.f), MLFloat16(14.f), MLFloat16(16.f), + MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f), MLFloat16(12.f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", + {kCudaExecutionProvider, kTensorrtExecutionProvider, kRocmExecutionProvider}); +} + +TEST(PoolFp16Test, MaxPool_Dilation_Ceil0_2d) { + // TODO: Unskip when fixed #41968513 + if (DefaultDmlExecutionProvider().get() != nullptr) { + GTEST_SKIP() << "Skipping because of the following error: MLOperatorAuthorImpl.cpp(2100): The parameter is incorrect."; + } + + OpTester test("MaxPool", 12); + + test.AddAttribute("auto_pad", ""); + test.AddAttribute("strides", std::vector{2, 1}); + test.AddAttribute("pads", std::vector{0, 0, 0, 0}); + test.AddAttribute("kernel_shape", std::vector{2, 2}); + test.AddAttribute("dilations", std::vector{2, 2}); + + std::vector x_vals = { + MLFloat16(1.f), MLFloat16(3.f), MLFloat16(2.f), MLFloat16(4.f), MLFloat16(-1.f), + MLFloat16(5.f), MLFloat16(7.f), MLFloat16(6.f), MLFloat16(8.f), MLFloat16(-2.f), + MLFloat16(9.f), MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(-3.f), + MLFloat16(13.f), MLFloat16(15.f), MLFloat16(14.f), MLFloat16(16.f), MLFloat16(-4.f)}; + std::vector x_dims = {1, 1, 4, 5}; + std::vector expected_dims = {1, 1, 1, 3}; + std::vector expected_vals = {MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider, kAclExecutionProvider}); +} + +TEST(PoolFp16Test, MaxPool_Dilation_Ceil1_2d) { + // TODO: Unskip when fixed #41968513 + if (DefaultDmlExecutionProvider().get() != nullptr) { + GTEST_SKIP() << "Skipping because of the following error: MLOperatorAuthorImpl.cpp(2100): The parameter is incorrect."; + } + + OpTester test("MaxPool", 12); + + test.AddAttribute("auto_pad", ""); + test.AddAttribute("strides", std::vector{2, 1}); + test.AddAttribute("pads", std::vector{0, 0, 0, 0}); + test.AddAttribute("kernel_shape", std::vector{2, 2}); + test.AddAttribute("dilations", std::vector{2, 2}); + test.AddAttribute("ceil_mode", (int64_t)1); + + std::vector x_vals = { + MLFloat16(1.f), MLFloat16(3.f), MLFloat16(2.f), MLFloat16(4.f), MLFloat16(-1.f), + MLFloat16(5.f), MLFloat16(7.f), MLFloat16(6.f), MLFloat16(8.f), MLFloat16(-2.f), + MLFloat16(9.f), MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(-3.f), + MLFloat16(13.f), MLFloat16(15.f), MLFloat16(14.f), MLFloat16(16.f), MLFloat16(-4.f)}; + std::vector x_dims = {1, 1, 4, 5}; + std::vector expected_dims = {1, 1, 2, 3}; + std::vector expected_vals = {MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider, kAclExecutionProvider}); +} + +TEST(PoolTest, MaxPool_DilationPadding_3d) { + OpTester test("MaxPool", 12); + + test.AddAttribute("auto_pad", ""); + test.AddAttribute("strides", std::vector{1, 1, 1}); + test.AddAttribute("pads", std::vector{1, 1, 1, 1, 1, 1}); + test.AddAttribute("kernel_shape", std::vector{2, 2, 2}); + test.AddAttribute("dilations", std::vector{2, 2, 2}); + + std::vector x_vals = { + MLFloat16(1.f), MLFloat16(3.f), MLFloat16(2.f), MLFloat16(4.f), MLFloat16(-1.f), + MLFloat16(5.f), MLFloat16(7.f), MLFloat16(6.f), MLFloat16(8.f), MLFloat16(-2.f), + MLFloat16(9.f), MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(-3.f), + MLFloat16(13.f), MLFloat16(15.f), MLFloat16(14.f), MLFloat16(16.f), MLFloat16(-4.f), + MLFloat16(1.f), MLFloat16(3.f), MLFloat16(2.f), MLFloat16(4.f), MLFloat16(-1.f), + MLFloat16(5.f), MLFloat16(7.f), MLFloat16(6.f), MLFloat16(8.f), MLFloat16(-2.f), + MLFloat16(9.f), MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(-3.f), + MLFloat16(13.f), MLFloat16(15.f), MLFloat16(14.f), MLFloat16(16.f), MLFloat16(-4.f)}; + std::vector x_dims = {1, 1, 2, 4, 5}; + std::vector expected_dims = {1, 1, 2, 4, 5}; + std::vector expected_vals = { + MLFloat16(7.f), MLFloat16(6.f), MLFloat16(8.f), MLFloat16(6.f), MLFloat16(8.f), + MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f), MLFloat16(12.f), + MLFloat16(15.f), MLFloat16(14.f), MLFloat16(16.f), MLFloat16(14.f), MLFloat16(16.f), + MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f), MLFloat16(12.f), + MLFloat16(7.f), MLFloat16(6.f), MLFloat16(8.f), MLFloat16(6.f), MLFloat16(8.f), + MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f), MLFloat16(12.f), + MLFloat16(15.f), MLFloat16(14.f), MLFloat16(16.f), MLFloat16(14.f), MLFloat16(16.f), + MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f), MLFloat16(12.f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", + {kCudaExecutionProvider, kTensorrtExecutionProvider, kRocmExecutionProvider}); +} + +TEST(PoolFp16Test, AveragePool) { + // TODO: Unskip when fixed #41968513 + if (DefaultDmlExecutionProvider().get() != nullptr) { + GTEST_SKIP() << "Skipping because of the following error: MLOperatorAuthorImpl.cpp(2100): The parameter is incorrect."; + } + + OpTester test("AveragePool", 11); + + test.AddAttribute("auto_pad", ""); + test.AddAttribute("strides", std::vector{1, 1}); + test.AddAttribute("pads", std::vector{0, 0, 0, 0}); + test.AddAttribute("kernel_shape", std::vector{8, 8}); + + std::vector x_vals = { + MLFloat16(0x1.55cp-2f), MLFloat16(0x1.c24p-1f), MLFloat16(0x1.598p-2f), + MLFloat16(0x1.554p-1f), MLFloat16(0x1.c54p-2f), MLFloat16(0x1.4b8p-1f), + MLFloat16(0x1.89p-1f), MLFloat16(0x1.c3cp-1f), MLFloat16(0x1.c54p-1f), + MLFloat16(0x1.7dcp-1f), MLFloat16(0x1.208p-2f), MLFloat16(0x1.bdcp-1f), + MLFloat16(0x1.14cp-1f), MLFloat16(0x1.4a4p-6f), MLFloat16(0x1.9cp-1f), + MLFloat16(0x1.35p-1f), MLFloat16(0x1.3f8p-1f), MLFloat16(0x1.418p-3f), + MLFloat16(0x1.91p-3f), MLFloat16(0x1.6ccp-1f), MLFloat16(0x1.1bp-4f), + MLFloat16(0x1.0cp-1f), MLFloat16(0x1.d18p-1f), MLFloat16(0x1.ba4p-2f), + MLFloat16(0x1.9ecp-1f), MLFloat16(0x1.a8cp-4f), MLFloat16(0x1.8e4p-2f), + MLFloat16(0x1.19cp-2f), MLFloat16(0x1.118p-2f), MLFloat16(0x1.60cp-5f), + MLFloat16(0x1.7ap-2f), MLFloat16(0x1.bccp-1f), MLFloat16(0x1.43p-1f), + MLFloat16(0x1.6c4p-1f), MLFloat16(0x1.03p-2f), MLFloat16(0x1.06cp-1f), + MLFloat16(0x1.c34p-4f), MLFloat16(0x1.9ccp-3f), MLFloat16(0x1.c04p-2f), + MLFloat16(0x1.838p-1f), MLFloat16(0x1.a08p-4f), MLFloat16(0x1.72cp-1f), + MLFloat16(0x1.fcp-2f), MLFloat16(0x1.d5cp-1f), MLFloat16(0x1.36p-1f), + MLFloat16(0x1.008p-2f), MLFloat16(0x1.e7p-2f), MLFloat16(0x1.cfp-1f), + MLFloat16(0x1.e3cp-2f), MLFloat16(0x1.b38p-1f), MLFloat16(0x1.1d8p-3f), + MLFloat16(0x1.f84p-1f), MLFloat16(0x1.57cp-1f), MLFloat16(0x1.cap-1f), + MLFloat16(0x1.9a4p-2f), MLFloat16(0x1.68cp-4f), MLFloat16(0x1.c98p-1f), + MLFloat16(0x1.61cp-2f), MLFloat16(0x1.9e4p-1f), MLFloat16(0x1.8acp-3f), + MLFloat16(0x1.43cp-5f), MLFloat16(0x1.bep-6f), MLFloat16(0x1.9f8p-1f), + MLFloat16(0x1.8acp-1f), MLFloat16(0x1.b8p-1f), MLFloat16(0x1.a1p-3f), + MLFloat16(0x1.918p-1f), MLFloat16(0x1.2bp-2f), MLFloat16(0x1.034p-1f), + MLFloat16(0x1.7bcp-1f), MLFloat16(0x1.b6p-4f), MLFloat16(0x1.2dcp-1f), + MLFloat16(0x1.b7p-3f), MLFloat16(0x1.404p-3f), MLFloat16(0x1.59p-3f), + MLFloat16(0x1.818p-1f), MLFloat16(0x1.2b4p-1f), MLFloat16(0x1.d38p-1f), + MLFloat16(0x1.5b4p-1f), MLFloat16(0x1.b1p-3f), MLFloat16(0x1.d8cp-5f), + MLFloat16(0x1.p-1f), MLFloat16(0x1.d9p-3f), MLFloat16(0x1.a8p-5f), + MLFloat16(0x1.64cp-1f), MLFloat16(0x1.e3cp-2f), MLFloat16(0x1.cf4p-4f), + MLFloat16(0x1.3ccp-1f), MLFloat16(0x1.cb4p-1f), MLFloat16(0x1.374p-1f), + MLFloat16(0x1.3acp-2f), MLFloat16(0x1.43cp-4f), MLFloat16(0x1.908p-5f), + MLFloat16(0x1.fc8p-3f), MLFloat16(0x1.f8p-1f), MLFloat16(0x1.cfp-2f), + MLFloat16(0x1.128p-2f), MLFloat16(0x1.84cp-1f), MLFloat16(0x1.834p-2f), + MLFloat16(0x1.3dp-2f), MLFloat16(0x1.c48p-1f), MLFloat16(0x1.7ecp-4f), + MLFloat16(0x1.84cp-2f), MLFloat16(0x1.93p-4f), MLFloat16(0x1.334p-1f), + MLFloat16(0x1.97p-1f), MLFloat16(0x1.d68p-2f), MLFloat16(0x1.1b8p-1f), + MLFloat16(0x1.8ep-2f), MLFloat16(0x1.a14p-2f), MLFloat16(0x1.8b8p-2f), + MLFloat16(0x1.c88p-1f), MLFloat16(0x1.bdp-3f), MLFloat16(0x1.57cp-1f), + MLFloat16(0x1.278p-1f), MLFloat16(0x1.f9p-1f), MLFloat16(0x1.3acp-5f), + MLFloat16(0x1.424p-3f), MLFloat16(0x1.7e8p-4f), MLFloat16(0x1.db8p-1f), + MLFloat16(0x1.49p-3f), MLFloat16(0x1.a64p-1f), MLFloat16(0x1.b1p-2f), + MLFloat16(0x1.f98p-1f), MLFloat16(0x1.e54p-1f), MLFloat16(0x1.d94p-1f), + MLFloat16(0x1.ff4p-1f), MLFloat16(0x1.50cp-2f), MLFloat16(0x1.85p-7f), + MLFloat16(0x1.f7cp-1f), MLFloat16(0x1.7f8p-4f), MLFloat16(0x1.56cp-2f), + MLFloat16(0x1.47p-1f), MLFloat16(0x1.f8cp-1f), MLFloat16(0x1.de8p-2f), + MLFloat16(0x1.e8p-1f), MLFloat16(0x1.458p-3f), MLFloat16(0x1.6f4p-1f), + MLFloat16(0x1.91cp-6f), MLFloat16(0x1.a4cp-1f), MLFloat16(0x1.274p-3f), + MLFloat16(0x1.cfp-2f), MLFloat16(0x1.c58p-2f), MLFloat16(0x1.fc8p-1f), + MLFloat16(0x1.b38p-1f), MLFloat16(0x1.0b4p-3f), MLFloat16(0x1.4p-4f), + MLFloat16(0x1.e3p-1f), MLFloat16(0x1.fb8p-6f), MLFloat16(0x1.bb4p-3f), + MLFloat16(0x1.e6p-1f), MLFloat16(0x1.258p-1f), MLFloat16(0x1.2f8p-1f), + MLFloat16(0x1.88p-1f), MLFloat16(0x1.2p-1f), MLFloat16(0x1.68cp-7f), + MLFloat16(0x1.75cp-1f), MLFloat16(0x1.8f4p-2f), MLFloat16(0x1.5d8p-4f), + MLFloat16(0x1.bbp-2f), MLFloat16(0x1.afcp-1f), MLFloat16(0x1.0f8p-1f), + MLFloat16(0x1.4a4p-1f), MLFloat16(0x1.518p-3f), MLFloat16(0x1.6fcp-2f), + MLFloat16(0x1.2d4p-5f), MLFloat16(0x1.23cp-1f), MLFloat16(0x1.2b4p-1f), + MLFloat16(0x1.ee4p-1f), MLFloat16(0x1.cf8p-1f), MLFloat16(0x1.2c4p-2f), + MLFloat16(0x1.0bcp-2f), MLFloat16(0x1.ee8p-2f), MLFloat16(0x1.21p-1f), + MLFloat16(0x1.ad4p-3f), MLFloat16(0x1.7f4p-2f), MLFloat16(0x1.f8p-2f), + MLFloat16(0x1.90cp-1f), MLFloat16(0x1.24p-2f), MLFloat16(0x1.dd8p-2f), + MLFloat16(0x1.974p-3f), MLFloat16(0x1.9dcp-3f), MLFloat16(0x1.46p-2f), + MLFloat16(0x1.cfcp-2f), MLFloat16(0x1.204p-2f), MLFloat16(0x1.a4p-1f), + MLFloat16(0x1.fc4p-2f), MLFloat16(0x1.dep-2f), MLFloat16(0x1.7b4p-1f), + MLFloat16(0x1.9b8p-2f), MLFloat16(0x1.b3p-3f), MLFloat16(0x1.e08p-2f)}; + std::vector x_dims = {1, 3, 8, 8}; + std::vector expected_dims = {1, 3, 1, 1}; + std::vector expected_vals = {MLFloat16(0.514681101f), MLFloat16(0.485104561f), MLFloat16(0.475683808f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); +} + +TEST(PoolFp16Test, AveragePool_IncludePadPixel) { + // TODO: Unskip when fixed #41968513 + if (DefaultDmlExecutionProvider().get() != nullptr) { + GTEST_SKIP() << "Skipping because of the following error: MLOperatorAuthorImpl.cpp(2100): The parameter is incorrect."; + } + + OpTester test("AveragePool", 11); + + test.AddAttribute("auto_pad", ""); + test.AddAttribute("strides", std::vector{1, 1}); + test.AddAttribute("pads", std::vector{1, 1, 1, 1}); + test.AddAttribute("kernel_shape", std::vector{2, 2}); + test.AddAttribute("count_include_pad", (int64_t)1); + std::vector x_vals = { + MLFloat16(0x1.55cp-2f), MLFloat16(0x1.c24p-1f), MLFloat16(0x1.598p-2f), + MLFloat16(0x1.554p-1f), MLFloat16(0x1.c54p-2f), MLFloat16(0x1.4b8p-1f), + MLFloat16(0x1.89p-1f), MLFloat16(0x1.c3cp-1f), MLFloat16(0x1.c54p-1f)}; + + std::vector x_dims = {1, 1, 3, 3}; + std::vector expected_dims = {1, 1, 4, 4}; + std::vector expected_vals = { + MLFloat16(0x1.55cp-4f), MLFloat16(0x1.369p-2f), MLFloat16(0x1.378p-2f), MLFloat16(0x1.598p-4f), + MLFloat16(0x1.001p-2f), MLFloat16(0x1.294p-1f), MLFloat16(0x1.2748p-1f), MLFloat16(0x1.f84p-3f), + MLFloat16(0x1.6f2p-2f), MLFloat16(0x1.6128p-1f), MLFloat16(0x1.6dc8p-1f), MLFloat16(0x1.886p-2f), + MLFloat16(0x1.89p-3f), MLFloat16(0x1.a66p-2f), MLFloat16(0x1.c48p-2f), MLFloat16(0x1.c54p-3f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); +} + +// test 'strides' attribute not specified +TEST(PoolFp16Test, AveragePool_DefaultStrides) { + OpTester test("AveragePool", 11); + test.AddAttribute("kernel_shape", std::vector{2}); + std::vector x_vals = { + MLFloat16(0.f), MLFloat16(1.f), MLFloat16(2.f), + MLFloat16(3.f), MLFloat16(4.f), MLFloat16(5.f), + MLFloat16(6.f), MLFloat16(7.f), MLFloat16(8.f)}; + + std::vector x_dims = {1, 3, 3}; + std::vector expected_dims = {1, 3, 2}; + std::vector expected_vals = { + MLFloat16(0.5f), MLFloat16(1.5f), + MLFloat16(3.5f), MLFloat16(4.5f), + MLFloat16(6.5f), MLFloat16(7.5f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); +} + +TEST(PoolFp16Test, AveragePool_10_ceil1_2d) { + // TODO: Unskip when fixed #41968513 + if (DefaultDmlExecutionProvider().get() != nullptr) { + GTEST_SKIP() << "Skipping because of the following error: MLOperatorAuthorImpl.cpp(2100): The parameter is incorrect."; + } + + OpTester test("AveragePool", 11); + + test.AddAttribute("auto_pad", ""); + test.AddAttribute("strides", std::vector{3, 1}); + test.AddAttribute("pads", std::vector{0, 0, 0, 0}); + test.AddAttribute("kernel_shape", std::vector{2, 2}); + test.AddAttribute("ceil_mode", (int64_t)1); + + std::vector x_vals = { + MLFloat16(1.f), MLFloat16(3.f), MLFloat16(2.f), MLFloat16(4.f), + MLFloat16(5.f), MLFloat16(7.f), MLFloat16(6.f), MLFloat16(8.f), + MLFloat16(9.f), MLFloat16(11.f), MLFloat16(10.f), MLFloat16(12.f), + MLFloat16(13.f), MLFloat16(15.f), MLFloat16(14.f), MLFloat16(16.f)}; + std::vector x_dims = {1, 1, 4, 4}; + std::vector expected_dims = {1, 1, 2, 3}; + std::vector expected_vals = { + MLFloat16(4.0f), MLFloat16(4.5f), MLFloat16(5.0f), MLFloat16(14.0f), MLFloat16(14.5f), MLFloat16(15.0f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider, kAclExecutionProvider}); +} + +TEST(PoolFp16Test, GlobalAveragePool) { + OpTester test("GlobalAveragePool"); + + std::vector x_vals = { + MLFloat16(0x1.55cp-2f), MLFloat16(0x1.c24p-1f), MLFloat16(0x1.598p-2f), + MLFloat16(0x1.554p-1f), MLFloat16(0x1.c54p-2f), MLFloat16(0x1.4b8p-1f), + MLFloat16(0x1.890p-1f), MLFloat16(0x1.c3cp-1f), MLFloat16(0x1.c54p-1f), + MLFloat16(0x1.7dcp-1f), MLFloat16(0x1.208p-2f), MLFloat16(0x1.bdcp-1f), + MLFloat16(0x1.14cp-1f), MLFloat16(0x1.4a4p-6f), MLFloat16(0x1.9c0p-1f), + MLFloat16(0x1.350p-1f), MLFloat16(0x1.3f8p-1f), MLFloat16(0x1.418p-3f), + MLFloat16(0x1.910p-3f), MLFloat16(0x1.6ccp-1f), MLFloat16(0x1.1b0p-4f), + MLFloat16(0x1.0c0p-1f), MLFloat16(0x1.d18p-1f), MLFloat16(0x1.ba4p-2f), + MLFloat16(0x1.9ecp-1f), MLFloat16(0x1.a8cp-4f), MLFloat16(0x1.8e4p-2f), + MLFloat16(0x1.19cp-2f), MLFloat16(0x1.118p-2f), MLFloat16(0x1.60cp-5f), + MLFloat16(0x1.7a0p-2f), MLFloat16(0x1.bccp-1f), MLFloat16(0x1.430p-1f), + MLFloat16(0x1.6c4p-1f), MLFloat16(0x1.030p-2f), MLFloat16(0x1.06cp-1f), + MLFloat16(0x1.c34p-4f), MLFloat16(0x1.9ccp-3f), MLFloat16(0x1.c04p-2f), + MLFloat16(0x1.838p-1f), MLFloat16(0x1.a08p-4f), MLFloat16(0x1.72cp-1f), + MLFloat16(0x1.fc0p-2f), MLFloat16(0x1.d5cp-1f), MLFloat16(0x1.360p-1f), + MLFloat16(0x1.008p-2f), MLFloat16(0x1.e70p-2f), MLFloat16(0x1.cf0p-1f), + MLFloat16(0x1.e3cp-2f), MLFloat16(0x1.b38p-1f), MLFloat16(0x1.1d8p-3f), + MLFloat16(0x1.f84p-1f), MLFloat16(0x1.57cp-1f), MLFloat16(0x1.ca0p-1f), + MLFloat16(0x1.9a4p-2f), MLFloat16(0x1.68cp-4f), MLFloat16(0x1.c98p-1f), + MLFloat16(0x1.61cp-2f), MLFloat16(0x1.9e4p-1f), MLFloat16(0x1.8acp-3f), + MLFloat16(0x1.43cp-5f), MLFloat16(0x1.be0p-6f), MLFloat16(0x1.9f8p-1f), + MLFloat16(0x1.8acp-1f), MLFloat16(0x1.b80p-1f), MLFloat16(0x1.a10p-3f), + MLFloat16(0x1.918p-1f), MLFloat16(0x1.2b0p-2f), MLFloat16(0x1.034p-1f), + MLFloat16(0x1.7bcp-1f), MLFloat16(0x1.b60p-4f), MLFloat16(0x1.2dcp-1f), + MLFloat16(0x1.b70p-3f), MLFloat16(0x1.404p-3f), MLFloat16(0x1.590p-3f), + MLFloat16(0x1.818p-1f), MLFloat16(0x1.2b4p-1f), MLFloat16(0x1.d38p-1f), + MLFloat16(0x1.5b4p-1f), MLFloat16(0x1.b10p-3f), MLFloat16(0x1.d8cp-5f), + MLFloat16(0x1.000p-1f), MLFloat16(0x1.d90p-3f), MLFloat16(0x1.a80p-5f), + MLFloat16(0x1.64cp-1f), MLFloat16(0x1.e3cp-2f), MLFloat16(0x1.cf4p-4f), + MLFloat16(0x1.3ccp-1f), MLFloat16(0x1.cb4p-1f), MLFloat16(0x1.374p-1f), + MLFloat16(0x1.3acp-2f), MLFloat16(0x1.43cp-4f), MLFloat16(0x1.908p-5f), + MLFloat16(0x1.fc8p-3f), MLFloat16(0x1.f80p-1f), MLFloat16(0x1.cf0p-2f), + MLFloat16(0x1.128p-2f), MLFloat16(0x1.84cp-1f), MLFloat16(0x1.834p-2f), + MLFloat16(0x1.3d0p-2f), MLFloat16(0x1.c48p-1f), MLFloat16(0x1.7ecp-4f), + MLFloat16(0x1.84cp-2f), MLFloat16(0x1.930p-4f), MLFloat16(0x1.334p-1f), + MLFloat16(0x1.970p-1f), MLFloat16(0x1.d68p-2f), MLFloat16(0x1.1b8p-1f), + MLFloat16(0x1.8e0p-2f), MLFloat16(0x1.a14p-2f), MLFloat16(0x1.8b8p-2f), + MLFloat16(0x1.c88p-1f), MLFloat16(0x1.bd0p-3f), MLFloat16(0x1.57cp-1f), + MLFloat16(0x1.278p-1f), MLFloat16(0x1.f90p-1f), MLFloat16(0x1.3acp-5f), + MLFloat16(0x1.424p-3f), MLFloat16(0x1.7e8p-4f), MLFloat16(0x1.db8p-1f), + MLFloat16(0x1.490p-3f), MLFloat16(0x1.a64p-1f), MLFloat16(0x1.b10p-2f), + MLFloat16(0x1.f98p-1f), MLFloat16(0x1.e54p-1f), MLFloat16(0x1.d94p-1f), + MLFloat16(0x1.ff4p-1f), MLFloat16(0x1.50cp-2f), MLFloat16(0x1.850p-7f), + MLFloat16(0x1.f7cp-1f), MLFloat16(0x1.7f8p-4f), MLFloat16(0x1.56cp-2f), + MLFloat16(0x1.470p-1f), MLFloat16(0x1.f8cp-1f), MLFloat16(0x1.de8p-2f), + MLFloat16(0x1.e80p-1f), MLFloat16(0x1.458p-3f), MLFloat16(0x1.6f4p-1f), + MLFloat16(0x1.91cp-6f), MLFloat16(0x1.a4cp-1f), MLFloat16(0x1.274p-3f), + MLFloat16(0x1.cf0p-2f), MLFloat16(0x1.c58p-2f), MLFloat16(0x1.fc8p-1f), + MLFloat16(0x1.b38p-1f), MLFloat16(0x1.0b4p-3f), MLFloat16(0x1.400p-4f), + MLFloat16(0x1.e30p-1f), MLFloat16(0x1.fb8p-6f), MLFloat16(0x1.bb4p-3f), + MLFloat16(0x1.e60p-1f), MLFloat16(0x1.258p-1f), MLFloat16(0x1.2f8p-1f), + MLFloat16(0x1.880p-1f), MLFloat16(0x1.200p-1f), MLFloat16(0x1.68cp-7f), + MLFloat16(0x1.75cp-1f), MLFloat16(0x1.8f4p-2f), MLFloat16(0x1.5d8p-4f), + MLFloat16(0x1.bb0p-2f), MLFloat16(0x1.afcp-1f), MLFloat16(0x1.0f8p-1f), + MLFloat16(0x1.4a4p-1f), MLFloat16(0x1.518p-3f), MLFloat16(0x1.6fcp-2f), + MLFloat16(0x1.2d4p-5f), MLFloat16(0x1.23cp-1f), MLFloat16(0x1.2b4p-1f), + MLFloat16(0x1.ee4p-1f), MLFloat16(0x1.cf8p-1f), MLFloat16(0x1.2c4p-2f), + MLFloat16(0x1.0bcp-2f), MLFloat16(0x1.ee8p-2f), MLFloat16(0x1.210p-1f), + MLFloat16(0x1.ad4p-3f), MLFloat16(0x1.7f4p-2f), MLFloat16(0x1.f80p-2f), + MLFloat16(0x1.90cp-1f), MLFloat16(0x1.240p-2f), MLFloat16(0x1.dd8p-2f), + MLFloat16(0x1.974p-3f), MLFloat16(0x1.9dcp-3f), MLFloat16(0x1.460p-2f), + MLFloat16(0x1.cfcp-2f), MLFloat16(0x1.204p-2f), MLFloat16(0x1.a40p-1f), + MLFloat16(0x1.fc4p-2f), MLFloat16(0x1.de0p-2f), MLFloat16(0x1.7b4p-1f), + MLFloat16(0x1.9b8p-2f), MLFloat16(0x1.b30p-3f), MLFloat16(0x1.e08p-2f)}; + std::vector x_dims = {1, 3, 8, 8}; + std::vector expected_dims = {1, 3, 1, 1}; + std::vector expected_vals = {MLFloat16(0x1.078448p-1f), MLFloat16(0x1.f0bf4p-2f), MLFloat16(0x1.e719a8p-2f)}; + + test.AddInput("X", x_dims, x_vals); + test.AddOutput("Y", expected_dims, expected_vals); + test.Run(); +} + +} // namespace test +} // namespace onnxruntime + +#endif \ No newline at end of file diff --git a/tools/ci_build/op_registration_validator.py b/tools/ci_build/op_registration_validator.py index 87da585443..a9ccb5c79d 100644 --- a/tools/ci_build/op_registration_validator.py +++ b/tools/ci_build/op_registration_validator.py @@ -47,7 +47,7 @@ class RegistrationValidator(op_registration_utils.RegistrationProcessor): key = domain + ":" + operator prev_start, prev_end = self.last_op_registrations[key] if key in self.last_op_registrations else (None, None) - if prev_start: + if prev_start and start_version > prev_start: # a typed registration where the to/from matches for each entry so nothing to update if prev_start == start_version and prev_end == end_version: return