#![cfg(feature = "simd")]
use core::f64::consts::{FRAC_2_PI, FRAC_PI_2, LN_2, LOG2_E, SQRT_2};
use oxiblas_core::simd::SimdRegister;
const LN2_HI: f64 = LN_2;
const LN2_LO: f64 = 2.319_046_813_846_3e-17_f64;
const LOG2E: f64 = LOG2_E;
const EXP_POLY: [f64; 12] = [
1.0_f64, 5e-1, 1.666_666_666_666_666_7e-1, 4.166_666_666_666_666_7e-2, 8.333_333_333_333_333e-3, 1.388_888_888_888_889e-3, 1.984_126_984_126_984e-4, 2.480_158_730_158_73e-5, 2.755_731_922_398_589_1e-6, 2.755_731_922_398_589_1e-7, 2.505_210_838_544_172e-8, 2.087_675_698_786_81e-9, ];
const EXP_MAX: f64 = 709.782_712_893_384_f64;
const EXP_MIN: f64 = -745.133_219_101_941_6_f64;
const LN_POLY: [f64; 10] = [
1.0_f64, 3.333_333_333_333_333_7e-1, 2e-1, 1.428_571_428_571_428_7e-1, 1.111_111_111_111_111e-1, 9.090_909_090_909_091e-2, 7.692_307_692_307_693e-2, 6.666_666_666_666_667e-2, 5.882_352_941_176_470_6e-2, 5.263_157_894_736_842e-2, ];
const SIN_POLY: [f64; 6] = [
1.0_f64, -1.666_666_666_666_666_7e-1, 8.333_333_333_333_333e-3, -1.984_126_984_126_984e-4, 2.755_731_922_398_589_1e-6, -2.505_210_838_544_172e-8, ];
const COS_POLY: [f64; 6] = [
1.0_f64, -5e-1, 4.166_666_666_666_666_7e-2, -1.388_888_888_888_889e-3, 2.480_158_730_158_73e-5, -2.755_731_922_398_589_1e-7, ];
const FRAC_PI_2_HI: f64 = FRAC_PI_2;
const FRAC_PI_2_LO: f64 = 6.123_233_995_736_766e-17_f64;
#[inline(always)]
pub(crate) fn simd_horner<R>(x: R, coeffs: &[f64]) -> R
where
R: SimdRegister<Scalar = f64>,
{
let n = coeffs.len();
if n == 0 {
return R::zero();
}
let mut acc = R::splat(coeffs[n - 1]);
for i in (0..n - 1).rev() {
acc = acc.mul_add(x, R::splat(coeffs[i]));
}
acc
}
pub fn simd_exp<R>(x: R) -> R
where
R: SimdRegister<Scalar = f64>,
{
let lanes = R::LANES;
let mut k_vals = [0_i64; 16];
let mut r_reg = R::zero();
let mut special_mask = [false; 16];
for lane in 0..lanes {
let xi = x.extract(lane);
if xi.is_nan() || xi >= EXP_MAX || xi <= EXP_MIN {
special_mask[lane] = true;
} else {
let k = (xi * LOG2E + 0.5).floor() as i64;
let kf = k as f64;
let r = xi - kf * LN2_HI - kf * LN2_LO;
k_vals[lane] = k;
r_reg = r_reg.insert(lane, r);
}
}
let p_reg = simd_horner(r_reg, &EXP_POLY);
let mut result = R::zero();
for lane in 0..lanes {
let val = if special_mask[lane] {
let xi = x.extract(lane);
if xi.is_nan() {
f64::NAN
} else if xi >= EXP_MAX {
f64::INFINITY
} else {
0.0
}
} else {
let r = r_reg.extract(lane);
let p = p_reg.extract(lane);
let exp_r = 1.0 + r * p;
let pow2k = reconstruct_pow2(k_vals[lane]);
exp_r * pow2k
};
result = result.insert(lane, val);
}
result
}
#[inline(always)]
fn scalar_exp(x: f64) -> f64 {
if x.is_nan() {
return f64::NAN;
}
if x >= EXP_MAX {
return f64::INFINITY;
}
if x <= EXP_MIN {
return 0.0;
}
let k = (x * LOG2E + 0.5).floor() as i64;
let kf = k as f64;
let r = x - kf * LN2_HI - kf * LN2_LO;
let p = horner_scalar(r, &EXP_POLY);
let exp_r = 1.0 + r * p;
exp_r * reconstruct_pow2(k)
}
#[inline(always)]
fn reconstruct_pow2(k: i64) -> f64 {
if (-1022..=1023_i64).contains(&k) {
f64::from_bits(((k + 1023) as u64) << 52)
} else if k > 1023 {
f64::INFINITY
} else {
let bit_pos = (52_i64 + k + 1022).clamp(0, 52) as u64;
f64::from_bits(1_u64 << bit_pos)
}
}
#[inline(always)]
fn horner_scalar(x: f64, coeffs: &[f64]) -> f64 {
let n = coeffs.len();
if n == 0 {
return 0.0;
}
let mut acc = coeffs[n - 1];
for i in (0..n - 1).rev() {
acc = acc * x + coeffs[i];
}
acc
}
pub fn simd_ln<R>(x: R) -> R
where
R: SimdRegister<Scalar = f64>,
{
let lanes = R::LANES;
let mut result = R::zero();
for lane in 0..lanes {
result = result.insert(lane, scalar_ln(x.extract(lane)));
}
result
}
#[inline(always)]
fn scalar_ln(x: f64) -> f64 {
if x.is_nan() || x < 0.0 {
return f64::NAN;
}
if x == 0.0 {
return f64::NEG_INFINITY;
}
if x.is_infinite() {
return f64::INFINITY;
}
let bits = x.to_bits();
let exp_bits = (bits >> 52) & 0x7FF;
let mantissa_bits = (bits & 0x000F_FFFF_FFFF_FFFF) | (1023_u64 << 52);
let mut m = f64::from_bits(mantissa_bits);
let mut e = exp_bits as i64 - 1023;
if exp_bits == 0 {
let scale = f64::from_bits(1075_u64 << 52);
let y = x * scale;
let b2 = y.to_bits();
let e2 = (b2 >> 52) & 0x7FF;
let m2_bits = (b2 & 0x000F_FFFF_FFFF_FFFF) | (1023_u64 << 52);
m = f64::from_bits(m2_bits);
e = e2 as i64 - 1023 - 52;
}
if m >= SQRT_2 {
m *= 0.5;
e += 1;
}
let s = (m - 1.0) / (m + 1.0);
let t = s * s;
let poly = horner_scalar(t, &LN_POLY);
let ln_m = 2.0 * s * poly;
(e as f64) * LN_2 + ln_m
}
pub fn simd_sin<R>(x: R) -> R
where
R: SimdRegister<Scalar = f64>,
{
let lanes = R::LANES;
let mut result = R::zero();
for lane in 0..lanes {
result = result.insert(lane, scalar_sin(x.extract(lane)));
}
result
}
pub fn simd_cos<R>(x: R) -> R
where
R: SimdRegister<Scalar = f64>,
{
let lanes = R::LANES;
let mut result = R::zero();
for lane in 0..lanes {
result = result.insert(lane, scalar_cos(x.extract(lane)));
}
result
}
#[inline(always)]
fn scalar_sin(x: f64) -> f64 {
if !x.is_finite() {
return f64::NAN;
}
let (r, quadrant) = reduce_pi2(x);
eval_sincos(r, quadrant, false)
}
#[inline(always)]
fn scalar_cos(x: f64) -> f64 {
if !x.is_finite() {
return f64::NAN;
}
let (r, quadrant) = reduce_pi2(x);
eval_sincos(r, quadrant, true)
}
fn reduce_pi2(x: f64) -> (f64, i32) {
let k = (x * FRAC_2_PI).round() as i64;
let kf = k as f64;
let r = x - kf * FRAC_PI_2_HI - kf * FRAC_PI_2_LO;
let quadrant = k.rem_euclid(4) as i32;
(r, quadrant)
}
fn eval_sincos(r: f64, quadrant: i32, want_cos: bool) -> f64 {
let r2 = r * r;
let sin_r = r * horner_scalar(r2, &SIN_POLY);
let cos_r = horner_scalar(r2, &COS_POLY);
let q = if want_cos {
(quadrant + 1).rem_euclid(4)
} else {
quadrant
};
match q {
0 => sin_r,
1 => cos_r,
2 => -sin_r,
3 => -cos_r,
_ => sin_r, }
}
pub fn simd_tanh<R>(x: R) -> R
where
R: SimdRegister<Scalar = f64>,
{
let lanes = R::LANES;
let mut result = R::zero();
for lane in 0..lanes {
result = result.insert(lane, scalar_tanh(x.extract(lane)));
}
result
}
#[inline(always)]
fn scalar_tanh(x: f64) -> f64 {
if x.is_nan() {
return f64::NAN;
}
if x.abs() >= 20.0 {
return if x > 0.0 { 1.0 } else { -1.0 };
}
let e_pos = scalar_exp(x);
let e_neg = scalar_exp(-x);
(e_pos - e_neg) / (e_pos + e_neg)
}
#[cfg(test)]
mod tests {
use super::*;
fn rel_err(a: f64, b: f64) -> f64 {
if b == 0.0 {
a.abs()
} else {
((a - b) / b).abs()
}
}
#[test]
fn test_scalar_exp_accuracy() {
let mut max_rel = 0.0_f64;
for i in 0..10_000_i32 {
let x = -20.0 + i as f64 * 40.0 / 10_000.0;
let got = scalar_exp(x);
let expected = x.exp();
if expected.is_finite() && expected != 0.0 {
let err = rel_err(got, expected);
if err > max_rel {
max_rel = err;
}
}
}
assert!(
max_rel < 1e-11,
"exp max relative error {max_rel} exceeds 1e-11"
);
}
#[test]
fn test_scalar_exp_special_values() {
assert!(scalar_exp(f64::NAN).is_nan());
assert_eq!(scalar_exp(f64::INFINITY), f64::INFINITY);
assert_eq!(scalar_exp(f64::NEG_INFINITY), 0.0);
assert_eq!(scalar_exp(710.0), f64::INFINITY);
assert_eq!(scalar_exp(-746.0), 0.0);
assert_eq!(scalar_exp(0.0), 1.0);
}
#[test]
fn test_scalar_ln_accuracy() {
let mut max_rel = 0.0_f64;
for i in 0..10_000_i32 {
let t = i as f64 / 10_000.0;
let x = 1e-6_f64 * (1e12_f64).powf(t);
let got = scalar_ln(x);
let expected = x.ln();
if expected.is_finite() {
let err = rel_err(got, expected);
if err > max_rel {
max_rel = err;
}
}
}
assert!(
max_rel < 1e-11,
"ln max relative error {max_rel} exceeds 1e-11"
);
}
#[test]
fn test_scalar_ln_special_values() {
assert!(scalar_ln(f64::NAN).is_nan());
assert_eq!(scalar_ln(f64::INFINITY), f64::INFINITY);
assert_eq!(scalar_ln(0.0), f64::NEG_INFINITY);
assert!(scalar_ln(-1.0).is_nan());
let got = scalar_ln(1.0);
assert!(got.abs() < 1e-15, "ln(1) = {got}");
let got_e = scalar_ln(core::f64::consts::E);
assert!((got_e - 1.0).abs() < 1e-14, "ln(e) = {got_e}");
}
#[test]
fn test_scalar_sin_accuracy() {
let mut max_rel = 0.0_f64;
for i in 0..10_000_i32 {
let x = -50.0 + i as f64 * 100.0 / 10_000.0;
let got = scalar_sin(x);
let expected = x.sin();
let err = if expected.abs() > 0.01 {
rel_err(got, expected)
} else {
(got - expected).abs()
};
if err > max_rel {
max_rel = err;
}
}
assert!(
max_rel < 1e-8,
"sin max relative error {max_rel} exceeds 1e-8 on [-50, 50]"
);
}
#[test]
fn test_scalar_sin_special_values() {
assert!(scalar_sin(f64::NAN).is_nan());
assert!(scalar_sin(f64::INFINITY).is_nan());
assert!(scalar_sin(f64::NEG_INFINITY).is_nan());
assert_eq!(scalar_sin(0.0), 0.0);
let pi_2 = core::f64::consts::PI / 2.0;
assert!(
(scalar_sin(pi_2) - 1.0).abs() < 1e-12,
"sin(pi/2) = {}",
scalar_sin(pi_2)
);
}
#[test]
fn test_scalar_cos_accuracy() {
let mut max_rel = 0.0_f64;
for i in 0..10_000_i32 {
let x = -50.0 + i as f64 * 100.0 / 10_000.0;
let got = scalar_cos(x);
let expected = x.cos();
let err = if expected.abs() > 0.01 {
rel_err(got, expected)
} else {
(got - expected).abs()
};
if err > max_rel {
max_rel = err;
}
}
assert!(
max_rel < 1e-8,
"cos max relative error {max_rel} exceeds 1e-8 on [-50, 50]"
);
}
#[test]
fn test_scalar_cos_special_values() {
assert!(scalar_cos(f64::NAN).is_nan());
assert!(scalar_cos(f64::INFINITY).is_nan());
assert_eq!(scalar_cos(0.0), 1.0);
let pi = core::f64::consts::PI;
assert!(
(scalar_cos(pi) + 1.0).abs() < 1e-12,
"cos(pi) = {}",
scalar_cos(pi)
);
}
#[test]
fn test_scalar_tanh_accuracy() {
let mut max_rel = 0.0_f64;
for i in 0..10_000_i32 {
let x = -20.0 + i as f64 * 40.0 / 10_000.0;
let got = scalar_tanh(x);
let expected = x.tanh();
let err = rel_err(got, expected);
if err > max_rel {
max_rel = err;
}
}
assert!(
max_rel < 1e-11,
"tanh max relative error {max_rel} exceeds 1e-11"
);
}
#[test]
fn test_scalar_tanh_special_values() {
assert!(scalar_tanh(f64::NAN).is_nan());
assert_eq!(scalar_tanh(f64::INFINITY), 1.0);
assert_eq!(scalar_tanh(f64::NEG_INFINITY), -1.0);
assert_eq!(scalar_tanh(0.0), 0.0);
}
#[test]
fn test_horner_scalar_matches() {
let coeffs = [1.0_f64, 2.0, 3.0];
let got = horner_scalar(2.0, &coeffs);
assert!(
(got - 17.0).abs() < 1e-15,
"horner_scalar gave {got}, expected 17"
);
}
#[test]
fn test_exp_extended_range() {
let cases = [
(-745.0_f64, (-745.0_f64).exp()),
(-700.0, (-700.0_f64).exp()),
(0.0, 1.0_f64),
(1.0, core::f64::consts::E),
(700.0, (700.0_f64).exp()),
];
for (x, expected) in cases {
let got = scalar_exp(x);
if expected.is_finite() && expected != 0.0 {
let err = rel_err(got, expected);
assert!(
err < 1e-11,
"exp({x}) = {got}, expected {expected}, rel_err={err}"
);
}
}
}
#[test]
fn test_ln_subnormal() {
let x = f64::from_bits(1); let got = scalar_ln(x);
let expected = x.ln();
let abs_err = (got - expected).abs();
assert!(
abs_err < 1.0,
"ln(subnormal) = {got}, expected approx {expected}"
);
}
}