use crate::activation_simd_avx2;
use crate::activation_simd_avx512;
use crate::math::constants::*;
use core::arch::x86_64::*;
#[inline]
#[target_feature(enable = "avx2,fma")]
pub unsafe fn simd_tanh_avx2(x: __m256) -> __m256 {
let clamp_lo = _mm256_set1_ps(-PADE_TANH_CLAMP);
let clamp_hi = _mm256_set1_ps(PADE_TANH_CLAMP);
let one = _mm256_set1_ps(1.0);
let neg_one = _mm256_set1_ps(-1.0);
let x = _mm256_max_ps(clamp_lo, _mm256_min_ps(clamp_hi, x));
let x_sq = _mm256_mul_ps(x, x);
let num_a = _mm256_set1_ps(PADE_TANH_NUM_A); let num_b = _mm256_set1_ps(PADE_TANH_NUM_B); let num = _mm256_add_ps(x_sq, num_a);
let num = _mm256_fmadd_ps(num, x_sq, num_b);
let num = _mm256_mul_ps(x, num);
let den_c4 = _mm256_set1_ps(PADE_TANH_DEN_C4); let den_c2 = _mm256_set1_ps(PADE_TANH_DEN_C2); let den_a = _mm256_set1_ps(PADE_TANH_DEN_A); let den = _mm256_fmadd_ps(x_sq, den_c4, den_c2);
let den = _mm256_fmadd_ps(den, x_sq, den_a);
let result = _mm256_div_ps(num, den);
_mm256_max_ps(neg_one, _mm256_min_ps(one, result))
}
#[inline]
#[target_feature(enable = "avx2,fma")]
pub unsafe fn simd_tanh_dual_avx2(x1: __m256, x2: __m256) -> (__m256, __m256) {
let clamp_lo = _mm256_set1_ps(-PADE_TANH_CLAMP);
let clamp_hi = _mm256_set1_ps(PADE_TANH_CLAMP);
let one = _mm256_set1_ps(1.0);
let neg_one = _mm256_set1_ps(-1.0);
let x1 = _mm256_max_ps(clamp_lo, _mm256_min_ps(clamp_hi, x1));
let x2 = _mm256_max_ps(clamp_lo, _mm256_min_ps(clamp_hi, x2));
let sq1 = _mm256_mul_ps(x1, x1);
let sq2 = _mm256_mul_ps(x2, x2);
let num_a = _mm256_set1_ps(PADE_TANH_NUM_A); let num_b = _mm256_set1_ps(PADE_TANH_NUM_B); let den_c4 = _mm256_set1_ps(PADE_TANH_DEN_C4); let den_c2 = _mm256_set1_ps(PADE_TANH_DEN_C2); let den_a = _mm256_set1_ps(PADE_TANH_DEN_A);
let num1 = _mm256_fmadd_ps(_mm256_add_ps(sq1, num_a), sq1, num_b);
let num1 = _mm256_mul_ps(x1, num1);
let den1 = _mm256_fmadd_ps(_mm256_fmadd_ps(sq1, den_c4, den_c2), sq1, den_a);
let num2 = _mm256_fmadd_ps(_mm256_add_ps(sq2, num_a), sq2, num_b);
let num2 = _mm256_mul_ps(x2, num2);
let den2 = _mm256_fmadd_ps(_mm256_fmadd_ps(sq2, den_c4, den_c2), sq2, den_a);
let res1 = _mm256_div_ps(num1, den1);
let res2 = _mm256_div_ps(num2, den2);
(
_mm256_max_ps(neg_one, _mm256_min_ps(one, res1)),
_mm256_max_ps(neg_one, _mm256_min_ps(one, res2)),
)
}
#[inline]
#[target_feature(enable = "avx512f,avx512vl,avx512dq")]
pub unsafe fn simd_tanh_avx512(x: __m512) -> __m512 {
let clamp_lo = _mm512_set1_ps(-PADE_TANH_CLAMP);
let clamp_hi = _mm512_set1_ps(PADE_TANH_CLAMP);
let one = _mm512_set1_ps(1.0);
let neg_one = _mm512_set1_ps(-1.0);
let x = _mm512_max_ps(clamp_lo, _mm512_min_ps(clamp_hi, x));
let x_sq = _mm512_mul_ps(x, x);
let num_a = _mm512_set1_ps(PADE_TANH_NUM_A); let num_b = _mm512_set1_ps(PADE_TANH_NUM_B); let num = _mm512_add_ps(x_sq, num_a);
let num = _mm512_fmadd_ps(num, x_sq, num_b);
let num = _mm512_mul_ps(x, num);
let den_c4 = _mm512_set1_ps(PADE_TANH_DEN_C4); let den_c2 = _mm512_set1_ps(PADE_TANH_DEN_C2); let den_a = _mm512_set1_ps(PADE_TANH_DEN_A); let den = _mm512_fmadd_ps(x_sq, den_c4, den_c2);
let den = _mm512_fmadd_ps(den, x_sq, den_a);
let result = _mm512_div_ps(num, den);
_mm512_max_ps(neg_one, _mm512_min_ps(one, result))
}
#[inline]
#[target_feature(enable = "avx2,fma")]
pub unsafe fn tanh_slice_avx2(slice: &mut [f32]) {
let mut i = 0;
let len = slice.len();
unsafe {
activation_simd_avx2!(
i,
len,
{
let x1 = _mm256_loadu_ps(slice.as_ptr().add(i));
let x2 = _mm256_loadu_ps(slice.as_ptr().add(i + 8));
let (y1, y2) = simd_tanh_dual_avx2(x1, x2);
_mm256_storeu_ps(slice.as_mut_ptr().add(i), y1);
_mm256_storeu_ps(slice.as_mut_ptr().add(i + 8), y2);
},
{
let x = _mm256_loadu_ps(slice.as_ptr().add(i));
let y = simd_tanh_avx2(x);
_mm256_storeu_ps(slice.as_mut_ptr().add(i), y);
}
);
}
for item in slice.iter_mut().skip(i) {
*item = scalar_pade_tanh(*item);
if item.abs() < f32::MIN_POSITIVE {
*item = 0.0;
}
}
}
#[inline]
#[target_feature(enable = "avx512f,avx512vl,avx512dq")]
pub unsafe fn tanh_slice_avx512(slice: &mut [f32]) {
let mut i = 0;
let len = slice.len();
unsafe {
activation_simd_avx512!(i, len, {
let x = _mm512_loadu_ps(slice.as_ptr().add(i));
let y = simd_tanh_avx512(x);
_mm512_storeu_ps(slice.as_mut_ptr().add(i), y);
});
}
for item in slice.iter_mut().skip(i) {
*item = scalar_pade_tanh(*item);
if item.abs() < f32::MIN_POSITIVE {
*item = 0.0;
}
}
}
#[inline]
pub fn scalar_pade_tanh(x: f32) -> f32 {
let x = x.clamp(-PADE_TANH_CLAMP, PADE_TANH_CLAMP);
let x2 = x * x;
let num = x * (x2 + PADE_TANH_NUM_A).mul_add(x2, PADE_TANH_NUM_B);
let den = (PADE_TANH_DEN_C4.mul_add(x2, PADE_TANH_DEN_C2)).mul_add(x2, PADE_TANH_DEN_A);
(num / den).clamp(-1.0, 1.0)
}
#[inline]
pub fn tanh(x: f32) -> f32 {
scalar_pade_tanh(x)
}