use super::ln_gamma;
const TINY: f64 = 1e-300;
const REL_EPS: f64 = 1e-16;
const MAX_ITERS: usize = 1000;
#[must_use]
pub fn betai(a: f64, b: f64, x: f64) -> f64 {
if x <= 0.0 {
return 0.0;
}
if x >= 1.0 {
return 1.0;
}
let log_prefactor = ln_gamma(a + b) - ln_gamma(a) - ln_gamma(b);
let bt = b
.mul_add((1.0 - x).ln(), a.mul_add(x.ln(), log_prefactor))
.exp();
if x < (a + 1.0) / (a + b + 2.0) {
bt * betacf(a, b, x) / a
} else {
1.0 - bt * betacf(b, a, 1.0 - x) / b
}
}
#[must_use]
pub fn ln_betai_lower(a: f64, b: f64, x: f64) -> f64 {
if x <= 0.0 {
return f64::NEG_INFINITY;
}
if x >= 1.0 {
return 0.0;
}
if x < (a + 1.0) / (a + b + 2.0) {
ln_bt(a, b, x) + (betacf(a, b, x) / a).ln()
} else {
let ln_upper = ln_bt(b, a, 1.0 - x) + (betacf(b, a, 1.0 - x) / b).ln();
(-ln_upper.exp()).ln_1p()
}
}
#[must_use]
pub fn ln_betai_upper(a: f64, b: f64, x: f64) -> f64 {
ln_betai_lower(b, a, 1.0 - x)
}
fn ln_bt(a: f64, b: f64, x: f64) -> f64 {
let log_prefactor = ln_gamma(a + b) - ln_gamma(a) - ln_gamma(b);
b.mul_add((1.0 - x).ln(), a.mul_add(x.ln(), log_prefactor))
}
#[allow(clippy::many_single_char_names)]
fn betacf(a: f64, b: f64, x: f64) -> f64 {
let qab = a + b;
let qap = a + 1.0;
let qam = a - 1.0;
let mut c = 1.0;
let mut d = lentz_guard(1.0 - qab * x / qap);
d = 1.0 / d;
let mut h = d;
let mut m = 0.0_f64;
for _ in 0..MAX_ITERS {
m += 1.0;
let m2 = 2.0 * m;
let aa = m * (b - m) * x / ((qam + m2) * (a + m2));
d = lentz_guard(aa.mul_add(d, 1.0));
c = lentz_guard(1.0 + aa / c);
d = 1.0 / d;
h *= d * c;
let aa = -(a + m) * (qab + m) * x / ((a + m2) * (qap + m2));
d = lentz_guard(aa.mul_add(d, 1.0));
c = lentz_guard(1.0 + aa / c);
d = 1.0 / d;
let del = d * c;
h *= del;
if (del - 1.0).abs() < REL_EPS {
break;
}
}
h
}
fn lentz_guard(v: f64) -> f64 {
if v.abs() < TINY { TINY } else { v }
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ln_betai_lower_matches_log_of_linear() {
for &(a, b, x) in &[(2.0, 3.0, 0.3), (0.5, 2.5, 0.9)] {
let want = betai(a, b, x).ln();
let got = ln_betai_lower(a, b, x);
assert!(
((got - want) / want.abs().max(1.0)).abs() < 1e-10,
"ln_betai_lower({a},{b},{x}) = {got}, want {want}"
);
}
}
#[test]
fn ln_betai_upper_matches_log_of_complement() {
for &(a, b, x) in &[(2.0, 3.0, 0.3), (1.5, 0.5, 0.001)] {
let want = (1.0 - betai(a, b, x)).ln();
let got = ln_betai_upper(a, b, x);
assert!(
((got - want) / want.abs().max(1.0)).abs() < 1e-10,
"ln_betai_upper({a},{b},{x}) = {got}, want {want}"
);
}
}
#[test]
fn ln_betai_lower_finite_in_deep_tail() {
let x = 5.0 / 50.0_f64.mul_add(50.0, 5.0); let got = ln_betai_lower(2.5, 0.5, x);
let want = -16.620_993_180_844_884;
assert!(got.is_finite(), "ln_betai_lower deep tail was {got}");
assert!(
((got - want) / want).abs() < 1e-9,
"ln_betai_lower deep tail rel error: got {got}, want {want}"
);
}
}