use super::exp::simd_exp;
use super::traits::SimdFloat;
#[inline]
pub fn simd_sigmoid<S: SimdFloat>(x: S) -> S {
let neg_abs_x = x.abs().neg(); let exp_neg = simd_exp(neg_abs_x); let denom = S::one().add(exp_neg);
let sig_pos = S::one().div(denom); let sig_neg = exp_neg.div(denom);
let mask = x.cmp_ge(S::zero());
S::blend(mask, sig_pos, sig_neg)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::simd::generic::ScalarFloat;
#[test]
fn test_sigmoid_accuracy() {
let n = 1_000_000;
let mut max_abs_err: f32 = 0.0;
for i in 0..n {
let x = -10.0 + 20.0 * (i as f32) / (n as f32);
let expected = 1.0 / (1.0 + (-x).exp());
let result = simd_sigmoid(ScalarFloat(x)).0;
let abs_err = (result - expected).abs();
max_abs_err = max_abs_err.max(abs_err);
}
assert!(
max_abs_err < 1e-5,
"sigmoid max abs error {max_abs_err} exceeds 1e-5"
);
}
#[test]
fn test_sigmoid_properties() {
let r = simd_sigmoid(ScalarFloat(0.0)).0;
assert!((r - 0.5).abs() < 1e-6, "sigmoid(0) = {r}");
for &x in &[1.0, 2.0, 5.0, -3.0] {
let s_pos = simd_sigmoid(ScalarFloat(x)).0;
let s_neg = simd_sigmoid(ScalarFloat(-x)).0;
assert!(
(s_pos + s_neg - 1.0).abs() < 1e-5,
"sigmoid({x}) + sigmoid({}) = {}",
-x,
s_pos + s_neg
);
}
let r = simd_sigmoid(ScalarFloat(50.0)).0;
assert!((r - 1.0).abs() < 1e-5, "sigmoid(50) = {r}");
let r = simd_sigmoid(ScalarFloat(-50.0)).0;
assert!(r.abs() < 1e-5, "sigmoid(-50) = {r}");
}
}