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(); } };