use num_traits::Float;
use strafe_type::{
FloatConstraint, LogProbability64, Positive64, Probability64, Rational64, Real64,
};
use crate::{
distribution::{
gamma::{dgamma, lgamma1p, log_pgamma, pgamma_raw},
norm::{log_qnorm, qnorm},
},
func::lgamma,
traits::DPQ,
};
pub fn qchisq_appr(p: f64, nu: f64, g: f64, lower_tail: bool, log: bool, tol: f64) -> f64
{
const C7: f64 = 4.67;
const C8: f64 = 6.66;
const C9: f64 = 6.73;
const C10: f64 = 13.32;
let mut alpha = 0.0;
let mut a = 0.0;
let mut c = 0.0;
let mut ch = 0.0;
let mut p1 = 0.0;
let mut p2 = 0.0;
let mut q = 0.0;
let mut t = 0.0;
let mut x = 0.0;
if p.is_nan() || nu.is_nan() {
return p + nu;
}
if nu <= 0.0 {
return f64::nan();
}
alpha = 0.5 * nu;
c = alpha - 1.0;
p1 = p.dt_log(lower_tail, log);
if nu < -1.24 * p1 {
let lgam1pa = if alpha < 0.5 {
lgamma1p(alpha)
} else {
(alpha.ln()) + g
};
ch = ((lgam1pa + p1) / alpha + std::f64::consts::LN_2).exp()
} else if nu > 0.32 {
x = if log {
log_qnorm(p, 0.0, 1.0, lower_tail).unwrap()
} else {
qnorm(p, 0.0, 1.0, lower_tail).unwrap()
};
p1 = 2.0 / (9.0 * nu);
ch = nu * (x * p1.sqrt() + 1.0 - p1).powf(3.0);
if ch > 2.2 * nu + 6.0 {
ch = -2.0 * (p.dt_clog(lower_tail, log) - c * (0.5 * ch).ln() + g)
}
} else {
ch = 0.4;
a = p.dt_clog(lower_tail, log) + g + c * std::f64::consts::LN_2;
loop {
q = ch;
p1 = 1.0 / (1.0 + ch * (C7 + ch));
p2 = ch * (C9 + ch * (C8 + ch));
t = -0.5 + (C7 + 2.0 * ch) * p1 - (C9 + ch * (C10 + 3.0 * ch)) / p2;
ch -= (1.0 - (a + 0.5 * ch).exp() * p2 * p1) / t;
if !((q - ch).abs() > tol * ch.abs()) {
break;
}
}
}
return ch;
}
pub fn qgamma<PR: Into<Probability64>, PO: Into<Positive64>, R: Into<Rational64>>(
p: PR,
alpha: PO,
scale: R,
lower_tail: bool,
) -> Real64 {
let p = p.into().unwrap();
qgamma_inner(p, alpha, scale, lower_tail, false)
}
pub fn log_qgamma<LP: Into<LogProbability64>, P: Into<Positive64>, R: Into<Rational64>>(
p: LP,
alpha: P,
scale: R,
lower_tail: bool,
) -> Real64 {
let p = p.into().unwrap();
qgamma_inner(p, alpha, scale, lower_tail, true)
}
pub fn qgamma_inner<P: Into<Positive64>, R: Into<Rational64>>(
mut p: f64,
alpha: P,
scale: R,
lower_tail: bool,
log: bool,
) -> Real64 {
let alpha = alpha.into().unwrap();
let scale = scale.into().unwrap();
const EPS1: f64 = 1e-2;
const EPS2: f64 = 5e-7;
const EPS_N: f64 = 1e-15;
const MAXIT: i32 = 1000;
const P_MIN: f64 = 1e-100;
const P_MAX: f64 = 1.0 - 1e-14;
static i420: f64 = 1.0 / 420.0;
static i2520: f64 = 1.0 / 2520.0;
static i5040: f64 = 1.0 / 5040.0;
let mut p_ = 0.0;
let mut a = 0.0;
let mut b = 0.0;
let mut c = 0.0;
let mut g = 0.0;
let mut ch = 0.0;
let mut ch0 = 0.0;
let mut p1 = 0.0;
let mut p2 = 0.0;
let mut q = 0.0;
let mut s1 = 0.0;
let mut s2 = 0.0;
let mut s3 = 0.0;
let mut s4 = 0.0;
let mut s5 = 0.0;
let mut s6 = 0.0;
let mut t = 0.0;
let mut x = 0.0;
let mut i = 0;
let mut max_it_Newton = 1;
if let Some(ret) = p.q_p01_boundaries(0.0, f64::infinity(), lower_tail, log) {
return ret.into();
}
if alpha == 0.0 {
return 0.0.into();
}
if alpha < 1e-10 {
max_it_Newton = 7
}
p_ = p.dt_qiv(lower_tail, log);
g = lgamma(alpha).unwrap();
ch = qchisq_appr(p, 2.0 * alpha, g, lower_tail, log, EPS1);
if !ch.is_finite() {
max_it_Newton = 0
} else if ch < 5e-7 || (log && (p_ > P_MAX || p_ < P_MIN)) {
max_it_Newton = 20
} else {
c = alpha - 1.0;
s6 = (120.0 + c * (346.0 + 127.0 * c)) * i5040;
ch0 = ch;
i = 1;
while i <= MAXIT {
q = ch;
p1 = 0.5 * ch;
if log && p1.is_nan() {
p2 = f64::nan();
} else {
p2 = p_ - pgamma_raw(p1, alpha, true, false);
}
if !p2.is_finite() || ch <= 0.0 {
ch = ch0;
max_it_Newton = 27;
break;
} else {
t = p2 * (alpha * std::f64::consts::LN_2 + g + p1 - c * ch.ln()).exp();
b = t / ch;
a = 0.5 * t - b * c;
s1 =
(210.0 + a * (140.0 + a * (105.0 + a * (84.0 + a * (70.0 + 60.0 * a))))) * i420;
s2 = (420.0 + a * (735.0 + a * (966.0 + a * (1141.0 + 1278.0 * a)))) * i2520;
s3 = (210.0 + a * (462.0 + a * (707.0 + 932.0 * a))) * i2520;
s4 = (252.0 + a * (672.0 + 1182.0 * a) + c * (294.0 + a * (889.0 + 1740.0 * a)))
* i5040;
s5 = (84.0 + 2264.0 * a + c * (1175.0 + 606.0 * a)) * i2520;
ch += t
* (1.0 + 0.5 * t * s1
- b * c * (s1 - b * (s2 - b * (s3 - b * (s4 - b * (s5 - b * s6))))));
if (q - ch).abs() < EPS2 * ch {
break;
}
if (q - ch).abs() > 0.1 * ch {
if ch < q {
ch = 0.9 * q
} else {
ch = 1.1 * q
}
}
i += 1
}
}
}
x = 0.5 * scale * ch;
if max_it_Newton != 0 {
if !log {
p = p.ln();
}
if x == 0.0 {
let _1_p = 1.0 + 1e-7;
let _1_m = 1.0 - 1e-7;
x = 2.2250738585072014e-308;
p_ = log_pgamma(x, alpha, scale, lower_tail).unwrap();
if lower_tail && p_ > p * _1_p || !lower_tail && p_ < p * _1_m {
return 0.0.into();
}
} else {
p_ = log_pgamma(x, alpha, scale, lower_tail).unwrap()
}
if p_ == f64::NEG_INFINITY {
return 0.0.into();
}
i = 1;
while i <= max_it_Newton {
p1 = p_ - p;
if p1.abs() < (EPS_N * p).abs() {
break;
}
g = dgamma(x, alpha, scale, true).unwrap();
if g == (f64::d_0(true)) {
break;
}
t = (p1) * (p_ - g).exp();
t = if lower_tail { (x) - t } else { (x) + t };
p_ = log_pgamma(t, alpha, scale, lower_tail).unwrap();
if (p_ - p).abs() > p1.abs() || i > 1 && (p_ - p).abs() == p1.abs() {
break;
} else {
x = t;
i += 1
}
}
}
x.into()
}