Simplify shrink, replace Eigne in Sign implemenation (#16975)

### Description
<!-- Describe your changes. -->
Simplify Shrink.
Replace Eigen code with the one that does not require fp16 conversion in
Sign.


### Motivation and Context
<!-- - Why is this change required? What problem does it solve?
- If it fixes an open issue, please link to the issue here. -->
This commit is contained in:
Dmitri Smirnov 2023-08-02 18:24:38 -07:00 committed by GitHub
parent 0df2e14038
commit 246cb3a197
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 14 additions and 39 deletions

View file

@ -58,12 +58,13 @@ struct CallSignImpl {
template <>
struct CallSignImpl<MLFloat16> {
void operator()(const Tensor* input, Tensor* output) const {
ConstEigenVectorMap<Eigen::half> input_data(
reinterpret_cast<const Eigen::half*>(input->Data<MLFloat16>()),
narrow<ptrdiff_t>(input->Shape().Size()));
EigenVectorMap<Eigen::half>(reinterpret_cast<Eigen::half*>(output->MutableData<MLFloat16>()),
narrow<ptrdiff_t>(output->Shape().Size())) = input_data.array().cwiseSign();
auto span = input->DataAsSpan<MLFloat16>();
auto output_data = output->MutableData<MLFloat16>();
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;
});
}
};

View file

@ -34,51 +34,25 @@ ONNX_CPU_OPERATOR_KERNEL(
#endif
namespace shrink_internal {
template <class T>
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<float>(t_val);
if (val < -lambd) {
return T(val + bias);
return static_cast<T>(val + bias);
}
if (val > lambd) {
return T(val - bias);
return static_cast<T>(val - bias);
} else {
return T(0);
return static_cast<T>(0.f);
}
}
template <class T>
Status ShrinkImpl(const Tensor* input, Tensor* output, float bias, float lambd) {
EigenMap<T>(*output) = EigenMap<T>(*input).unaryExpr([bias, lambd](const T& val) { return ShrinkCore<T>(val, bias, lambd); });
return Status::OK();
}
template <>
Status ShrinkImpl<MLFloat16>(const Tensor* input, Tensor* output, float bias, float lambd) {
const auto span = input->DataAsSpan<MLFloat16>();
auto* output_data = output->MutableData<MLFloat16>();
std::transform(span.begin(), span.end(), output_data, [bias, lambd](const MLFloat16& val) {
float fl = val.ToFloat();
return MLFloat16(ShrinkCore<float>(fl, bias, lambd));
});
return Status::OK();
}
template <>
Status ShrinkImpl<BFloat16>(const Tensor* input, Tensor* output, float bias, float lambd) {
const auto span = input->DataAsSpan<BFloat16>();
auto* output_data = output->MutableData<BFloat16>();
std::transform(span.begin(), span.end(), output_data, [bias, lambd](const BFloat16& val) {
float fl = val.ToFloat();
return BFloat16(ShrinkCore<float>(fl, bias, lambd));
});
return Status::OK();
}
template <class T>
struct CallShrinkImpl {
Status operator()(const Tensor* input, Tensor* output, float bias, float lambd) const {
return ShrinkImpl<T>(input, output, bias, lambd);
EigenMap<T>(*output) = EigenMap<T>(*input).unaryExpr([bias, lambd](const T& val) { return ShrinkCore<T>(val, bias, lambd); });
return Status::OK();
}
};