Improve accuracy of fast approximate tanh and the logistic functions in Eigen, such that they preserve relative accuracy to within a few ULPs where their function values tend to zero (around x=0 for tanh, and for large negative x for the logistic function).

This change re-instates the fast rational approximation of the logistic function for float32 in Eigen (removed in 66f07efeae), but uses the more accurate approximation 1/(1+exp(-1)) ~= exp(x) below -9. The exponential is only calculated on the vectorized path if at least one element in the SIMD input vector is less than -9.

This change also contains a few improvements to speed up the original float specialization of logistic:
  - Introduce EIGEN_PREDICT_{FALSE,TRUE} for __builtin_predict and use it to predict that the logistic-only path is most likely (~2-3% speedup for the common case).
  - Carefully set the upper clipping point to the smallest x where the approximation evaluates to exactly 1. This saves the explicit clamping of the output (~7% speedup).

The increased accuracy for tanh comes at a cost of 10-20% depending on instruction set.

The benchmarks below repeated calls

   u = v.logistic()  (u = v.tanh(), respectively)

where u and v are of type Eigen::ArrayXf, have length 8k, and v contains random numbers in [-1,1].

Benchmark numbers for logistic:

Before:
Benchmark                  Time(ns)        CPU(ns)     Iterations
-----------------------------------------------------------------
SSE
BM_eigen_logistic_float        4467           4468         155835  model_time: 4827
AVX
BM_eigen_logistic_float        2347           2347         299135  model_time: 2926
AVX+FMA
BM_eigen_logistic_float        1467           1467         476143  model_time: 2926
AVX512
BM_eigen_logistic_float         805            805         858696  model_time: 1463

After:
Benchmark                  Time(ns)        CPU(ns)     Iterations
-----------------------------------------------------------------
SSE
BM_eigen_logistic_float        2589           2590         270264  model_time: 4827
AVX
BM_eigen_logistic_float        1428           1428         489265  model_time: 2926
AVX+FMA
BM_eigen_logistic_float        1059           1059         662255  model_time: 2926
AVX512
BM_eigen_logistic_float         673            673        1000000  model_time: 1463

Benchmark numbers for tanh:

Before:
Benchmark                  Time(ns)        CPU(ns)     Iterations
-----------------------------------------------------------------
SSE
BM_eigen_tanh_float        2391           2391         292624  model_time: 4242
AVX
BM_eigen_tanh_float        1256           1256         554662  model_time: 2633
AVX+FMA
BM_eigen_tanh_float         823            823         866267  model_time: 1609
AVX512
BM_eigen_tanh_float         443            443        1578999  model_time: 805

After:
Benchmark                  Time(ns)        CPU(ns)     Iterations
-----------------------------------------------------------------
SSE
BM_eigen_tanh_float        2588           2588         273531  model_time: 4242
AVX
BM_eigen_tanh_float        1536           1536         452321  model_time: 2633
AVX+FMA
BM_eigen_tanh_float        1007           1007         694681  model_time: 1609
AVX512
BM_eigen_tanh_float         471            471        1472178  model_time: 805
This commit is contained in:
Rasmus Munk Larsen
2019-12-16 21:33:42 +00:00
parent 8e5da71466
commit a566074480
9 changed files with 191 additions and 23 deletions

View File

