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.
This commit is contained in:
Chen Fu 2023-04-25 14:01:47 -07:00 committed by GitHub
parent 9bf08bdb52
commit 2fa10fb803
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
10 changed files with 903 additions and 49 deletions

View file

@ -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<void>, // default entry to avoid the list become empty after ops-reducing
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, MLFloat16, NhwcFusedConv)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSInternalNHWCDomain, 11, MLFloat16, MaxPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSInternalNHWCDomain, 12, MLFloat16, MaxPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSInternalNHWCDomain, 11, MLFloat16, AveragePool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSInternalNHWCDomain, 1, MLFloat16, GlobalAveragePool)>,
};

View file

@ -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);

View file

@ -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<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MaxUnpool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 17, LpPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Conv)>,
#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MLFloat16, Conv)>,
#endif
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, ConvTranspose)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, 12, If)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, SequenceLength)>,
@ -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<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, MLFloat16, GlobalAveragePool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MLFloat16, Conv)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MLFloat16, AveragePool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 8, 11, MLFloat16, MaxPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, MLFloat16, MaxPool)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 12, MLFloat16, Relu)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 13, 13, MLFloat16, Relu)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 14, MLFloat16, Relu)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 15, MLFloat16, LeakyRelu)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 16, MLFloat16, LeakyRelu)>,
};
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

View file

