Skip to main content

probl_engine/
conjugate.rs

1//! Exact updates for conjugate priors (docs/semantics.md, section 14): the
2//! probability of an observation when its parameter is unknown, and the
3//! parameter's distribution after it.
4//!
5//! Probabilities are computed as logarithms. A run's weight has an extended
6//! exponent, but one observation's probability can already be far below the
7//! smallest `f64`: `0` successes out of 100,000 with a rate near 50% has a
8//! probability of e^−4234.
9
10use crate::continuous::{Family, ln_beta};
11
12/// What an observation saw, with its parameters besides the variable.
13#[derive(Clone, Copy, Debug, PartialEq)]
14pub enum Seen {
15    /// `k` successes from `binomial(trials, x)`.
16    Binomial { trials: u64, k: f64 },
17    /// A fact from `bernoulli(x)`.
18    Bernoulli(bool),
19    /// `k` from `poisson(x)`.
20    Poisson { k: f64 },
21    /// `y` from `normal(x, sd)`.
22    Normal { y: f64, sd: f64 },
23}
24
25/// Whether a draw from `family` may be delayed: it's the prior of a
26/// conjugate pair.
27pub fn is_prior(family: &Family) -> bool {
28    matches!(
29        family,
30        Family::Beta { .. } | Family::Gamma { .. } | Family::Normal { .. }
31    )
32}
33
34/// Observing `seen`, with the variable distributed as `prior`: the natural
35/// logarithm of the observation's probability (a density for `Normal`), and
36/// the variable's distribution after it. `None` if the two aren't a
37/// conjugate pair. A logarithm of −∞ means the observation is impossible,
38/// like a count above the number of trials; the prior is returned then.
39pub fn update(prior: &Family, seen: Seen) -> Option<(f64, Family)> {
40    let impossible = Some((f64::NEG_INFINITY, *prior));
41    let count = |k: f64| k >= 0.0 && k.fract() == 0.0;
42    match (*prior, seen) {
43        (Family::Beta { a, b }, Seen::Binomial { trials, k }) => {
44            let n = trials as f64;
45            if !count(k) || k > n {
46                return impossible;
47            }
48            // C(n, k) B(a + k, b + n − k) / B(a, b), where
49            // C(n, k) = 1 / ((n + 1) B(k + 1, n − k + 1)).
50            let ln = -libm::log1p(n) - ln_beta(k + 1.0, n - k + 1.0) + ln_beta(a + k, b + n - k) - ln_beta(a, b);
51            Some((ln, Family::Beta { a: a + k, b: b + n - k }))
52        }
53        (Family::Beta { a, b }, Seen::Bernoulli(yes)) => Some(if yes {
54            (libm::log(a / (a + b)), Family::Beta { a: a + 1.0, b })
55        } else {
56            (libm::log(b / (a + b)), Family::Beta { a, b: b + 1.0 })
57        }),
58        (Family::Gamma { shape, scale }, Seen::Poisson { k }) => {
59            if !count(k) || !k.is_finite() {
60                return impossible;
61            }
62            // Γ(shape + k) / (Γ(shape) k!) · (scale / (1 + scale))^k
63            // · (1 + scale)^−shape, where the gammas and the factorial make
64            // 1 / (k B(k, shape)) for k ≥ 1.
65            let mut ln = -shape * libm::log1p(scale);
66            if k > 0.0 {
67                // ln(scale / (1 + scale)), without cancelling for a large scale.
68                let odds = if scale > 1.0 {
69                    -libm::log1p(1.0 / scale)
70                } else {
71                    libm::log(scale) - libm::log1p(scale)
72                };
73                ln += -libm::log(k) - ln_beta(k, shape) + k * odds;
74            }
75            let posterior = Family::Gamma {
76                shape: shape + k,
77                scale: scale / (1.0 + scale),
78            };
79            Some((ln, posterior))
80        }
81        (Family::Normal { mean, sd }, Seen::Normal { y, sd: noise }) => {
82            // y is normal with variance sd² + noise²; √ of it without
83            // overflow, and the posterior in terms of shares of it.
84            let spread = libm::hypot(sd, noise);
85            let z = (y - mean) / spread;
86            let ln = -libm::log(spread) - 0.5 * libm::log(2.0 * std::f64::consts::PI) - z * z / 2.0;
87            let (s, t) = (sd / spread, noise / spread);
88            let posterior = Family::Normal {
89                mean: mean + (y - mean) * s * s,
90                sd: sd * t,
91            };
92            Some((if ln.is_nan() { f64::NEG_INFINITY } else { ln }, posterior))
93        }
94        _ => None,
95    }
96}
97
98#[cfg(test)]
99mod tests {
100    use super::*;
101
102    fn close(a: f64, b: f64, tol: f64) {
103        assert!((a - b).abs() <= tol * (1.0 + b.abs()), "{a} vs {b}");
104    }
105
106    fn beta(a: f64, b: f64) -> Family {
107        Family::beta(a, b).unwrap()
108    }
109
110    /// ∫ f(x) dx over (lo, hi), by the midpoint rule on a fine grid.
111    fn integrate(f: impl Fn(f64) -> f64, lo: f64, hi: f64) -> f64 {
112        let n = 200_000;
113        let h = (hi - lo) / n as f64;
114        (0..n).map(|i| f(lo + (i as f64 + 0.5) * h) * h).sum()
115    }
116
117    #[test]
118    fn log_beta_is_accurate_for_small_and_large_arguments() {
119        // Values from mpmath, with 50 digits.
120        close(ln_beta(2.0, 3.0), -2.484_906_649_788_000_3, 1e-15);
121        close(ln_beta(0.5, 0.5), 1.144_729_885_849_400_2, 1e-15);
122        close(ln_beta(1000.0, 1000.0), -1_388.482_601_635_902_3, 1e-15);
123        close(ln_beta(1000.0, 101_000.0), -5_622.584_683_645_091, 1e-15);
124        close(ln_beta(3.0, 1e12), -82.199_916_167_228_7, 1e-15);
125        close(ln_beta(1e10, 1e10), -13_862_943_621.446_32, 1e-15);
126        close(ln_beta(12.5, 9.0), -14.512_975_445_993_465, 1e-15);
127    }
128
129    /// The marginal probability is the prior average of the likelihood, and
130    /// the posterior is the prior times the likelihood, normalized.
131    #[test]
132    fn updates_agree_with_integrating_over_the_prior() {
133        let cases = [
134            (beta(2.0, 50.0), Seen::Binomial { trials: 400, k: 12.0 }),
135            (beta(1.0, 1.0), Seen::Binomial { trials: 10, k: 0.0 }),
136            (beta(3.0, 2.0), Seen::Bernoulli(true)),
137            (beta(3.0, 2.0), Seen::Bernoulli(false)),
138            (Family::gamma(2.0, 3.0).unwrap(), Seen::Poisson { k: 4.0 }),
139            (Family::gamma(1.5, 0.2).unwrap(), Seen::Poisson { k: 0.0 }),
140            (Family::normal(1.0, 2.0).unwrap(), Seen::Normal { y: 1.5, sd: 0.5 }),
141            (Family::normal(-3.0, 0.1).unwrap(), Seen::Normal { y: 4.0, sd: 3.0 }),
142        ];
143        for (prior, seen) in cases {
144            let likelihood = |x: f64| match seen {
145                Seen::Binomial { trials, k } => crate::dist::Counts::Binomial { n: trials, p: x }.pmf(k),
146                Seen::Bernoulli(yes) => {
147                    if yes {
148                        x
149                    } else {
150                        1.0 - x
151                    }
152                }
153                Seen::Poisson { k } => crate::dist::Counts::Poisson { rate: x }.pmf(k),
154                Seen::Normal { y, sd } => Family::normal(x, sd).unwrap().pdf(y),
155            };
156            let (lo, hi) = (prior.quantile(1e-13), prior.quantile(1.0 - 1e-13));
157            let joint = |x: f64| prior.pdf(x) * likelihood(x);
158            let marginal = integrate(joint, lo, hi);
159            let (ln, posterior) = update(&prior, seen).unwrap();
160            close(ln.exp(), marginal, 1e-6);
161            // The posterior's density and moments.
162            for q in [0.1, 0.5, 0.9] {
163                let x = posterior.quantile(q);
164                close(posterior.pdf(x), joint(x) / marginal, 1e-5);
165            }
166            close(posterior.mean(), integrate(|x| x * joint(x), lo, hi) / marginal, 1e-6);
167        }
168    }
169
170    #[test]
171    fn impossible_observations_have_no_probability() {
172        let prior = beta(2.0, 3.0);
173        for k in [11.0, -1.0, 2.5, f64::NAN] {
174            let (ln, after) = update(&prior, Seen::Binomial { trials: 10, k }).unwrap();
175            assert_eq!(ln, f64::NEG_INFINITY);
176            assert_eq!(after, prior);
177        }
178        let gamma = Family::gamma(2.0, 1.0).unwrap();
179        for k in [-1.0, 0.5, f64::INFINITY, f64::NAN] {
180            assert_eq!(update(&gamma, Seen::Poisson { k }).unwrap().0, f64::NEG_INFINITY);
181        }
182        let normal = Family::normal(0.0, 1.0).unwrap();
183        for y in [f64::INFINITY, f64::NAN] {
184            assert_eq!(
185                update(&normal, Seen::Normal { y, sd: 1.0 }).unwrap().0,
186                f64::NEG_INFINITY
187            );
188        }
189        // Pairs that aren't conjugate.
190        assert!(update(&normal, Seen::Poisson { k: 1.0 }).is_none());
191        assert!(update(&gamma, Seen::Bernoulli(true)).is_none());
192        assert!(update(&Family::lognormal(0.0, 1.0).unwrap(), Seen::Poisson { k: 1.0 }).is_none());
193    }
194
195    #[test]
196    fn extreme_observations_keep_their_probability() {
197        // The review's case: far below the smallest f64, but exact.
198        let (ln, after) = update(
199            &beta(1000.0, 1000.0),
200            Seen::Binomial {
201                trials: 100_000,
202                k: 0.0,
203            },
204        )
205        .unwrap();
206        close(ln, -4_234.102_082_009_189, 1e-15);
207        assert_eq!(after, beta(1000.0, 101_000.0));
208        // A million trials, and a count far out in the tail.
209        let (ln, _) = update(
210            &beta(2.0, 50.0),
211            Seen::Binomial {
212                trials: 1_000_000,
213                k: 900_000.0,
214            },
215        )
216        .unwrap();
217        close(ln, -118.892_768_879_054_9, 1e-12);
218        // Priors whose densities are infinite at an end.
219        let (ln, _) = update(&beta(0.5, 0.5), Seen::Binomial { trials: 10, k: 0.0 }).unwrap();
220        close(ln, -1.736_152_296_596_451_7, 1e-15);
221        let (ln, _) = update(&Family::gamma(0.7, 0.2).unwrap(), Seen::Poisson { k: 5.0 }).unwrap();
222        close(ln, -9.850_813_770_178_176, 1e-15);
223        // A density above 1.
224        let (ln, _) = update(&Family::normal(0.0, 0.001).unwrap(), Seen::Normal { y: 0.0, sd: 0.001 }).unwrap();
225        close(ln, 5.642_243_155_497_492, 1e-15);
226        // A huge scale, and a tiny one.
227        let (ln, _) = update(&Family::gamma(2.0, 1e12).unwrap(), Seen::Poisson { k: 3.0 }).unwrap();
228        close(ln, -53.875_747_870_742_21, 1e-14);
229        let (ln, _) = update(&Family::gamma(2.0, 1e-12).unwrap(), Seen::Poisson { k: 3.0 }).unwrap();
230        close(ln, -81.506_768_986_670_75, 1e-14);
231    }
232
233    /// Updating one observation at a time gives the probability of all of
234    /// them together, constants included.
235    #[test]
236    fn sequential_updates_give_the_whole_evidence() {
237        let data = [(400u64, 12.0), (400, 9.0), (350, 15.0), (410, 0.0), (1, 1.0)];
238        let (a, b) = (2.0, 50.0);
239        let mut prior = beta(a, b);
240        let mut total = 0.0;
241        for (trials, k) in data {
242            let (ln, after) = update(&prior, Seen::Binomial { trials, k }).unwrap();
243            total += ln;
244            prior = after;
245        }
246        let (n, k): (f64, f64) = data.iter().fold((0.0, 0.0), |(n, s), &(t, k)| (n + t as f64, s + k));
247        let choose: f64 = data
248            .iter()
249            .map(|&(t, k)| -libm::log1p(t as f64) - ln_beta(k + 1.0, t as f64 - k + 1.0))
250            .sum();
251        close(total, choose + ln_beta(a + k, b + n - k) - ln_beta(a, b), 1e-13);
252        assert_eq!(prior, beta(a + k, b + n - k));
253    }
254}