use num_traits::{Float, FloatConst};
use strafe_type::{FloatConstraint, LogProbability64, Probability64, Rational64, Real64};
use crate::{
distribution::{
norm::{log_qnorm, qnorm},
t::{dt, pt},
},
traits::{TrigPI, DPQ},
};
pub fn qt<P: Into<Probability64>, R: Into<Rational64>>(p: P, ndf: R, lower_tail: bool) -> Real64 {
let p = p.into().unwrap();
qt_inner(p, ndf, lower_tail, false)
}
pub fn log_qt<LP: Into<LogProbability64>, R: Into<Rational64>>(
p: LP,
ndf: R,
lower_tail: bool,
) -> Real64 {
let p = p.into().unwrap();
qt_inner(p, ndf, lower_tail, true)
}
fn qt_inner<R: Into<Rational64>>(mut p: f64, ndf: R, lower_tail: bool, log: bool) -> Real64 {
let ndf = ndf.into().unwrap();
let eps = 1.0e-12;
let mut P;
let mut q;
if let Some(ret) = p.q_p01_boundaries(f64::neg_infinity(), f64::infinity(), lower_tail, log) {
return ret.into();
}
if ndf < 1.0 {
let accu = 1e-13;
let Eps = 1e-11;
let mut ux;
let mut lx;
let mut nx;
let mut pp;
let mut iter = 0;
p = p.dt_qiv(lower_tail, log);
if p > 1.0 - f64::EPSILON {
return f64::infinity().into();
}
pp = (1.0 - f64::EPSILON).min(p * (1.0 + Eps));
ux = 1.0;
while ux < f64::max_value() && pt(ux, ndf, true).unwrap() < pp {
ux *= 2.0
}
pp = p * (1.0 - Eps);
lx = -1.0;
while lx > -f64::max_value() && pt(lx, ndf, true).unwrap() > pp {
lx *= 2.0
}
loop {
nx = 0.5 * (lx + ux);
if pt(nx, ndf, true).unwrap() > p {
ux = nx;
} else {
lx = nx;
}
iter += 1;
if (ux - lx) / nx.abs() < accu || iter > 1000 {
break;
}
}
return (0.5 * (lx + ux)).into();
}
if ndf > 1e20 {
return if log {
log_qnorm(p, 0.0, 1.0, lower_tail)
} else {
qnorm(p, 0.0, 1.0, lower_tail)
};
}
P = p.d_qiv(log);
let neg = (!lower_tail || P < 0.5) && (lower_tail || P > 0.5);
let is_neg_lower = lower_tail == neg;
if neg {
P = 2.0
* if log {
if lower_tail {
P
} else {
-p.exp_m1()
}
} else {
p.d_lval(lower_tail)
}
} else {
P = 2.0
* if log {
if lower_tail {
-p.exp_m1()
} else {
P
}
} else {
p.d_cval(lower_tail)
};
}
if (ndf - 2.0).abs() < eps {
if P > f64::min_value() {
if 3.0 * P < f64::epsilon() {
q = 1.0 / P.sqrt();
} else if P > 0.9 {
q = (1.0 - P) * (2.0 / (P * (2.0 - P))).sqrt();
} else {
q = (2.0 / (P * (2.0 - P)) - 2.0).sqrt();
}
} else {
if log {
q = if is_neg_lower {
(-p / 2.0).exp() / f64::SQRT_2()
} else {
1.0 / (-p.exp_m1()).sqrt()
};
} else {
q = f64::infinity();
}
}
} else if ndf < 1.0 + eps {
if P == 1.0 {
q = 0.0; } else if P > 0.0 {
q = 1.0 / (P / 2.0).tan_pi();
} else {
if log {
q = if is_neg_lower {
f64::FRAC_1_PI() * (-p).exp()
} else {
-1.0 / (f64::PI() * p.exp_m1())
};
} else {
q = f64::infinity()
}
}
} else {
let mut x = 0.0;
let mut y = f64::nan();
let mut log_P2 = 0.0;
let a = 1.0 / (ndf - 0.5);
let b = 48.0 / (a * a);
let mut c = ((20700.0 * a / b - 98.0) * a - 16.0) * a + 96.36;
let d = ((94.5 / (b + c) - 3.0) / b + 1.0) * (a * f64::FRAC_PI_2()).sqrt() * ndf;
let P_ok1 = P > f64::min_value() || !log;
let mut P_ok = P_ok1; if P_ok1 {
y = (d * P).powf(2.0 / ndf);
P_ok = y >= f64::epsilon();
}
if !P_ok {
log_P2 = if is_neg_lower {
p.d_log(log)
} else {
p.d_lexp(log)
};
x = (d.ln() + f64::LN_2() + log_P2) / ndf;
y = (2.0 * x).exp();
}
if (ndf < 2.1 && P > 0.5) || y > 0.05 + a {
if P_ok {
x = qnorm(0.5 * P, 0.0, 1.0, true).unwrap();
} else {
x = log_qnorm(log_P2, 0.0, 1.0, lower_tail).unwrap();
}
y = x * x;
if ndf < 5.0 {
c += 0.3 * (ndf - 4.5) * (x + 0.6);
}
c = (((0.05 * d * x - 5.0) * x - 7.0) * x - 2.0) * x + b + c;
y = (((((0.4 * y + 6.3) * y + 36.0) * y + 94.5) / c - y - 3.0) / b + 1.0) * x;
y = (a * y * y).exp_m1();
q = (ndf * y).sqrt();
} else if !P_ok && x < -f64::LN_2() * f64::MANTISSA_DIGITS as f64 {
q = ndf.sqrt() * (-x).exp();
} else {
y = ((1.0 / (((ndf + 6.0) / (ndf * y) - 0.089 * d - 0.822) * (ndf + 2.0) * 3.0)
+ 0.5 / (ndf + 4.0))
* y
- 1.0)
* (ndf + 1.0)
/ (ndf + 2.0)
+ 1.0 / y;
q = (ndf * y).sqrt();
}
if P_ok1 {
let M = ((f64::max_value() / 2.0).sqrt() - ndf).abs();
let mut it = 0;
loop {
it += 1;
y = dt(q, ndf, false).unwrap();
x = (pt(q, ndf, false).unwrap() - P / 2.0) / y;
if !(it < 10 && y > 0.0 && x.is_finite() && x.abs() > 1e-14 * q.abs()) {
break;
}
let F = if q.abs() < M {
q * (ndf + 1.0) / (2.0 * (q * q + ndf))
} else {
(ndf + 1.0) / (2.0 * (q + ndf / q))
};
let del_q = x * (1.0 + x * F);
if del_q.is_finite() && (q + del_q).is_finite() {
q += del_q;
} else if x.is_finite() && (q + x).is_finite() {
q += x;
} else {
break; }
}
}
}
if neg { -q } else { q }.into()
}