use crate::errors::QlResult;
use crate::fail;
use crate::math::gammafunction::log_gamma;
use crate::types::Real;
const ACCURACY: Real = 1e-16;
const MAX_ITERATIONS: u32 = 100;
pub fn beta_function(z: Real, w: Real) -> QlResult<Real> {
Ok((log_gamma(z)? + log_gamma(w)? - log_gamma(z + w)?).exp())
}
pub fn incomplete_beta(a: Real, b: Real, x: Real) -> QlResult<Real> {
if !a.is_finite() || a <= 0.0 {
fail!("incomplete_beta requires a finite a > 0, got a={a}");
}
if !b.is_finite() || b <= 0.0 {
fail!("incomplete_beta requires a finite b > 0, got b={b}");
}
if !(0.0..=1.0).contains(&x) {
fail!("incomplete_beta requires x in [0, 1], got x={x}");
}
if x == 0.0 {
return Ok(0.0);
}
if x == 1.0 {
return Ok(1.0);
}
let prefactor =
(log_gamma(a + b)? - log_gamma(a)? - log_gamma(b)? + a * x.ln() + b * (1.0 - x).ln()).exp();
if x < (a + 1.0) / (a + b + 2.0) {
Ok(prefactor * beta_continued_fraction(a, b, x)? / a)
} else {
Ok(1.0 - prefactor * beta_continued_fraction(b, a, 1.0 - x)? / b)
}
}
fn beta_continued_fraction(a: Real, b: Real, x: Real) -> QlResult<Real> {
let eps = Real::EPSILON;
let qab = a + b;
let qap = a + 1.0;
let qam = a - 1.0;
let mut c = 1.0;
let mut d = 1.0 - qab * x / qap;
if d.abs() < eps {
d = eps;
}
d = 1.0 / d;
let mut result = d;
for iter in 1..=MAX_ITERATIONS {
let m = iter as Real;
let m2 = 2.0 * m;
let aa = m * (b - m) * x / ((qam + m2) * (a + m2));
d = 1.0 + aa * d;
if d.abs() < eps {
d = eps;
}
c = 1.0 + aa / c;
if c.abs() < eps {
c = eps;
}
d = 1.0 / d;
result *= d * c;
let aa = -(a + m) * (qab + m) * x / ((a + m2) * (qap + m2));
d = 1.0 + aa * d;
if d.abs() < eps {
d = eps;
}
c = 1.0 + aa / c;
if c.abs() < eps {
c = eps;
}
d = 1.0 / d;
let del = d * c;
result *= del;
if (del - 1.0).abs() < ACCURACY {
return Ok(result);
}
}
fail!("incomplete_beta continued fraction did not converge (a={a}, b={b})");
}
#[cfg(test)]
mod tests {
use super::*;
const TOL: Real = 1e-12;
fn assert_close(got: Real, expected: Real) {
let tol = TOL * (1.0 + expected.abs());
assert!(
(got - expected).abs() <= tol,
"got {got}, expected {expected}, diff {}",
(got - expected).abs()
);
}
#[test]
fn beta_function_known_values() {
assert_close(beta_function(2.0, 3.0).unwrap(), 1.0 / 12.0);
assert_close(beta_function(0.5, 0.5).unwrap(), std::f64::consts::PI);
}
#[test]
fn boundaries_and_uniform() {
assert_eq!(incomplete_beta(2.0, 3.0, 0.0).unwrap(), 0.0);
assert_eq!(incomplete_beta(2.0, 3.0, 1.0).unwrap(), 1.0);
assert_close(incomplete_beta(1.0, 1.0, 0.3).unwrap(), 0.3);
assert_close(incomplete_beta(1.0, 1.0, 0.5).unwrap(), 0.5);
}
#[test]
fn known_values() {
assert_close(incomplete_beta(2.0, 2.0, 0.5).unwrap(), 0.5);
assert_close(incomplete_beta(2.0, 3.0, 0.5).unwrap(), 0.6875);
assert_close(incomplete_beta(2.0, 3.0, 0.3).unwrap(), 0.3483);
}
#[test]
fn symmetry_identity() {
for &(a, b, x) in &[(2.0, 3.0, 0.3), (0.5, 2.5, 0.7), (4.0, 1.5, 0.2)] {
let lhs = incomplete_beta(a, b, x).unwrap();
let rhs = 1.0 - incomplete_beta(b, a, 1.0 - x).unwrap();
assert!(
(lhs - rhs).abs() < TOL,
"symmetry failed at a={a}, b={b}, x={x}"
);
}
}
#[test]
fn invalid_args_rejected() {
assert!(incomplete_beta(0.0, 1.0, 0.5).is_err()); assert!(incomplete_beta(1.0, -1.0, 0.5).is_err()); assert!(incomplete_beta(1.0, 1.0, 1.5).is_err()); assert!(incomplete_beta(1.0, 1.0, Real::NAN).is_err()); assert!(incomplete_beta(Real::NAN, 1.0, 0.5).is_err()); assert!(incomplete_beta(1.0, Real::NAN, 0.5).is_err()); assert!(incomplete_beta(Real::INFINITY, 1.0, 0.5).is_err()); assert!(incomplete_beta(1.0, Real::INFINITY, 0.5).is_err()); assert!(beta_function(Real::INFINITY, 1.0).is_err());
}
}