use num_traits::{Float, FloatConst};
use strafe_type::{FloatConstraint, LogProbability64, Positive64, Probability64, Real64};
use crate::{distribution::beta::pbeta_raw, func::lbeta, traits::DPQ};
const USE_LOG_X_CUTOFF: f64 = -5.0;
const N_NEWTON_FREE: i32 = 4;
pub fn qbeta<PR: Into<Probability64>, PO1: Into<Positive64>, PO2: Into<Positive64>>(
alpha: PR,
p: PO1,
q: PO2,
lower_tail: bool,
) -> Real64 {
let alpha = alpha.into().unwrap();
qbeta_inner(alpha, p, q, lower_tail, false)
}
pub fn log_qbeta<LP: Into<LogProbability64>, P1: Into<Positive64>, P2: Into<Positive64>>(
alpha: LP,
p: P1,
q: P2,
lower_tail: bool,
) -> Real64 {
let alpha = alpha.into().unwrap();
qbeta_inner(alpha, p, q, lower_tail, true)
}
pub fn qbeta_inner<P1: Into<Positive64>, P2: Into<Positive64>>(
alpha: f64,
p: P1,
q: P2,
lower_tail: bool,
log: bool,
) -> Real64 {
let p = p.into().unwrap();
let q = q.into().unwrap();
let mut qbet: [f64; 2] = [0.0; 2]; qbeta_raw(
alpha,
p,
q,
lower_tail,
log,
None,
USE_LOG_X_CUTOFF,
N_NEWTON_FREE,
&mut qbet,
);
qbet[0].into()
}
static DBL_very_MIN: f64 = 2.2250738585072014e-308 / 4.0;
static DBL_log_v_MIN: f64 = std::f64::consts::LN_2 * (f64::MIN_EXP as f64 - 2.0);
static DBL_1__eps: f64 = 0.9999999999999999;
const FPU: f64 = 3e-308;
const ACU_MIN: f64 = 1e-300;
const P_LO: f64 = FPU;
const P_HI: f64 = 1.0 - 2.22e-16;
const CONST1: f64 = 2.30753;
const CONST2: f64 = 0.27061;
const CONST3: f64 = 0.99229;
const CONST4: f64 = 0.04481;
#[derive(Copy, Clone, Default, Debug)]
pub struct QBetaState {
swap_choose: bool,
swap_tail: bool,
log_: bool,
give_log_q: bool,
use_log_x: bool,
warned: bool,
add_N_step: bool,
a: f64,
la: f64,
logbeta: f64,
g: f64,
h: f64,
pp: f64,
p_: f64,
qq: f64,
r: f64,
s: f64,
t: f64,
w: f64,
y: f64,
u: f64,
xinbta: f64,
n_maybe_swaps: usize,
u_n: f64,
acu: f64,
tx: f64,
exit: bool,
}
fn qbeta_raw(
alpha: f64,
p: f64,
q: f64,
lower_tail: bool,
log: bool,
swap_01: Option<bool>,
log_q_cut: f64,
n_N: i32,
qb: &mut [f64],
) {
let mut state = QBetaState::default();
state.swap_choose = swap_01.is_none();
state.give_log_q = log_q_cut == f64::infinity();
state.use_log_x = state.give_log_q; state.warned = false;
state.add_N_step = true;
state.y = -1.0;
let return_q_0 = |qb: &mut [f64]| {
if state.give_log_q {
qb[0] = f64::neg_infinity();
qb[1] = 0.0;
} else {
qb[0] = 0.0;
qb[1] = 1.0;
}
};
let return_q_1 = |qb: &mut [f64]| {
if state.give_log_q {
qb[0] = 0.0;
qb[1] = f64::neg_infinity();
} else {
qb[0] = 1.0;
qb[1] = 0.0;
}
};
if alpha == f64::dt_0(lower_tail, log) {
return_q_0(qb);
return;
}
if alpha == f64::dt_1(lower_tail, log) {
return_q_1(qb);
return;
}
if (log && alpha > 0.0) || (!log && (alpha < 0.0 || alpha > 1.0)) {
qb[0] = f64::nan();
qb[1] = f64::nan();
return;
}
if p == 0.0 || q == 0.0 || !p.is_finite() || !q.is_finite() {
let return_q_half = |qb: &mut [f64]| {
if state.give_log_q {
qb[0] = -f64::LN_2();
qb[1] = -f64::LN_2();
} else {
qb[0] = 0.5;
qb[1] = 0.5;
}
};
if p == 0.0 && q == 0.0 {
if alpha < f64::d_half(log) {
return_q_0(qb);
} else if alpha > f64::d_half(log) {
return_q_1(qb);
} else {
return_q_half(qb);
}
} else if p == 0.0 || p / q == 0.0 {
return_q_0(qb);
} else if q == 0.0 || q / p == 0.0 {
return_q_1(qb);
} else {
return_q_half(qb);
}
return;
}
state.p_ = alpha.dt_qiv(lower_tail, log);
state.logbeta = lbeta(p, q).unwrap();
state.swap_tail = if state.swap_choose {
state.p_ > 0.5
} else {
swap_01.unwrap()
};
maybe_swap(
&mut state, alpha, p, q, lower_tail, log, swap_01, log_q_cut, n_N, qb,
);
if state.exit {
return;
}
l_newton(
&mut state, alpha, p, q, lower_tail, log, swap_01, log_q_cut, n_N, qb,
);
if state.exit {
return;
}
l_converged(&mut state, log);
if state.exit {
return;
}
l_return(&mut state, log, qb);
}
fn maybe_swap(
state: &mut QBetaState,
alpha: f64,
p: f64,
q: f64,
lower_tail: bool,
log: bool,
swap_01: Option<bool>,
log_q_cut: f64,
n_N: i32,
qb: &mut [f64],
) {
if state.swap_tail {
state.a = alpha.dt_civ(lower_tail, log);
state.la = alpha.dt_clog(lower_tail, log);
state.pp = q;
state.qq = p;
} else {
state.a = state.p_;
state.la = alpha.dt_log(lower_tail, log);
state.pp = p;
state.qq = q;
}
state.n_maybe_swaps += 1;
state.acu =
ACU_MIN.max(10.0.powf(-13.0 - 2.5 / (state.pp * state.pp) - 0.5 / (state.a * state.a)));
let u0 = (state.la + state.pp.ln() + state.logbeta) / state.pp; let mut rp = state.pp * (1.0 - state.qq) / (state.pp + 1.0);
let log_eps_c = f64::LN_2() * (1.0 - f64::MANTISSA_DIGITS as f64);
state.t = 0.2;
let u0_maybe = (f64::LN_2() * f64::MIN_EXP as f64) < u0 && u0 < -0.01;
state.u_n = 1.0; if u0_maybe &&
u0 < (state.t*log_eps_c - ((state.pp*(1.0-state.qq)*(2.0-state.qq)/(2.*(state.pp+2.0))).abs()).ln())/2.0
{
rp = rp * u0.exp(); if rp > -1.0 {
state.u = u0 - rp.ln_1p() / state.pp;
} else {
state.u = u0;
}
state.tx = state.u.exp();
state.xinbta = state.u.exp();
state.use_log_x = true; l_newton(
state, alpha, p, q, lower_tail, log, swap_01, log_q_cut, n_N, qb,
);
if state.exit {
return;
}
}
state.r = (-2.0 * state.la).sqrt();
state.y = state.r - (CONST1 + CONST2 * state.r) / (1. + (CONST3 + CONST4 * state.r) * state.r);
if state.pp > 1.0 && state.qq > 1.0 {
state.r = (state.y * state.y - 3.) / 6.;
state.s = 1.0 / (state.pp + state.pp - 1.0);
state.t = 1.0 / (state.qq + state.qq - 1.0);
state.h = 2.0 / (state.s + state.t);
state.w = state.y * (state.h + state.r).sqrt() / state.h
- (state.t - state.s) * (state.r + 5.0 / 6.0 - 2.0 / (3.0 * state.h));
if state.w > 300.0 {
state.t = state.w + state.w + state.qq.ln() - state.pp.ln(); state.u = if state.t <= 18.0 {-state.t.exp().ln_1p()} else {-state.t - (-state.t).exp()};
state.xinbta = state.u.exp();
} else {
state.xinbta = state.pp / (state.pp + state.qq * (state.w + state.w).exp());
state.u = - (state.qq/state.pp * (state.w+state.w).exp()).ln_1p();
}
} else {
state.r = state.qq + state.qq;
state.t = 1.0 / (3.0 * state.qq.sqrt()); state.t = state.r * (1.0 + state.t * (-state.t + state.y)).pow_di(3); state.s = 4.0 * state.pp + state.r - 2.0; if state.t == 0.0 || (state.t < 0.0 && state.s >= state.t) {
let l1ma =
if state.swap_tail {alpha.dt_log(lower_tail, log) } else { alpha.dt_clog(lower_tail, log)};
let xx = (l1ma + state.qq.ln() + state.logbeta) / state.qq;
if xx <= 0.0 {
state.xinbta = -xx.exp_m1();
state.u = xx.log1_exp(); } else {
let r_ = rp * u0.exp();
if r_ > -1.0 {
state.u = u0 - r_.ln_1p() / state.pp; } else {
state.u = u0; }
state.xinbta = state.u.exp();
}
} else {
state.t = state.s / state.t;
if state.t <= 1.0 {
state.u = u0;
state.xinbta = state.u.exp();
} else {
state.xinbta = 1.0 - 2.0 / (state.t + 1.0);
state.u = (-2.0 / (state.t + 1.0)).ln_1p();
}
}
}
if state.swap_choose && (( state.swap_tail && state.u >= -log_q_cut.exp()) || (!state.swap_tail && state.u >= -(4.0*log_q_cut).exp() && state.pp / state.qq < 1000.0) )
{
state.swap_tail = !state.swap_tail;
if state.swap_tail {
state.a = alpha.dt_civ(lower_tail, log); state.la = alpha.dt_clog(lower_tail, log);
state.pp = q;
state.qq = p;
} else {
state.a = state.p_;
state.la = alpha.dt_log(lower_tail, log);
state.pp = p;
state.qq = q;
}
state.u = state.u.log1_exp();
state.xinbta = state.u.exp();
}
if !state.use_log_x {
state.use_log_x = state.u < log_q_cut; }
let bad_u = !state.u.is_finite();
let bad_init = bad_u || state.xinbta > P_HI;
state.tx = state.xinbta;
if bad_u || state.u < log_q_cut {
state.w = pbeta_raw(DBL_very_MIN, state.pp, state.qq, true, log);
if state.w > if log { state.la } else { state.a } {
if log || (state.w - state.a).abs() < (0.0 - state.a).abs() {
state.tx = DBL_very_MIN;
state.u_n = DBL_log_v_MIN; } else {
state.tx = 0.0;
state.u_n = f64::neg_infinity();
}
state.use_log_x = log;
state.add_N_step = false;
l_return(state, log, qb);
state.exit = true;
return;
} else {
if state.u < DBL_log_v_MIN {
state.u = DBL_log_v_MIN; state.xinbta = DBL_very_MIN;
}
}
}
if bad_init && !(state.use_log_x && state.tx > 0.0) {
if state.u == f64::neg_infinity() {
state.u = f64::LN_2() * f64::MIN_EXP as f64;
state.xinbta = f64::min_value();
} else {
state.xinbta = if state.xinbta > 1.1 {
0.5 } else if state.xinbta < P_LO {
state.u.exp()
} else {
P_HI
};
if bad_u {
state.u = state.xinbta.ln();
}
}
}
}
fn l_newton(
state: &mut QBetaState,
alpha: f64,
p: f64,
q: f64,
lower_tail: bool,
log: bool,
swap_01: Option<bool>,
log_q_cut: f64,
n_N: i32,
qb: &mut [f64],
) {
state.r = 1.0 - state.pp;
state.t = 1.0 - state.qq;
let mut wprev = 0.0;
let mut prev = 1.0;
let mut adj = 1.0;
if state.use_log_x {
let mut i_pb = 0;
while i_pb < 1000 {
state.y = pbeta_raw(
state.xinbta,
state.pp,
state.qq,
true,
true,
);
state.w = if state.y == f64::neg_infinity() {
0.0
} else {
(state.y - state.la)
* (state.y - state.u
+ state.logbeta
+ state.r * state.u
+ state.t * state.u.log1_exp())
.exp()
};
if !state.w.is_finite() {
if state.n_maybe_swaps <= 1 {
maybe_swap(
state, alpha, p, q, lower_tail, log, swap_01, log_q_cut, n_N, qb,
);
if state.exit {
return;
}
}
qb[0] = f64::nan();
qb[1] = f64::nan();
state.exit = true;
return;
}
if i_pb >= n_N && state.w * wprev <= 0.0 {
prev = adj.abs().max(FPU);
}
state.g = 1.0;
let mut i_inn = 0;
while i_inn < 1000 {
adj = state.g * state.w;
if adj.abs() < prev {
state.u_n = state.u - adj; if state.u_n <= 0.0 {
if prev <= state.acu || state.w.abs() <= state.acu {
l_converged(state, log);
if state.exit {
return;
}
}
break;
}
}
state.g /= 3.0;
i_inn += 1;
}
let D = adj.abs().min((state.u_n - state.u).abs());
if D <= 4e-16 * (state.u_n + state.u).abs() {
l_converged(state, log);
if state.exit {
return;
}
}
state.u = state.u_n;
state.xinbta = state.u.exp();
wprev = state.w;
i_pb += 1;
} } else {
let mut i_pb = 0;
while i_pb < 1000 {
state.y = pbeta_raw(
state.xinbta,
state.pp,
state.qq,
true,
log,
);
state.w = if log {
(state.y - state.la)
* (state.y
+ state.logbeta
+ state.r * state.xinbta.ln()
+ state.t * (-state.xinbta).ln_1p())
.exp()
} else {
(state.y - state.a)
* (state.logbeta
+ state.r * state.xinbta.ln()
+ state.t * (-state.xinbta).ln_1p())
.exp()
};
if !state.w.is_finite() {
if state.n_maybe_swaps <= 2 {
if !log && state.n_maybe_swaps == 2 {
state.use_log_x = true;
} if !log || state.n_maybe_swaps <= 1 {
maybe_swap(
state, alpha, p, q, lower_tail, log, swap_01, log_q_cut, n_N, qb,
);
if state.exit {
return;
}
}
}
qb[0] = f64::nan();
qb[1] = f64::nan();
state.exit = true;
return;
}
if i_pb >= n_N && state.w * wprev <= 0.0 {
prev = adj.abs().max(FPU);
}
state.g = 1.0;
let mut i_inn = 0;
while i_inn < 1000 {
adj = state.g * state.w;
if i_pb < n_N || adj.abs() < prev {
state.tx = state.xinbta - adj; if 0.0 <= state.tx && state.tx <= 1.0 {
if prev <= state.acu || state.w.abs() <= state.acu {
l_converged(state, log);
if state.exit {
return;
}
}
if state.tx != 0.0 && state.tx != 1.0 {
break;
}
}
}
state.g /= 3.0;
i_inn += 1;
}
if (state.tx - state.xinbta).abs() <= 4e-16 * (state.tx + state.xinbta) {
l_converged(state, log);
if state.exit {
return;
}
}
state.xinbta = state.tx;
if state.tx == 0.0 {
break;
}
wprev = state.w;
i_pb += 1;
} }
state.warned = true;
}
fn l_converged(state: &mut QBetaState, log: bool) {
state.log_ = log || state.use_log_x;
if (state.log_ && state.y == f64::neg_infinity()) || (!state.log_ && state.y == 0.0) {
state.w = pbeta_raw(DBL_very_MIN, state.pp, state.qq, true, state.log_);
if state.log_ || (state.w - state.a).abs() <= (state.y - state.a).abs() {
state.tx = DBL_very_MIN;
state.u_n = DBL_log_v_MIN; }
state.add_N_step = false; } else if !state.warned
&& if state.log_ {
(state.y - state.la).abs() > 3.0
} else {
(state.y - state.a).abs() > 1e-4
}
&& !(state.log_ && state.y == f64::neg_infinity() &&
pbeta_raw(DBL_1__eps, state.pp, state.qq, true, true) > state.la + 2.0)
{
}
}
fn l_return(state: &mut QBetaState, log: bool, qb: &mut [f64]) {
if state.give_log_q {
if !state.use_log_x { }
let r = state.u_n.log1_exp();
if state.swap_tail {
qb[0] = r;
qb[1] = state.u_n;
} else {
qb[0] = state.u_n;
qb[1] = r;
}
} else {
if state.use_log_x {
if state.add_N_step {
if state.u_n != 1.0 {
state.xinbta = state.u_n.exp();
}
state.y = pbeta_raw(
state.xinbta,
state.pp,
state.qq,
true,
log,
);
state.w = if log {
(state.y - state.la)
* (state.y
+ state.logbeta
+ state.r * state.xinbta.ln()
+ state.t * (-state.xinbta).ln_1p())
.exp()
} else {
(state.y - state.a)
* (state.logbeta
+ state.r * state.xinbta.ln()
+ state.t * (-state.xinbta).ln_1p())
.exp()
};
if state.w.is_finite() {
state.tx = state.xinbta - state.w;
} else {
state.tx = state.xinbta;
}
} else {
if state.swap_tail {
qb[0] = -state.u_n.exp_m1();
qb[1] = state.u_n.exp();
} else {
qb[0] = state.u_n.exp();
qb[1] = -state.u_n.exp_m1();
}
return;
}
}
if state.swap_tail {
qb[0] = 1.0 - state.tx;
qb[1] = state.tx;
} else {
qb[0] = state.tx;
qb[1] = 1.0 - state.tx;
}
}
return;
}