use std::f64::consts;
pub fn erf(x: f64) -> f64 {
if x.is_nan() {
return f64::NAN;
}
if x.is_infinite() {
return x.signum();
}
let ax = x.abs();
let result = if ax < 1e-8 {
ax * consts::FRAC_2_SQRT_PI
} else if ax < 5.0 {
let xsq = ax * ax;
let mut term = ax; let mut sum = ax / 1.0_f64; for n in 1..100usize {
term *= -xsq / n as f64;
let contrib = term / (2 * n + 1) as f64;
sum += contrib;
if contrib.abs() < sum.abs().max(1e-300) * f64::EPSILON {
break;
}
}
consts::FRAC_2_SQRT_PI * sum
} else {
let erfc_val = erfc(ax);
1.0 - erfc_val
};
if x < 0.0 { -result } else { result }
}
pub fn erfc(x: f64) -> f64 {
if x.is_nan() {
return f64::NAN;
}
if x < 0.0 {
return 2.0 - erfc(-x);
}
if x > 27.0 {
return 0.0;
}
if x < 0.5 {
return 1.0 - erf(x);
}
let t = 1.0 / (1.0 + 0.5 * x);
let poly = t
* (-x * x - 1.265_512_23
+ t * (1.000_023_68
+ t * (0.374_091_96
+ t * (0.096_784_18
+ t * (-0.186_288_06
+ t * (0.278_868_07
+ t * (-1.135_203_98
+ t * (1.488_515_87
+ t * (-0.822_152_23 + t * 0.170_872_94)))))))));
poly.exp() * t
}
pub fn lgamma(x: f64) -> f64 {
if x.is_nan() {
return f64::NAN;
}
if x <= 0.0 {
if x.fract() == 0.0 {
return f64::INFINITY;
}
let lnpi = consts::PI.ln();
let sinpix = (consts::PI * x).sin().abs();
return lnpi - sinpix.ln() - lgamma(1.0 - x);
}
if x < 0.5 {
let lnpi = consts::PI.ln();
let sinpix = (consts::PI * x).sin().abs();
return lnpi - sinpix.ln() - lgamma(1.0 - x);
}
const G: f64 = 7.0;
const C: [f64; 9] = [
0.999_999_999_999_809_9,
676.520_368_121_885_1,
-1_259.139_216_722_402_9,
771.323_428_777_653_1,
-176.615_029_162_140_6,
12.507_343_278_686_9,
-0.138_571_095_265_72,
9.984_369_578_019_57e-6,
1.505_632_735_149_31e-7,
];
let xm1 = x - 1.0;
let mut a = C[0];
for (k, &c) in C[1..].iter().enumerate() {
a += c / (xm1 + (k as f64) + 1.0);
}
let t = xm1 + G + 0.5;
(2.0 * consts::PI).sqrt().ln() + a.ln() + (xm1 + 0.5) * t.ln() - t
}
pub fn digamma(x: f64) -> f64 {
if x.is_nan() {
return f64::NAN;
}
if x <= 0.0 && x.fract() == 0.0 {
return f64::NEG_INFINITY;
}
if x < 0.5 {
return digamma(1.0 - x) - consts::PI / (consts::PI * x).tan();
}
let mut x = x;
let mut result = 0.0;
while x < 7.0 {
result -= 1.0 / x;
x += 1.0;
}
let xinv = 1.0 / x;
let xinv2 = xinv * xinv;
result += x.ln()
- 0.5 * xinv
- xinv2
* (1.0 / 12.0
- xinv2
* (1.0 / 120.0
- xinv2 * (1.0 / 252.0 - xinv2 * (1.0 / 240.0 - xinv2 * (1.0 / 132.0)))));
result
}
pub fn trigamma(x: f64) -> f64 {
if x.is_nan() {
return f64::NAN;
}
if x <= 0.0 && x == x.floor() {
return f64::INFINITY;
}
if x < 0.5 {
let pi_x = std::f64::consts::PI * x;
let s = pi_x.sin();
return std::f64::consts::PI * std::f64::consts::PI / (s * s) - trigamma(1.0 - x);
}
if x < 6.0 {
return trigamma(x + 1.0) + 1.0 / (x * x);
}
let inv_x = 1.0 / x;
let inv_x2 = inv_x * inv_x;
inv_x
+ 0.5 * inv_x2
+ inv_x2
* inv_x
* (1.0 / 6.0
+ inv_x2
* (-1.0 / 30.0
+ inv_x2 * (1.0 / 42.0 + inv_x2 * (-1.0 / 30.0 + inv_x2 * (5.0 / 66.0)))))
}
pub fn ei(x: f64) -> f64 {
if x.is_nan() {
return f64::NAN;
}
if x == 0.0 {
return f64::NEG_INFINITY;
}
if x < 0.0 {
return -e1(-x);
}
if x > 40.0 {
let ex_over_x = x.exp() / x;
let mut series = 1.0_f64;
let mut term = 1.0_f64;
for n in 1..30usize {
term *= n as f64 / x;
if term.abs() > 1e50 {
break;
}
series += term;
if term.abs() < series.abs() * 1e-15 {
break;
}
}
return ex_over_x * series;
}
const EULER: f64 = 0.577_215_664_901_532_9;
let mut series = 0.0_f64;
let mut term = x;
for n in 1..200usize {
series += term / n as f64;
term *= x / (n + 1) as f64;
if term.abs() < (series.abs().max(1.0)) * 1e-16 {
break;
}
}
EULER + x.abs().ln() + series
}
fn e1(x: f64) -> f64 {
if x <= 0.0 {
return f64::NAN;
}
if x <= 1.0 {
const EULER: f64 = 0.577_215_664_901_532_9;
let mut series = 0.0_f64;
let mut term = -x;
for n in 1..200usize {
series += term / n as f64;
term *= -x / (n + 1) as f64;
if term.abs() < (series.abs().max(1e-100)) * 1e-16 {
break;
}
}
return -EULER - x.ln() - series;
}
e1_cf(x)
}
fn e1_cf(x: f64) -> f64 {
let tiny = 1e-300_f64;
let mut f;
let mut c;
let mut d;
d = 1.0 / x.max(tiny);
f = d;
c = 1.0 / tiny;
for n in 1i64..200 {
let half_n = (n + 1) / 2;
let a = half_n as f64;
let b = if n % 2 == 1 { 1.0 } else { x };
d = 1.0 / (b + a * d);
c = b + a / c;
f *= c * d;
if (c * d - 1.0).abs() < 1e-15 {
break;
}
}
(-x).exp() * f
}
pub fn si(x: f64) -> f64 {
if x.is_nan() {
return f64::NAN;
}
let (si_val, _) = si_ci_internal(x.abs());
if x < 0.0 { -si_val } else { si_val }
}
pub fn ci(x: f64) -> f64 {
if x.is_nan() || x <= 0.0 {
return f64::NAN;
}
let (_, ci_val) = si_ci_internal(x);
ci_val
}
fn si_ci_internal(x: f64) -> (f64, f64) {
const EULER: f64 = 0.577_215_664_901_532_9;
if x < 1e-8 {
let ci_val = if x > 0.0 {
EULER + x.ln()
} else {
f64::NEG_INFINITY
};
return (x, ci_val);
}
if x < 4.0 {
let xsq = x * x;
let mut si_t = x; let mut si_acc = x; for n in 1..100usize {
si_t *= -xsq / ((2 * n) as f64 * (2 * n + 1) as f64);
let contrib = si_t / (2 * n + 1) as f64;
si_acc += contrib;
if contrib.abs() < si_acc.abs().max(1e-300) * f64::EPSILON {
break;
}
}
let mut u_n = -xsq / 2.0; let mut ci_sum = u_n / 2.0; for n in 2..100usize {
u_n *= -xsq / ((2 * n - 1) as f64 * (2 * n) as f64);
let contrib = u_n / (2 * n) as f64;
ci_sum += contrib;
if contrib.abs() < ci_sum.abs().max(1e-300) * f64::EPSILON {
break;
}
}
let ci_val = EULER + x.ln() + ci_sum;
return (si_acc, ci_val);
}
let (f_val, g_val) = si_ci_aux(x);
let sinx = x.sin();
let cosx = x.cos();
let si_val = consts::FRAC_PI_2 - f_val * cosx - g_val * sinx;
let ci_val = f_val * sinx - g_val * cosx;
(si_val, ci_val)
}
fn si_ci_aux(x: f64) -> (f64, f64) {
let xsq = x * x;
let mut f_sum = 1.0_f64;
let mut g_sum = 1.0_f64;
let mut f_term = 1.0_f64;
let mut g_term = 1.0_f64;
for n in 1..40usize {
let even = (2 * n) as f64;
let odd = (2 * n - 1) as f64;
f_term *= -odd * even / xsq;
g_term *= -(even) * (even + 1.0) / xsq;
if f_term.abs() >= f_sum.abs() {
break;
}
f_sum += f_term;
if g_term.abs() < g_sum.abs() {
g_sum += g_term;
}
}
(f_sum / x, g_sum / xsq)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_erf_known_values() {
assert!(erf(0.0).abs() < 1e-15);
assert!((erf(1.0) - 0.842_700_792_949_715).abs() < 1e-10);
assert!((erf(-1.0) + 0.842_700_792_949_715).abs() < 1e-10);
assert!((erf(2.0) - 0.995_322_265_004_329).abs() < 1e-8);
}
#[test]
fn test_erf_odd() {
for &x in &[0.1_f64, 0.5, 1.0, 2.0, 3.0] {
assert!((erf(x) + erf(-x)).abs() < 1e-12, "erf not odd at x={x}");
}
}
#[test]
fn test_lgamma_known() {
assert!(lgamma(1.0).abs() < 1e-12);
assert!(lgamma(2.0).abs() < 1e-12);
assert!((lgamma(3.0) - 2.0_f64.ln()).abs() < 1e-10);
assert!((lgamma(0.5) - 0.5 * consts::PI.ln()).abs() < 1e-10);
}
#[test]
fn test_lgamma_recurrence() {
for &x in &[1.0_f64, 2.0, 3.0, 5.0] {
let lhs = lgamma(x + 1.0);
let rhs = lgamma(x) + x.ln();
assert!((lhs - rhs).abs() < 1e-10, "lgamma recurrence at x={x}");
}
}
#[test]
fn test_digamma_known() {
const EULER: f64 = 0.577_215_664_901_532_9;
assert!((digamma(1.0) + EULER).abs() < 1e-10);
}
#[test]
fn test_digamma_recurrence() {
for &x in &[1.0_f64, 2.0, 3.0, 0.5] {
let diff = digamma(x + 1.0) - digamma(x);
assert!(
(diff - 1.0 / x).abs() < 1e-10,
"digamma recurrence at x={x}: got {diff}"
);
}
}
#[test]
fn test_ei_known() {
assert!((ei(1.0) - 1.895_117_816_355_937).abs() < 1e-8);
}
#[test]
fn test_si_known() {
assert!((si(1.0) - 0.946_083_070_367_183).abs() < 1e-8);
}
#[test]
fn test_si_odd() {
for &x in &[0.5_f64, 1.0, 2.0] {
assert!((si(x) + si(-x)).abs() < 1e-10, "Si not odd at x={x}");
}
}
#[test]
fn test_ci_positive() {
assert!((ci(1.0) - 0.337_403_922_900_968).abs() < 1e-8);
assert!(ci(-1.0).is_nan());
assert!(ci(0.0).is_nan());
}
#[test]
fn test_trigamma_known_values() {
use std::f64::consts::PI;
let expected_1 = PI * PI / 6.0;
assert!(
(trigamma(1.0) - expected_1).abs() < 1e-10,
"trigamma(1) = {}, expected {}",
trigamma(1.0),
expected_1
);
let expected_2 = PI * PI / 6.0 - 1.0;
assert!(
(trigamma(2.0) - expected_2).abs() < 1e-10,
"trigamma(2) = {}, expected {}",
trigamma(2.0),
expected_2
);
let expected_half = PI * PI / 2.0;
assert!(
(trigamma(0.5) - expected_half).abs() < 1e-10,
"trigamma(0.5) = {}, expected {}",
trigamma(0.5),
expected_half
);
}
#[test]
fn test_trigamma_recurrence() {
for &x in &[1.0f64, 2.0, 3.0, 0.5, 1.5, 2.5] {
let diff = trigamma(x) - trigamma(x + 1.0);
let expected = 1.0 / (x * x);
assert!(
(diff - expected).abs() < 1e-10,
"recurrence failed at x={}: {} != {}",
x,
diff,
expected
);
}
}
#[test]
fn test_trigamma_reflection() {
use std::f64::consts::PI;
for &x in &[0.3f64, 0.25, 0.1, 0.4] {
let lhs = trigamma(x) + trigamma(1.0 - x);
let s = (PI * x).sin();
let rhs = PI * PI / (s * s);
assert!(
(lhs - rhs).abs() < 1e-9,
"reflection failed at x={}: {} != {}",
x,
lhs,
rhs
);
}
}
#[test]
fn test_trigamma_pole() {
assert!(trigamma(0.0).is_infinite());
assert!(trigamma(-1.0).is_infinite());
assert!(trigamma(-2.0).is_infinite());
}
#[test]
fn test_trigamma_nan() {
assert!(trigamma(f64::NAN).is_nan());
}
}