qualia_core_db/solvers/learning/variational/
gaussian.rs1use crate::solvers::learning::LearningError;
10
11#[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 pub fn mean(&self) -> f64 {
25 self.mu_n
26 }
27 pub fn precision_mean(&self) -> f64 {
29 self.a_n / self.b_n
30 }
31 pub fn variance_mean(&self) -> f64 {
33 self.b_n / self.a_n
34 }
35}
36
37pub 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 let mu_n = (lambda0 * mu0 + nf * xbar) / (lambda0 + nf);
57 let a_n = a0 + (nf + 1.0) / 2.0;
58
59 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 lambda_n = (lambda0 + nf) * e_tau;
71 let var_mu = 1.0 / lambda_n;
72 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 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 assert!(
112 (q.mean() - xbar).abs() < 1e-3,
113 "mean {} vs xbar {xbar}",
114 q.mean()
115 );
116 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}