Three issues from two downstream consumers, all small, all sharing a theme: the crate had the information and would not hand it over. #44 — `UnknownKey { team: 0, member: 0 }` did not say which key. A consumer upgrading 0.1.2 -> 0.4.1 had every one of 5591 predictions return this error, fell back to a neutral 0.5, and lost its entire metadata model for a day. Nothing crashed and nothing logged; it was found by sweeping an unrelated parameter and noticing the output did not move. The 0.4.0 change that made unknown keys an error was right — the error was just too anonymous to act on. It now carries the key's `Debug` rendering, and its `Display` says what to do about it. The precondition is documented on every prediction entry point, which the reporter said would alone have saved the day. #43 — `cdf` was `pub(crate)`, so a consumer asking "is this competitor below the cutoff" approximated it with a `mu + z * sigma` band and had no way to say what confidence any `z` bought. Adds `Gaussian::probability_below` / `probability_above`. The second is separate on purpose: `1 - cdf` collapses to exactly zero past ~8.3 sigma, and a stopping rule is evaluated precisely there. Both route through the survival function added in 0.4.1, so this is visibility rather than new numerics. #50 — `ConvergenceReport` was not `#[must_use]`, so the one signal that a fit stopped short was trivially discarded. It now is, and that immediately found 78 sites doing exactly that — including this crate's own ATP example, which was capped at 10 sweeps when the history needs 30. The example now reads the report and says so. `ITERATIONS = 30` is documented as the floor it is, with the three measurements to hand: 400 events over 100 competitors already stops there at ~7e-3 against a 1e-6 tolerance, the ATP example needs 30 at a much looser one, and a consumer's 2000-node model needs 76 to 161. BREAKING CHANGE: `InferenceError::UnknownKey` gains a `key` field, and the prediction methods now require `K: Debug` in order to fill it. Closes #43, #50. Refs #44 — its third ask, an opt-in `UnknownKeys::Skip` mode, is a live API question and deliberately not answered here. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_011hcFjNDmHXZF8URGLku5zZ
447 lines
15 KiB
Rust
447 lines
15 KiB
Rust
use std::ops;
|
||
|
||
use crate::{MU, N_INF, SIGMA};
|
||
|
||
/// A Gaussian distribution stored in natural parameters.
|
||
///
|
||
/// `pi = 1 / sigma^2` (precision)
|
||
/// `tau = mu * pi` (precision-adjusted mean)
|
||
///
|
||
/// Multiplication and division in message passing become pure adds/subs of
|
||
/// the stored fields with no `sqrt` or reciprocal in the hot path. `mu()` and
|
||
/// `sigma()` are accessors computed on demand.
|
||
#[derive(Clone, Copy, PartialEq, Debug)]
|
||
pub struct Gaussian {
|
||
pi: f64,
|
||
tau: f64,
|
||
}
|
||
|
||
impl Gaussian {
|
||
/// Construct from mean and standard deviation.
|
||
#[must_use]
|
||
pub const fn from_ms(mu: f64, sigma: f64) -> Self {
|
||
if sigma == f64::INFINITY {
|
||
Self { pi: 0.0, tau: 0.0 }
|
||
} else if sigma == 0.0 {
|
||
// Point mass at mu. tau = mu * pi = mu * inf.
|
||
// For mu == 0 this is 0; for mu != 0 it is inf * mu = inf (IEEE).
|
||
// Only N00 (mu=0, sigma=0) is used in practice.
|
||
Self {
|
||
pi: f64::INFINITY,
|
||
tau: if mu == 0.0 { 0.0 } else { f64::INFINITY },
|
||
}
|
||
} else {
|
||
let pi = 1.0 / (sigma * sigma);
|
||
Self { pi, tau: mu * pi }
|
||
}
|
||
}
|
||
|
||
/// Construct from mean and *variance*, skipping the square-root round trip.
|
||
///
|
||
/// `from_ms(mu, var.sqrt())` immediately squares the root away again to
|
||
/// recover `pi = 1/var`. Variance-combining operations (`Add`, `Sub`,
|
||
/// `exclude`, `forget`) work in variance space throughout, so they go
|
||
/// through here instead and never take a root.
|
||
#[inline]
|
||
pub(crate) fn from_mv(mu: f64, var: f64) -> Self {
|
||
if var == f64::INFINITY {
|
||
Self { pi: 0.0, tau: 0.0 }
|
||
} else if var == 0.0 {
|
||
// Point mass at mu; see `from_ms` for the tau convention.
|
||
Self {
|
||
pi: f64::INFINITY,
|
||
tau: if mu == 0.0 { 0.0 } else { f64::INFINITY },
|
||
}
|
||
} else {
|
||
let pi = 1.0 / var;
|
||
Self { pi, tau: mu * pi }
|
||
}
|
||
}
|
||
|
||
/// Construct directly from natural parameters.
|
||
#[inline]
|
||
pub(crate) const fn from_natural(pi: f64, tau: f64) -> Self {
|
||
Self { pi, tau }
|
||
}
|
||
|
||
#[inline]
|
||
#[must_use]
|
||
pub fn pi(&self) -> f64 {
|
||
self.pi
|
||
}
|
||
|
||
#[inline]
|
||
#[must_use]
|
||
pub fn tau(&self) -> f64 {
|
||
self.tau
|
||
}
|
||
|
||
#[inline]
|
||
#[must_use]
|
||
pub fn mu(&self) -> f64 {
|
||
// A non-positive precision is an improper (uninformative) Gaussian — its mean is
|
||
// undefined. Treat it like `pi == 0` and return 0. EP message cancellation can land
|
||
// `pi` on a tiny negative value (round-off of exactly zero); without this guard
|
||
// `tau / pi` would yield a spurious finite mean.
|
||
if self.pi <= 0.0 {
|
||
0.0
|
||
} else {
|
||
self.tau / self.pi
|
||
}
|
||
}
|
||
|
||
/// Variance, `1 / pi`, without the root-and-square of `sigma().powi(2)`.
|
||
///
|
||
/// Mirrors `sigma()`'s treatment of the improper (`pi <= 0`) and point-mass
|
||
/// (`pi == inf`) cases.
|
||
#[inline]
|
||
pub(crate) fn variance(&self) -> f64 {
|
||
if self.pi <= 0.0 {
|
||
f64::INFINITY
|
||
} else if self.pi.is_infinite() {
|
||
0.0
|
||
} else {
|
||
1.0 / self.pi
|
||
}
|
||
}
|
||
|
||
#[inline]
|
||
#[must_use]
|
||
pub fn sigma(&self) -> f64 {
|
||
// A non-positive precision is improper → infinite standard deviation. Guarding
|
||
// `pi <= 0.0` (not just `== 0.0`) keeps `1.0 / pi.sqrt()` from returning NaN when EP
|
||
// cancellation produces a tiny negative precision (round-off of exactly zero).
|
||
if self.pi <= 0.0 {
|
||
f64::INFINITY
|
||
} else if self.pi.is_infinite() {
|
||
0.0
|
||
} else {
|
||
1.0 / self.pi.sqrt()
|
||
}
|
||
}
|
||
|
||
pub(crate) fn delta(&self, other: Gaussian) -> (f64, f64) {
|
||
(
|
||
(self.mu() - other.mu()).abs(),
|
||
(self.sigma() - other.sigma()).abs(),
|
||
)
|
||
}
|
||
|
||
pub(crate) fn exclude(&self, other: Gaussian) -> Self {
|
||
let var = self.variance() - other.variance();
|
||
if var <= 0.0 {
|
||
// When sigma_self ≈ sigma_other (including ULP-level rounding differences
|
||
// from the pi→sigma accessor round-trip), the excluded contribution is N00.
|
||
// Computing from_ms(tiny_mu, 0.0) would give {pi:inf, tau:inf}, whose
|
||
// mu() = inf/inf = NaN. Returning N00 is correct: when both Gaussians
|
||
// carry the same variance, the residual is a point mass at 0.
|
||
return Gaussian::from_mv(0.0, 0.0);
|
||
}
|
||
|
||
Self::from_mv(self.mu() - other.mu(), var)
|
||
}
|
||
|
||
pub(crate) fn forget(&self, variance_delta: f64) -> Self {
|
||
Self::from_mv(self.mu(), self.variance() + variance_delta)
|
||
}
|
||
|
||
/// `P(X < x)` under this Gaussian.
|
||
///
|
||
/// The question a stopping rule asks: *how sure am I that this competitor's
|
||
/// true skill is below the cutoff?* Expressing that as a probability keeps
|
||
/// its meaning as sigma changes, where a `mu + z * sigma` band silently
|
||
/// means different confidence at different uncertainties — which is exactly
|
||
/// the regime a stopping rule operates in.
|
||
///
|
||
/// Accurate in the *lower* tail. For the upper tail use
|
||
/// [`Gaussian::probability_above`] rather than `1.0 - probability_below(x)`,
|
||
/// which cancels away every significant digit once the result is small.
|
||
///
|
||
/// An improper Gaussian (non-positive precision) has no defined mean, so
|
||
/// this returns `0.5` — the same convention `mu()` and `sigma()` follow.
|
||
#[must_use]
|
||
pub fn probability_below(&self, x: f64) -> f64 {
|
||
if self.pi <= 0.0 {
|
||
return 0.5;
|
||
}
|
||
crate::cdf(x, self.mu(), self.sigma())
|
||
}
|
||
|
||
/// `P(X > x)` under this Gaussian.
|
||
///
|
||
/// Computed as a survival function rather than `1 - cdf`, so it keeps full
|
||
/// relative precision in the upper tail: `1 - cdf` returns exactly zero
|
||
/// past about 8.3 sigma, where the true value is still 1e-19 and perfectly
|
||
/// representable. A stopping rule is evaluated precisely there — the
|
||
/// interesting cases are the ones near certainty.
|
||
///
|
||
/// An improper Gaussian returns `0.5`, as [`Gaussian::probability_below`].
|
||
#[must_use]
|
||
pub fn probability_above(&self, x: f64) -> f64 {
|
||
if self.pi <= 0.0 {
|
||
return 0.5;
|
||
}
|
||
crate::sf(x, self.mu(), self.sigma())
|
||
}
|
||
|
||
/// EP damping in natural-parameter space: `α·new + (1−α)·self`.
|
||
///
|
||
/// Used by within-game inference to stabilise oscillating fixed-point
|
||
/// loops on hard graphs. `alpha = 1.0` returns `new` exactly;
|
||
/// `alpha < 1.0` shrinks each per-step update.
|
||
#[must_use]
|
||
pub fn damp_natural(self, new: Gaussian, alpha: f64) -> Gaussian {
|
||
Gaussian::from_natural(
|
||
alpha * new.pi() + (1.0 - alpha) * self.pi(),
|
||
alpha * new.tau() + (1.0 - alpha) * self.tau(),
|
||
)
|
||
}
|
||
}
|
||
|
||
impl Default for Gaussian {
|
||
fn default() -> Self {
|
||
Self::from_ms(MU, SIGMA)
|
||
}
|
||
}
|
||
|
||
impl ops::Add<Gaussian> for Gaussian {
|
||
type Output = Gaussian;
|
||
/// Variance addition: (mu1 + mu2, sqrt(σ1² + σ2²)).
|
||
/// Used for combining performance and noise; rare relative to mul/div.
|
||
fn add(self, rhs: Gaussian) -> Self::Output {
|
||
Self::from_mv(self.mu() + rhs.mu(), self.variance() + rhs.variance())
|
||
}
|
||
}
|
||
|
||
impl ops::Sub<Gaussian> for Gaussian {
|
||
type Output = Gaussian;
|
||
/// (mu1 - mu2, sqrt(σ1² + σ2²)). Same sigma combination as Add.
|
||
fn sub(self, rhs: Gaussian) -> Self::Output {
|
||
Self::from_mv(self.mu() - rhs.mu(), self.variance() + rhs.variance())
|
||
}
|
||
}
|
||
|
||
impl ops::Mul<Gaussian> for Gaussian {
|
||
type Output = Gaussian;
|
||
/// Factor product: nat-param add. Hot path — two f64 additions, no sqrt.
|
||
fn mul(self, rhs: Gaussian) -> Self::Output {
|
||
Self::from_natural(self.pi + rhs.pi, self.tau + rhs.tau)
|
||
}
|
||
}
|
||
|
||
impl ops::Mul<f64> for Gaussian {
|
||
type Output = Gaussian;
|
||
fn mul(self, scalar: f64) -> Self::Output {
|
||
if !scalar.is_finite() {
|
||
return N_INF;
|
||
}
|
||
if scalar == 0.0 {
|
||
// Scaling by 0 collapses to a point mass at 0 (sigma' = 0, mu' = 0).
|
||
// This is N00, the additive identity, NOT N_INF.
|
||
return Gaussian::from_mv(0.0, 0.0);
|
||
}
|
||
// sigma' = sigma * |scalar| => pi' = pi / scalar²
|
||
// mu' = mu * scalar => tau' = tau / scalar
|
||
Self::from_natural(self.pi / (scalar * scalar), self.tau / scalar)
|
||
}
|
||
}
|
||
|
||
impl ops::Div<Gaussian> for Gaussian {
|
||
type Output = Gaussian;
|
||
/// Cavity: nat-param sub. Hot path — two f64 subtractions, no sqrt.
|
||
fn div(self, rhs: Gaussian) -> Self::Output {
|
||
Self::from_natural(self.pi - rhs.pi, self.tau - rhs.tau)
|
||
}
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
|
||
#[test]
|
||
fn non_positive_precision_is_improper_not_nan() {
|
||
// EP message cancellation can leave `pi` a tiny negative (round-off of exactly zero).
|
||
// Such a Gaussian is improper/uninformative: mu() must be 0 and sigma() infinite, not
|
||
// NaN. A NaN here propagates through the moment-space `Sub` in the game chain and
|
||
// poisons every skill in the slice.
|
||
let tiny_neg = Gaussian::from_natural(-5.55e-17, -8.88e-16);
|
||
assert_eq!(tiny_neg.mu(), 0.0);
|
||
assert!(tiny_neg.sigma().is_infinite());
|
||
|
||
// A frankly-negative precision is treated the same way.
|
||
let neg = Gaussian::from_natural(-1.0, 2.0);
|
||
assert_eq!(neg.mu(), 0.0);
|
||
assert!(neg.sigma().is_infinite());
|
||
|
||
// Subtracting such a message must not produce NaN (the original failure path).
|
||
let proper = Gaussian::from_ms(9.75, 1.256);
|
||
let diff = proper - tiny_neg;
|
||
assert!(diff.pi().is_finite() && !diff.pi().is_nan());
|
||
assert!(diff.tau().is_finite() && !diff.tau().is_nan());
|
||
}
|
||
|
||
#[test]
|
||
fn test_add() {
|
||
let n = Gaussian::from_ms(25.0, 25.0 / 3.0);
|
||
let m = Gaussian::from_ms(0.0, 1.0);
|
||
let r = n + m;
|
||
assert!((r.mu() - 25.0).abs() < 1e-12);
|
||
assert!((r.sigma() - 8.393118874676116).abs() < 1e-10);
|
||
}
|
||
|
||
#[test]
|
||
fn test_sub() {
|
||
let n = Gaussian::from_ms(25.0, 25.0 / 3.0);
|
||
let m = Gaussian::from_ms(1.0, 1.0);
|
||
let r = n - m;
|
||
assert!((r.mu() - 24.0).abs() < 1e-12);
|
||
assert!((r.sigma() - 8.393118874676116).abs() < 1e-10);
|
||
}
|
||
|
||
#[test]
|
||
fn test_mul() {
|
||
let n = Gaussian::from_ms(25.0, 25.0 / 3.0);
|
||
let m = Gaussian::from_ms(0.0, 1.0);
|
||
let r = n * m;
|
||
assert!((r.mu() - 0.35488958990536273).abs() < 1e-10);
|
||
assert!((r.sigma() - 0.992876838486922).abs() < 1e-10);
|
||
}
|
||
|
||
#[test]
|
||
fn test_div() {
|
||
let n = Gaussian::from_ms(25.0, 25.0 / 3.0);
|
||
let m = Gaussian::from_ms(0.0, 1.0);
|
||
let r = m / n;
|
||
assert!((r.mu() - (-0.3652597402597402)).abs() < 1e-10);
|
||
assert!((r.sigma() - 1.0072787050317253).abs() < 1e-10);
|
||
}
|
||
|
||
#[test]
|
||
fn test_n00_is_add_identity() {
|
||
// N00 (sigma=0) is the additive identity for the variance-convolution Add op.
|
||
// N_INF (sigma=inf) is the identity for the EP-product Mul op.
|
||
let g = Gaussian::from_ms(3.0, 2.0);
|
||
let n00 = Gaussian::from_ms(0.0, 0.0);
|
||
let r = n00 + g;
|
||
assert!((r.mu() - g.mu()).abs() < 1e-12);
|
||
assert!((r.sigma() - g.sigma()).abs() < 1e-12);
|
||
}
|
||
|
||
#[test]
|
||
fn test_mul_is_factor_product() {
|
||
// n * m in nat-params should be pi_n + pi_m, tau_n + tau_m
|
||
let n = Gaussian::from_ms(2.0, 3.0);
|
||
let m = Gaussian::from_ms(1.0, 2.0);
|
||
let r = n * m;
|
||
let expected_pi = n.pi() + m.pi();
|
||
let expected_tau = n.tau() + m.tau();
|
||
assert!((r.pi() - expected_pi).abs() < 1e-15);
|
||
assert!((r.tau() - expected_tau).abs() < 1e-15);
|
||
}
|
||
|
||
#[test]
|
||
fn test_div_is_cavity() {
|
||
let n = Gaussian::from_ms(2.0, 1.0);
|
||
let m = Gaussian::from_ms(1.0, 2.0);
|
||
let r = n / m;
|
||
let expected_pi = n.pi() - m.pi();
|
||
let expected_tau = n.tau() - m.tau();
|
||
assert!((r.pi() - expected_pi).abs() < 1e-15);
|
||
assert!((r.tau() - expected_tau).abs() < 1e-15);
|
||
}
|
||
|
||
#[test]
|
||
fn damp_natural_alpha_one_returns_new() {
|
||
let old = Gaussian::from_ms(1.0, 2.0);
|
||
let new = Gaussian::from_ms(5.0, 0.5);
|
||
let damped = old.damp_natural(new, 1.0);
|
||
assert_eq!(damped.pi(), new.pi());
|
||
assert_eq!(damped.tau(), new.tau());
|
||
}
|
||
|
||
#[test]
|
||
fn damp_natural_alpha_zero_returns_self() {
|
||
let old = Gaussian::from_ms(1.0, 2.0);
|
||
let new = Gaussian::from_ms(5.0, 0.5);
|
||
let damped = old.damp_natural(new, 0.0);
|
||
assert_eq!(damped.pi(), old.pi());
|
||
assert_eq!(damped.tau(), old.tau());
|
||
}
|
||
|
||
#[test]
|
||
fn damp_natural_alpha_half_is_midpoint_in_natural_params() {
|
||
let old = Gaussian::from_ms(1.0, 2.0);
|
||
let new = Gaussian::from_ms(5.0, 0.5);
|
||
let damped = old.damp_natural(new, 0.5);
|
||
let expected_pi = 0.5 * new.pi() + 0.5 * old.pi();
|
||
let expected_tau = 0.5 * new.tau() + 0.5 * old.tau();
|
||
assert!((damped.pi() - expected_pi).abs() < 1e-12);
|
||
assert!((damped.tau() - expected_tau).abs() < 1e-12);
|
||
}
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tail_probability_tests {
|
||
use super::*;
|
||
|
||
#[test]
|
||
fn probability_below_matches_published_quantiles() {
|
||
let g = Gaussian::from_ms(0.0, 1.0);
|
||
for (x, expected) in [
|
||
(-1.959_963_984_540_054, 0.025),
|
||
(0.0, 0.5),
|
||
(1.281_551_565_544_6, 0.9),
|
||
(1.959_963_984_540_054, 0.975),
|
||
] {
|
||
let got = g.probability_below(x);
|
||
assert!(
|
||
(got - expected).abs() < 1e-12,
|
||
"P(X < {x}) = {got}, expected {expected}"
|
||
);
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn the_two_tails_partition_the_mass() {
|
||
let g = Gaussian::from_ms(3.0, 2.0);
|
||
for x in [-4.0f64, 0.0, 3.0, 7.5] {
|
||
let total = g.probability_below(x) + g.probability_above(x);
|
||
assert!((total - 1.0).abs() < 1e-15, "at {x}: {total}");
|
||
}
|
||
}
|
||
|
||
/// The reason `probability_above` exists rather than `1 - probability_below`.
|
||
#[test]
|
||
fn probability_above_keeps_precision_where_the_complement_collapses() {
|
||
let g = Gaussian::from_ms(0.0, 1.0);
|
||
for (x, expected) in [(9.0f64, 1.128_588e-19), (20.0, 2.753_624e-89)] {
|
||
let got = g.probability_above(x);
|
||
assert!(
|
||
(got - expected).abs() / expected < 1e-6,
|
||
"P(X > {x}) = {got}, expected ~{expected}"
|
||
);
|
||
assert_eq!(
|
||
1.0 - g.probability_below(x),
|
||
0.0,
|
||
"the complement should still collapse at {x}"
|
||
);
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn a_scaled_gaussian_shifts_and_stretches() {
|
||
let g = Gaussian::from_ms(25.0, 6.0);
|
||
assert!((g.probability_below(25.0) - 0.5).abs() < 1e-15);
|
||
// One sigma either side of the mean.
|
||
assert!((g.probability_below(31.0) - 0.841_344_746_068_543).abs() < 1e-12);
|
||
assert!((g.probability_above(19.0) - 0.841_344_746_068_543).abs() < 1e-12);
|
||
}
|
||
|
||
#[test]
|
||
fn an_improper_gaussian_is_uninformative_rather_than_nan() {
|
||
let improper = Gaussian::from_ms(0.0, f64::INFINITY);
|
||
assert_eq!(improper.probability_below(5.0), 0.5);
|
||
assert_eq!(improper.probability_above(5.0), 0.5);
|
||
}
|
||
}
|