use std::f64::consts::LN_10;
use nonstdfloat::f128;
use num_traits::{Float, ToPrimitive, Zero};
use strafe_type::{
FloatConstraint, LogProbability64, Positive64, Probability64, Rational64, Real64,
};
use crate::{
distribution::{
chisq::{log_pchisq, pchisq},
func::logspace_add,
},
func::lgamma,
traits::DPQ,
};
static _dbl_min_exp: f64 = std::f64::consts::LN_2 * f64::MIN_EXP as f64;
pub fn pnchisq<RE: Into<Real64>, RA: Into<Rational64>, P: Into<Positive64>>(
x: RE,
df: RA,
ncp: P,
lower_tail: bool,
) -> Probability64 {
pnchisq_inner(x, df, ncp, lower_tail, false).into()
}
pub fn log_pnchisq<RE: Into<Real64>, RA: Into<Rational64>, P: Into<Positive64>>(
x: RE,
df: RA,
ncp: P,
lower_tail: bool,
) -> LogProbability64 {
pnchisq_inner(x, df, ncp, lower_tail, true).into()
}
fn pnchisq_inner<RE: Into<Real64>, RA: Into<Rational64>, P: Into<Positive64>>(
x: RE,
df: RA,
ncp: P,
lower_tail: bool,
log: bool,
) -> f64 {
let mut x = x.into().unwrap();
let df = df.into().unwrap();
let ncp = ncp.into().unwrap();
if x == 0.0 {
x = f64::min_positive_value();
}
let mut ans = 0.0;
if !df.is_finite() || !ncp.is_finite() {
return f64::nan();
}
ans = pnchisq_raw(
x,
df,
ncp,
1e-12,
8.0 * 2.2204460492503131e-16,
1000000,
lower_tail,
log,
);
if x <= 0. || x == f64::infinity() {
return ans; }
if ncp >= 80.0 {
if lower_tail {
ans = ans.min(f64::d_1(log))
} else {
if ans < if log { -10. * LN_10 } else { 1e-10 } {
warn!("full precision may not have been achieved in pnchisq");
}
if !log && ans < 0.0 {
ans = 0.0;
}
}
}
if !log || ans < -1e-8 {
ans
} else {
ans = pnchisq_raw(
x,
df,
ncp,
1e-12,
8.0 * 2.2204460492503131e-16,
1000000,
!lower_tail,
false,
);
(-ans).ln_1p()
}
}
pub fn pnchisq_raw(
x: f64,
f: f64,
theta: f64,
errmax: f64,
reltol: f64,
itrmax: i32,
lower_tail: bool,
log: bool,
) -> f64 {
let mut lam = 0.0;
let mut x2 = 0.0;
let mut f2 = 0.0;
let mut term = 0.0;
let mut bound = 0.0;
let mut f_x_2n = 0.0;
let mut f_2n = 0.0;
let mut l_lam = -1.0;
let mut l_x = -1.0;
let mut n = 0;
let mut lamSml = false;
let mut tSml = false;
let mut is_r = false;
let mut is_b = false;
let mut is_it = false;
let mut ans = f128::zero();
let mut u = f128::zero();
let mut v = f128::zero();
let mut t = f128::zero();
let mut lt = f128::zero();
let mut lu = f128::new(-(1.0));
if x <= 0.0 {
if x == 0.0 && f == 0.0 {
let _L = (-0.5) * theta;
return if lower_tail {
_L.d_exp(log)
} else if log {
_L.log1_exp()
} else {
-_L.exp_m1()
};
}
return f64::dt_0(lower_tail, log);
}
if !x.is_finite() {
return f64::dt_1(lower_tail, log);
}
if theta < 80.0 {
let mut ans_0 = f128::zero();
let mut i = 0;
return if lower_tail
&& f > 0.0
&& x.ln()
< std::f64::consts::LN_2 + 2.0 / f * (lgamma(f / 2.0 + 1.0).unwrap() + _dbl_min_exp)
{
let lambda = 0.5 * theta; let mut sum = 0.0;
let mut sum2 = 0.0;
let mut pr = -lambda;
let log_lam = lambda.ln();
sum2 = f64::NEG_INFINITY;
sum = sum2;
i = 0; while i < 110 {
sum2 = logspace_add(sum2, pr);
sum = logspace_add(
sum,
pr + log_pchisq(x, f + (2 * i) as f64, lower_tail).unwrap(),
);
if sum2 >= -1e-15 {
break;
}
i += 1;
pr += log_lam - (i as f64).ln()
}
ans_0 = f128::new(sum - sum2);
if log { ans_0 } else { ans_0.exp() }.to_f64().unwrap()
} else {
let lambda_0 = f128::new(0.5 * theta); let mut sum_0 = f128::new(0.0);
let mut sum2_0 = f128::new(0.0);
let mut pr_0 = (-lambda_0).exp();
i = 0;
while i < 110 {
sum2_0 += pr_0;
sum_0 += pr_0 * f128::new(pchisq(x, f + (2 * i) as f64, lower_tail));
if sum2_0 >= f128::new(1.0 - 1e-15) {
break;
}
i += 1;
pr_0 *= lambda_0 / f128::new(i as f64)
}
ans_0 = sum_0 / sum2_0;
if log { ans_0.ln() } else { ans_0 }.to_f64().unwrap()
};
}
lam = 0.5 * theta; lamSml = -lam < _dbl_min_exp;
if lamSml {
u = f128::new(0.0);
lu = f128::new(-lam);
l_lam = lam.ln()
} else {
u = f128::new((-lam).exp())
}
v = u;
x2 = 0.5 * x;
f2 = 0.5 * f;
f_x_2n = f - x;
if f2 * 2.2204460492503131e-16 > 0.125 && {
t = f128::new(x2 - f2);
(t.abs()) < f128::new((2.2204460492503131e-16).sqrt() * f2)
} {
lt = (f128::new(1.0) - t) * (f128::new(2.0) - t / f128::new(f2 + 1.0))
- f128::new(strafe_consts::LN_SQRT_2TPI)
- f128::new(0.5 * (f2 + 1.0).ln())
} else {
lt = f128::new(f2 * x2.ln() - x2 - lgamma(f2 + 1.0).unwrap())
}
tSml = lt < f128::new(_dbl_min_exp);
if tSml {
if x > f + theta + 5.0 * (2.0 * (f + 2.0 * theta)).sqrt() {
return f64::dt_1(lower_tail, log);
}
l_x = x.ln();
term = 0.0;
ans = f128::new(term);
t = f128::new(0.0)
} else {
t = lt.exp();
term = (v * t).to_f64().unwrap();
ans = f128::new(term)
}
n = 1;
f_2n = f + 2.0;
f_x_2n += 2.0;
loop {
if f_x_2n > 0.0 {
bound = (t * f128::new(x) / f128::new(f_x_2n)).to_f64().unwrap();
is_it = false;
is_r = is_it;
is_b = bound <= errmax;
if is_b && {
is_r = f128::new(term) <= f128::new(reltol) * ans;
is_r
} || {
is_it = n > itrmax;
is_it
} {
break;
}
}
if lamSml {
lu += f128::new(l_lam - (n as f64).ln());
if lu >= f128::new(_dbl_min_exp) {
u = lu.exp();
v = u;
lamSml = false
}
} else {
u *= f128::new(lam / n as f64);
v += u
}
if tSml {
lt += f128::new(l_x - f_2n.ln());
if lt >= f128::new(_dbl_min_exp) {
t = lt.exp();
tSml = false
}
} else {
t *= f128::new(x / f_2n)
}
if lamSml as u64 == 0 && tSml as u64 == 0 {
term = (v * t).to_f64().unwrap();
ans += f128::new(term)
}
n += 1;
f_2n += 2.0;
f_x_2n += 2.0
}
if is_it {
warn!("pnchisq(x={}, ..): not converged in {} iter.", x, itrmax);
}
let dans = ans.to_f64().unwrap();
return dans.dt_val(lower_tail, log);
}