From 55b10ec08e2cc6160b927c745899d3ebb5eef3f1 Mon Sep 17 00:00:00 2001 From: Dmitri Smirnov Date: Wed, 6 Feb 2019 17:09:34 -0800 Subject: [PATCH 1/2] Implement Sign operator. --- .../providers/cpu/cpu_execution_provider.cc | 26 ++- onnxruntime/core/providers/cpu/math/sign.cc | 168 ++++++++++++++++ onnxruntime/test/onnx/main.cc | 1 - .../test/providers/cpu/math/sing_test.cc | 182 ++++++++++++++++++ .../test/python/onnx_backend_test_series.py | 1 - 5 files changed, 375 insertions(+), 3 deletions(-) create mode 100644 onnxruntime/core/providers/cpu/math/sign.cc create mode 100644 onnxruntime/test/providers/cpu/math/sing_test.cc diff --git a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc index 955d545300c17..127ac79ffb929 100644 --- a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc +++ b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc @@ -236,6 +236,18 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, EyeLike); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, IsNaN); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MLFloat16, IsNaN); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MLFloat16, Sign); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, BFloat16, Sign); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, Sign); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, double, Sign); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int8_t, Sign); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int16_t, Sign); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, Sign); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t, Sign); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, uint8_t, Sign); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, uint16_t, Sign); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, uint32_t, Sign); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, uint64_t, Sign); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Erf); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t_int64_t_int64_t, OneHot); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float_int64_t_int64_t, OneHot); @@ -459,7 +471,7 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); @@ -486,6 +498,18 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); diff --git a/onnxruntime/core/providers/cpu/math/sign.cc b/onnxruntime/core/providers/cpu/math/sign.cc new file mode 100644 index 0000000000000..fdd6da8813652 --- /dev/null +++ b/onnxruntime/core/providers/cpu/math/sign.cc @@ -0,0 +1,168 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/common/common.h" +#include "core/framework/data_types.h" +#include "core/framework/op_kernel.h" +#include "core/util/math.h" + +#include "gsl/span" +#include + +using namespace ::onnxruntime::common; +using namespace ONNX_NAMESPACE; +namespace onnxruntime { + +class Sign final : public OpKernel { + public: + explicit Sign(const OpKernelInfo& info) : OpKernel(info) {} + + Status Compute(OpKernelContext* ctx) const override; +}; + +#define ADD_TYPED_SIGN_OP(data_type) \ + ONNX_CPU_OPERATOR_TYPED_KERNEL( \ + Sign, \ + 9, \ + data_type, \ + KernelDefBuilder() \ + .TypeConstraint("T1", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("T2", DataTypeImpl::GetTensorType()), \ + Sign); + +ADD_TYPED_SIGN_OP(MLFloat16); +ADD_TYPED_SIGN_OP(BFloat16); +ADD_TYPED_SIGN_OP(float); +ADD_TYPED_SIGN_OP(double); + +ADD_TYPED_SIGN_OP(int8_t); +ADD_TYPED_SIGN_OP(int16_t); +ADD_TYPED_SIGN_OP(int32_t); +ADD_TYPED_SIGN_OP(int64_t); + +ADD_TYPED_SIGN_OP(uint8_t); +ADD_TYPED_SIGN_OP(uint16_t); +ADD_TYPED_SIGN_OP(uint32_t); +ADD_TYPED_SIGN_OP(uint64_t); + +namespace sign_internal { +// Unsigned types can only be eq or gt zero +// Signed can be lt, gt or eq to zero +// float, float16 and double will require special handling bc +// - float16 requires unpacking +// - all of then require care for comparing to zeros + +// Unsigned Integer types +template +static void SignUnsigned(const Tensor* input, Tensor* output) { + static_assert(std::numeric_limits::is_integer && + !std::numeric_limits::is_signed, + "Expect a unsigned integer type"); + auto span = gsl::make_span(input->Data(), input->Shape().Size()); + auto output_data = output->template MutableData(); + std::transform(span.cbegin(), span.cend(), output_data, [](T val) { + return (val == T(0)) ? T(0) : T(1); + }); +} + +// Signed types +template +void SignSignedInteger(const Tensor* input, Tensor* output) { + static_assert(std::numeric_limits::is_integer && + std::numeric_limits::is_signed, + "Expect a signed type"); + auto span = gsl::make_span(input->Data(), input->Shape().Size()); + auto output_data = output->template MutableData(); + std::transform(span.cbegin(), span.cend(), output_data, [](T val) { + if (val > T(0)) { + return T(1); + } else if (val < T(0)) { + return T(-1); + } + return T(0); + }); +} + +// The spec does not specify how NaN is +// treated but we have to treat it somehow. We choose +// to return 0 for NaN as TF does. +template +inline T FloatingImpl(T val) { + if (std::isnan(val) || val == T(0)) { + return T(0); + } else if (val > T(0)) { + return T(1); + } else { + return T(-1); + } +} + +template +void SignFloat(const Tensor* input, Tensor* output) { + static_assert((std::is_same::value || std::is_same::value), + "Expect a signed type"); + auto span = gsl::make_span(input->Data(), input->Shape().Size()); + auto output_data = output->template MutableData(); + std::transform(span.cbegin(), span.cend(), output_data, [](T val) { + return FloatingImpl(val); + }); +} + +void SignMLFloat16(const Tensor* input, Tensor* output) { + auto span = gsl::make_span(input->Data(), input->Shape().Size()); + auto output_data = output->template MutableData(); + std::transform(span.cbegin(), span.cend(), output_data, [](const MLFloat16& val) { + float fl = math::halfToFloat(val.val); + return MLFloat16(math::floatToHalf(FloatingImpl(fl))); + }); +} + +void SignBFloat16(const Tensor* input, Tensor* output) { + auto span = gsl::make_span(input->Data(), input->Shape().Size()); + auto output_data = output->template MutableData(); + std::transform(span.cbegin(), span.cend(), output_data, [](const BFloat16& val) { + float fl = val.ToFloat(); + return BFloat16(FloatingImpl(fl)); + }); +} + +} // namespace sign_internal + +Status Sign::Compute(OpKernelContext* ctx) const { + using namespace sign_internal; + + auto input = ctx->Input(0); + auto output = ctx->Output(0, input->Shape()); + + auto dtype = input->DataType(); + if (dtype == DataTypeImpl::GetType()) { + SignSignedInteger(input, output); + } else if (dtype == DataTypeImpl::GetType()) { + SignSignedInteger(input, output); + } else if (dtype == DataTypeImpl::GetType()) { + SignSignedInteger(input, output); + } else if (dtype == DataTypeImpl::GetType()) { + SignSignedInteger(input, output); + } else if (dtype == DataTypeImpl::GetType()) { + SignUnsigned(input, output); + } else if (dtype == DataTypeImpl::GetType()) { + SignUnsigned(input, output); + } else if (dtype == DataTypeImpl::GetType()) { + SignUnsigned(input, output); + } else if (dtype == DataTypeImpl::GetType()) { + SignUnsigned(input, output); + } else if (dtype == DataTypeImpl::GetType()) { + SignFloat(input, output); + } else if (dtype == DataTypeImpl::GetType()) { + SignFloat(input, output); + } else if (dtype == DataTypeImpl::GetType()) { + SignMLFloat16(input, output); + } else if (dtype == DataTypeImpl::GetType()) { + SignBFloat16(input, output); + } else { + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Unsupported input datatype"); + } + return Status::OK(); +} + +} // namespace onnxruntime diff --git a/onnxruntime/test/onnx/main.cc b/onnxruntime/test/onnx/main.cc index faa2527e20f73..8a231615ae302 100644 --- a/onnxruntime/test/onnx/main.cc +++ b/onnxruntime/test/onnx/main.cc @@ -309,7 +309,6 @@ int real_main(int argc, char* argv[]) { {"acosh_example", "opset 9 not supported yet"}, {"atanh_example", "opset 9 not supported yet"}, {"sign_model", "opset 9 not supported yet"}, - {"sign", "opset 9 not supported yet"}, {"scatter_with_axis", "opset 9 not supported yet"}, {"scatter_without_axis", "opset 9 not supported yet"}, {"scan_sum", "opset 9 not supported yet"}, diff --git a/onnxruntime/test/providers/cpu/math/sing_test.cc b/onnxruntime/test/providers/cpu/math/sing_test.cc new file mode 100644 index 0000000000000..ea6cd1993a7f1 --- /dev/null +++ b/onnxruntime/test/providers/cpu/math/sing_test.cc @@ -0,0 +1,182 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "gtest/gtest.h" +#include "test/providers/provider_test_utils.h" +#include "core/util/math.h" + +namespace onnxruntime { +namespace test { + +namespace test_sign_internal { + +template +struct make_type { + static T make(A v) { + return T(v); + } +}; + +template +struct make_type { + static MLFloat16 make(A v) { + return MLFloat16(math::floatToHalf(float(v))); + } +}; + +template +struct make_type { + static BFloat16 make(A v) { + return BFloat16(float(v)); + } +}; + +template +typename std::enable_if::is_signed>::type +GenerateSequence(OutputIter out) { + for (int i = 0; i < 7; ++i) { + *out = make_type::make(i); + ++out; + } +} + +template +typename std::enable_if::is_signed>::type +GenerateSequence(OutputIter out) { + for (int i = -5; i < 2; ++i) { + *out = make_type::make(i); + ++out; + } +} + +template +inline auto to_testable_type(T v) { + return v; +} + +template <> +inline auto to_testable_type(MLFloat16 v) { + return math::halfToFloat(v.val); +} + +template <> +inline auto to_testable_type(BFloat16 v) { + return v.ToFloat(); +} + +template +typename std::enable_if::is_signed && + !std::is_same::value && + !std::is_same::value>::type +TestImpl(ForwardIter first, ForwardIter last, OutputIter out) { + std::transform(first, last, out, [](T v) { + auto t = to_testable_type(v); + if (t == 0) { + t = 0; + } else { + t = 1; + } + return make_type::make(t); + }); +} + +template +typename std::enable_if::is_signed || + std::is_same::value || + std::is_same::value>::type +TestImpl(ForwardIter first, ForwardIter last, OutputIter out) { + std::transform(first, last, out, [](T v) { + auto t = to_testable_type(v); + if (t == 0) { + t = 0; + } else if (t > 0) { + t = 1; + } else { + t = -1; + } + return make_type::make(t); + }); +} +} // namespace test_sign_internal + +TEST(MathOpTest, Sign_uint64) { + using namespace test_sign_internal; + OpTester test("Sign", 9); + + std::vector input_dims{7}; + std::vector input; + GenerateSequence(std::back_inserter(input)); + ASSERT_EQ(input.size(), 7U); + test.AddInput("input", input_dims, input); + + std::vector output; + TestImpl(input.cbegin(), input.cend(), std::back_inserter(output)); + test.AddOutput("output", input_dims, output); + test.Run(OpTester::ExpectResult::kExpectSuccess); +} + +TEST(MathOpTest, Sign_int64) { + using namespace test_sign_internal; + OpTester test("Sign", 9); + + std::vector input_dims{7}; + std::vector input; + GenerateSequence(std::back_inserter(input)); + ASSERT_EQ(input.size(), 7U); + test.AddInput("input", input_dims, input); + + std::vector output; + TestImpl(input.cbegin(), input.cend(), std::back_inserter(output)); + test.AddOutput("output", input_dims, output); + test.Run(OpTester::ExpectResult::kExpectSuccess); +} + +TEST(MathOpTest, Sign_float) { + using namespace test_sign_internal; + OpTester test("Sign", 9); + + std::vector input_dims{7}; + std::vector input; + GenerateSequence(std::back_inserter(input)); + ASSERT_EQ(input.size(), 7U); + test.AddInput("input", input_dims, input); + + std::vector output; + TestImpl(input.cbegin(), input.cend(), std::back_inserter(output)); + test.AddOutput("output", input_dims, output); + test.Run(OpTester::ExpectResult::kExpectSuccess); +} + +TEST(MathOpTest, Sign_double) { + using namespace test_sign_internal; + OpTester test("Sign", 9); + + std::vector input_dims{7}; + std::vector input; + GenerateSequence(std::back_inserter(input)); + ASSERT_EQ(input.size(), 7U); + test.AddInput("input", input_dims, input); + + std::vector output; + TestImpl(input.cbegin(), input.cend(), std::back_inserter(output)); + test.AddOutput("output", input_dims, output); + test.Run(OpTester::ExpectResult::kExpectSuccess); +} +TEST(MathOpTest, Sign_MLFloat16) { + using namespace test_sign_internal; + OpTester test("Sign", 9); + + std::vector input_dims{7}; + std::vector input; + GenerateSequence(std::back_inserter(input)); + ASSERT_EQ(input.size(), 7U); + test.AddInput("input", input_dims, input); + + std::vector output; + TestImpl(input.cbegin(), input.cend(), std::back_inserter(output)); + test.AddOutput("output", input_dims, output); + test.Run(OpTester::ExpectResult::kExpectSuccess); +} + +} // namespace test +} // namespace onnxruntime diff --git a/onnxruntime/test/python/onnx_backend_test_series.py b/onnxruntime/test/python/onnx_backend_test_series.py index 829975e641bd9..cfc137e54ce56 100644 --- a/onnxruntime/test/python/onnx_backend_test_series.py +++ b/onnxruntime/test/python/onnx_backend_test_series.py @@ -31,7 +31,6 @@ '|^test_scatter_without_axis_cpu.*' '|^test_shrink_hard_cpu.*' '|^test_shrink_soft_cpu.*' -'|^test_sign_cpu.*' '|^test_where_example_cpu.*' '|^test_AvgPool1d_cpu.*' '|^test_AvgPool1d_stride_cpu.*' From 4f967975fbd505fb5388fdf4a4bbaff2aa167e69 Mon Sep 17 00:00:00 2001 From: Dmitri Smirnov Date: Fri, 8 Feb 2019 16:26:26 -0800 Subject: [PATCH 2/2] Re-implement using Eigen coefficient-wise sign except for MLFloat16 and BFloat16. Change kernel registration to non-typed. --- .../providers/cpu/cpu_execution_provider.cc | 26 +--- onnxruntime/core/providers/cpu/math/sign.cc | 116 +++++------------- 2 files changed, 32 insertions(+), 110 deletions(-) diff --git a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc index 127ac79ffb929..5d7a46313da54 100644 --- a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc +++ b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc @@ -236,18 +236,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, EyeLike); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, IsNaN); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MLFloat16, IsNaN); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MLFloat16, Sign); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, BFloat16, Sign); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, Sign); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, double, Sign); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int8_t, Sign); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int16_t, Sign); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, Sign); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t, Sign); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, uint8_t, Sign); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, uint16_t, Sign); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, uint32_t, Sign); -class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, uint64_t, Sign); +class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Sign); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Erf); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int64_t_int64_t_int64_t, OneHot); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float_int64_t_int64_t, OneHot); @@ -498,18 +487,7 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); - kernel_registry.Register(BuildKernelCreateInfo()); + kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); kernel_registry.Register(BuildKernelCreateInfo()); diff --git a/onnxruntime/core/providers/cpu/math/sign.cc b/onnxruntime/core/providers/cpu/math/sign.cc index fdd6da8813652..fbfa4ce117375 100644 --- a/onnxruntime/core/providers/cpu/math/sign.cc +++ b/onnxruntime/core/providers/cpu/math/sign.cc @@ -5,6 +5,7 @@ #include "core/framework/data_types.h" #include "core/framework/op_kernel.h" #include "core/util/math.h" +#include "core/util/math_cpuonly.h" #include "gsl/span" #include @@ -20,69 +21,24 @@ class Sign final : public OpKernel { Status Compute(OpKernelContext* ctx) const override; }; -#define ADD_TYPED_SIGN_OP(data_type) \ - ONNX_CPU_OPERATOR_TYPED_KERNEL( \ - Sign, \ - 9, \ - data_type, \ - KernelDefBuilder() \ - .TypeConstraint("T1", DataTypeImpl::GetTensorType()) \ - .TypeConstraint("T2", DataTypeImpl::GetTensorType()), \ - Sign); - -ADD_TYPED_SIGN_OP(MLFloat16); -ADD_TYPED_SIGN_OP(BFloat16); -ADD_TYPED_SIGN_OP(float); -ADD_TYPED_SIGN_OP(double); - -ADD_TYPED_SIGN_OP(int8_t); -ADD_TYPED_SIGN_OP(int16_t); -ADD_TYPED_SIGN_OP(int32_t); -ADD_TYPED_SIGN_OP(int64_t); - -ADD_TYPED_SIGN_OP(uint8_t); -ADD_TYPED_SIGN_OP(uint16_t); -ADD_TYPED_SIGN_OP(uint32_t); -ADD_TYPED_SIGN_OP(uint64_t); +ONNX_CPU_OPERATOR_KERNEL( + Sign, + 9, + KernelDefBuilder().TypeConstraint("T", {DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType()}), + Sign); namespace sign_internal { -// Unsigned types can only be eq or gt zero -// Signed can be lt, gt or eq to zero -// float, float16 and double will require special handling bc -// - float16 requires unpacking -// - all of then require care for comparing to zeros - -// Unsigned Integer types -template -static void SignUnsigned(const Tensor* input, Tensor* output) { - static_assert(std::numeric_limits::is_integer && - !std::numeric_limits::is_signed, - "Expect a unsigned integer type"); - auto span = gsl::make_span(input->Data(), input->Shape().Size()); - auto output_data = output->template MutableData(); - std::transform(span.cbegin(), span.cend(), output_data, [](T val) { - return (val == T(0)) ? T(0) : T(1); - }); -} - -// Signed types -template -void SignSignedInteger(const Tensor* input, Tensor* output) { - static_assert(std::numeric_limits::is_integer && - std::numeric_limits::is_signed, - "Expect a signed type"); - auto span = gsl::make_span(input->Data(), input->Shape().Size()); - auto output_data = output->template MutableData(); - std::transform(span.cbegin(), span.cend(), output_data, [](T val) { - if (val > T(0)) { - return T(1); - } else if (val < T(0)) { - return T(-1); - } - return T(0); - }); -} - // The spec does not specify how NaN is // treated but we have to treat it somehow. We choose // to return 0 for NaN as TF does. @@ -97,17 +53,6 @@ inline T FloatingImpl(T val) { } } -template -void SignFloat(const Tensor* input, Tensor* output) { - static_assert((std::is_same::value || std::is_same::value), - "Expect a signed type"); - auto span = gsl::make_span(input->Data(), input->Shape().Size()); - auto output_data = output->template MutableData(); - std::transform(span.cbegin(), span.cend(), output_data, [](T val) { - return FloatingImpl(val); - }); -} - void SignMLFloat16(const Tensor* input, Tensor* output) { auto span = gsl::make_span(input->Data(), input->Shape().Size()); auto output_data = output->template MutableData(); @@ -125,7 +70,6 @@ void SignBFloat16(const Tensor* input, Tensor* output) { return BFloat16(FloatingImpl(fl)); }); } - } // namespace sign_internal Status Sign::Compute(OpKernelContext* ctx) const { @@ -135,26 +79,26 @@ Status Sign::Compute(OpKernelContext* ctx) const { auto output = ctx->Output(0, input->Shape()); auto dtype = input->DataType(); - if (dtype == DataTypeImpl::GetType()) { - SignSignedInteger(input, output); + if (dtype == DataTypeImpl::GetType()) { + EigenMap(*output) = EigenMap(*input).array().cwiseSign(); + } else if (dtype == DataTypeImpl::GetType()) { + EigenMap(*output) = EigenMap(*input).array().cwiseSign(); + } else if (dtype == DataTypeImpl::GetType()) { + EigenMap(*output) = EigenMap(*input).array().cwiseSign(); } else if (dtype == DataTypeImpl::GetType()) { - SignSignedInteger(input, output); + EigenMap(*output) = EigenMap(*input).array().cwiseSign(); } else if (dtype == DataTypeImpl::GetType()) { - SignSignedInteger(input, output); + EigenMap(*output) = EigenMap(*input).array().cwiseSign(); } else if (dtype == DataTypeImpl::GetType()) { - SignSignedInteger(input, output); + EigenMap(*output) = EigenMap(*input).array().cwiseSign(); } else if (dtype == DataTypeImpl::GetType()) { - SignUnsigned(input, output); + EigenMap(*output) = EigenMap(*input).array().cwiseSign(); } else if (dtype == DataTypeImpl::GetType()) { - SignUnsigned(input, output); + EigenMap(*output) = EigenMap(*input).array().cwiseSign(); } else if (dtype == DataTypeImpl::GetType()) { - SignUnsigned(input, output); + EigenMap(*output) = EigenMap(*input).array().cwiseSign(); } else if (dtype == DataTypeImpl::GetType()) { - SignUnsigned(input, output); - } else if (dtype == DataTypeImpl::GetType()) { - SignFloat(input, output); - } else if (dtype == DataTypeImpl::GetType()) { - SignFloat(input, output); + EigenMap(*output) = EigenMap(*input).array().cwiseSign(); } else if (dtype == DataTypeImpl::GetType()) { SignMLFloat16(input, output); } else if (dtype == DataTypeImpl::GetType()) {