Skip to main content

oxicuda_cs/sbl/
sparse_bayesian.rs

1//! Sparse Bayesian Learning (Tipping 2001) for sparse recovery.
2//!
3//! Bayesian model: y = Φ x + ε, with prior `p(x_j) = N(0, γ_j)` and noise `ε ~ N(0, σ² I)`.
4//! Expectation-maximisation alternates:
5//!   - Posterior mean μ = Σ Φᵀ y / σ², Σ = (diag(1/γ) + ΦᵀΦ/σ²)⁻¹.
6//!   - γ_j ← μ_j² + Σ_{jj}.
7//!   - σ² ← ||y − Φμ||² / (m − Σ_j (1 − γ_j⁻¹ Σ_{jj})).
8//!
9//! Variables with γ_j → 0 are pruned (sparsity emerges naturally).
10
11use 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
16/// Sparse Bayesian Learning EM iteration.
17pub 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        // Build (diag(1/γ) + ΦᵀΦ/σ²).
41        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        // μ = Σ Φᵀ y / σ²; rhs = Φᵀ y / σ².
62        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        // Σ_{jj}: solve L L^T s_j = e_j; we read diag(Σ) = diag(L^{-T} L^{-1}).
69        // Cheap approximation: compute the diagonal by solving for each unit vector.
70        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        // γ_j update.
78        for j in 0..n {
79            gamma[j] = mu_new[j] * mu_new[j] + sigma_diag[j];
80        }
81        // σ² update.
82        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        // Convergence check.
101        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        // Should recover sparse x = [1, 0].
130        assert!(r.x[0].abs() > r.x[1].abs());
131    }
132}