use crate::math::constants::*;
use core::arch::x86_64::*;
#[inline]
#[target_feature(enable = "avx2,fma")]
pub unsafe fn simd_tanh_pade_nr1_avx2(x: __m256) -> __m256 {
let clamp_lo = _mm256_set1_ps(-PADE_TANH_CLAMP);
let clamp_hi = _mm256_set1_ps(PADE_TANH_CLAMP);
let two = _mm256_set1_ps(2.0);
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 mut r = _mm256_rcp_ps(den);
r = _mm256_mul_ps(r, _mm256_fnmadd_ps(den, r, two));
let result = _mm256_mul_ps(num, r);
_mm256_max_ps(neg_one, _mm256_min_ps(one, result))
}
#[inline]
#[target_feature(enable = "avx2,fma")]
pub unsafe fn simd_tanh_pade_nr1_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 two = _mm256_set1_ps(2.0);
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 mut r1 = _mm256_rcp_ps(den1);
r1 = _mm256_mul_ps(r1, _mm256_fnmadd_ps(den1, r1, two));
let mut r2 = _mm256_rcp_ps(den2);
r2 = _mm256_mul_ps(r2, _mm256_fnmadd_ps(den2, r2, two));
let res1 = _mm256_mul_ps(num1, r1);
let res2 = _mm256_mul_ps(num2, r2);
(
_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_pade_nr1_avx512(x: __m512) -> __m512 {
let clamp_lo = _mm512_set1_ps(-PADE_TANH_CLAMP);
let clamp_hi = _mm512_set1_ps(PADE_TANH_CLAMP);
let two = _mm512_set1_ps(2.0);
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 mut r = _mm512_rcp14_ps(den);
r = _mm512_mul_ps(r, _mm512_fnmadd_ps(den, r, two));
let result = _mm512_mul_ps(num, r);
_mm512_max_ps(neg_one, _mm512_min_ps(one, result))
}
#[inline]
#[target_feature(enable = "avx2,fma")]
pub unsafe fn simd_tanh_pade_nr2_avx2(x: __m256) -> __m256 {
let clamp_lo = _mm256_set1_ps(-PADE_TANH_CLAMP);
let clamp_hi = _mm256_set1_ps(PADE_TANH_CLAMP);
let two = _mm256_set1_ps(2.0);
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 mut r = _mm256_rcp_ps(den);
r = _mm256_mul_ps(r, _mm256_fnmadd_ps(den, r, two));
r = _mm256_mul_ps(r, _mm256_fnmadd_ps(den, r, two));
let result = _mm256_mul_ps(num, r);
_mm256_max_ps(neg_one, _mm256_min_ps(one, result))
}
#[inline]
#[target_feature(enable = "avx512f,avx512vl,avx512dq")]
pub unsafe fn simd_tanh_pade_nr2_avx512(x: __m512) -> __m512 {
let clamp_lo = _mm512_set1_ps(-PADE_TANH_CLAMP);
let clamp_hi = _mm512_set1_ps(PADE_TANH_CLAMP);
let two = _mm512_set1_ps(2.0);
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 mut r = _mm512_rcp14_ps(den);
r = _mm512_mul_ps(r, _mm512_fnmadd_ps(den, r, two));
r = _mm512_mul_ps(r, _mm512_fnmadd_ps(den, r, two));
let result = _mm512_mul_ps(num, r);
_mm512_max_ps(neg_one, _mm512_min_ps(one, result))
}
#[cfg(test)]
#[path = "reference_test.rs"]
mod reference_test;