1use crate::continuous::{Family, ln_beta};
11
12#[derive(Clone, Copy, Debug, PartialEq)]
14pub enum Seen {
15 Binomial { trials: u64, k: f64 },
17 Bernoulli(bool),
19 Poisson { k: f64 },
21 Normal { y: f64, sd: f64 },
23}
24
25pub fn is_prior(family: &Family) -> bool {
28 matches!(
29 family,
30 Family::Beta { .. } | Family::Gamma { .. } | Family::Normal { .. }
31 )
32}
33
34pub 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 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 let mut ln = -shape * libm::log1p(scale);
66 if k > 0.0 {
67 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 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 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 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 #[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 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 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 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 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 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 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 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 #[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}