use crate::base::errors::SymplexError;
use crate::output::codegen::numeric_rt::{erfc, erfcinv, lgamma};
const LN_SQRT_2PI: f64 = 0.918_938_533_204_672_7;
const SQRT_2PI: f64 = 2.506_628_274_631_000_2;
const EPS: f64 = f64::EPSILON;
const GAMMA_MAX_ITER: usize = 10_000_000;
#[derive(Clone, Copy, Debug, PartialEq)]
struct Tails {
lower: f64,
upper: f64,
}
impl Tails {
const NAN: Tails = Tails {
lower: f64::NAN,
upper: f64::NAN,
};
const ZERO: Tails = Tails {
lower: 0.0,
upper: 1.0,
};
const ONE: Tails = Tails {
lower: 1.0,
upper: 0.0,
};
fn flipped(self) -> Tails {
Tails {
lower: self.upper,
upper: self.lower,
}
}
}
fn stirlerr(z: f64) -> f64 {
if z >= 10.0 {
let inv = 1.0 / z;
let inv2 = inv * inv;
inv * (1.0 / 12.0
- inv2
* (1.0 / 360.0
- inv2
* (1.0 / 1260.0
- inv2
* (1.0 / 1680.0
- inv2
* (1.0 / 1188.0
- inv2 * (691.0 / 360_360.0 - inv2 / 156.0))))))
} else {
lgamma(z) - ((z - 0.5) * z.ln() - z + LN_SQRT_2PI)
}
}
fn lbeta(a: f64, b: f64) -> f64 {
let (p, q) = if a < b { (a, b) } else { (b, a) };
if p >= 10.0 {
let corr = stirlerr(p) + stirlerr(q) - stirlerr(p + q);
-0.5 * q.ln()
+ LN_SQRT_2PI
+ corr
+ (p - 0.5) * (p / (p + q)).ln()
+ q * (-p / (p + q)).ln_1p()
} else if q >= 10.0 {
let corr = stirlerr(q) - stirlerr(p + q);
lgamma(p) + corr + p - p * (p + q).ln() + (q - 0.5) * (-p / (p + q)).ln_1p()
} else {
lgamma(p) + lgamma(q) - lgamma(p + q)
}
}
fn rlog1(x: f64) -> f64 {
if !(-0.5..=0.5).contains(&x) {
return x - x.ln_1p();
}
let r = x / (2.0 + x);
let t = r * r;
let mut w = 1.0 / 3.0;
let mut term = 1.0;
let mut k = 1.0;
loop {
term *= t;
k += 1.0;
let add = term / (2.0 * k + 1.0);
w += add;
if add <= EPS * w {
break;
}
}
2.0 * t * (1.0 / (1.0 - r) - r * w)
}
fn bd0_with(k: f64, m: f64, diff: f64, ln_ratio: f64) -> f64 {
if diff.abs() < 0.1 * (k + m) {
let v = diff / (k + m);
let v2 = v * v;
let mut s = diff * v;
let mut ej = 2.0 * k * v;
let mut j = 1.0;
loop {
ej *= v2;
let s1 = s + ej / (2.0 * j + 1.0);
if s1 == s || j > 1000.0 {
return s1;
}
s = s1;
j += 1.0;
}
}
k * ln_ratio - diff
}
fn bd0(k: f64, m: f64) -> f64 {
bd0_with(k, m, k - m, (k / m).ln())
}
fn log_beta_pref(a: f64, b: f64, x: f64, y: f64) -> f64 {
if x <= 0.0 || y <= 0.0 {
return f64::NEG_INFINITY;
}
let n = a + b;
let (s, small, large) = if y <= x { (y, b, a) } else { (x, a, b) };
let ln_l = (-s).ln_1p();
let ns = n * s;
let d = ns - small;
let t_small = bd0_with(small, ns, -d, (small / n).ln() - s.ln());
let nl = n - ns;
let t_large = bd0_with(large, nl, d, (large / n).ln() - ln_l);
-t_small - t_large - stirlerr(a) - stirlerr(b) + stirlerr(n) + 0.5 * (a * b / n).ln()
- LN_SQRT_2PI
}
fn log_gamma_pref(a: f64, x: f64) -> f64 {
-stirlerr(a) - bd0(a, x) + 0.5 * a.ln() - LN_SQRT_2PI
}
fn gamma_p_series_over_r(a: f64, x: f64) -> f64 {
let mut term = 1.0 / a;
let mut sum = term;
let mut ap = a;
for _ in 0..GAMMA_MAX_ITER {
ap += 1.0;
term *= x / ap;
sum += term;
if term <= sum * EPS {
break;
}
}
sum
}
fn gamma_q_cf_over_r(a: f64, x: f64) -> f64 {
const TINY: f64 = 1e-300;
let mut b = x + 1.0 - a;
let mut c = 1.0 / TINY;
let mut d = 1.0 / b;
let mut h = d;
let mut i = 1.0;
for _ in 0..GAMMA_MAX_ITER {
let an = -i * (i - a);
b += 2.0;
d = an * d + b;
if d.abs() < TINY {
d = TINY;
}
c = b + an / c;
if c.abs() < TINY {
c = TINY;
}
d = 1.0 / d;
let del = d * c;
h *= del;
if (del - 1.0).abs() < EPS {
break;
}
i += 1.0;
}
h
}
const GAMMA_TEMME_MIN_A: f64 = 1e6;
fn temme_eta(d: f64) -> f64 {
let phi = if d.abs() < 0.1 {
let mut term = d * d; let mut sum = 0.0;
let mut k = 2.0;
let mut positive = true; loop {
let contribution = term / k;
let next = if positive {
sum + contribution
} else {
sum - contribution
};
positive = !positive;
if next == sum || k > 60.0 {
break next;
}
sum = next;
term *= d;
k += 1.0;
}
} else {
d - d.ln_1p()
};
let eta = (2.0 * phi).sqrt();
if d < 0.0 { -eta } else { eta }
}
fn temme_c0_c1(eta: f64, d: f64) -> (f64, f64) {
if eta.abs() < 0.05 {
let e = eta;
let c0 = -1.0 / 3.0
+ e * (1.0 / 12.0
+ e * (-2.0 / 135.0
+ e * (1.0 / 864.0
+ e * (1.0 / 2835.0
+ e * (-139.0 / 777_600.0
+ e * (1.0 / 25_515.0
+ e * (-571.0 / 261_273_600.0
+ e * (-281.0 / 151_559_100.0))))))));
let c1 = -1.0 / 540.0
+ e * (-1.0 / 288.0
+ e * (1.0 / 378.0
+ e * (-77.0 / 77_760.0
+ e * (1.0 / 4860.0
+ e * (-1.0 / 2_488_320.0 + e * (-2743.0 / 151_559_100.0))))));
return (c0, c1);
}
let lm1 = d;
let c0 = 1.0 / lm1 - 1.0 / eta;
let dlambda = eta * (1.0 + d) / lm1;
let dc0 = -dlambda / (lm1 * lm1) + 1.0 / (eta * eta);
let c1 = dc0 / eta - (1.0 / 12.0) / lm1;
(c0, c1)
}
fn gammainc_tails_temme(a: f64, x: f64) -> Tails {
let d = (x - a) / a;
let eta = temme_eta(d);
let z = eta * (a / 2.0).sqrt();
let (c0, c1) = temme_c0_c1(eta, d);
let r = (-0.5 * a * eta * eta).exp() / (2.0 * std::f64::consts::PI * a).sqrt() * (c0 + c1 / a);
Tails {
lower: 0.5 * erfc(-z) - r,
upper: 0.5 * erfc(z) + r,
}
}
fn gammainc_tails(a: f64, x: f64) -> Tails {
if a.is_nan() || x.is_nan() || a <= 0.0 || a.is_infinite() || x < 0.0 {
return Tails::NAN;
}
if x == 0.0 {
return Tails::ZERO;
}
if x.is_infinite() {
return Tails::ONE;
}
if a >= GAMMA_TEMME_MIN_A {
let sd = a.sqrt();
if (x - a).abs() < 40.0 * sd {
return gammainc_tails_temme(a, x);
}
}
let r = log_gamma_pref(a, x).exp();
if x < a + 1.0 {
let p = r * gamma_p_series_over_r(a, x);
Tails {
lower: p,
upper: 1.0 - p,
}
} else {
let q = r * gamma_q_cf_over_r(a, x);
Tails {
lower: 1.0 - q,
upper: q,
}
}
}
pub fn gammainc_lower_regularized_f64(a: f64, x: f64) -> f64 {
gammainc_tails(a, x).lower
}
pub fn gammainc_upper_regularized_f64(a: f64, x: f64) -> f64 {
gammainc_tails(a, x).upper
}
fn bpser(a: f64, b: f64, x: f64) -> f64 {
if x == 0.0 {
return 0.0;
}
let lead = (a * x.ln() - lbeta(a, b)).exp() / a;
if lead == 0.0 {
return 0.0;
}
let tol = EPS / a;
let mut n = 0.0;
let mut sum = 0.0;
let mut c = 1.0;
loop {
n += 1.0;
c *= (1.0 - b / n) * x;
let w = c / (a + n);
sum += w;
if w.abs() <= tol || n >= 1e7 {
break;
}
}
lead * (1.0 + a * sum)
}
fn bup(a: f64, b: f64, x: f64, y: f64, n: usize) -> f64 {
let log_t0 = log_beta_pref(a, b, x, y) - a.ln();
if log_t0 == f64::NEG_INFINITY {
return 0.0;
}
let apb = a + b;
let ap1 = a + 1.0;
let lnx = x.ln();
let mut log_t = log_t0;
let mut max = log_t;
let mut sum = 1.0;
for j in 1..n {
let jf = (j - 1) as f64;
log_t += lnx + ((apb + jf) / (ap1 + jf)).ln();
if log_t > max {
sum = sum * (max - log_t).exp() + 1.0;
max = log_t;
} else {
let term = (log_t - max).exp();
sum += term;
if term <= EPS * sum {
break;
}
}
}
max.exp() * sum
}
fn gam1(a: f64) -> f64 {
let lg = lgamma(a + 1.0);
-lg.exp_m1() * (-lg).exp()
}
fn grat_r(a: f64, x: f64, log_r: f64) -> f64 {
if a * x == 0.0 {
return if x <= a { (-log_r).exp() } else { 0.0 };
}
if x >= 1.1 {
return gamma_q_cf_over_r(a, x);
}
let mut an = 3.0;
let mut c = x;
let mut sum = x / (a + 3.0);
let tol = 0.1 * EPS / (a + 1.0);
loop {
an += 1.0;
c *= -(x / an);
let t = c / (a + an);
sum += t;
if t.abs() <= tol {
break;
}
}
let j = a * x * ((sum / 6.0 - 0.5 / (a + 2.0)) * x + 1.0 / (a + 1.0));
let z = a * x.ln();
let h = gam1(a);
let g = h + 1.0;
if (x >= 0.25 && a < x / 2.59) || z > -0.13394 {
let l = z.exp_m1();
let q = ((l + 1.0) * j - l) * g - h;
if q <= 0.0 { 0.0 } else { q * (-log_r).exp() }
} else {
let p = z.exp() * g * (1.0 - j);
(1.0 - p) * (-log_r).exp()
}
}
fn bgrat(a: f64, b: f64, x: f64, y: f64) -> Option<f64> {
const N_TERMS: usize = 30;
let bm1 = b - 1.0;
let nu = a + 0.5 * bm1;
let lnx = if y > 0.375 { x.ln() } else { (-y).ln_1p() };
let z = -nu * lnx;
if b * z == 0.0 {
return None;
}
let log_r = b * z.ln() - z - lgamma(b);
let log_u = log_r - ((lbeta(a, b) - lgamma(b)) + b * nu.ln());
if log_u == f64::NEG_INFINITY {
return Some(0.0);
}
let u = log_u.exp();
let v = 0.25 / (nu * nu);
let t2 = lnx * 0.25 * lnx;
let mut j = grat_r(b, z, log_r);
let mut sum = j;
let mut t = 1.0;
let mut cn = 1.0;
let mut n2 = 0.0;
let mut c = [0.0; N_TERMS];
let mut d = [0.0; N_TERMS];
for n in 1..=N_TERMS {
let bp2n = b + n2;
j = (bp2n * (bp2n + 1.0) * j + (z + bp2n + 1.0) * t) * v;
n2 += 2.0;
t *= t2;
cn /= n2 * (n2 + 1.0);
c[n - 1] = cn;
let mut s = 0.0;
if n > 1 {
let mut coef = b - n as f64;
for i in 1..n {
s += coef * c[i - 1] * d[n - 1 - i];
coef += b;
}
}
d[n - 1] = bm1 * cn + s / n as f64;
let dj = d[n - 1] * j;
sum += dj;
if sum <= 0.0 {
return None;
}
if dj.abs() <= 15.0 * EPS * sum {
break;
}
}
Some(if u == 0.0 {
(log_u + sum.ln()).exp()
} else {
u * sum
})
}
fn bfrac(a: f64, b: f64, x: f64, y: f64, lambda: f64) -> f64 {
let brc = log_beta_pref(a, b, x, y).exp();
if brc == 0.0 {
return 0.0;
}
let c = lambda + 1.0;
let c0 = b / a;
let c1 = 1.0 / a + 1.0;
let yp1 = y + 1.0;
let mut n = 0.0;
let mut p = 1.0;
let mut s = a + 1.0;
let mut an = 0.0;
let mut bn = 1.0;
let mut anp1 = 1.0;
let mut bnp1 = c / c1;
let mut r = c1 / c;
for _ in 0..10_000 {
n += 1.0;
let w = n * x * (b - n);
let t = n / a;
let e = a / s;
let alpha = p * (p + c0) * e * e * (w * x);
let e = (t + 1.0) / (c1 + t + t);
let beta = w / s + n + e * (c + n * yp1);
p = t + 1.0;
s += 2.0;
let t = alpha * an + beta * anp1;
an = anp1;
anp1 = t;
let t = alpha * bn + beta * bnp1;
bn = bnp1;
bnp1 = t;
let r0 = r;
r = anp1 / bnp1;
if (r - r0).abs() <= 15.0 * EPS * r {
break;
}
an /= bnp1;
bn /= bnp1;
anp1 = r;
bnp1 = 1.0;
}
brc * r
}
fn erfcx(x: f64) -> f64 {
(x * x).exp() * erfc(x)
}
fn basym(a: f64, b: f64, lambda: f64) -> f64 {
const NUM: usize = 20;
const E0: f64 = std::f64::consts::FRAC_2_SQRT_PI;
const E1: f64 = 0.353_553_390_593_273_7;
let f = a * rlog1(-lambda / a) + b * rlog1(lambda / b);
let t = (-f).exp();
if t == 0.0 {
return 0.0;
}
let z0 = f.sqrt();
let z = 0.5 * z0 / E1;
let z2 = f + f;
let (h, r0, r1, w0) = if a < b {
let h = a / b;
(
h,
1.0 / (h + 1.0),
(b - a) / b,
1.0 / (a * (h + 1.0)).sqrt(),
)
} else {
let h = b / a;
(
h,
1.0 / (h + 1.0),
(b - a) / a,
1.0 / (b * (h + 1.0)).sqrt(),
)
};
let mut a0 = [0.0; NUM + 1];
let mut b0 = [0.0; NUM + 1];
let mut c = [0.0; NUM + 1];
let mut d = [0.0; NUM + 1];
a0[0] = r1 * 2.0 / 3.0;
c[0] = -0.5 * a0[0];
d[0] = -c[0];
let mut j0 = 0.5 / E0 * erfcx(z0);
let mut j1 = E1;
let mut sum = j0 + d[0] * w0 * j1;
let mut s = 1.0;
let h2 = h * h;
let mut hn = 1.0;
let mut w = w0;
let mut znm1 = z;
let mut zn = z2;
let mut n = 2;
while n <= NUM {
hn *= h2;
a0[n - 1] = r0 * 2.0 * (h * hn + 1.0) / (n as f64 + 2.0);
let np1 = n + 1;
s += hn;
a0[np1 - 1] = r1 * 2.0 * s / (n as f64 + 3.0);
for i in n..=np1 {
let r = -0.5 * (i as f64 + 1.0);
b0[0] = r * a0[0];
for m in 2..=i {
let mut bsum = 0.0;
for j in 1..m {
let mmj = m - j;
bsum += (j as f64 * r - mmj as f64) * a0[j - 1] * b0[mmj - 1];
}
b0[m - 1] = r * a0[m - 1] + bsum / m as f64;
}
c[i - 1] = b0[i - 1] / (i as f64 + 1.0);
let mut dsum = 0.0;
for j in 1..i {
dsum += d[i - j - 1] * c[j - 1];
}
d[i - 1] = -(dsum + c[i - 1]);
}
j0 = E1 * znm1 + (n as f64 - 1.0) * j0;
j1 = E1 * zn + n as f64 * j1;
znm1 *= z2;
zn *= z2;
w *= w0;
let t0 = d[n - 1] * w * j0;
w *= w0;
let t1 = d[np1 - 1] * w * j1;
sum += t0 + t1;
if t0.abs() + t1.abs() <= 100.0 * EPS * sum {
break;
}
n += 2;
}
let bcorr = stirlerr(a) + stirlerr(b) - stirlerr(a + b);
E0 * t * (-bcorr).exp() * sum
}
#[derive(Clone, Copy, Debug)]
enum Side {
Lower(f64),
Upper(f64),
}
fn bratio_route(a: f64, b: f64, x: f64, y: f64) -> Option<(Side, bool)> {
let (side, swap) = if a.min(b) <= 1.0 {
let swap = x > 0.5;
let (a0, b0, x0, y0) = if swap { (b, a, y, x) } else { (a, b, x, y) };
let side = if a0.max(b0) > 1.0 {
if b0 <= 1.0 {
Side::Lower(bpser(a0, b0, x0))
} else if x0 >= 0.29 {
Side::Upper(bpser(b0, a0, y0))
} else if x0 < 0.1 && (x0 * b0).powf(a0) <= 0.7 {
Side::Lower(bpser(a0, b0, x0))
} else if b0 > 15.0 {
Side::Upper(bgrat(b0, a0, y0, x0)?)
} else {
Side::Upper(bup(b0, a0, y0, x0, 20) + bgrat(b0 + 20.0, a0, y0, x0)?)
}
} else if a0 >= 0.2f64.min(b0) || x0.powf(a0) <= 0.9 {
Side::Lower(bpser(a0, b0, x0))
} else if x0 >= 0.3 {
Side::Upper(bpser(b0, a0, y0))
} else {
Side::Upper(bup(b0, a0, y0, x0, 20) + bgrat(b0 + 20.0, a0, y0, x0)?)
};
(side, swap)
} else {
let lambda = if a > b {
(a + b) * y - b
} else {
a - (a + b) * x
};
let swap = lambda < 0.0;
let (a0, b0, x0, y0, lambda) = if swap {
(b, a, y, x, -lambda)
} else {
(a, b, x, y, lambda)
};
let side = if b0 < 40.0 {
if b0 * x0 <= 0.7 {
Side::Lower(bpser(a0, b0, x0))
} else {
let mut n = b0.floor();
let mut bf = b0 - n;
if bf == 0.0 {
n -= 1.0;
bf = 1.0;
}
let mut w = bup(bf, a0, y0, x0, n as usize);
if x0 <= 0.7 {
w += bpser(a0, bf, x0);
} else {
let mut aa = a0;
if aa <= 15.0 {
w += bup(aa, bf, x0, y0, 20);
aa += 20.0;
}
w += bgrat(aa, bf, x0, y0)?;
}
Side::Lower(w)
}
} else if a0 > b0 {
if b0 <= 100.0 || lambda > 0.03 * b0 {
Side::Lower(bfrac(a0, b0, x0, y0, lambda))
} else {
Side::Lower(basym(a0, b0, lambda))
}
} else if a0 <= 100.0 || lambda > 0.03 * a0 {
Side::Lower(bfrac(a0, b0, x0, y0, lambda))
} else {
Side::Lower(basym(a0, b0, lambda))
};
(side, swap)
};
Some((side, swap))
}
fn bratio(a: f64, b: f64, x: f64, y: f64) -> Tails {
if a.is_nan()
|| b.is_nan()
|| x.is_nan()
|| y.is_nan()
|| a <= 0.0
|| b <= 0.0
|| !(0.0..=1.0).contains(&x)
|| !(0.0..=1.0).contains(&y)
{
return Tails::NAN;
}
if x == 0.0 {
return Tails::ZERO;
}
if y == 0.0 {
return Tails::ONE;
}
let Some((side, swap)) = bratio_route(a, b, x, y) else {
return Tails::NAN;
};
let tails = match side {
Side::Lower(w) => Tails {
lower: w,
upper: 1.0 - w,
},
Side::Upper(w1) => Tails {
lower: 1.0 - w1,
upper: w1,
},
};
if swap { tails.flipped() } else { tails }
}
pub fn betainc_regularized_f64(a: f64, b: f64, x: f64) -> f64 {
bratio(a, b, x, 1.0 - x).lower
}
fn invalid(op: &'static str, reason: String) -> SymplexError {
SymplexError::invalid_argument(op, reason)
}
fn check_level(op: &'static str, name: &str, p: f64) -> Result<(), SymplexError> {
if p > 0.0 && p < 1.0 {
Ok(())
} else {
Err(invalid(
op,
format!("{name} must lie strictly between 0 and 1, got {p}"),
))
}
}
fn check_positive(op: &'static str, name: &str, v: f64) -> Result<(), SymplexError> {
if v.is_finite() && v > 0.0 {
Ok(())
} else {
Err(invalid(
op,
format!("{name} must be finite and positive, got {v}"),
))
}
}
fn nearest_tail(p: f64) -> Side {
if p <= 0.5 {
Side::Lower(p)
} else {
Side::Upper(1.0 - p)
}
}
fn nearest_tail_upper(q: f64) -> Side {
if q <= 0.5 {
Side::Upper(q)
} else {
Side::Lower(1.0 - q)
}
}
struct Eval {
g: f64,
dg: f64,
}
fn solve_increasing(
op: &'static str,
g: impl Fn(f64) -> Eval,
v0: f64,
lo: f64,
hi: f64,
) -> Result<f64, SymplexError> {
const MAX_ITER: usize = 200;
const TOL: f64 = 16.0 * EPS;
const QUADRATIC_PHASE: f64 = 1e-6;
if v0.is_nan() {
return Err(SymplexError::computation_failed(
op,
"the starting point of the quantile iteration is not a number",
));
}
let (mut lo, mut hi) = (lo, hi);
let mut v = v0.clamp(lo, hi);
let mut step = 1.0;
let mut width = f64::INFINITY;
let mut best = (v, f64::INFINITY);
for _ in 0..MAX_ITER {
let Eval { g: gv, dg } = g(v);
if gv.is_nan() {
return Err(SymplexError::computation_failed(
op,
format!("the distribution function is not a number at {v}"),
));
}
if gv == 0.0 {
return Ok(v);
}
if gv.abs() < best.1 {
best = (v, gv.abs());
} else if width <= QUADRATIC_PHASE * (1.0 + v.abs()) {
return Ok(best.0);
}
if gv < 0.0 {
lo = v;
} else {
hi = v;
}
let newton = v - gv / dg;
let newton_ok = gv.is_finite()
&& dg.is_finite()
&& dg > 0.0
&& newton > lo
&& newton < hi
&& 2.0 * gv.abs() <= (width * dg).abs();
let next = if newton_ok {
newton
} else if gv < 0.0 {
if hi.is_finite() {
0.5 * (lo + hi)
} else {
step *= 2.0;
v + step
}
} else if lo.is_finite() {
0.5 * (lo + hi)
} else {
step *= 2.0;
v - step
};
let tol = TOL * (1.0 + v.abs());
if (next - v).abs() <= tol {
return Ok(next);
}
if hi - lo <= tol {
return Ok(0.5 * (lo + hi));
}
width = (next - v).abs();
v = next;
}
Err(SymplexError::computation_failed(
op,
format!("the quantile iteration did not converge in {MAX_ITER} steps"),
))
}
fn discrete_ppf(
op: &'static str,
cdf: impl Fn(f64) -> f64,
p: f64,
k0: f64,
kmin: f64,
kmax: f64,
) -> Result<f64, SymplexError> {
const MAX_STEPS: usize = 2000;
let nan = || SymplexError::computation_failed(op, "the distribution function is not a number");
let k0 = k0.round().clamp(kmin, kmax);
let c0 = cdf(k0);
if c0.is_nan() {
return Err(nan());
}
let mut lo;
let mut hi;
let mut step = 1.0;
if c0 >= p {
hi = k0;
lo = k0 - 1.0;
for _ in 0..MAX_STEPS {
if lo < kmin {
lo = kmin - 1.0;
break;
}
let c = cdf(lo);
if c.is_nan() {
return Err(nan());
}
if c < p {
break;
}
hi = lo;
lo -= step;
step *= 2.0;
}
} else {
lo = k0;
hi = k0 + 1.0;
for _ in 0..MAX_STEPS {
if hi >= kmax {
hi = kmax;
break;
}
let c = cdf(hi);
if c.is_nan() {
return Err(nan());
}
if c >= p {
break;
}
lo = hi;
hi += step;
step *= 2.0;
}
}
while hi - lo > 1.0 {
let mid = (0.5 * (lo + hi)).floor();
let c = cdf(mid);
if c.is_nan() {
return Err(nan());
}
if c >= p {
hi = mid;
} else {
lo = mid;
}
}
Ok(hi)
}
pub mod norm {
use super::{SQRT_2PI, check_level, erfc, erfcinv};
use crate::base::errors::SymplexError;
use std::f64::consts::SQRT_2;
pub fn cdf(x: f64) -> f64 {
0.5 * erfc(-x / SQRT_2)
}
pub fn sf(x: f64) -> f64 {
0.5 * erfc(x / SQRT_2)
}
pub fn pdf(x: f64) -> f64 {
(-0.5 * x * x).exp() / SQRT_2PI
}
pub fn ppf(p: f64) -> Result<f64, SymplexError> {
check_level("numdist::norm::ppf", "p", p)?;
Ok(-SQRT_2 * erfcinv(2.0 * p))
}
pub fn isf(q: f64) -> Result<f64, SymplexError> {
check_level("numdist::norm::isf", "q", q)?;
Ok(SQRT_2 * erfcinv(2.0 * q))
}
}
pub mod t {
use super::{
Eval, Side, Tails, bratio, check_level, check_positive, lbeta, log_beta_pref, nearest_tail,
nearest_tail_upper, norm, solve_increasing,
};
use crate::base::errors::SymplexError;
fn tails(x: f64, df: f64) -> Tails {
if x.is_nan() || df.is_nan() || df <= 0.0 {
return Tails::NAN;
}
if df.is_infinite() {
return Tails {
lower: norm::cdf(x),
upper: norm::sf(x),
};
}
let x2 = x * x;
if x2.is_infinite() {
return if x > 0.0 { Tails::ONE } else { Tails::ZERO };
}
if x == 0.0 {
return Tails {
lower: 0.5,
upper: 0.5,
};
}
let i = bratio(0.5 * df, 0.5, df / (df + x2), x2 / (df + x2));
let tail = 0.5 * i.lower;
let body = 0.5 + 0.5 * i.upper;
if x > 0.0 {
Tails {
lower: body,
upper: tail,
}
} else {
Tails {
lower: tail,
upper: body,
}
}
}
pub fn cdf(x: f64, df: f64) -> f64 {
tails(x, df).lower
}
pub fn sf(x: f64, df: f64) -> f64 {
tails(x, df).upper
}
pub fn pdf(x: f64, df: f64) -> f64 {
if x.is_nan() || df.is_nan() || df <= 0.0 {
return f64::NAN;
}
if df.is_infinite() {
return norm::pdf(x);
}
if x == 0.0 {
return (-lbeta(0.5 * df, 0.5)).exp() / df.sqrt();
}
let x2 = x * x;
if x2.is_infinite() {
return 0.0;
}
log_beta_pref(0.5 * df, 0.5, df / (df + x2), x2 / (df + x2)).exp() / x.abs()
}
fn upper_quantile(op: &'static str, q: f64, df: f64) -> Result<f64, SymplexError> {
let a = 0.5 * df;
let z = norm::isf(q)?;
let t0 = z
+ (z * z * z + z) / (4.0 * df)
+ (5.0 * z.powi(5) + 16.0 * z.powi(3) + 3.0 * z) / (96.0 * df * df);
let mut u0 = t0.ln();
if q < 0.05 || df < 3.0 {
let ln_tail = ((a - 1.0) * df.ln() - lbeta(a, 0.5) - q.ln()) / df;
u0 = u0.max(ln_tail);
}
let u0 = u0.min(700.0);
let lnq = q.ln();
let g = |u: f64| {
let t = u.exp();
let x2 = t * t;
let sf = tails(t, df).upper;
let ln_sf = sf.ln();
let lp = log_beta_pref(a, 0.5, df / (df + x2), x2 / (df + x2));
Eval {
g: lnq - ln_sf,
dg: (lp - ln_sf).exp(),
}
};
solve_increasing(op, g, u0, f64::NEG_INFINITY, f64::INFINITY).map(f64::exp)
}
fn quantile(op: &'static str, side: Side, df: f64) -> Result<f64, SymplexError> {
if df.is_infinite() && df > 0.0 {
return match side {
Side::Lower(p) => norm::ppf(p),
Side::Upper(q) => norm::isf(q),
};
}
check_positive(op, "df", df)?;
Ok(match side {
Side::Lower(0.5) => 0.0,
Side::Lower(pl) => -upper_quantile(op, pl, df)?,
Side::Upper(q) => upper_quantile(op, q, df)?,
})
}
pub fn ppf(p: f64, df: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::t::ppf";
check_level(OP, "p", p)?;
quantile(OP, nearest_tail(p), df)
}
pub fn isf(q: f64, df: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::t::isf";
check_level(OP, "q", q)?;
quantile(OP, nearest_tail_upper(q), df)
}
}
pub mod gamma {
use super::{
Eval, Side, Tails, check_level, check_positive, gammainc_tails, lgamma, log_gamma_pref,
nearest_tail, nearest_tail_upper, norm, solve_increasing,
};
use crate::base::errors::SymplexError;
pub(super) fn tails(x: f64, shape: f64, scale: f64) -> Tails {
if x.is_nan() || !(shape.is_finite() && shape > 0.0) || !(scale.is_finite() && scale > 0.0)
{
return Tails::NAN;
}
if x <= 0.0 {
return Tails::ZERO;
}
gammainc_tails(shape, x / scale)
}
pub fn cdf(x: f64, shape: f64, scale: f64) -> f64 {
tails(x, shape, scale).lower
}
pub fn sf(x: f64, shape: f64, scale: f64) -> f64 {
tails(x, shape, scale).upper
}
pub(super) fn standard_quantile(
op: &'static str,
side: Side,
shape: f64,
) -> Result<f64, SymplexError> {
let (z, target, lower) = match side {
Side::Lower(pl) => (norm::ppf(pl)?, pl, true),
Side::Upper(q) => (norm::isf(q)?, q, false),
};
let mut x0 = shape * (1.0 - 1.0 / (9.0 * shape) + z / (3.0 * shape.sqrt())).powi(3);
if !(x0.is_finite() && x0 > 0.0) {
x0 = if lower {
((target.ln() + lgamma(shape + 1.0)) / shape).exp()
} else {
(-target.ln()).max(shape)
};
}
let ln_target = target.ln();
let g = |u: f64| {
let x = u.exp();
let tl = gammainc_tails(shape, x);
let ln_tail = if lower { tl.lower.ln() } else { tl.upper.ln() };
let gv = if lower {
ln_tail - ln_target
} else {
ln_target - ln_tail
};
Eval {
g: gv,
dg: (log_gamma_pref(shape, x) - ln_tail).exp(),
}
};
solve_increasing(op, g, x0.ln(), f64::NEG_INFINITY, f64::INFINITY).map(f64::exp)
}
fn quantile(op: &'static str, side: Side, shape: f64, scale: f64) -> Result<f64, SymplexError> {
check_positive(op, "the shape", shape)?;
check_positive(op, "the scale", scale)?;
standard_quantile(op, side, shape).map(|x| scale * x)
}
pub fn ppf(p: f64, shape: f64, scale: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::gamma::ppf";
check_level(OP, "p", p)?;
quantile(OP, nearest_tail(p), shape, scale)
}
pub fn isf(q: f64, shape: f64, scale: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::gamma::isf";
check_level(OP, "q", q)?;
quantile(OP, nearest_tail_upper(q), shape, scale)
}
}
pub mod chi2 {
use super::{Side, check_level, check_positive, gamma, nearest_tail, nearest_tail_upper};
use crate::base::errors::SymplexError;
pub fn cdf(x: f64, df: f64) -> f64 {
gamma::tails(x, 0.5 * df, 2.0).lower
}
pub fn sf(x: f64, df: f64) -> f64 {
gamma::tails(x, 0.5 * df, 2.0).upper
}
fn quantile(op: &'static str, side: Side, df: f64) -> Result<f64, SymplexError> {
check_positive(op, "df", df)?;
gamma::standard_quantile(op, side, 0.5 * df).map(|x| 2.0 * x)
}
pub fn ppf(p: f64, df: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::chi2::ppf";
check_level(OP, "p", p)?;
quantile(OP, nearest_tail(p), df)
}
pub fn isf(q: f64, df: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::chi2::isf";
check_level(OP, "q", q)?;
quantile(OP, nearest_tail_upper(q), df)
}
}
pub mod beta {
use super::{
Eval, Side, Tails, bratio, check_level, check_positive, lbeta, log_beta_pref, nearest_tail,
nearest_tail_upper, norm, solve_increasing,
};
use crate::base::errors::SymplexError;
fn tails(x: f64, alpha: f64, beta: f64) -> Tails {
if x.is_nan() || !(alpha.is_finite() && alpha > 0.0) || !(beta.is_finite() && beta > 0.0) {
return Tails::NAN;
}
if x <= 0.0 {
return Tails::ZERO;
}
if x >= 1.0 {
return Tails::ONE;
}
bratio(alpha, beta, x, 1.0 - x)
}
pub fn cdf(x: f64, alpha: f64, beta: f64) -> f64 {
tails(x, alpha, beta).lower
}
pub fn sf(x: f64, alpha: f64, beta: f64) -> f64 {
tails(x, alpha, beta).upper
}
pub(super) fn start_logit(side: Side, alpha: f64, beta: f64) -> Result<f64, SymplexError> {
let s = alpha + beta;
let mean = alpha / s;
let sd = (alpha * beta / (s * s * (s + 1.0))).sqrt();
let z = match side {
Side::Lower(pl) => norm::ppf(pl)?,
Side::Upper(q) => norm::isf(q)?,
};
let x0 = mean + z * sd;
Ok(match side {
Side::Lower(pl) => {
let x = if x0 > 0.0 && x0 < 1.0 && pl >= 1e-3 {
x0
} else {
((pl.ln() + alpha.ln() + lbeta(alpha, beta)) / alpha)
.exp()
.min(0.5)
};
x.ln() - (-x).ln_1p()
}
Side::Upper(q) => {
let y = if x0 > 0.0 && x0 < 1.0 && q >= 1e-3 {
1.0 - x0
} else {
((q.ln() + beta.ln() + lbeta(alpha, beta)) / beta)
.exp()
.min(0.5)
};
(-y).ln_1p() - y.ln()
}
})
}
fn quantile_logit(
op: &'static str,
side: Side,
alpha: f64,
beta: f64,
) -> Result<f64, SymplexError> {
let v0 = start_logit(side, alpha, beta)?;
let (target, lower) = match side {
Side::Lower(pl) => (pl, true),
Side::Upper(q) => (q, false),
};
let ln_target = target.ln();
let g = |v: f64| {
let x = 1.0 / (1.0 + (-v).exp());
let y = 1.0 / (1.0 + v.exp());
let tl = bratio(alpha, beta, x, y);
let ln_tail = if lower { tl.lower.ln() } else { tl.upper.ln() };
let gv = if lower {
ln_tail - ln_target
} else {
ln_target - ln_tail
};
Eval {
g: gv,
dg: (log_beta_pref(alpha, beta, x, y) - ln_tail).exp(),
}
};
solve_increasing(op, g, v0, f64::NEG_INFINITY, f64::INFINITY)
}
fn quantile(op: &'static str, side: Side, alpha: f64, beta: f64) -> Result<f64, SymplexError> {
check_positive(op, "α", alpha)?;
check_positive(op, "β", beta)?;
quantile_logit(op, side, alpha, beta).map(|v| 1.0 / (1.0 + (-v).exp()))
}
pub fn ppf(p: f64, alpha: f64, beta: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::beta::ppf";
check_level(OP, "p", p)?;
quantile(OP, nearest_tail(p), alpha, beta)
}
pub fn isf(q: f64, alpha: f64, beta: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::beta::isf";
check_level(OP, "q", q)?;
quantile(OP, nearest_tail_upper(q), alpha, beta)
}
}
pub mod f {
use super::{
Eval, Side, Tails, beta, bratio, check_level, check_positive, log_beta_pref, nearest_tail,
nearest_tail_upper, solve_increasing,
};
use crate::base::errors::SymplexError;
fn tails(x: f64, d1: f64, d2: f64) -> Tails {
if x.is_nan() || !(d1.is_finite() && d1 > 0.0) || !(d2.is_finite() && d2 > 0.0) {
return Tails::NAN;
}
if x <= 0.0 {
return Tails::ZERO;
}
let num = d1 * x;
let den = num + d2;
if den.is_infinite() {
return Tails::ONE;
}
bratio(0.5 * d1, 0.5 * d2, num / den, d2 / den)
}
pub fn cdf(x: f64, d1: f64, d2: f64) -> f64 {
tails(x, d1, d2).lower
}
pub fn sf(x: f64, d1: f64, d2: f64) -> f64 {
tails(x, d1, d2).upper
}
fn quantile(op: &'static str, side: Side, d1: f64, d2: f64) -> Result<f64, SymplexError> {
check_positive(op, "d1", d1)?;
check_positive(op, "d2", d2)?;
let (a, b) = (0.5 * d1, 0.5 * d2);
let u0 = beta::start_logit(side, a, b)? + (d2 / d1).ln();
let (target, lower) = match side {
Side::Lower(pl) => (pl, true),
Side::Upper(q) => (q, false),
};
let ln_target = target.ln();
let g = |u: f64| {
let x = u.exp();
let num = d1 * x;
let den = num + d2;
let (xb, yb) = (num / den, d2 / den);
let tl = bratio(a, b, xb, yb);
let ln_tail = if lower { tl.lower.ln() } else { tl.upper.ln() };
let gv = if lower {
ln_tail - ln_target
} else {
ln_target - ln_tail
};
Eval {
g: gv,
dg: (log_beta_pref(a, b, xb, yb) - ln_tail).exp(),
}
};
solve_increasing(op, g, u0, f64::NEG_INFINITY, f64::INFINITY).map(f64::exp)
}
pub fn ppf(p: f64, d1: f64, d2: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::f::ppf";
check_level(OP, "p", p)?;
quantile(OP, nearest_tail(p), d1, d2)
}
pub fn isf(q: f64, d1: f64, d2: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::f::isf";
check_level(OP, "q", q)?;
quantile(OP, nearest_tail_upper(q), d1, d2)
}
}
pub mod binom {
use super::{Tails, bratio, check_level, discrete_ppf, invalid, norm};
use crate::base::errors::SymplexError;
fn tails(k: f64, n: f64, p: f64) -> Tails {
if k.is_nan() || !(n.is_finite() && n >= 0.0) || !(0.0..=1.0).contains(&p) {
return Tails::NAN;
}
let kf = k.floor();
if kf < 0.0 {
return Tails::ZERO;
}
if kf >= n {
return Tails::ONE;
}
bratio(n - kf, kf + 1.0, 1.0 - p, p)
}
pub fn cdf(k: f64, n: f64, p: f64) -> f64 {
tails(k, n, p).lower
}
pub fn sf(k: f64, n: f64, p: f64) -> f64 {
tails(k, n, p).upper
}
pub fn ppf(q: f64, n: f64, p: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::binom::ppf";
check_level(OP, "q", q)?;
if !(n.is_finite() && n >= 0.0) {
return Err(invalid(
OP,
format!("n must be a non-negative count, got {n}"),
));
}
if !(0.0..=1.0).contains(&p) {
return Err(invalid(OP, format!("p must lie in [0, 1], got {p}")));
}
let n = n.floor();
let k0 = n * p + norm::ppf(q)? * (n * p * (1.0 - p)).sqrt();
discrete_ppf(OP, |k| cdf(k, n, p), q, k0, 0.0, n)
}
}
pub mod poisson {
use super::{Tails, check_level, check_positive, discrete_ppf, gammainc_tails, norm};
use crate::base::errors::SymplexError;
fn tails(k: f64, rate: f64) -> Tails {
if k.is_nan() || !(rate.is_finite() && rate > 0.0) {
return Tails::NAN;
}
let kf = k.floor();
if kf < 0.0 {
return Tails::ZERO;
}
gammainc_tails(kf + 1.0, rate).flipped()
}
pub fn cdf(k: f64, rate: f64) -> f64 {
tails(k, rate).lower
}
pub fn sf(k: f64, rate: f64) -> f64 {
tails(k, rate).upper
}
pub fn ppf(q: f64, rate: f64) -> Result<f64, SymplexError> {
const OP: &str = "numdist::poisson::ppf";
check_level(OP, "q", q)?;
check_positive(OP, "the rate", rate)?;
let k0 = rate + norm::ppf(q)? * rate.sqrt();
discrete_ppf(OP, |k| cdf(k, rate), q, k0, 0.0, f64::INFINITY)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn rel(a: f64, b: f64) -> f64 {
((a - b) / b).abs()
}
#[test]
fn stirling_pieces() {
assert!((stirlerr(10.0) - 0.00833056343336287).abs() < 1e-16);
assert!((stirlerr(10.0) - stirlerr(9.999_999_999)).abs() < 1e-12);
let direct = lgamma(3.5) + lgamma(12.0) - lgamma(15.5);
assert!((lbeta(3.5, 12.0) - direct).abs() < 1e-13);
assert!((rlog1(0.7) - (0.7 - 1.7f64.ln())).abs() < 1e-16);
assert!((rlog1(1e-3) - (1e-3 - 1e-3f64.ln_1p())).abs() < 1e-16);
assert_eq!(bd0(5.0, 5.0), 0.0);
assert!((bd0(2.0, 5.0) - (2.0 * (2.0f64 / 5.0).ln() + 3.0)).abs() < 1e-15);
}
#[test]
fn prefactors_match_definitions() {
let (a, b, x): (f64, f64, f64) = (2.5, 4.0, 0.3);
let direct = a * x.ln() + b * (1.0f64 - x).ln() - (lgamma(a) + lgamma(b) - lgamma(a + b));
assert!((log_beta_pref(a, b, x, 1.0 - x) - direct).abs() < 1e-14);
let direct = a * x.ln() - x - lgamma(a);
assert!((log_gamma_pref(a, x) - direct).abs() < 1e-14);
}
#[test]
fn grat_r_is_continuous_at_the_branch_point() {
for a in [0.05, 0.5, 0.9] {
let lo = grat_r(a, 1.099_999_999, log_gamma_pref(a, 1.099_999_999));
let hi = grat_r(a, 1.1, log_gamma_pref(a, 1.1));
assert!(rel(lo, hi) < 1e-9, "a = {a}: {lo} vs {hi}");
let q = gammainc_tails(a, 1.1).upper / log_gamma_pref(a, 1.1).exp();
assert!(rel(hi, q) < 1e-13, "a = {a}: {hi} vs {q}");
}
}
#[test]
fn incomplete_beta_identities() {
assert!((betainc_regularized_f64(1.0, 1.0, 0.37) - 0.37).abs() < 1e-15);
let t = bratio(3.0, 7.0, 0.2, 0.8);
assert!((t.lower + t.upper - 1.0).abs() < 1e-15);
assert!((betainc_regularized_f64(7.0, 3.0, 0.8) - t.upper).abs() < 1e-15);
assert!((betainc_regularized_f64(2.0, 3.0, 0.4) - 0.5248).abs() < 1e-15);
assert!(betainc_regularized_f64(-1.0, 3.0, 0.4).is_nan());
assert!(betainc_regularized_f64(1.0, 3.0, 1.4).is_nan());
}
#[test]
fn incomplete_gamma_identities() {
assert!((gammainc_lower_regularized_f64(1.0, 0.7) - (1.0 - (-0.7f64).exp())).abs() < 1e-15);
assert!(
rel(
gammainc_upper_regularized_f64(0.5, 3.0),
erfc(3.0f64.sqrt())
) < 1e-14
);
assert_eq!(gammainc_lower_regularized_f64(2.0, 0.0), 0.0);
assert!(gammainc_lower_regularized_f64(0.0, 1.0).is_nan());
}
#[test]
fn quantiles_round_trip() -> Result<(), SymplexError> {
for (p, df) in [
(1e-10, 1.0),
(1e-3, 2.5),
(0.3, 30.0),
(0.999, 1e4),
(0.5, 7.0),
] {
let x = t::ppf(p, df)?;
assert!(
rel(t::cdf(x, df), p) < 1e-13 || p == 0.5,
"t: p = {p}, df = {df}"
);
}
for (p, df) in [(1e-10, 1.0), (1e-3, 2.5), (0.3, 30.0), (0.999, 1e4)] {
let x = chi2::ppf(p, df)?;
assert!(rel(chi2::cdf(x, df), p) < 1e-13, "chi2: p = {p}, df = {df}");
}
for (p, a, b) in [
(1e-10, 0.5, 0.5),
(1e-3, 2.0, 30.0),
(0.3, 5.0, 9.0),
(0.999, 0.05, 3.0),
] {
let x = beta::ppf(p, a, b)?;
assert!(
rel(beta::cdf(x, a, b), p) < 1e-13,
"beta: p = {p}, a = {a}, b = {b}"
);
let x = f::ppf(p, 2.0 * a, 2.0 * b)?;
assert!(rel(f::cdf(x, 2.0 * a, 2.0 * b), p) < 1e-13, "f: p = {p}");
}
assert!(t::ppf(0.0, 3.0).is_err());
assert!(t::ppf(0.5, -3.0).is_err());
assert!(chi2::ppf(1.0, 3.0).is_err());
assert!(beta::ppf(0.5, 0.0, 1.0).is_err());
Ok(())
}
#[test]
fn discrete_quantiles() -> Result<(), SymplexError> {
assert_eq!(binom::ppf(0.3, 9.0, 3.0 / 7.0)?, 3.0);
assert_eq!(binom::ppf(0.9, 9.0, 3.0 / 7.0)?, 6.0);
assert_eq!(poisson::ppf(0.1, 7.0 / 3.0)?, 1.0);
assert_eq!(poisson::ppf(0.3, 7.0 / 3.0)?, 1.0);
assert_eq!(poisson::ppf(0.95, 7.0 / 3.0)?, 5.0);
let k = poisson::ppf(1e-10, 1e6)?;
assert!(poisson::cdf(k, 1e6) >= 1e-10 && poisson::cdf(k - 1.0, 1e6) < 1e-10);
assert_eq!(binom::ppf(0.5, 0.0, 0.3)?, 0.0);
Ok(())
}
}