Skip to main content

qualia_core_db/solvers/learning/variational/
gaussian.rs

1//! Mean-field variational inference for a univariate Gaussian (PRML ch 10.1.3) —
2//! the canonical CAVI example. Data `xₙ ~ N(μ, τ⁻¹)` with a Normal-Gamma prior; the
3//! variational posterior is factorized `q(μ,τ) = q(μ)·q(τ)` (`q(μ)` Gaussian, `q(τ)`
4//! Gamma), and the coordinate-ascent updates iterate to a fixed point.
5//!
6//! This is the worked instance of the general principle: approximate an intractable
7//! posterior by the closest factorized distribution, maximizing the ELBO.
8
9use crate::solvers::learning::LearningError;
10
11/// The fitted factorized posterior `q(μ) = N(μ_n, λ_n⁻¹)`, `q(τ) = Gamma(a_n, b_n)`.
12#[derive(Debug, Clone, Copy)]
13pub struct VariationalGaussian {
14    pub mu_n: f64,
15    pub lambda_n: f64,
16    pub a_n: f64,
17    pub b_n: f64,
18    pub n_iter: usize,
19    pub converged: bool,
20}
21
22impl VariationalGaussian {
23    /// Posterior mean of μ.
24    pub fn mean(&self) -> f64 {
25        self.mu_n
26    }
27    /// Posterior mean of the precision τ, `E[τ] = a_n/b_n`.
28    pub fn precision_mean(&self) -> f64 {
29        self.a_n / self.b_n
30    }
31    /// Implied posterior-mean variance `1/E[τ]`.
32    pub fn variance_mean(&self) -> f64 {
33        self.b_n / self.a_n
34    }
35}
36
37/// Run CAVI for the univariate Gaussian. Priors: `μ ~ N(μ0, (λ0·τ)⁻¹)` (`mu0`,
38/// `lambda0`), `τ ~ Gamma(a0, b0)`. Fails closed on too little data.
39pub fn fit(
40    data: &[f64],
41    mu0: f64,
42    lambda0: f64,
43    a0: f64,
44    b0: f64,
45    max_iter: usize,
46    tol: f64,
47) -> Result<VariationalGaussian, LearningError> {
48    let n = data.len();
49    if n < 2 || lambda0 < 0.0 || !(a0 > 0.0) || !(b0 > 0.0) {
50        return Err(LearningError::InsufficientData);
51    }
52    let nf = n as f64;
53    let xbar = data.iter().sum::<f64>() / nf;
54
55    // q(μ) mean is fixed by the data + prior; only its precision depends on E[τ].
56    let mu_n = (lambda0 * mu0 + nf * xbar) / (lambda0 + nf);
57    let a_n = a0 + (nf + 1.0) / 2.0;
58
59    // Initialise E[τ] from the sample variance.
60    let s2 = data.iter().map(|&x| (x - xbar).powi(2)).sum::<f64>() / nf;
61    let mut e_tau = 1.0 / s2.max(1e-9);
62    let mut lambda_n = (lambda0 + nf) * e_tau;
63    let mut b_n = b0;
64    let mut converged = false;
65    let mut iters = 0;
66
67    for it in 1..=max_iter.max(1) {
68        iters = it;
69        // q(μ): precision λ_n = (λ0 + N)·E[τ].
70        lambda_n = (lambda0 + nf) * e_tau;
71        let var_mu = 1.0 / lambda_n;
72        // q(τ): b_n = b0 + ½·E_μ[ Σ(xₙ−μ)² + λ0(μ−μ0)² ].
73        let sum_sq: f64 = data.iter().map(|&x| (x - mu_n).powi(2)).sum::<f64>() + nf * var_mu;
74        let prior_term = lambda0 * ((mu_n - mu0).powi(2) + var_mu);
75        let new_b = b0 + 0.5 * (sum_sq + prior_term);
76        let new_e_tau = a_n / new_b;
77        let delta = (new_e_tau - e_tau).abs();
78        b_n = new_b;
79        e_tau = new_e_tau;
80        if delta < tol {
81            converged = true;
82            break;
83        }
84    }
85
86    Ok(VariationalGaussian {
87        mu_n,
88        lambda_n,
89        a_n,
90        b_n,
91        n_iter: iters,
92        converged,
93    })
94}
95
96#[cfg(test)]
97mod tests {
98    use super::*;
99
100    #[test]
101    fn recovers_mean_and_precision() {
102        // Data spread around 5 with variance ~1 → E[μ]≈5, E[τ]≈1.
103        let data: Vec<f64> = (0..200)
104            .map(|i| 5.0 + ((i * 37 % 100) as f64 - 50.0) / 29.0)
105            .collect();
106        let q = fit(&data, 0.0, 1e-3, 1e-3, 1e-3, 100, 1e-10).unwrap();
107        assert!(q.converged);
108        let xbar = data.iter().sum::<f64>() / data.len() as f64;
109        // mu_n is shrunk toward the prior mean (0) by lambda0 — very close to xbar
110        // for a vague prior, but not exactly equal.
111        assert!(
112            (q.mean() - xbar).abs() < 1e-3,
113            "mean {} vs xbar {xbar}",
114            q.mean()
115        );
116        // E[τ] ≈ 1/sample_variance.
117        let s2 = data.iter().map(|&x| (x - xbar).powi(2)).sum::<f64>() / data.len() as f64;
118        assert!(
119            (q.precision_mean() - 1.0 / s2).abs() / (1.0 / s2) < 0.05,
120            "prec {}",
121            q.precision_mean()
122        );
123    }
124
125    #[test]
126    fn tighter_data_gives_higher_precision() {
127        let tight: Vec<f64> = (0..100)
128            .map(|i| 2.0 + ((i % 5) as f64 - 2.0) * 0.05)
129            .collect();
130        let loose: Vec<f64> = (0..100)
131            .map(|i| 2.0 + ((i % 5) as f64 - 2.0) * 1.0)
132            .collect();
133        let qt = fit(&tight, 0.0, 1e-3, 1e-3, 1e-3, 100, 1e-12).unwrap();
134        let ql = fit(&loose, 0.0, 1e-3, 1e-3, 1e-3, 100, 1e-12).unwrap();
135        assert!(qt.precision_mean() > ql.precision_mean());
136    }
137
138    #[test]
139    fn guards() {
140        assert_eq!(
141            fit(&[1.0], 0.0, 1.0, 1.0, 1.0, 10, 1e-6).unwrap_err(),
142            LearningError::InsufficientData
143        );
144    }
145}