use num_traits::{One, Signed, ToPrimitive, Zero};
use crate::api::context::Context;
use crate::api::expr::Ex;
use crate::base::errors::SymplexError;
use super::data::Q;
const OP: &str = "stats::sequential";
fn invalid(reason: impl Into<String>) -> SymplexError {
SymplexError::invalid_argument(OP, reason)
}
fn failed(reason: impl Into<String>) -> SymplexError {
SymplexError::computation_failed(OP, reason)
}
fn q_f64(q: &Q) -> f64 {
q.to_f64().unwrap_or(f64::NAN)
}
fn check_rates(alpha: f64, beta: f64) -> Result<(), SymplexError> {
if !(alpha > 0.0 && alpha < 1.0) {
return Err(invalid(format!("alpha must lie in (0, 1), got {alpha}")));
}
if !(beta > 0.0 && beta < 1.0) {
return Err(invalid(format!("beta must lie in (0, 1), got {beta}")));
}
if alpha + beta >= 1.0 {
return Err(invalid(format!(
"alpha + beta must be below 1 for the boundaries to be ordered, got {alpha} + {beta}"
)));
}
Ok(())
}
fn check_probability(p: &Q, what: &str) -> Result<(), SymplexError> {
if !p.is_positive() || p >= &Q::one() {
return Err(invalid(format!(
"{what} must lie strictly in (0, 1), got {p}"
)));
}
Ok(())
}
pub fn wald_boundaries(alpha: f64, beta: f64) -> Result<(f64, f64), SymplexError> {
check_rates(alpha, beta)?;
Ok(((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,
lower: f64,
upper: f64,
observations: usize,
successes: usize,
sum: Q,
}
impl Sprt {
pub fn bernoulli(p0: Q, p1: Q, alpha: f64, beta: f64) -> Result<Self, SymplexError> {
check_probability(&p0, "p0")?;
check_probability(&p1, "p1")?;
if p0 == p1 {
return Err(invalid("p0 and p1 must differ"));
}
let (lower, upper) = wald_boundaries(alpha, beta)?;
Ok(Sprt {
model: Model::Bernoulli { p0, p1 },
alpha,
beta,
lower,
upper,
observations: 0,
successes: 0,
sum: Q::zero(),
})
}
pub fn normal_mean(
mu0: Q,
mu1: Q,
sigma: Q,
alpha: f64,
beta: f64,
) -> Result<Self, SymplexError> {
if !sigma.is_positive() {
return Err(invalid(format!("sigma must be positive, got {sigma}")));
}
if mu0 == mu1 {
return Err(invalid("mu0 and mu1 must differ"));
}
let (lower, upper) = wald_boundaries(alpha, beta)?;
Ok(Sprt {
model: Model::NormalMean { mu0, mu1, sigma },
alpha,
beta,
lower,
upper,
observations: 0,
successes: 0,
sum: Q::zero(),
})
}
pub fn update(&mut self, success: bool) -> Decision {
self.observations += 1;
if success {
self.successes += 1;
self.sum += Q::one();
}
self.decision()
}
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(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.decision())
}
}
}
pub fn decision(&self) -> Decision {
let llr = self.log_likelihood_ratio_f64();
if llr >= self.upper {
Decision::AcceptH1
} else if llr <= self.lower {
Decision::AcceptH0
} else {
Decision::Continue
}
}
pub fn reset(&mut self) {
self.observations = 0;
self.successes = 0;
self.sum = Q::zero();
}
pub fn boundaries(&self) -> (f64, f64) {
(self.lower, self.upper)
}
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 = ctx.int(usize_to_i64(self.successes));
let f = ctx.int(usize_to_i64(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_f64(&(p1 / p0)).ln();
let lf = q_f64(&((&one - p1) / (&one - p0))).ln();
self.successes as f64 * ls + self.failures() as f64 * lf
}
Model::NormalMean { .. } => q_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 usize_to_i64(n: usize) -> i64 {
i64::try_from(n).unwrap_or(i64::MAX)
}
fn bernoulli_increments(p: f64, p0: &Q, p1: &Q) -> Result<(f64, f64, f64), SymplexError> {
check_probability(p0, "p0")?;
check_probability(p1, "p1")?;
if p0 == p1 {
return Err(invalid("p0 and p1 must differ"));
}
if !(0.0..=1.0).contains(&p) {
return Err(invalid(format!("p must lie in [0, 1], got {p}")));
}
let one = Q::one();
let ls = q_f64(&(p1 / p0)).ln();
let lf = q_f64(&((&one - p1) / (&one - p0))).ln();
Ok((ls, lf, p * ls + (1.0 - p) * lf))
}
fn wald_h(p: f64, ls: f64, lf: f64, drift: f64) -> Result<Option<f64>, SymplexError> {
if drift.abs() < 1e-13 {
return Ok(None);
}
let g = |h: f64| p * (h * ls).exp() + (1.0 - p) * (h * lf).exp() - 1.0;
let side = if drift < 0.0 { 1.0 } else { -1.0 };
let mut hi = side;
let mut grown = 0;
while g(hi) <= 0.0 {
hi *= 2.0;
grown += 1;
if grown > 60 {
return Err(failed(
"operating characteristic: could not bracket the root of Wald's identity",
));
}
}
let mut lo = 0.0;
for _ in 0..200 {
let mid = 0.5 * (lo + hi);
if g(mid) > 0.0 {
hi = mid;
} else {
lo = mid;
}
}
Ok(Some(0.5 * (lo + hi)))
}
pub fn operating_characteristic_bernoulli(
p: f64,
p0: &Q,
p1: &Q,
alpha: f64,
beta: f64,
) -> Result<f64, SymplexError> {
let (ls, lf, drift) = bernoulli_increments(p, p0, p1)?;
let (a, b) = wald_boundaries(alpha, beta)?;
match wald_h(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 expected_sample_size_bernoulli(
p: f64,
p0: &Q,
p1: &Q,
alpha: f64,
beta: f64,
) -> Result<f64, SymplexError> {
let (ls, lf, drift) = bernoulli_increments(p, p0, p1)?;
let (a, b) = wald_boundaries(alpha, beta)?;
match wald_h(p, ls, lf, drift)? {
Some(h) => {
let l = ((b * h).exp() - 1.0) / ((b * h).exp() - (a * h).exp());
Ok((l * a + (1.0 - l) * b) / drift)
}
None => {
let second_moment = p * ls * ls + (1.0 - p) * lf * lf;
Ok(-a * b / second_moment)
}
}
}