mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
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:
parent
9bf08bdb52
commit
2fa10fb803
10 changed files with 903 additions and 49 deletions
|
|
@ -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)>,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
83
onnxruntime/core/providers/cpu/fp16/fp16_activations.h
Normal file
83
onnxruntime/core/providers/cpu/fp16/fp16_activations.h
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
@ -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_);
|
||||
|
|
|
|||
|
|
@ -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>(
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
570
onnxruntime/test/providers/cpu/nn/pool_fp16_op_test.cc
Normal file
570
onnxruntime/test/providers/cpu/nn/pool_fp16_op_test.cc
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in a new issue