use crate::special::ln_beta;
#[derive(Debug, Clone, Default)]
pub struct Welford {
n: u64,
mean: f64,
m2: f64,
}
impl Welford {
pub fn new() -> Self {
Self::default()
}
pub fn push(&mut self, x: f64) {
self.n += 1;
let delta = x - self.mean;
self.mean += delta / self.n as f64;
self.m2 += delta * (x - self.mean);
}
pub fn count(&self) -> u64 {
self.n
}
pub fn mean(&self) -> f64 {
self.mean
}
pub fn sample_variance(&self) -> f64 {
if self.n < 2 {
0.0
} else {
self.m2 / (self.n - 1) as f64
}
}
pub fn population_variance(&self) -> f64 {
if self.n == 0 {
0.0
} else {
self.m2 / self.n as f64
}
}
pub fn sample_std(&self) -> f64 {
self.sample_variance().sqrt()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConfidenceLevel {
P90,
P95,
P99,
}
impl ConfidenceLevel {
pub fn z(&self) -> f64 {
match self {
Self::P90 => 1.644_853_626_951_472_2,
Self::P95 => 1.959_963_984_540_054,
Self::P99 => 2.575_829_303_548_900_4,
}
}
pub fn alpha(&self) -> f64 {
match self {
Self::P90 => 0.10,
Self::P95 => 0.05,
Self::P99 => 0.01,
}
}
pub fn as_percent(&self) -> u32 {
match self {
Self::P90 => 90,
Self::P95 => 95,
Self::P99 => 99,
}
}
}
pub fn wilson_interval(successes: u64, trials: u64, level: ConfidenceLevel) -> (f64, f64) {
if trials == 0 {
return (0.0, 1.0);
}
let n = trials as f64;
let p = (successes as f64 / n).min(1.0);
let z = level.z();
let z2 = z * z;
let denom = 1.0 + z2 / n;
let center = (p + z2 / (2.0 * n)) / denom;
let spread = (z / denom) * (p * (1.0 - p) / n + z2 / (4.0 * n * n)).sqrt();
((center - spread).max(0.0), (center + spread).min(1.0))
}
const BRACKET_EPS: f64 = 1e-15;
const BISECTION_ITERS: u32 = 200;
fn ln_mixture(ln_b: f64, successes: f64, failures: f64, p: f64) -> f64 {
let mut ln_m = ln_b;
if successes > 0.0 {
ln_m -= successes * p.ln();
}
if failures > 0.0 {
ln_m -= failures * (1.0 - p).ln();
}
ln_m
}
fn bisect_endpoint(
ln_b: f64,
successes: f64,
failures: f64,
threshold: f64,
mut outside: f64,
mut inside: f64,
) -> f64 {
for _ in 0..BISECTION_ITERS {
let mid = 0.5 * (outside + inside);
if ln_mixture(ln_b, successes, failures, mid) > threshold {
outside = mid;
} else {
inside = mid;
}
}
0.5 * (outside + inside)
}
#[derive(Debug, Clone)]
pub struct BernoulliConfidenceSequence {
successes: u64,
trials: u64,
alpha: f64,
}
impl BernoulliConfidenceSequence {
pub fn new(level: ConfidenceLevel) -> Self {
Self {
successes: 0,
trials: 0,
alpha: level.alpha(),
}
}
pub fn update(&mut self, hit: bool) {
self.trials = self.trials.saturating_add(1);
if hit {
self.successes = self.successes.saturating_add(1);
}
}
pub fn update_batch(&mut self, hits: u64, trials: u64) {
self.trials = self.trials.saturating_add(trials);
self.successes = self.successes.saturating_add(hits.min(trials));
}
pub fn successes(&self) -> u64 {
self.successes
}
pub fn trials(&self) -> u64 {
self.trials
}
pub fn bounds(&self) -> (f64, f64) {
if self.trials == 0 {
return (0.0, 1.0);
}
let n = self.trials as f64;
let successes = self.successes.min(self.trials) as f64;
let failures = n - successes;
let p_hat = successes / n;
let threshold = -self.alpha.ln(); let ln_b = ln_beta(successes + 1.0, failures + 1.0);
let lower = if p_hat <= BRACKET_EPS
|| ln_mixture(ln_b, successes, failures, BRACKET_EPS) <= threshold
{
0.0
} else {
bisect_endpoint(ln_b, successes, failures, threshold, BRACKET_EPS, p_hat)
};
let upper = if p_hat >= 1.0 - BRACKET_EPS
|| ln_mixture(ln_b, successes, failures, 1.0 - BRACKET_EPS) <= threshold
{
1.0
} else {
bisect_endpoint(
ln_b,
successes,
failures,
threshold,
1.0 - BRACKET_EPS,
p_hat,
)
};
(lower, upper)
}
pub fn half_width(&self) -> f64 {
let (lo, hi) = self.bounds();
(hi - lo) / 2.0
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn welford_matches_two_pass_moments_on_a_hostile_fixture() {
let xs: Vec<f64> = (0..1000).map(|i| 1.0e9 + (i % 7) as f64 * 0.25).collect();
let mut w = Welford::new();
for &x in &xs { w.push(x); }
let n = xs.len() as f64;
let mean = xs.iter().sum::<f64>() / n;
let m2 = xs.iter().map(|x| (x - mean).powi(2)).sum::<f64>();
assert_eq!(w.count(), 1000);
assert!((w.mean() - mean).abs() < 1e-6);
assert!((w.population_variance() - m2 / n).abs() < 1e-6);
assert!((w.sample_variance() - m2 / (n - 1.0)).abs() < 1e-6);
}
#[test]
fn welford_empty_and_single_are_defined() {
let mut w = Welford::new();
assert_eq!(w.count(), 0);
assert_eq!(w.mean(), 0.0);
assert_eq!(w.sample_variance(), 0.0);
assert_eq!(w.population_variance(), 0.0);
w.push(42.0);
assert_eq!(w.mean(), 42.0);
assert_eq!(w.sample_variance(), 0.0); assert_eq!(w.population_variance(), 0.0);
}
#[test]
fn confidence_level_accessors_are_pinned_exactly() {
assert_eq!(ConfidenceLevel::P90.z().to_bits(), 1.644_853_626_951_472_2_f64.to_bits());
assert_eq!(ConfidenceLevel::P95.z().to_bits(), 1.959_963_984_540_054_f64.to_bits());
assert_eq!(ConfidenceLevel::P99.z().to_bits(), 2.575_829_303_548_900_4_f64.to_bits());
assert_eq!(ConfidenceLevel::P90.alpha(), 0.10);
assert_eq!(ConfidenceLevel::P95.alpha(), 0.05);
assert_eq!(ConfidenceLevel::P99.alpha(), 0.01);
assert_eq!(ConfidenceLevel::P90.as_percent(), 90);
assert_eq!(ConfidenceLevel::P95.as_percent(), 95);
assert_eq!(ConfidenceLevel::P99.as_percent(), 99);
}
#[test]
fn wilson_matches_the_newcombe_canonical_example() {
let (lo, hi) = wilson_interval(81, 263, ConfidenceLevel::P95);
assert!((lo - 0.2553).abs() < 5e-4, "lo = {lo}");
assert!((hi - 0.3662).abs() < 5e-4, "hi = {hi}");
}
#[test]
fn wilson_two_algebraic_forms_agree() {
for &(k, n) in &[(0_u64, 10_u64), (10, 10), (1, 10), (8, 10), (500, 1000), (3, 7)] {
for level in [ConfidenceLevel::P90, ConfidenceLevel::P95, ConfidenceLevel::P99] {
let (lo, hi) = wilson_interval(k, n, level);
let z = level.z();
let (kf, nf) = (k as f64, n as f64);
let p_hat = kf / nf;
let z2 = z * z;
let a = nf + z2;
let b = 2.0 * nf * p_hat + z2;
let root_term = z * (z2 + 4.0 * nf * p_hat * (1.0 - p_hat)).sqrt();
let lo_root = (b - root_term) / (2.0 * a);
let hi_root = (b + root_term) / (2.0 * a);
assert!((lo - lo_root).abs() < 1e-12, "lo mismatch at k={k} n={n}: {lo} vs {lo_root}");
assert!((hi - hi_root).abs() < 1e-12, "hi mismatch at k={k} n={n}: {hi} vs {hi_root}");
}
}
}
#[test]
#[allow(clippy::manual_range_contains)]
fn wilson_edges_and_ordering_properties() {
assert_eq!(wilson_interval(0, 0, ConfidenceLevel::P95), (0.0, 1.0));
for n in 1_u64..=200 {
for level in [ConfidenceLevel::P90, ConfidenceLevel::P95, ConfidenceLevel::P99] {
let (lo0, _) = wilson_interval(0, n, level);
assert!(lo0 >= 0.0 && lo0 < 1e-9, "k=0 n={n} {level:?}: lo = {lo0:e}, want [0, 1e-9)");
let (_, hi_n) = wilson_interval(n, n, level);
assert!(hi_n <= 1.0 && hi_n > 1.0 - 1e-9, "k=n n={n} {level:?}: hi = {hi_n:e}, want (1-1e-9, 1]");
}
}
let w = |k, n, l| { let (a, b) = wilson_interval(k, n, l); b - a };
assert!(w(40, 100, ConfidenceLevel::P99) > w(40, 100, ConfidenceLevel::P95));
assert!(w(40, 100, ConfidenceLevel::P95) > w(400, 1000, ConfidenceLevel::P95));
let (lo, hi) = wilson_interval(3, 7, ConfidenceLevel::P90);
let p = 3.0 / 7.0;
assert!(lo < p && p < hi);
}
#[test]
fn wilson_interval_saturates_when_successes_exceeds_trials() {
for level in [ConfidenceLevel::P90, ConfidenceLevel::P95, ConfidenceLevel::P99] {
assert_eq!(wilson_interval(37, 20, level), wilson_interval(20, 20, level));
assert_eq!(wilson_interval(u64::MAX, 20, level), wilson_interval(20, 20, level));
}
}
#[test]
fn cs_starts_ignorant_and_shrinks_monotonically_in_n() {
let mut cs = BernoulliConfidenceSequence::new(ConfidenceLevel::P95);
assert_eq!(cs.bounds(), (0.0, 1.0));
let mut prev = 1.0_f64;
for i in 0..2000 {
cs.update(i % 2 == 0);
if i % 100 == 99 {
let hw = cs.half_width();
assert!(hw <= prev + 1e-12, "half-width grew at n={}: {hw} > {prev}", i + 1);
prev = hw;
}
}
let (lo, hi) = cs.bounds();
assert!(lo > 0.0 && hi < 1.0 && lo < 0.5 && 0.5 < hi);
}
#[test]
fn cs_edges_saturate_exactly() {
let mut cs = BernoulliConfidenceSequence::new(ConfidenceLevel::P95);
cs.update_batch(0, 50);
let (lo, hi) = cs.bounds();
assert_eq!(lo, 0.0);
assert!(hi < 0.25 && hi > 0.0);
let mut cs2 = BernoulliConfidenceSequence::new(ConfidenceLevel::P95);
cs2.update_batch(50, 50);
let (lo2, hi2) = cs2.bounds();
assert_eq!(hi2, 1.0);
assert!(lo2 > 0.75 && lo2 < 1.0);
}
#[test]
fn cs_is_wider_than_wilson_at_the_same_n() {
for &(k, n) in &[(10_u64, 40_u64), (81, 263), (500, 1000)] {
let mut cs = BernoulliConfidenceSequence::new(ConfidenceLevel::P95);
cs.update_batch(k, n);
let (clo, chi) = cs.bounds();
let (wlo, whi) = wilson_interval(k, n, ConfidenceLevel::P95);
assert!(chi - clo > whi - wlo, "CS not wider than Wilson at k={k}, n={n}");
}
}
#[test]
fn cs_bounds_are_deterministic_and_batch_order_invariant() {
let mut a = BernoulliConfidenceSequence::new(ConfidenceLevel::P99);
let mut b = BernoulliConfidenceSequence::new(ConfidenceLevel::P99);
a.update_batch(30, 100);
for i in 0..100 { b.update(i < 30); }
assert_eq!(a.bounds(), b.bounds()); }
#[test]
fn cs_brackets_the_mle_and_nests_by_confidence_level() {
for n in 1_u64..=60 {
for s in 0..=n {
let p_hat = s as f64 / n as f64;
let mut widths = Vec::new();
for level in [ConfidenceLevel::P90, ConfidenceLevel::P95, ConfidenceLevel::P99] {
let mut cs = BernoulliConfidenceSequence::new(level);
cs.update_batch(s, n);
let (lo, hi) = cs.bounds();
assert!(lo.is_finite() && hi.is_finite(), "non-finite bound at s={s} n={n}");
assert!((0.0..=1.0).contains(&lo), "lo={lo} out of [0,1] at s={s} n={n}");
assert!((0.0..=1.0).contains(&hi), "hi={hi} out of [0,1] at s={s} n={n}");
assert!(lo <= p_hat, "lo={lo} above p_hat={p_hat} at s={s} n={n} {level:?}");
assert!(hi >= p_hat, "hi={hi} below p_hat={p_hat} at s={s} n={n} {level:?}");
assert!((hi - lo - 2.0 * cs.half_width()).abs() < 1e-15);
widths.push(hi - lo);
}
assert!(widths[0] <= widths[1], "P90 wider than P95 at s={s} n={n}");
assert!(widths[1] <= widths[2], "P95 wider than P99 at s={s} n={n}");
}
}
let mut bad = BernoulliConfidenceSequence::new(ConfidenceLevel::P95);
bad.update_batch(u64::MAX, 20);
assert_eq!(bad.successes(), 20);
assert_eq!(bad.trials(), 20);
let (lo, hi) = bad.bounds();
assert!(lo > 0.0 && !lo.is_nan(), "lo={lo}");
assert_eq!(hi, 1.0);
}
fn binomial_cdf(k: usize, n: usize, p: f64) -> f64 {
let ratio = p / (1.0 - p);
let mut term = (1.0 - p).powi(n as i32);
let mut total = term;
for i in 1..=k.min(n) {
term *= ratio * (n - i + 1) as f64 / i as f64;
total += term;
}
total.min(1.0)
}
#[test]
fn cs_coverage_survives_optional_stopping() {
use rand::{rngs::StdRng, RngExt, SeedableRng};
const TRIALS: usize = 200;
const P_TRUE: f64 = 0.30;
let mut rng = StdRng::seed_from_u64(0x1352_C0FF_EE00_0001);
let mut covered = 0_usize;
let mut final_half_widths = 0.0_f64;
for _ in 0..TRIALS {
let mut cs = BernoulliConfidenceSequence::new(ConfidenceLevel::P95);
for _ in 0..5000 {
cs.update(rng.random::<f64>() < P_TRUE);
if cs.trials() >= 50 && cs.half_width() <= 0.08 { break; }
}
let (lo, hi) = cs.bounds();
if lo <= P_TRUE && P_TRUE <= hi { covered += 1; }
final_half_widths += cs.half_width();
}
let mean_hw = final_half_widths / TRIALS as f64;
eprintln!("MBA-1352 CS coverage: {covered}/{TRIALS}, mean final half-width {mean_hw:.5}");
const COVERAGE_FLOOR: usize = 180;
let false_failure_rate = binomial_cdf(COVERAGE_FLOOR - 1, TRIALS, 0.95);
assert!(
(false_failure_rate - 1.1599e-3).abs() < 1e-6,
"exact tail P(X <= {}) = {false_failure_rate:.6e}, expected ~1.1599e-3",
COVERAGE_FLOOR - 1
);
assert!(
covered >= COVERAGE_FLOOR,
"coverage {covered}/{TRIALS} below the exact-tail floor {COVERAGE_FLOOR}"
);
assert!(mean_hw < 0.12, "mean final half-width {mean_hw} — intervals not converging");
}
}