const L2UF: f32 = 0.693145751953125;
const L2LF: f32 = 1.428606765330187045e-06;
const R_LN2F: f32 = 1.442_695_040_888_963_4_f32;
const POLY0: f32 = 0.000_198_527_617_612_853_64;
const POLY1: f32 = 0.001_393_043_552_525_341_5;
const POLY2: f32 = 0.008_333_360_776_305_198_7;
const POLY3: f32 = 0.041_666_485_369_205_475;
const POLY4: f32 = 0.166_666_671_633_720_4;
const POLY5: f32 = 0.5;
const EXPF_UNDERFLOW: f32 = -104.0;
const EXPF_OVERFLOW: f32 = 100.0;
#[inline]
pub fn sleef_expf(d: f32) -> f32 {
let q_int = (d * R_LN2F).round_ties_even() as i32;
let qf = q_int as f32;
let s = qf.mul_add(-L2UF, d);
let s = qf.mul_add(-L2LF, s);
let mut u: f32 = POLY0;
u = u.mul_add(s, POLY1);
u = u.mul_add(s, POLY2);
u = u.mul_add(s, POLY3);
u = u.mul_add(s, POLY4);
u = u.mul_add(s, POLY5);
let u = 1.0_f32 + (s * s).mul_add(u, s);
let q1 = q_int >> 1;
let q2 = q_int - q1;
let two_q1 = f32::from_bits(((127_i32 + q1) as u32) << 23);
let two_q2 = f32::from_bits(((127_i32 + q2) as u32) << 23);
let y = u * two_q1 * two_q2;
if d < EXPF_UNDERFLOW {
0.0
} else if d > EXPF_OVERFLOW {
f32::INFINITY
} else {
y
}
}
#[cfg(target_arch = "aarch64")]
#[inline]
pub fn sleef_expf_inplace_neon(data: &mut [f32]) {
use std::arch::aarch64::*;
let n = data.len();
let n4 = n - (n % 4);
let mut i = 0;
unsafe {
while i < n4 {
let v_d = vld1q_f32(data.as_ptr().add(i));
let v_q_f = vmulq_f32(v_d, vdupq_n_f32(R_LN2F));
let v_q_i = vcvtnq_s32_f32(v_q_f);
let v_qf = vcvtq_f32_s32(v_q_i);
let v_neg_l2u = vdupq_n_f32(-L2UF);
let v_s = vfmaq_f32(v_d, v_qf, v_neg_l2u);
let v_neg_l2l = vdupq_n_f32(-L2LF);
let v_s = vfmaq_f32(v_s, v_qf, v_neg_l2l);
let v_u = vdupq_n_f32(POLY0);
let v_u = vfmaq_f32(vdupq_n_f32(POLY1), v_u, v_s);
let v_u = vfmaq_f32(vdupq_n_f32(POLY2), v_u, v_s);
let v_u = vfmaq_f32(vdupq_n_f32(POLY3), v_u, v_s);
let v_u = vfmaq_f32(vdupq_n_f32(POLY4), v_u, v_s);
let v_u = vfmaq_f32(vdupq_n_f32(POLY5), v_u, v_s);
let v_s_sq = vmulq_f32(v_s, v_s);
let v_inner = vfmaq_f32(v_s, v_s_sq, v_u);
let v_u = vaddq_f32(vdupq_n_f32(1.0), v_inner);
let v_q1 = vshrq_n_s32(v_q_i, 1);
let v_q2 = vsubq_s32(v_q_i, v_q1);
let v_127 = vdupq_n_s32(127);
let v_e1 = vshlq_n_s32(vaddq_s32(v_q1, v_127), 23);
let v_e2 = vshlq_n_s32(vaddq_s32(v_q2, v_127), 23);
let v_two_q1 = vreinterpretq_f32_s32(v_e1);
let v_two_q2 = vreinterpretq_f32_s32(v_e2);
let v_y = vmulq_f32(vmulq_f32(v_u, v_two_q1), v_two_q2);
let v_neg104 = vdupq_n_f32(EXPF_UNDERFLOW);
let v_p100 = vdupq_n_f32(EXPF_OVERFLOW);
let v_zero = vdupq_n_f32(0.0);
let v_inf = vreinterpretq_f32_u32(vdupq_n_u32(0x7f800000));
let m_under = vcltq_f32(v_d, v_neg104);
let m_over = vcltq_f32(v_p100, v_d);
let v_y = vbslq_f32(m_under, v_zero, v_y);
let v_y = vbslq_f32(m_over, v_inf, v_y);
vst1q_f32(data.as_mut_ptr().add(i), v_y);
i += 4;
}
}
while i < n {
data[i] = sleef_expf(data[i]);
i += 1;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn matches_torch_on_known_divergence_point() {
let x = -3.796875_f32;
let got = sleef_expf(x);
let expected: u32 = 0x3cb7d5c0;
assert_eq!(
got.to_bits(),
expected,
"sleef_expf(-3.796875): expected 0x{:08x} (torch), got 0x{:08x}",
expected,
got.to_bits()
);
}
#[test]
fn differs_from_libm_at_divergence_point() {
let x = -3.796875_f32;
let libm_bits = x.exp().to_bits();
let sleef_bits = sleef_expf(x).to_bits();
assert_ne!(
libm_bits, sleef_bits,
"libm and sleef_expf should differ at -3.796875 (1-ULP gap); both gave 0x{:08x}",
libm_bits
);
}
#[test]
fn matches_standard_reference_values() {
assert_eq!(sleef_expf(0.0).to_bits(), 0x3f800000);
assert_eq!(sleef_expf(1.0).to_bits(), 0x402df854);
assert_eq!(sleef_expf(-1.0).to_bits(), 0x3ebc5ab2);
}
#[test]
fn underflow_and_overflow_thresholds() {
assert_eq!(sleef_expf(-200.0), 0.0);
assert_eq!(sleef_expf(200.0), f32::INFINITY);
assert_eq!(sleef_expf(f32::NEG_INFINITY), 0.0);
assert!(sleef_expf(f32::NAN).is_nan());
}
#[test]
fn matches_torch_in_subnormal_and_overflow_ranges() {
let r = sleef_expf(-100.0);
assert_eq!(
r.to_bits(),
0x0000001b,
"sleef_expf(-100): expected torch 0x0000001b (subnormal), got 0x{:08x}",
r.to_bits()
);
let r = sleef_expf(-90.0);
assert_eq!(
r.to_bits(),
0x0008ec28,
"sleef_expf(-90): expected torch 0x0008ec28 (subnormal), got 0x{:08x}",
r.to_bits()
);
let r = sleef_expf(90.0);
assert_eq!(
r.to_bits(),
0x7f800000,
"sleef_expf(90): expected torch +inf 0x7f800000, got 0x{:08x}",
r.to_bits()
);
let r = sleef_expf(-104.0);
assert_eq!(r, 0.0, "sleef_expf(-104) post-guard → 0");
let r = sleef_expf(100.0);
assert_eq!(
r.to_bits(),
0x7f800000,
"sleef_expf(100): expected +inf (natural IEEE overflow), got 0x{:08x}",
r.to_bits()
);
}
#[cfg(target_arch = "aarch64")]
#[test]
fn neon_matches_scalar_on_known_divergence_point() {
let mut data4 = vec![-3.796875_f32; 4];
sleef_expf_inplace_neon(&mut data4);
for (i, &x) in data4.iter().enumerate() {
assert_eq!(
x.to_bits(),
0x3cb7d5c0,
"NEON lane {} mismatch: expected 0x3cb7d5c0 (torch), got 0x{:08x}",
i,
x.to_bits()
);
}
let mut data7: Vec<f32> = vec![
-3.796875, 0.0, 1.0, -1.0, -100.0, 90.0, -3.796875, ];
sleef_expf_inplace_neon(&mut data7);
let expected = [
0x3cb7d5c0_u32, 0x3f800000, 0x402df854, 0x3ebc5ab2, 0x0000001b, 0x7f800000, 0x3cb7d5c0, ];
for (i, (&x, &e)) in data7.iter().zip(expected.iter()).enumerate() {
assert_eq!(
x.to_bits(),
e,
"lane {} (mixed SIMD/tail) mismatch: expected 0x{:08x}, got 0x{:08x}",
i,
e,
x.to_bits()
);
}
}
#[cfg(target_arch = "aarch64")]
#[test]
fn neon_matches_scalar_on_sweep() {
let mut x: u32 = 0xCAFEBABE;
const N: usize = 1024;
let mut input: Vec<f32> = Vec::with_capacity(N);
for _ in 0..N {
x ^= x << 13;
x ^= x >> 17;
x ^= x << 5;
input.push((x as f32 / u32::MAX as f32) * 15.0 - 10.0);
}
let mut scalar_out = vec![0.0_f32; N];
for (i, &v) in input.iter().enumerate() {
scalar_out[i] = sleef_expf(v);
}
let mut neon_out = input.clone();
sleef_expf_inplace_neon(&mut neon_out);
let mut differ = 0;
for (i, (&s, &nu)) in scalar_out.iter().zip(neon_out.iter()).enumerate() {
if s.to_bits() != nu.to_bits() {
if differ < 5 {
eprintln!(
" divergence at idx {}: input={}, scalar=0x{:08x}, neon=0x{:08x}",
i,
input[i],
s.to_bits(),
nu.to_bits()
);
}
differ += 1;
}
}
assert_eq!(
differ, 0,
"NEON sleef_expf must bit-match scalar on every input; {differ}/{N} differed"
);
}
#[cfg(target_arch = "aarch64")]
#[test]
#[ignore]
fn sleef_expf_neon_microbench() {
const N: usize = 1_000_000;
let mut data: Vec<f32> = Vec::with_capacity(N);
let mut x: u32 = 0xCAFEBABE;
for _ in 0..N {
x ^= x << 13;
x ^= x >> 17;
x ^= x << 5;
data.push((x as f32 / u32::MAX as f32) * 15.0 - 10.0);
}
const ITERS: usize = 5;
let mut libm_buf = vec![0f32; N];
let mut libm_ns = Vec::with_capacity(ITERS);
for _ in 0..ITERS {
let t0 = std::time::Instant::now();
for i in 0..N {
libm_buf[i] = data[i].exp();
}
libm_ns.push(t0.elapsed().as_nanos() as u64);
std::hint::black_box(&libm_buf);
}
libm_ns.sort();
let mut scalar_buf = vec![0f32; N];
let mut scalar_ns = Vec::with_capacity(ITERS);
for _ in 0..ITERS {
let t0 = std::time::Instant::now();
for i in 0..N {
scalar_buf[i] = sleef_expf(data[i]);
}
scalar_ns.push(t0.elapsed().as_nanos() as u64);
std::hint::black_box(&scalar_buf);
}
scalar_ns.sort();
let mut neon_ns = Vec::with_capacity(ITERS);
for _ in 0..ITERS {
let mut neon_buf = data.clone();
let t0 = std::time::Instant::now();
sleef_expf_inplace_neon(&mut neon_buf);
neon_ns.push(t0.elapsed().as_nanos() as u64);
std::hint::black_box(&neon_buf);
}
neon_ns.sort();
let libm_med = libm_ns[ITERS / 2] as f64;
let scalar_med = scalar_ns[ITERS / 2] as f64;
let neon_med = neon_ns[ITERS / 2] as f64;
eprintln!(
"\nsleef NEON microbench (N = {} f32, {} iters):\n \
libm f32::exp: {:.3} ms {:.1} M elem/s [1.00× ref]\n \
scalar sleef_expf: {:.3} ms {:.1} M elem/s [{:.2}× vs libm]\n \
NEON sleef_expf_inplace: {:.3} ms {:.1} M elem/s [{:.2}× vs libm]",
N,
ITERS,
libm_med / 1e6,
N as f64 / (libm_med / 1e9) / 1e6,
scalar_med / 1e6,
N as f64 / (scalar_med / 1e9) / 1e6,
scalar_med / libm_med,
neon_med / 1e6,
N as f64 / (neon_med / 1e9) / 1e6,
neon_med / libm_med,
);
}
#[test]
#[ignore]
fn sleef_expf_microbench() {
const N: usize = 1_000_000;
let mut data: Vec<f32> = Vec::with_capacity(N);
let mut x: u32 = 0xCAFEBABE;
for _ in 0..N {
x ^= x << 13;
x ^= x >> 17;
x ^= x << 5;
data.push((x as f32 / u32::MAX as f32) * 15.0 - 10.0);
}
const ITERS: usize = 5;
let mut libm_buf = vec![0f32; N];
let mut libm_ns = Vec::with_capacity(ITERS);
for _ in 0..ITERS {
let t0 = std::time::Instant::now();
for i in 0..N {
libm_buf[i] = data[i].exp();
}
libm_ns.push(t0.elapsed().as_nanos() as u64);
std::hint::black_box(&libm_buf);
}
libm_ns.sort();
let libm_median = libm_ns[ITERS / 2];
let mut sleef_buf = vec![0f32; N];
let mut sleef_ns = Vec::with_capacity(ITERS);
for _ in 0..ITERS {
let t0 = std::time::Instant::now();
for i in 0..N {
sleef_buf[i] = sleef_expf(data[i]);
}
sleef_ns.push(t0.elapsed().as_nanos() as u64);
std::hint::black_box(&sleef_buf);
}
sleef_ns.sort();
let sleef_median = sleef_ns[ITERS / 2];
let libm_mhz = (N as f64 / (libm_median as f64 / 1e9)) / 1e6;
let sleef_mhz = (N as f64 / (sleef_median as f64 / 1e9)) / 1e6;
let ratio = sleef_median as f64 / libm_median as f64;
eprintln!(
"\nsleef_expf vs libm exp microbench (N = {} f32, {} iters):\n \
libm f32::exp: {:.3} ms median ({:.1} M elem/s)\n \
sleef_expf (port): {:.3} ms median ({:.1} M elem/s)\n \
cost ratio: {:.2}× (1.0 = same speed)\n \
libm all iters: {:?}\n \
sleef all iters: {:?}",
N,
ITERS,
libm_median as f64 / 1e6,
libm_mhz,
sleef_median as f64 / 1e6,
sleef_mhz,
ratio,
libm_ns,
sleef_ns,
);
}
#[test]
fn sweep_shows_libm_divergence() {
let mut x: u32 = 0xCAFEBABE;
let mut differ = 0usize;
let mut total = 0usize;
for _ in 0..1024 {
x ^= x << 13;
x ^= x >> 17;
x ^= x << 5;
let f = (x as f32 / u32::MAX as f32) * 15.0 - 10.0;
let lb = f.exp().to_bits();
let sb = sleef_expf(f).to_bits();
if lb != sb {
differ += 1;
}
total += 1;
}
assert!(
differ * 100 >= total * 2,
"sleef_expf appears to shadow libm: only {}/{} diverged (expected ~8%)",
differ,
total
);
}
}