@@ -905,14 +905,106 @@ struct scalar_logistic_op {
}
};
/** \internal
* \brief Template specialization of the logistic function for float.
*
* Uses just a 9/10-degree rational interpolant which
* interpolates 1/(1+exp(-x)) - 0.5 up to a couple of ulps in the range
* [-9, 18]. Below -9 we use the more accurate approximation
* 1/(1+exp(-x)) ~= exp(x), and above 18 the logistic function is 1 withing
* one ulp. The shifted logistic is interpolated because it was easier to
* make the fit converge.
*
*/
template <>
struct scalar_logistic_op<float> {
EIGEN_EMPTY_STRUCT_CTOR(scalar_logistic_op)
EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float operator()(const float& x) const {
// The upper cut-off is the smallest x for which the rational approximation evaluates to 1.
// Choosing this value saves us a few instructions clamping the results at the end.
#ifdef EIGEN_VECTORIZE_FMA
const float cutoff_upper = 16.285715103149414062f;
#else
const float cutoff_upper = 16.619047164916992188f;
#endif
const float cutoff_lower = -9.f;
if (x > cutoff_upper) return 1.0f;
else if (x < cutoff_lower) return numext::exp(x);
else return 1.0f / (1.0f + numext::exp(-x));
}
template <typename Packet> EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
Packet packetOp(const Packet& _x) const {
const Packet cutoff_lower = pset1<Packet>(-9.f);
const Packet lt_mask = pcmp_lt<Packet>(_x, cutoff_lower);
const bool any_small = predux(lt_mask);
// Clamp the input to be at most 'cutoff_upper'.
#ifdef EIGEN_VECTORIZE_FMA
const Packet cutoff_upper = pset1<Packet>(16.285715103149414062f);
#else
const Packet cutoff_upper = pset1<Packet>(16.619047164916992188f);
#endif
const Packet x = pmin(_x, cutoff_upper);
// The monomial coefficients of the numerator polynomial (odd).
const Packet alpha_1 = pset1<Packet>(2.48287947061529e-01f);
const Packet alpha_3 = pset1<Packet>(8.51377133304701e-03f);
const Packet alpha_5 = pset1<Packet>(6.08574864600143e-05f);
const Packet alpha_7 = pset1<Packet>(1.15627324459942e-07f);
const Packet alpha_9 = pset1<Packet>(4.37031012579801e-11f);
// The monomial coefficients of the denominator polynomial (even).
const Packet beta_0 = pset1<Packet>(9.93151921023180e-01f);
const Packet beta_2 = pset1<Packet>(1.16817656904453e-01f);
const Packet beta_4 = pset1<Packet>(1.70198817374094e-03f);
const Packet beta_6 = pset1<Packet>(6.29106785017040e-06f);
const Packet beta_8 = pset1<Packet>(5.76102136993427e-09f);
const Packet beta_10 = pset1<Packet>(6.10247389755681e-13f);
// Since the polynomials are odd/even, we need x^2.
const Packet x2 = pmul(x, x);
// Evaluate the numerator polynomial p.
Packet p = pmadd(x2, alpha_9, alpha_7);
p = pmadd(x2, p, alpha_5);
p = pmadd(x2, p, alpha_3);
p = pmadd(x2, p, alpha_1);
p = pmul(x, p);
// Evaluate the denominator polynomial q.
Packet q = pmadd(x2, beta_10, beta_8);
q = pmadd(x2, q, beta_6);
q = pmadd(x2, q, beta_4);
q = pmadd(x2, q, beta_2);
q = pmadd(x2, q, beta_0);
// Divide the numerator by the denominator and shift it up.
const Packet logistic = padd(pdiv(p, q), pset1<Packet>(0.5f));
if (EIGEN_PREDICT_FALSE(any_small)) {
const Packet exponential = pexp(_x);
return pselect(lt_mask, exponential, logistic);
} else {
return logistic;
}
}
};
template <typename T>
struct functor_traits<scalar_logistic_op<T> > {
enum {
// The cost estimate for float here here is for the common(?) case where
// all arguments are greater than -9.
Cost = scalar_div_cost<T, packet_traits<T>::HasDiv>::value +
NumTraits<T>::AddCost * 2 + functor_traits<scalar_exp_op<T> >::Cost,
(internal::is_same<T, float>::value
? NumTraits<T>::AddCost * 15 + NumTraits<T>::MulCost * 11
: NumTraits<T>::AddCost * 2 +
functor_traits<scalar_exp_op<T> >::Cost),
PacketAccess =
packet_traits<T>::HasAdd && packet_traits<T>::HasDiv &&
packet_traits<T>::HasNegate && packet_traits<T>::HasExp
(internal::is_same<T, float>::value
? packet_traits<T>::HasMul && packet_traits<T>::HasMax &&
packet_traits<T>::HasMin
: packet_traits<T>::HasNegate && packet_traits<T>::HasExp)
};
};