@ -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<MLFloat16> : public ElementWiseRangedTransform<MLFloat16> {
MLAS_ACTIVATION Activation;
Status Init(const onnxruntime::NodeAttributes&) {
Activation.ActivationKind = MlasReluActivation;
return Status::OK();
}
GSL_SUPPRESS(r .11)
ElementWiseRangedTransform<MLFloat16>* Copy() const final {
using T1 = typename std::remove_pointer<decltype(this)>::type;
using T2 = typename std::remove_const<T1>::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<MLFloat16> : public ElementWiseRangedTransform<MLFloat16> {
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<MLFloat16>* Copy() const final {
using T1 = typename std::remove_pointer<decltype(this)>::type;
using T2 = typename std::remove_const<T1>::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

View file

@ -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<Tensor>(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<size_t>(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<MLFloat16> padding_data;
@ -106,16 +122,51 @@ Status NhwcPoolFp16::Compute(OpKernelContext* context) const {
}
const auto* Xdata = X->Data<MLFloat16>();
auto* Y = context->Output(0, output_dims);
auto* Ydata = Y->MutableData<MLFloat16>();
// 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<size_t>(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<size_t>(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<size_t>(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<MLFloat16*>(transpose_input_buffer.get()),
static_cast<size_t>(C),
static_cast<size_t>(input_image_size));
input_data = static_cast<MLFloat16*>(transpose_input_buffer.get());
output_data = static_cast<MLFloat16*>(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<MLFloat16, StorageOrder::NHWC>()(
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<MLFloat16 const**>(col_buffer.get()),
Ydata,
outputptr,
static_cast<size_t>(C),
static_cast<size_t>(output_count),
static_cast<size_t>(kernel_size));
} else {
MlasNhwcAvgPool(
static_cast<MLFloat16 const**>(col_buffer.get()),
Ydata,
outputptr,
static_cast<size_t>(C),
static_cast<size_t>(output_count),
static_cast<size_t>(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<size_t>(output_image_size),
static_cast<size_t>(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<MLFloat16>()),
PoolFp16);
ONNX_CPU_OPERATOR_TYPED_KERNEL(
MaxPool,
12,
MLFloat16,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<MLFloat16>()),
PoolFp16);
ONNX_CPU_OPERATOR_TYPED_KERNEL(
AveragePool,
11,
MLFloat16,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<MLFloat16>()),
PoolFp16);
ONNX_CPU_OPERATOR_TYPED_KERNEL(
GlobalAveragePool,
1,
MLFloat16,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<MLFloat16>()),
PoolFp16);
#ifndef DISABLE_CONTRIB_OPS
namespace contrib {
ONNX_OPERATOR_TYPED_KERNEL_EX(
MaxPool,
kMSInternalNHWCDomain,
11,
12,
MLFloat16,
kCpuExecutionProvider,
KernelDefBuilder()
.TypeConstraint("T", DataTypeImpl::GetTensorType<MLFloat16>()),
NhwcPoolFp16);
PoolFp16);
ONNX_OPERATOR_TYPED_KERNEL_EX(
AveragePool,
@ -169,7 +262,7 @@ ONNX_OPERATOR_TYPED_KERNEL_EX(
kCpuExecutionProvider,
KernelDefBuilder()
.TypeConstraint("T", DataTypeImpl::GetTensorType<MLFloat16>()),
NhwcPoolFp16);
PoolFp16);
ONNX_OPERATOR_TYPED_KERNEL_EX(
GlobalAveragePool,
@ -179,9 +272,10 @@ ONNX_OPERATOR_TYPED_KERNEL_EX(
kCpuExecutionProvider,
KernelDefBuilder()
.TypeConstraint("T", DataTypeImpl::GetTensorType<MLFloat16>()),
NhwcPoolFp16);
#endif
PoolFp16);
} // namespace contrib
#endif // DISABLE_CONTRIB_OPS
} // namespace onnxruntime
#endif // MLAS_F16VEC_INTRINSICS_SUPPORTED

View file

@ -168,7 +168,7 @@ class NhwcFp16PoolOpTester {
std::vector<int64_t> 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<MLFloat16>("x", X_shape_, X_data_);
test.AddOutput<MLFloat16>("y", Y_shape, Y_data);
test.AddAttribute("kernel_shape", kernel_shape_);

View file

@ -121,6 +121,18 @@ TEST_F(ActivationOpTest, Relu) {
{},
/*is_tensorrt_supported=*/false,
/*opset_version= */ 14);
#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED
TestActivationOp<MLFloat16>(
"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<float> X = input_values.front();
std::vector<float> Y;
for (unsigned i = 0; i < X.size(); i++)
Y.push_back(formula(X[i]));
std::vector<int64_t> dims{(int64_t)X.size()};
std::vector<MLFloat16> bf_X(X.size());
ConvertFloatToMLFloat16(X.data(), bf_X.data(), (int)X.size());
std::vector<MLFloat16> bf_Y(Y.size());
ConvertFloatToMLFloat16(Y.data(), bf_Y.data(), (int)Y.size());
test.AddInput<MLFloat16>("X", dims, bf_X);
test.AddOutput<MLFloat16>("Y", dims, bf_Y);
test.Run();
}
#endif // MLAS_F16VEC_INTRINSICS_SUPPORTED
TEST_F(ActivationOpTest, ThresholdedRelu) {
float alpha = 0.1f;
TestActivationOp<float>(

View file

@ -7,6 +7,7 @@
#include <unordered_map>
#include <functional>
#include <random>
#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<double>::infinity()}}; // max, -max, inf
std::vector<std::vector<int8_t>> input_values_int8{{-1, -5, 0, 1, 5, 100, -100, // normal input values for activation
std::numeric_limits<int8_t>::min(), std::numeric_limits<int8_t>::max()}}; // min, max
#ifdef MLAS_F16VEC_INTRINSICS_SUPPORTED
std::vector<std::vector<MLFloat16>> 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;

View file

@ -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<int64_t>{1, 1});
test.AddAttribute("pads", std::vector<int64_t>{0, 0, 0, 0});
test.AddAttribute("kernel_shape", std::vector<int64_t>{8, 8});
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 3, 8, 8};
std::vector<int64_t> expected_dims = {1, 3, 1, 1};
std::vector<MLFloat16> expected_vals = {MLFloat16(0x1.fcp-1f), MLFloat16(0x1.f6cp-1f), MLFloat16(0x1.fd4p-1f)};
test.AddInput<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{1});
test.AddAttribute("pads", std::vector<int64_t>{0, 0});
test.AddAttribute("kernel_shape", std::vector<int64_t>{3});
test.AddAttribute("dilations", std::vector<int64_t>{3});
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 1, 12};
std::vector<int64_t> expected_dims = {1, 1, 6};
std::vector<MLFloat16> expected_vals = {
MLFloat16(4.f), MLFloat16(3.f), MLFloat16(2.f),
MLFloat16(4.f), MLFloat16(-1.f), MLFloat16(-2.f)};
test.AddInput<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{2});
std::vector<int64_t> x_dims = {1, 3, 3};
std::vector<MLFloat16> 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<int64_t> expected_dims = {1, 3, 2};
std::vector<MLFloat16> expected_vals = {
MLFloat16(1.f), MLFloat16(2.f),
MLFloat16(4.f), MLFloat16(5.f),
MLFloat16(7.f), MLFloat16(8.f)};
test.AddInput<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{1});
test.AddAttribute("pads", std::vector<int64_t>{1, 1});
test.AddAttribute("kernel_shape", std::vector<int64_t>{3});
test.AddAttribute("dilations", std::vector<int64_t>{3});
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 1, 12};
std::vector<int64_t> expected_dims = {1, 1, 8};
std::vector<MLFloat16> 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<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{1, 1});
test.AddAttribute("pads", std::vector<int64_t>{0, 0, 0, 0});
test.AddAttribute("kernel_shape", std::vector<int64_t>{2, 2});
test.AddAttribute("dilations", std::vector<int64_t>{2, 2});
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 1, 4, 5};
std::vector<int64_t> expected_dims = {1, 1, 2, 3};
std::vector<MLFloat16> expected_vals = {
MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f), MLFloat16(14.f), MLFloat16(16.f), MLFloat16(14.f)};
test.AddInput<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{1, 1});
test.AddAttribute("pads", std::vector<int64_t>{1, 1, 1, 1});
test.AddAttribute("kernel_shape", std::vector<int64_t>{2, 2});
test.AddAttribute("dilations", std::vector<int64_t>{2, 2});
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 1, 4, 5};
std::vector<int64_t> expected_dims = {1, 1, 4, 5};
std::vector<MLFloat16> 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<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{2, 1});
test.AddAttribute("pads", std::vector<int64_t>{0, 0, 0, 0});
test.AddAttribute("kernel_shape", std::vector<int64_t>{2, 2});
test.AddAttribute("dilations", std::vector<int64_t>{2, 2});
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 1, 4, 5};
std::vector<int64_t> expected_dims = {1, 1, 1, 3};
std::vector<MLFloat16> expected_vals = {MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f)};
test.AddInput<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{2, 1});
test.AddAttribute("pads", std::vector<int64_t>{0, 0, 0, 0});
test.AddAttribute("kernel_shape", std::vector<int64_t>{2, 2});
test.AddAttribute("dilations", std::vector<int64_t>{2, 2});
test.AddAttribute("ceil_mode", (int64_t)1);
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 1, 4, 5};
std::vector<int64_t> expected_dims = {1, 1, 2, 3};
std::vector<MLFloat16> expected_vals = {MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f), MLFloat16(10.f), MLFloat16(12.f), MLFloat16(10.f)};
test.AddInput<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{1, 1, 1});
test.AddAttribute("pads", std::vector<int64_t>{1, 1, 1, 1, 1, 1});
test.AddAttribute("kernel_shape", std::vector<int64_t>{2, 2, 2});
test.AddAttribute("dilations", std::vector<int64_t>{2, 2, 2});
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 1, 2, 4, 5};
std::vector<int64_t> expected_dims = {1, 1, 2, 4, 5};
std::vector<MLFloat16> 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<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{1, 1});
test.AddAttribute("pads", std::vector<int64_t>{0, 0, 0, 0});
test.AddAttribute("kernel_shape", std::vector<int64_t>{8, 8});
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 3, 8, 8};
std::vector<int64_t> expected_dims = {1, 3, 1, 1};
std::vector<MLFloat16> expected_vals = {MLFloat16(0.514681101f), MLFloat16(0.485104561f), MLFloat16(0.475683808f)};
test.AddInput<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{1, 1});
test.AddAttribute("pads", std::vector<int64_t>{1, 1, 1, 1});
test.AddAttribute("kernel_shape", std::vector<int64_t>{2, 2});
test.AddAttribute("count_include_pad", (int64_t)1);
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 1, 3, 3};
std::vector<int64_t> expected_dims = {1, 1, 4, 4};
std::vector<MLFloat16> 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<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{2});
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 3, 3};
std::vector<int64_t> expected_dims = {1, 3, 2};
std::vector<MLFloat16> expected_vals = {
MLFloat16(0.5f), MLFloat16(1.5f),
MLFloat16(3.5f), MLFloat16(4.5f),
MLFloat16(6.5f), MLFloat16(7.5f)};
test.AddInput<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("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<int64_t>{3, 1});
test.AddAttribute("pads", std::vector<int64_t>{0, 0, 0, 0});
test.AddAttribute("kernel_shape", std::vector<int64_t>{2, 2});
test.AddAttribute("ceil_mode", (int64_t)1);
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 1, 4, 4};
std::vector<int64_t> expected_dims = {1, 1, 2, 3};
std::vector<MLFloat16> expected_vals = {
MLFloat16(4.0f), MLFloat16(4.5f), MLFloat16(5.0f), MLFloat16(14.0f), MLFloat16(14.5f), MLFloat16(15.0f)};
test.AddInput<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("Y", expected_dims, expected_vals);
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider, kAclExecutionProvider});
}
TEST(PoolFp16Test, GlobalAveragePool) {
OpTester test("GlobalAveragePool");
std::vector<MLFloat16> 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<int64_t> x_dims = {1, 3, 8, 8};
std::vector<int64_t> expected_dims = {1, 3, 1, 1};
std::vector<MLFloat16> expected_vals = {MLFloat16(0x1.078448p-1f), MLFloat16(0x1.f0bf4p-2f), MLFloat16(0x1.e719a8p-2f)};
test.AddInput<MLFloat16>("X", x_dims, x_vals);
test.AddOutput<MLFloat16>("Y", expected_dims, expected_vals);
test.Run();
}
} // namespace test
} // namespace onnxruntime
#endif

View file

@ -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