oxicuda_cs/sbl/
sparse_bayesian.rs1use crate::error::{CsError, CsResult};
12use crate::linalg::cholesky::{cholesky_factor, cholesky_solve};
13use crate::linalg::{mat_t_vec, mat_vec, norm2};
14use crate::sbl::SblResult;
15
16pub fn sparse_bayesian(
18 phi: &[f64],
19 m: usize,
20 n: usize,
21 y: &[f64],
22 max_iter: usize,
23 tol: f64,
24) -> CsResult<SblResult> {
25 if phi.len() != m * n {
26 return Err(CsError::ShapeMismatch {
27 expected: vec![m, n],
28 got: vec![phi.len()],
29 });
30 }
31 if y.len() != m {
32 return Err(CsError::DimensionMismatch { a: y.len(), b: m });
33 }
34 let mut gamma = vec![1.0_f64; n];
35 let mut sigma2 = (norm2(y) / (m as f64).sqrt()).max(1.0e-6);
36 sigma2 *= sigma2;
37 let mut mu = vec![0.0_f64; n];
38 let mut iter = 0usize;
39 for _ in 0..max_iter {
40 let mut g = vec![0.0_f64; n * n];
42 let inv_sigma2 = 1.0 / sigma2;
43 for k in 0..m {
44 let row = k * n;
45 for i in 0..n {
46 let pki = phi[row + i] * inv_sigma2;
47 for j in 0..n {
48 g[i * n + j] += pki * phi[row + j];
49 }
50 }
51 }
52 for j in 0..n {
53 let g_inv = if gamma[j] > 1.0e-300 {
54 1.0 / gamma[j]
55 } else {
56 1.0e12
57 };
58 g[j * n + j] += g_inv;
59 }
60 let l = cholesky_factor(&g, n)?;
61 let phi_t_y = mat_t_vec(phi, m, n, y)?;
63 let mut rhs = vec![0.0_f64; n];
64 for j in 0..n {
65 rhs[j] = phi_t_y[j] * inv_sigma2;
66 }
67 let mu_new = cholesky_solve(&l, n, &rhs)?;
68 let mut sigma_diag = vec![0.0_f64; n];
71 for j in 0..n {
72 let mut ej = vec![0.0_f64; n];
73 ej[j] = 1.0;
74 let s_j = cholesky_solve(&l, n, &ej)?;
75 sigma_diag[j] = s_j[j];
76 }
77 for j in 0..n {
79 gamma[j] = mu_new[j] * mu_new[j] + sigma_diag[j];
80 }
81 let phi_mu = mat_vec(phi, m, n, &mu_new)?;
83 let mut resid_sq = 0.0_f64;
84 for i in 0..m {
85 let d = y[i] - phi_mu[i];
86 resid_sq += d * d;
87 }
88 let mut denom = m as f64;
89 for j in 0..n {
90 let frac = if gamma[j] > 1.0e-300 {
91 sigma_diag[j] / gamma[j]
92 } else {
93 1.0
94 };
95 denom -= 1.0 - frac;
96 }
97 if denom > 1.0e-6 {
98 sigma2 = (resid_sq / denom).max(1.0e-12);
99 }
100 let mut delta = 0.0_f64;
102 for j in 0..n {
103 let d = mu_new[j] - mu[j];
104 delta += d * d;
105 }
106 mu = mu_new;
107 iter += 1;
108 if delta.sqrt() / norm2(&mu).max(1.0e-300) < tol {
109 break;
110 }
111 }
112 Ok(SblResult {
113 x: mu,
114 gamma,
115 sigma2,
116 iterations: iter,
117 })
118}
119
120#[cfg(test)]
121mod tests {
122 use super::*;
123
124 #[test]
125 fn sbl_runs() {
126 let phi = vec![1.0, 0.0, 0.0, 1.0];
127 let y = vec![1.0, 0.0];
128 let r = sparse_bayesian(&phi, 2, 2, &y, 30, 1.0e-7).expect("ok");
129 assert!(r.x[0].abs() > r.x[1].abs());
131 }
132}