use num_traits::{One, Signed, Zero};
use crate::api::context::Context;
use crate::api::expr::Ex;
use crate::base::errors::SymplexError;
use crate::base::interval::{Bounds, Interval};
use crate::domains::optimize::{RootOpts, bisect, grow_bracket};
use super::common::{check_alpha, check_unit_open, ex_usize, invalid, q_to_f64};
use super::data::Q;
fn check_rates(op: &'static str, alpha: f64, beta: f64) -> Result<(), SymplexError> {
check_alpha(op, alpha)?;
check_unit_open(op, "beta", beta)?;
if alpha + beta >= 1.0 {
return Err(invalid(
op,
format!(
"alpha + beta must be below 1 for the boundaries to be ordered, got {alpha} + {beta}"
),
));
}
Ok(())
}
fn check_probability(op: &'static str, p: &Q, what: &str) -> Result<(), SymplexError> {
if !p.is_positive() || p >= &Q::one() {
return Err(invalid(
op,
format!("{what} must lie strictly in (0, 1), got {p}"),
));
}
Ok(())
}
pub fn wald_boundaries(alpha: f64, beta: f64) -> Result<Interval<f64>, SymplexError> {
check_rates("wald_boundaries", alpha, beta)?;
Ok(Interval::open(
(beta / (1.0 - alpha)).ln(),
((1.0 - beta) / alpha).ln(),
))
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum Decision {
Continue,
AcceptH0,
AcceptH1,
}
#[derive(Clone, Debug, PartialEq)]
enum Model {
Bernoulli { p0: Q, p1: Q },
NormalMean { mu0: Q, mu1: Q, sigma: Q },
}
#[derive(Clone, Debug, PartialEq)]
pub struct Sprt {
model: Model,
alpha: f64,
beta: f64,
boundaries: Interval<f64>,
observations: usize,
successes: usize,
sum: Q,
stopped: Option<Stop>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct Stop {
decision: Decision,
at: usize,
}
impl Sprt {
pub fn bernoulli(p0: &Q, p1: &Q, alpha: f64, beta: f64) -> Result<Self, SymplexError> {
const OP: &str = "Sprt::bernoulli";
check_probability(OP, p0, "p0")?;
check_probability(OP, p1, "p1")?;
if p0 == p1 {
return Err(invalid(OP, "p0 and p1 must differ"));
}
check_rates(OP, alpha, beta)?;
let boundaries = wald_boundaries(alpha, beta)?;
Ok(Sprt {
model: Model::Bernoulli {
p0: p0.clone(),
p1: p1.clone(),
},
alpha,
beta,
boundaries,
observations: 0,
successes: 0,
sum: Q::zero(),
stopped: None,
})
}
pub fn normal_mean(
mu0: &Q,
mu1: &Q,
sigma: &Q,
alpha: f64,
beta: f64,
) -> Result<Self, SymplexError> {
const OP: &str = "Sprt::normal_mean";
if !sigma.is_positive() {
return Err(invalid(OP, format!("sigma must be positive, got {sigma}")));
}
if mu0 == mu1 {
return Err(invalid(OP, "mu0 and mu1 must differ"));
}
check_rates(OP, alpha, beta)?;
let boundaries = wald_boundaries(alpha, beta)?;
Ok(Sprt {
model: Model::NormalMean {
mu0: mu0.clone(),
mu1: mu1.clone(),
sigma: sigma.clone(),
},
alpha,
beta,
boundaries,
observations: 0,
successes: 0,
sum: Q::zero(),
stopped: None,
})
}
pub fn update(&mut self, success: bool) -> Decision {
self.observations += 1;
if success {
self.successes += 1;
self.sum += Q::one();
}
self.settle()
}
pub fn observe(&mut self, x: &Q) -> Result<Decision, SymplexError> {
match &self.model {
Model::Bernoulli { .. } => {
if x.is_one() {
Ok(self.update(true))
} else if x.is_zero() {
Ok(self.update(false))
} else {
Err(invalid(
"Sprt::observe",
format!("a Bernoulli observation must be 0 or 1, got {x}"),
))
}
}
Model::NormalMean { .. } => {
self.observations += 1;
if x.is_positive() {
self.successes += 1;
}
self.sum += x;
Ok(self.settle())
}
}
}
fn settle(&mut self) -> Decision {
if let Some(stop) = self.stopped {
return stop.decision;
}
let decision = self.current();
if decision != Decision::Continue {
self.stopped = Some(Stop {
decision,
at: self.observations,
});
}
decision
}
fn current(&self) -> Decision {
let llr = self.log_likelihood_ratio_f64();
if llr >= self.boundaries.upper {
Decision::AcceptH1
} else if llr <= self.boundaries.lower {
Decision::AcceptH0
} else {
Decision::Continue
}
}
pub fn decision(&self) -> Decision {
match self.stopped {
Some(stop) => stop.decision,
None => self.current(),
}
}
pub fn is_decided(&self) -> bool {
self.stopped.is_some()
}
pub fn stopped_at(&self) -> Option<usize> {
self.stopped.map(|s| s.at)
}
pub fn reset(&mut self) {
self.observations = 0;
self.successes = 0;
self.sum = Q::zero();
self.stopped = None;
}
pub fn boundaries(&self) -> Interval<f64> {
self.boundaries
}
pub fn alpha(&self) -> f64 {
self.alpha
}
pub fn beta(&self) -> f64 {
self.beta
}
pub fn observations(&self) -> usize {
self.observations
}
pub fn successes(&self) -> usize {
self.successes
}
pub fn failures(&self) -> usize {
self.observations - self.successes
}
pub fn sum(&self) -> &Q {
&self.sum
}
pub fn log_likelihood_ratio(&self, ctx: &Context) -> Ex {
match &self.model {
Model::Bernoulli { p0, p1 } => {
let s = ex_usize(ctx, self.successes);
let f = ex_usize(ctx, self.failures());
let one = Q::one();
s * ctx.from_ratio(p1 / p0).ln()
+ f * ctx.from_ratio((&one - p1) / (&one - p0)).ln()
}
Model::NormalMean { .. } => ctx.from_ratio(self.normal_llr_exact()),
}
}
pub fn log_likelihood_ratio_f64(&self) -> f64 {
match &self.model {
Model::Bernoulli { p0, p1 } => {
let one = Q::one();
let ls = q_to_f64(&(p1 / p0)).ln();
let lf = q_to_f64(&((&one - p1) / (&one - p0))).ln();
self.successes as f64 * ls + self.failures() as f64 * lf
}
Model::NormalMean { .. } => q_to_f64(&self.normal_llr_exact()),
}
}
fn normal_llr_exact(&self) -> Q {
match &self.model {
Model::NormalMean { mu0, mu1, sigma } => {
let var = sigma * sigma;
let n = Q::from_integer(self.observations.into());
let two = Q::from_integer(2.into());
(mu1 - mu0) / &var * &self.sum - n * (mu1 * mu1 - mu0 * mu0) / (&two * &var)
}
Model::Bernoulli { .. } => Q::zero(),
}
}
}
fn bernoulli_increments(
op: &'static str,
p: f64,
p0: &Q,
p1: &Q,
) -> Result<(f64, f64, f64), SymplexError> {
check_probability(op, p0, "p0")?;
check_probability(op, p1, "p1")?;
if p0 == p1 {
return Err(invalid(op, "p0 and p1 must differ"));
}
if !(0.0..=1.0).contains(&p) {
return Err(invalid(op, format!("p must lie in [0, 1], got {p}")));
}
let one = Q::one();
let ls = q_to_f64(&(p1 / p0)).ln();
let lf = q_to_f64(&((&one - p1) / (&one - p0))).ln();
Ok((ls, lf, p * ls + (1.0 - p) * lf))
}
fn wald_h(
op: &'static str,
p: f64,
ls: f64,
lf: f64,
drift: f64,
) -> Result<Option<f64>, SymplexError> {
if drift.abs() < 1e-13 {
return Ok(None);
}
let phi = |h: f64| {
if h == 0.0 {
drift
} else {
(p * (h * ls).exp_m1() + (1.0 - p) * (h * lf).exp_m1()) / h
}
};
let side = if drift < 0.0 { 1.0 } else { -1.0 };
let (a, b, bounds) = if side > 0.0 {
(0.0, side, Bounds::at_least(0.0))
} else {
(side, 0.0, Bounds::at_most(0.0))
};
let bracket = grow_bracket(phi, a, b, bounds, 60).map_err(|_| {
SymplexError::computation_failed(
op,
"operating characteristic: could not bracket the root of Wald's identity",
)
})?;
let tight = RootOpts {
xtol: 0.0,
max_iter: 200,
..RootOpts::default()
};
bisect(phi, bracket.lower, bracket.upper, &tight)
.map(Some)
.map_err(|e| SymplexError::computation_failed(op, e.to_string()))
}
fn wald_oc(
op: &'static str,
p: f64,
ls: f64,
lf: f64,
drift: f64,
a: f64,
b: f64,
) -> Result<f64, SymplexError> {
if p <= 0.0 || p >= 1.0 {
return Ok(if drift < 0.0 { 1.0 } else { 0.0 });
}
match wald_h(op, p, ls, lf, drift)? {
Some(h) => Ok(((b * h).exp() - 1.0) / ((b * h).exp() - (a * h).exp())),
None => Ok(b / (b - a)),
}
}
pub fn operating_characteristic_bernoulli(
p: f64,
p0: &Q,
p1: &Q,
alpha: f64,
beta: f64,
) -> Result<f64, SymplexError> {
const OP: &str = "operating_characteristic_bernoulli";
let (ls, lf, drift) = bernoulli_increments(OP, p, p0, p1)?;
check_rates(OP, alpha, beta)?;
let (a, b) = wald_boundaries(alpha, beta)?.into_pair();
wald_oc(OP, p, ls, lf, drift, a, b)
}
pub fn expected_sample_size_bernoulli(
p: f64,
p0: &Q,
p1: &Q,
alpha: f64,
beta: f64,
) -> Result<f64, SymplexError> {
const OP: &str = "expected_sample_size_bernoulli";
let (ls, lf, drift) = bernoulli_increments(OP, p, p0, p1)?;
check_rates(OP, alpha, beta)?;
let (a, b) = wald_boundaries(alpha, beta)?.into_pair();
if drift.abs() < 1e-13 {
let second_moment = p * ls * ls + (1.0 - p) * lf * lf;
return Ok(-a * b / second_moment);
}
let l = wald_oc(OP, p, ls, lf, drift, a, b)?;
Ok((l * a + (1.0 - l) * b) / drift)
}