qualia_core_db/solvers/learning/regression/
bayesian.rs1use 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#[derive(Debug, Clone)]
19pub struct BayesianLinear {
20 pub mean: Vec<f64>, pub cov: Vec<f64>, 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 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 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 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 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 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 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 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 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 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 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 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]); let (_, var_far) = bl.predict_row(&[100.0]); assert!(
175 var_far > var_center,
176 "var_far {var_far} should exceed var_center {var_center}"
177 );
178 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 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}