Skip to main content

qualia_core_db/solvers/learning/regression/
bayesian.rs

1//! Bayesian linear regression (PRML ch 3.3) — a conjugate Gaussian model that
2//! returns a **predictive distribution** (mean + variance), not just a point
3//! estimate. This is the mission-aligned payoff: a model that can say "ŷ, and here
4//! is how sure" — calibrated uncertainty, with the predictive variance widening
5//! away from the data.
6//!
7//! Prior `w ~ N(0, α⁻¹I)`, noise precision `β = 1/σ²`. Posterior (PRML 3.53–3.54):
8//! `S_N⁻¹ = αI + β ΦᵀΦ`, `m_N = β S_N Φᵀy`. Predictive (3.58–3.59):
9//! `mean = m_Nᵀφ`, `var = 1/β + φᵀ S_N φ`. The `k×k` solves reuse
10//! `linear_algebra::cholesky` (no new solver). Kernel-class `DenseLinear`.
11
12use crate::solvers::learning::LearningError;
13use crate::solvers::linear_algebra::cholesky::{cholesky_factor, cholesky_solve};
14use crate::solvers::linear_algebra::gemm::{gemm, matvec, Transpose};
15
16/// A fitted Bayesian linear model: the posterior over weights (`mean` = `m_N`,
17/// `cov` = `S_N`) and the noise precision `beta`.
18#[derive(Debug, Clone)]
19pub struct BayesianLinear {
20    pub mean: Vec<f64>, // posterior mean m_N (k)
21    pub cov: Vec<f64>,  // posterior covariance S_N (k×k)
22    pub beta: f64,
23    pub fit_intercept: bool,
24    pub p: usize,
25}
26
27fn design_row(x_row: &[f64], fit_intercept: bool, out: &mut [f64]) {
28    if fit_intercept {
29        out[0] = 1.0;
30        out[1..].copy_from_slice(x_row);
31    } else {
32        out.copy_from_slice(x_row);
33    }
34}
35
36impl BayesianLinear {
37    /// Fit with prior precision `alpha` (> 0) and noise precision `beta` (> 0).
38    /// Fails closed on shape mismatch / non-positive hyper-parameters.
39    pub fn fit(
40        x: &[f64],
41        y: &[f64],
42        n: usize,
43        p: usize,
44        alpha: f64,
45        beta: f64,
46        fit_intercept: bool,
47    ) -> Result<Self, LearningError> {
48        if n == 0 || p == 0 || x.len() != n * p || y.len() != n {
49            return Err(LearningError::InvalidDimension);
50        }
51        if !(alpha > 0.0) || !(beta > 0.0) {
52            return Err(LearningError::InsufficientData);
53        }
54        let k = p + usize::from(fit_intercept);
55
56        // Design matrix Φ (n × k).
57        let mut phi = vec![0.0; n * k];
58        for i in 0..n {
59            design_row(
60                &x[i * p..(i + 1) * p],
61                fit_intercept,
62                &mut phi[i * k..(i + 1) * k],
63            );
64        }
65
66        // A = αI + β ΦᵀΦ.
67        let mut a = vec![0.0; k * k];
68        gemm(
69            Transpose::Yes,
70            Transpose::No,
71            k,
72            k,
73            n,
74            beta,
75            &phi,
76            &phi,
77            0.0,
78            &mut a,
79        )?;
80        for j in 0..k {
81            a[j * k + j] += alpha;
82        }
83        // b = β Φᵀy.
84        let mut b = vec![0.0; k];
85        matvec(Transpose::Yes, k, n, &phi, y, &mut b)?;
86        for v in b.iter_mut() {
87            *v *= beta;
88        }
89
90        // Posterior mean: solve A m_N = b.  Posterior covariance S_N = A⁻¹.
91        let mut l = vec![0.0; k * k];
92        cholesky_factor(k, &a, &mut l).map_err(|_| LearningError::Singular)?;
93        let mut mean = vec![0.0; k];
94        cholesky_solve(k, &l, &b, &mut mean)?;
95        // S_N columns via A⁻¹ = solving A·s_j = e_j.
96        let mut cov = vec![0.0; k * k];
97        let mut ej = vec![0.0; k];
98        let mut sj = vec![0.0; k];
99        for j in 0..k {
100            ej.iter_mut().for_each(|v| *v = 0.0);
101            ej[j] = 1.0;
102            cholesky_solve(k, &l, &ej, &mut sj)?;
103            for i in 0..k {
104                cov[i * k + j] = sj[i];
105            }
106        }
107
108        Ok(Self {
109            mean,
110            cov,
111            beta,
112            fit_intercept,
113            p,
114        })
115    }
116
117    /// Predictive distribution at one row: `(mean, variance)`.
118    /// `variance = 1/β + φᵀ S_N φ` — the noise floor plus the model uncertainty,
119    /// which grows away from the training data.
120    pub fn predict_row(&self, x_row: &[f64]) -> (f64, f64) {
121        let k = self.mean.len();
122        let mut phi = vec![0.0; k];
123        design_row(x_row, self.fit_intercept, &mut phi);
124        let mean: f64 = phi.iter().zip(&self.mean).map(|(p, m)| p * m).sum();
125        // φᵀ S_N φ.
126        let mut sphi = vec![0.0; k];
127        for i in 0..k {
128            sphi[i] = (0..k).map(|j| self.cov[i * k + j] * phi[j]).sum();
129        }
130        let model_var: f64 = phi.iter().zip(&sphi).map(|(p, s)| p * s).sum();
131        (mean, 1.0 / self.beta + model_var)
132    }
133
134    /// Predictive `(mean, variance)` for each row of a row-major `m × p` matrix.
135    pub fn predict(&self, x: &[f64], m: usize) -> Vec<(f64, f64)> {
136        (0..m)
137            .map(|i| self.predict_row(&x[i * self.p..(i + 1) * self.p]))
138            .collect()
139    }
140}
141
142#[cfg(test)]
143mod tests {
144    use super::*;
145    use crate::solvers::learning::regression::linear;
146
147    #[test]
148    fn posterior_mean_approaches_ols_with_weak_prior() {
149        // y = 1 + 2x; a weak prior + precise likelihood → posterior mean ≈ OLS.
150        let x: Vec<f64> = (0..20).map(|i| i as f64 / 2.0).collect();
151        let y: Vec<f64> = x.iter().map(|&xi| 1.0 + 2.0 * xi).collect();
152        let bl = BayesianLinear::fit(&x, &y, 20, 1, 1e-6, 1e6, true).unwrap();
153        let ols = linear::fit(&x, &y, 20, 1, true).unwrap();
154        assert!(
155            (bl.mean[0] - ols.coefficients[0]).abs() < 1e-2,
156            "intercept {}",
157            bl.mean[0]
158        );
159        assert!(
160            (bl.mean[1] - ols.coefficients[1]).abs() < 1e-2,
161            "slope {}",
162            bl.mean[1]
163        );
164    }
165
166    #[test]
167    fn predictive_variance_grows_away_from_data() {
168        // Train on x∈[0,10]; predictive std should be larger when extrapolating.
169        let x: Vec<f64> = (0..11).map(|i| i as f64).collect();
170        let y: Vec<f64> = x.iter().map(|&xi| 0.5 * xi).collect();
171        let bl = BayesianLinear::fit(&x, &y, 11, 1, 1.0, 25.0, true).unwrap();
172        let (_, var_center) = bl.predict_row(&[5.0]); // middle of the data
173        let (_, var_far) = bl.predict_row(&[100.0]); // far extrapolation
174        assert!(
175            var_far > var_center,
176            "var_far {var_far} should exceed var_center {var_center}"
177        );
178        // Variance never drops below the noise floor 1/β.
179        assert!(var_center >= 1.0 / 25.0 - 1e-12);
180    }
181
182    #[test]
183    fn stronger_prior_shrinks_weights() {
184        let x: Vec<f64> = (0..10).map(|i| i as f64).collect();
185        let y: Vec<f64> = x.iter().map(|&xi| 3.0 * xi).collect();
186        let weak = BayesianLinear::fit(&x, &y, 10, 1, 0.01, 10.0, false).unwrap();
187        let strong = BayesianLinear::fit(&x, &y, 10, 1, 100.0, 10.0, false).unwrap();
188        // A stronger zero-mean prior pulls the slope toward 0.
189        assert!(strong.mean[0].abs() < weak.mean[0].abs());
190    }
191
192    #[test]
193    fn guards() {
194        assert_eq!(
195            BayesianLinear::fit(&[1.0, 2.0], &[1.0], 2, 1, 1.0, 1.0, true).unwrap_err(),
196            LearningError::InvalidDimension
197        );
198        assert_eq!(
199            BayesianLinear::fit(&[1.0, 2.0], &[1.0, 2.0], 2, 1, 0.0, 1.0, true).unwrap_err(),
200            LearningError::InsufficientData
201        );
202    }
203}