#![allow(clippy::excessive_precision)]
use symplex::prelude::*;
fn rel(got: f64, want: f64) -> f64 {
if want == 0.0 {
got.abs()
} else {
((got - want) / want).abs()
}
}
fn assert_rel(got: f64, want: f64, tol: f64, what: &str) {
assert!(
rel(got, want) <= tol,
"{what}: got {got:.17e}, want {want:.17e}, rel err {:.3e} > {tol:.0e}",
rel(got, want)
);
}
fn assert_poly_eq(a: &Ex, b: &Ex, label: &str) {
let d = (a - b).expand().simplify();
assert!(
d.is_zero_structural(),
"{label}: `{a}` ≠ `{b}` (difference `{d}`)"
);
}
#[test]
fn compiled_erfinv_matches_mpmath() {
let ctx = Context::new();
let x = ctx.symbol("x");
let f = x.erfinv().compile(&["x"]).unwrap();
let cases: [(f64, f64); 12] = [
(0.1, 0.088855990494257691974),
(0.3, 0.27246271472675434502),
(0.46875, 0.44271885732435436322),
(0.5, 0.47693627620446987338),
(0.6, 0.59511608144999482198),
(0.75, 0.81341984759761854169),
(0.9, 1.1630871536766741628),
(0.99, 1.8213863677184494559),
(0.999, 2.3267537655135244939),
(0.99999, 3.123413274341570864),
(1.0 - 1e-9, 4.3200053881053620459),
(1.0 - 1e-12, 5.0420318985726961301),
];
for &(xv, want) in &cases {
assert_rel(f(&[xv]), want, 1e-15, &format!("erfinv({xv})"));
assert_rel(f(&[-xv]), -want, 1e-15, &format!("erfinv(-{xv})"));
}
assert_rel(
f(&[1.0 - 1e-15]),
5.6759157397447131788,
1e-14,
"erfinv(1-1e-15)",
);
assert_rel(
f(&[1.0 - f64::EPSILON]),
5.8050186831934533002,
1e-14,
"erfinv(1-2^-52)",
);
assert_rel(
f(&[1e-300]),
8.8622692545275803586e-301,
1e-15,
"erfinv(1e-300)",
);
assert_eq!(f(&[0.0]), 0.0);
assert_eq!(f(&[1.0]), f64::INFINITY);
assert_eq!(f(&[-1.0]), f64::NEG_INFINITY);
assert!(f(&[1.5]).is_nan());
assert!(f(&[-1.0000001]).is_nan());
assert!(f(&[f64::NAN]).is_nan());
}
#[test]
fn compiled_erfcinv_matches_mpmath() {
let ctx = Context::new();
let y = ctx.symbol("y");
let f = y.erfcinv().compile(&["y"]).unwrap();
let cases: [(f64, f64); 10] = [
(0.5, 0.47693627620446987338),
(1.5, -0.47693627620446987338),
(0.1, 1.1630871536766740677),
(1.9, -1.1630871536766737823),
(0.999, 0.00088622715746655289169),
(1e-5, 3.1234132743408750177),
(1e-10, 4.5728249673894852748),
(1e-20, 6.6015806223551425656),
(1e-100, 15.065574702592645704),
(1e-300, 26.209469960516123886),
];
for &(yv, want) in &cases {
assert_rel(f(&[yv]), want, 1e-15, &format!("erfcinv({yv})"));
}
assert_rel(
f(&[5e-324]),
27.213293210812948815,
1e-15,
"erfcinv(5e-324)",
);
assert_eq!(f(&[1.0]), 0.0);
assert_eq!(f(&[0.0]), f64::INFINITY);
assert_eq!(f(&[2.0]), f64::NEG_INFINITY);
assert!(f(&[-0.5]).is_nan());
assert!(f(&[2.5]).is_nan());
}
#[test]
fn compiled_normal_quantile() {
let ctx = Context::new();
let p = ctx.symbol("p");
let mu = ctx.symbol("mu");
let sigma = ctx.symbol("sigma");
let q = &mu + &sigma * ctx.int(2).sqrt() * (ctx.int(2) * &p - 1).erfinv();
let f = q.compile(&["mu", "sigma", "p"]).unwrap();
assert_eq!(f.arity(), 3);
assert_rel(
f(&[0.0, 1.0, 0.975]),
1.9599639845400538556,
1e-15,
"z_{0.975}",
);
assert_rel(
f(&[1.0, 2.0, 0.975]),
4.9199279690801077112,
1e-15,
"N(1,2) quantile",
);
assert_rel(
f(&[0.0, 1.0, 0.999]),
3.0902323061678132778,
1e-15,
"z_{0.999}",
);
assert_eq!(f(&[0.0, 1.0, 0.5]), 0.0);
assert_eq!(f(&[0.0, 1.0, 0.025]), -f(&[0.0, 1.0, 0.975]));
assert_eq!(f(&[0.0, 1.0, 1.0]), f64::INFINITY);
assert_eq!(f(&[0.0, 1.0, 0.0]), f64::NEG_INFINITY);
let g = q.exp().compile(&["mu", "sigma", "p"]).unwrap();
assert_rel(
g(&[1.0, 2.0, 0.975]),
4.9199279690801077112f64.exp(),
1e-15,
"LogNormal(1,2) quantile",
);
}
#[test]
fn erfinv_rust_codegen_uses_the_shared_runtime() {
let ctx = Context::new();
let p = ctx.symbol("p");
let q = ctx.int(2).sqrt() * (ctx.int(2) * &p - 1).erfinv();
let code = q.to_rust_fn("quantile", &["p"]).unwrap();
assert!(code.contains("symplex_rt::erfinv("), "{code}");
assert!(
code.contains("pub fn erfinv("),
"runtime section not embedded:\n{code}"
);
assert!(
code.contains("fn erfcx_large("),
"erf section dependency missing:\n{code}"
);
let code = p.erfcinv().to_rust_fn("f", &["p"]).unwrap();
assert!(code.contains("symplex_rt::erfcinv("), "{code}");
let code = p.erfinv().to_c_fn("f", &["p"]).unwrap();
assert!(code.contains("symplex_erfinv("), "{code}");
}
#[test]
fn negative_binomial_series_numeric_c() {
let ctx = Context::new();
let k = ctx.symbol("k");
let zero = ctx.int(0);
let inf = ctx.infinity();
let third = ctx.rational(1, 3);
let bin = (&k + ctx.int(3)).binomial(&k);
let s = (&bin * third.pow(&k)).summation(&k, &zero, &inf);
assert_eq!(s, ctx.rational(81, 16), "{s}");
let bin3 = (&k + ctx.int(3)).binomial(&ctx.int(3));
let s = (&bin3 * third.pow(&k)).summation(&k, &zero, &inf);
assert_eq!(s, ctx.rational(81, 16), "{s}");
let s = (&k * &bin * third.pow(&k)).summation(&k, &zero, &inf);
assert_eq!(s, ctx.rational(81, 8), "{s}");
let s = (k.powi(2) * &bin * third.pow(&k)).summation(&k, &zero, &inf);
assert_eq!(s, ctx.rational(567, 16), "{s}");
}
#[test]
fn negative_binomial_pmf_total_mass_and_moments() {
let ctx = Context::new();
let k = ctx.symbol("k");
let zero = ctx.int(0);
let inf = ctx.infinity();
let half = ctx.rational(1, 2);
let pmf = (&k + ctx.int(2)).binomial(&k) * half.powi(3) * half.pow(&k);
assert_eq!(pmf.summation(&k, &zero, &inf), ctx.int(1));
assert_eq!((&k * &pmf).summation(&k, &zero, &inf), ctx.int(3));
assert_eq!((k.powi(2) * &pmf).summation(&k, &zero, &inf), ctx.int(15));
assert_eq!(pmf.summation(&k, &ctx.int(3), &inf), ctx.rational(1, 2));
let r = ctx.symbol("r");
let pmf_r = (&k + &r - 1).binomial(&k) * half.pow(&r) * half.pow(&k);
assert_eq!(pmf_r.summation(&k, &zero, &inf).simplify(), ctx.int(1));
}
#[test]
fn negative_binomial_series_symbolic_c() {
let ctx = Context::new();
let k = ctx.symbol("k");
let c = ctx.symbol("c");
let zero = ctx.int(0);
let inf = ctx.infinity();
let third = ctx.rational(1, 3);
let s = ((&k + &c).binomial(&k) * third.pow(&k)).summation(&k, &zero, &inf);
assert!(!s.has_unevaluated(), "{s}");
let expected = ctx.rational(2, 3).pow(&(-&c - 1));
assert_eq!(s, expected, "{s}");
let s2 = ((&k + &c).binomial(&c) * third.pow(&k)).summation(&k, &zero, &inf);
assert_eq!(s2, expected, "{s2}");
let s3 = (&k * (&k + &c).binomial(&k) * third.pow(&k)).summation(&k, &zero, &inf);
assert!(!s3.has_unevaluated(), "{s3}");
assert_eq!(s3.subs_i64(&c, 5).eval(), ctx.rational(2187, 64), "{s3}");
assert_eq!(s.subs_i64(&c, 2).eval(), ctx.rational(27, 8));
}
#[test]
fn negative_binomial_series_symbolic_ratio_stays_formal() {
let ctx = Context::new();
let k = ctx.symbol("k");
let c = ctx.symbol("c");
let x = ctx.symbol("x");
let p = ctx.symbol("p");
let r = ctx.symbol("r");
let zero = ctx.int(0);
let inf = ctx.infinity();
let s = ((&k + &c).binomial(&k) * x.pow(&k)).summation(&k, &zero, &inf);
assert!(s.has_unevaluated(), "{s}");
let s = ((&k + ctx.int(3)).binomial(&k) * x.pow(&k)).summation(&k, &zero, &inf);
assert!(s.has_unevaluated(), "{s}");
let pmf = (&k + &r - 1).binomial(&k) * p.pow(&r) * (ctx.one() - &p).pow(&k);
assert!(pmf.summation(&k, &zero, &inf).has_unevaluated());
let s = ((&k + ctx.int(3)).binomial(&k) * ctx.int(3).pow(&k)).summation(&k, &zero, &inf);
assert_eq!(s, ctx.infinity(), "{s}");
}
#[test]
fn negative_binomial_finite_sum_via_gosper() {
let ctx = Context::new();
let k = ctx.symbol("k");
let n = ctx.symbol("n");
let half = ctx.rational(1, 2);
let body = (&k + ctx.int(2)).binomial(&k) * half.pow(&k);
let s = body.summation(&k, &ctx.int(0), &n);
assert!(!s.has_unevaluated(), "{s}");
for (nv, want) in [
(0, ctx.int(1)),
(1, ctx.rational(5, 2)),
(4, ctx.rational(99, 16)),
(7, ctx.rational(121, 16)),
] {
let v = s.subs_i64(&n, nv).eval().simplify();
assert_eq!(v, want, "n = {nv}: {v}");
}
assert_eq!(
body.summation(&k, &ctx.int(0), &ctx.int(4)),
ctx.rational(99, 16)
);
}
#[test]
fn poisson_total_mass_and_moments_symbolic_rate() {
let ctx = Context::new();
let k = ctx.symbol("k");
let lam = ctx.symbol("lambda");
let zero = ctx.int(0);
let inf = ctx.infinity();
let pmf = lam.pow(&k) * (-&lam).exp() / k.factorial();
let total = pmf.summation(&k, &zero, &inf);
assert!(!total.has_unevaluated(), "{total}");
assert_eq!(total.simplify(), ctx.int(1), "{total}");
let mean = (&k * &pmf).summation(&k, &zero, &inf);
assert_eq!(mean.simplify(), lam, "{mean}");
let second = (k.powi(2) * &pmf).summation(&k, &zero, &inf);
assert_poly_eq(&second.simplify(), &(lam.powi(2) + &lam), "Poisson E[X²]");
}
#[test]
fn geometric_mean_numeric_p_closes_symbolic_p_stays_formal() {
let ctx = Context::new();
let k = ctx.symbol("k");
let one = ctx.int(1);
let inf = ctx.infinity();
let pmf = ctx.rational(2, 3).pow(&(&k - 1)) * ctx.rational(1, 3);
assert_eq!(pmf.summation(&k, &one, &inf), ctx.int(1));
assert_eq!((&k * &pmf).summation(&k, &one, &inf), ctx.int(3));
assert_eq!((k.powi(2) * &pmf).summation(&k, &one, &inf), ctx.int(15));
let p = ctx.symbol("p");
let pmf_p = (ctx.one() - &p).pow(&(&k - 1)) * &p;
assert!(pmf_p.summation(&k, &one, &inf).has_unevaluated());
assert!((&k * &pmf_p).summation(&k, &one, &inf).has_unevaluated());
let r = ctx.symbol("r");
assert!(r.pow(&k).summation(&k, &ctx.int(0), &inf).has_unevaluated());
}
#[test]
fn binomial_theorem_symbolic_n_and_p() {
let ctx = Context::new();
let k = ctx.symbol("k");
let n = ctx.symbol("n");
let p = ctx.symbol("p");
let zero = ctx.int(0);
let pmf = n.binomial(&k) * p.pow(&k) * (ctx.one() - &p).pow(&(&n - &k));
let total = pmf.summation(&k, &zero, &n);
assert_eq!(total, ctx.int(1), "{total}");
let mean = (&k * &pmf).summation(&k, &zero, &n);
assert_eq!(mean.simplify(), &n * &p, "{mean}");
let second = (k.powi(2) * &pmf).summation(&k, &zero, &n);
assert!(!second.has_unevaluated(), "{second}");
let want = n.powi(2) * p.powi(2) - &n * p.powi(2) + &n * &p;
assert_poly_eq(&second, &want, "Binomial E[X²]");
let pmf_num = n.binomial(&k) * ctx.rational(1, 3).pow(&k) * ctx.rational(2, 3).pow(&(&n - &k));
assert_eq!(pmf_num.summation(&k, &zero, &n), ctx.int(1));
}
#[test]
fn binomial_coefficient_with_k_greater_than_n_is_zero_on_every_route() {
let ctx = Context::new();
let c12 = ctx.int(1).binomial(&ctx.int(2));
assert_eq!(c12.eval(), ctx.int(0));
assert_eq!(c12.eval_f64().unwrap(), 0.0);
assert_eq!(ctx.int(2).binomial(&ctx.int(5)).eval(), ctx.int(0));
assert_eq!(ctx.int(0).binomial(&ctx.int(1)).eval(), ctx.int(0));
assert_eq!(ctx.int(5).binomial(&ctx.int(2)).eval(), ctx.int(10));
let n = ctx.symbol("n");
let k = ctx.symbol("k");
let f = n.binomial(&k).compile(&["n", "k"]).unwrap();
assert_eq!(f(&[1.0, 2.0]), 0.0);
assert_eq!(f(&[2.0, 5.0]), 0.0);
assert_eq!(f(&[0.0, 1.0]), 0.0);
}