From 246cb3a197007ecc632f26550138f501dd890979 Mon Sep 17 00:00:00 2001 From: Dmitri Smirnov Date: Wed, 2 Aug 2023 18:24:38 -0700 Subject: [PATCH] Simplify shrink, replace Eigne in Sign implemenation (#16975) ### Description Simplify Shrink. Replace Eigen code with the one that does not require fp16 conversion in Sign. ### Motivation and Context --- onnxruntime/core/providers/cpu/math/sign.cc | 13 +++---- onnxruntime/core/providers/cpu/nn/shrink.cc | 40 ++++----------------- 2 files changed, 14 insertions(+), 39 deletions(-) diff --git a/onnxruntime/core/providers/cpu/math/sign.cc b/onnxruntime/core/providers/cpu/math/sign.cc index abc3bdee16..60080135bb 100644 --- a/onnxruntime/core/providers/cpu/math/sign.cc +++ b/onnxruntime/core/providers/cpu/math/sign.cc @@ -58,12 +58,13 @@ struct CallSignImpl { template <> struct CallSignImpl { void operator()(const Tensor* input, Tensor* output) const { - ConstEigenVectorMap input_data( - reinterpret_cast(input->Data()), - narrow(input->Shape().Size())); - - EigenVectorMap(reinterpret_cast(output->MutableData()), - narrow(output->Shape().Size())) = input_data.array().cwiseSign(); + auto span = input->DataAsSpan(); + auto output_data = output->MutableData(); + std::transform(span.begin(), span.end(), output_data, [](const MLFloat16& val) { + // Return 0 as TF does for NaN. + if (val.IsNaNOrZero()) return MLFloat16::Zero; + return (val.IsNegative()) ? MLFloat16::MinusOne : MLFloat16::One; + }); } }; diff --git a/onnxruntime/core/providers/cpu/nn/shrink.cc b/onnxruntime/core/providers/cpu/nn/shrink.cc index bbfbb0eb11..406e8870b2 100644 --- a/onnxruntime/core/providers/cpu/nn/shrink.cc +++ b/onnxruntime/core/providers/cpu/nn/shrink.cc @@ -34,51 +34,25 @@ ONNX_CPU_OPERATOR_KERNEL( #endif namespace shrink_internal { template -inline T ShrinkCore(const T& val, float bias, float lambd) { +inline T ShrinkCore(const T& t_val, float bias, float lambd) { // The ONNX spec doesn't take numeric overflow and underflow into account // Implementing the spec as is for now + float val = static_cast(t_val); if (val < -lambd) { - return T(val + bias); + return static_cast(val + bias); } if (val > lambd) { - return T(val - bias); + return static_cast(val - bias); } else { - return T(0); + return static_cast(0.f); } } -template -Status ShrinkImpl(const Tensor* input, Tensor* output, float bias, float lambd) { - EigenMap(*output) = EigenMap(*input).unaryExpr([bias, lambd](const T& val) { return ShrinkCore(val, bias, lambd); }); - return Status::OK(); -} - -template <> -Status ShrinkImpl(const Tensor* input, Tensor* output, float bias, float lambd) { - const auto span = input->DataAsSpan(); - auto* output_data = output->MutableData(); - std::transform(span.begin(), span.end(), output_data, [bias, lambd](const MLFloat16& val) { - float fl = val.ToFloat(); - return MLFloat16(ShrinkCore(fl, bias, lambd)); - }); - return Status::OK(); -} - -template <> -Status ShrinkImpl(const Tensor* input, Tensor* output, float bias, float lambd) { - const auto span = input->DataAsSpan(); - auto* output_data = output->MutableData(); - std::transform(span.begin(), span.end(), output_data, [bias, lambd](const BFloat16& val) { - float fl = val.ToFloat(); - return BFloat16(ShrinkCore(fl, bias, lambd)); - }); - return Status::OK(); -} - template struct CallShrinkImpl { Status operator()(const Tensor* input, Tensor* output, float bias, float lambd) const { - return ShrinkImpl(input, output, bias, lambd); + EigenMap(*output) = EigenMap(*input).unaryExpr([bias, lambd](const T& val) { return ShrinkCore(val, bias, lambd); }); + return Status::OK(); } };