mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
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:
parent
0df2e14038
commit
246cb3a197
2 changed files with 14 additions and 39 deletions
|
|
@ -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;
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
};
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue