use nonstdfloat::f128;
use num_traits::{ToPrimitive, Zero};
use strafe_type::{FloatConstraint, LogProbability64, Probability64, Rational64, Real64};
use crate::{
distribution::{
beta::pbeta,
norm::{log_pnorm, pnorm},
t::{log_pt, pt},
},
func::lgamma,
traits::DPQ,
};
pub fn pnt<RE1: Into<Real64>, RA: Into<Rational64>, RE2: Into<Real64>>(
t: RE1,
df: RA,
ncp: RE2,
lower_tail: bool,
) -> Probability64 {
pnt_inner(t, df, ncp, lower_tail, false).into()
}
pub fn log_pnt<RE1: Into<Real64>, RA: Into<Rational64>, RE2: Into<Real64>>(
t: RE1,
df: RA,
ncp: RE2,
lower_tail: bool,
) -> LogProbability64 {
pnt_inner(t, df, ncp, lower_tail, true).into()
}
fn pnt_inner<RE1: Into<Real64>, RA: Into<Rational64>, RE2: Into<Real64>>(
t: RE1,
df: RA,
ncp: RE2,
mut lower_tail: bool,
log: bool,
) -> f64 {
let t = t.into().unwrap();
let df = df.into().unwrap();
let ncp = ncp.into().unwrap();
let current_block: u64;
let mut albeta = 0.0;
let mut a = 0.0;
let mut b = 0.0;
let mut del = 0.0;
let mut errbd = 0.0;
let mut lambda = 0.0;
let mut rxb = 0.0;
let mut tt = 0.0;
let mut x = 0.0;
let mut geven = f128::zero();
let mut godd = f128::zero();
let mut p = f128::zero();
let mut q = f128::zero();
let mut s = f128::zero();
let mut tnc = f128::zero();
let mut xeven = f128::zero();
let mut xodd = f128::zero();
let mut it = 0;
let mut negdel = false;
let itrmax = 1000;
static errmax: f64 = 1.0e-12;
if ncp == 0.0 {
return if log {
log_pt(t, df, lower_tail).unwrap()
} else {
pt(t, df, lower_tail).unwrap()
};
}
if !t.is_finite() {
return if t < 0.0 {
f64::dt_0(lower_tail, log)
} else {
f64::dt_1(lower_tail, log)
};
}
if t >= 0.0 {
negdel = false;
tt = t;
del = ncp
} else {
if ncp > 40.0 && !lower_tail {
return f64::dt_0(lower_tail, log);
}
negdel = true;
tt = -t;
del = -ncp
}
if df > 4e5 || del * del > 2.0 * std::f64::consts::LN_2 * -f64::MIN_EXP as f64 {
s = f128::new(1.0 / (4.0 * df));
let x = (f128::new(tt) * (f128::new(1.0) - s)).to_f64().unwrap();
let mu = del;
let sigma = (f128::new(1.0) + f128::new(tt * tt * 2.0) * s)
.to_f64()
.unwrap()
.sqrt();
let lower_tail = lower_tail != negdel;
return if log {
log_pnorm(x, mu, sigma, lower_tail).unwrap()
} else {
pnorm(x, mu, sigma, lower_tail).unwrap()
};
}
x = t * t;
rxb = df / (x + df);
x = x / (x + df);
if x > 0.0 {
lambda = del * del;
p = f128::new(0.5 * (-0.5 * lambda).exp());
if p == f128::new(0.0) {
warn!("underflow occurred in pnt");
warn!("value out of range in pnt");
return f64::dt_0(lower_tail, log);
}
q = f128::new(strafe_consts::SQRT_2DPI) * p * f128::new(del);
s = f128::new(0.5) - p;
if s < f128::new(1e-7) {
s = f128::new(-0.5 * (-0.5 * lambda).exp_m1())
}
a = 0.5;
b = 0.5 * df;
rxb = rxb.powf(b);
albeta = strafe_consts::LN_SQRT_PI + lgamma(b).unwrap() - lgamma(0.5 + b).unwrap();
xodd = f128::new(pbeta(x, a, b, true).unwrap());
godd = f128::new(2.0 * rxb * (a * x.ln() - albeta).exp());
tnc = f128::new(b * x);
xeven = if tnc < f128::new(2.2204460492503131e-16) {
tnc
} else {
f128::new(1.0 - rxb)
};
geven = tnc * f128::new(rxb);
tnc = p * xodd + q * xeven;
it = 1;
loop {
if !(it <= itrmax) {
current_block = 8304106758420804164;
break;
}
a += 1.0;
xodd -= godd;
xeven -= geven;
godd *= f128::new(x * (a + b - 1.0) / a);
geven *= f128::new(x * (a + b - 0.5) / (a + 0.5));
p *= f128::new(lambda / (2 * it) as f64);
q *= f128::new(lambda / (2 * it + 1) as f64);
tnc += p * xodd + q * xeven;
s -= p;
if s < f128::new(-1.0e-10) {
warn!("full precision may not have been achieved in pnt");
current_block = 7175283761034176179;
break;
} else {
if s <= f128::new(0.0) && it > 1 {
current_block = 7175283761034176179;
break;
}
errbd = (f128::new(2.0) * s * (xodd - godd)).to_f64().unwrap();
if errbd.abs() < errmax {
current_block = 7175283761034176179;
break;
}
it += 1
}
}
match current_block {
7175283761034176179 => {}
_ => {
warn!("convergence failed in pnt");
}
}
} else {
tnc = f128::new(0.0)
}
tnc += f128::new(pnorm(-del, 0.0, 1.0, true));
lower_tail = lower_tail != negdel;
if tnc > f128::new(1.0 - 1e-10) && lower_tail {
warn!("full precision may not have been achieved in pnt{{final}}");
}
tnc.to_f64().unwrap().min(1.0).dt_val(lower_tail, log)
}