Skip to main content

solow_decomposition/
ica.rs

1//! FastICA (Hyvärinen 1999) with symmetric decorrelation.
2
3use ndarray::{Array1, Array2, ArrayView2};
4use solow_core::{Error, Result};
5use solow_manifold::isomap::jacobi_symmetric;
6
7/// Non-Gaussianity contrast function.
8#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
9#[derive(Copy, Clone, Debug, PartialEq, Eq)]
10pub enum IcaFun {
11    /// `G(u) = (1/α) log cosh(α u)` (default, α = 1).
12    LogCosh,
13    /// `G(u) = -exp(-u²/2)`.
14    Exp,
15}
16
17impl IcaFun {
18    fn g_and_gprime(&self, u: f64) -> (f64, f64) {
19        match self {
20            IcaFun::LogCosh => {
21                let t = u.tanh();
22                (t, 1.0 - t * t)
23            }
24            IcaFun::Exp => {
25                let e = (-0.5 * u * u).exp();
26                (u * e, (1.0 - u * u) * e)
27            }
28        }
29    }
30}
31
32/// Fitted FastICA decomposition.
33#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
34#[derive(Clone, Debug, PartialEq)]
35pub struct FastIca {
36    /// Un-mixing matrix `W` of shape `(n_components, d)` (applied to
37    /// the centered data).
38    pub components: Array2<f64>,
39    /// Whitening matrix `K` of shape `(n_components, d)`.
40    pub whitening: Array2<f64>,
41    /// Data mean subtracted before whitening.
42    pub mean: Array1<f64>,
43    /// Number of components.
44    pub n_components: usize,
45    /// Number of iterations run.
46    pub n_iter: usize,
47}
48
49impl FastIca {
50    /// Fit with defaults (`fun = LogCosh`, `max_iter = 200`, `tol = 1e-4`).
51    pub fn fit(x: ArrayView2<'_, f64>, n_components: usize, seed: u64) -> Result<Self> {
52        Self::fit_with(x, n_components, IcaFun::LogCosh, 200, 1e-4, seed)
53    }
54
55    /// Full-configuration fit.
56    pub fn fit_with(
57        x: ArrayView2<'_, f64>,
58        n_components: usize,
59        fun: IcaFun,
60        max_iter: usize,
61        tol: f64,
62        seed: u64,
63    ) -> Result<Self> {
64        if x.nrows() < 2 || x.ncols() == 0 {
65            return Err(Error::Value(
66                "FastIca::fit_with: need ≥ 2 samples and ≥ 1 feature".into(),
67            ));
68        }
69        if n_components == 0 || n_components > x.ncols() {
70            return Err(Error::Value(format!(
71                "FastIca::fit_with: n_components in [1, d] (got {n_components})"
72            )));
73        }
74        let (n, d) = (x.nrows(), x.ncols());
75        // Centre.
76        let mut mean = Array1::<f64>::zeros(d);
77        for j in 0..d {
78            mean[j] = x.column(j).iter().sum::<f64>() / n as f64;
79        }
80        let mut xc = Array2::<f64>::zeros((n, d));
81        for i in 0..n {
82            for j in 0..d {
83                xc[[i, j]] = x[[i, j]] - mean[j];
84            }
85        }
86        // Whitening: eigendecompose XᵀX/n.
87        let mut cov = Array2::<f64>::zeros((d, d));
88        for i in 0..d {
89            for j in 0..d {
90                let mut s = 0.0_f64;
91                for k in 0..n {
92                    s += xc[[k, i]] * xc[[k, j]];
93                }
94                cov[[i, j]] = s / n as f64;
95            }
96        }
97        let (eigvals, eigvecs) = jacobi_symmetric(&cov, 300, 1e-12);
98        let mut order: Vec<usize> = (0..d).collect();
99        order.sort_by(|&a, &b| eigvals[b].partial_cmp(&eigvals[a]).unwrap());
100        // Build whitening K of shape (n_components, d): rows = eigvecs[:, i] / sqrt(λ_i).
101        let mut k = Array2::<f64>::zeros((n_components, d));
102        for c in 0..n_components {
103            let idx = order[c];
104            let s = eigvals[idx].max(1e-12).sqrt();
105            for j in 0..d {
106                k[[c, j]] = eigvecs[[j, idx]] / s;
107            }
108        }
109        // Whitened data X1 = xc · Kᵀ  (shape n × n_components).
110        let mut x1 = Array2::<f64>::zeros((n, n_components));
111        for i in 0..n {
112            for c in 0..n_components {
113                let mut s = 0.0_f64;
114                for j in 0..d {
115                    s += xc[[i, j]] * k[[c, j]];
116                }
117                x1[[i, c]] = s;
118            }
119        }
120        // Initialise W (n_components × n_components) with a deterministic random matrix.
121        let mut state = seed.wrapping_add(0xABCD_EF01_2345_6789);
122        let mut w = Array2::<f64>::zeros((n_components, n_components));
123        for i in 0..n_components {
124            for j in 0..n_components {
125                w[[i, j]] = uniform_symmetric(&mut state, 1.0);
126            }
127        }
128        symmetric_decorrelate(&mut w);
129        // FastICA symmetric loop.
130        let mut n_iter_used = 0usize;
131        for it in 0..max_iter {
132            n_iter_used = it + 1;
133            let mut new_w = Array2::<f64>::zeros(w.dim());
134            for c in 0..n_components {
135                let (mut mean_g, mut mean_gp) = (0.0_f64, 0.0_f64);
136                let mut acc = Array1::<f64>::zeros(n_components);
137                for i in 0..n {
138                    let mut u = 0.0_f64;
139                    for kk in 0..n_components {
140                        u += w[[c, kk]] * x1[[i, kk]];
141                    }
142                    let (g, gp) = fun.g_and_gprime(u);
143                    mean_g += g;
144                    mean_gp += gp;
145                    for kk in 0..n_components {
146                        acc[kk] += x1[[i, kk]] * g;
147                    }
148                }
149                mean_g /= n as f64;
150                mean_gp /= n as f64;
151                for kk in 0..n_components {
152                    new_w[[c, kk]] = acc[kk] / n as f64 - mean_gp * w[[c, kk]];
153                }
154                let _ = mean_g;
155            }
156            symmetric_decorrelate(&mut new_w);
157            // Convergence.
158            let mut delta = 0.0_f64;
159            for c in 0..n_components {
160                let mut s = 0.0_f64;
161                for kk in 0..n_components {
162                    s += new_w[[c, kk]] * w[[c, kk]];
163                }
164                delta = delta.max((s.abs() - 1.0).abs());
165            }
166            w = new_w;
167            if delta < tol {
168                break;
169            }
170        }
171        Ok(Self {
172            components: w,
173            whitening: k,
174            mean,
175            n_components,
176            n_iter: n_iter_used,
177        })
178    }
179
180    /// Project `x` into the independent-component space.
181    pub fn transform(&self, x: ArrayView2<'_, f64>) -> Array2<f64> {
182        let (n, d) = (x.nrows(), x.ncols());
183        // (x - mean) · Kᵀ · Wᵀ
184        let mut out = Array2::<f64>::zeros((n, self.n_components));
185        for i in 0..n {
186            // Centre.
187            let mut centered = vec![0.0_f64; d];
188            for j in 0..d {
189                centered[j] = x[[i, j]] - self.mean[j];
190            }
191            // Whiten.
192            let mut w1 = vec![0.0_f64; self.n_components];
193            for c in 0..self.n_components {
194                let mut s = 0.0_f64;
195                for j in 0..d {
196                    s += centered[j] * self.whitening[[c, j]];
197                }
198                w1[c] = s;
199            }
200            // Un-mix.
201            for c in 0..self.n_components {
202                let mut s = 0.0_f64;
203                for kk in 0..self.n_components {
204                    s += self.components[[c, kk]] * w1[kk];
205                }
206                out[[i, c]] = s;
207            }
208        }
209        out
210    }
211}
212
213fn symmetric_decorrelate(w: &mut Array2<f64>) {
214    let n = w.nrows();
215    // A = W · Wᵀ
216    let mut a = Array2::<f64>::zeros((n, n));
217    for i in 0..n {
218        for j in 0..n {
219            let mut s = 0.0_f64;
220            for k in 0..n {
221                s += w[[i, k]] * w[[j, k]];
222            }
223            a[[i, j]] = s;
224        }
225    }
226    let (eigvals, eigvecs) = jacobi_symmetric(&a, 300, 1e-12);
227    // A^{-1/2} = V · diag(1/√λ) · Vᵀ.
228    let mut d_half = Array2::<f64>::zeros((n, n));
229    for i in 0..n {
230        let inv = if eigvals[i] > 1e-12 {
231            1.0 / eigvals[i].sqrt()
232        } else {
233            0.0
234        };
235        d_half[[i, i]] = inv;
236    }
237    // vt = eigvecs.T · W, then result = eigvecs · d_half · vt
238    // Efficient: (eigvecs · d_half · eigvecsᵀ) · W
239    let mut tmp = Array2::<f64>::zeros((n, n));
240    for i in 0..n {
241        for j in 0..n {
242            let mut s = 0.0_f64;
243            for k in 0..n {
244                s += eigvecs[[i, k]] * d_half[[k, k]] * eigvecs[[j, k]];
245            }
246            tmp[[i, j]] = s;
247        }
248    }
249    // W_new = tmp · W.
250    let mut new_w = Array2::<f64>::zeros((n, n));
251    for i in 0..n {
252        for j in 0..n {
253            let mut s = 0.0_f64;
254            for k in 0..n {
255                s += tmp[[i, k]] * w[[k, j]];
256            }
257            new_w[[i, j]] = s;
258        }
259    }
260    *w = new_w;
261}
262
263fn lcg_next(state: &mut u64) -> u64 {
264    *state = state
265        .wrapping_mul(6_364_136_223_846_793_005)
266        .wrapping_add(1_442_695_040_888_963_407);
267    *state
268}
269
270fn uniform_f64(state: &mut u64) -> f64 {
271    (lcg_next(state) >> 11) as f64 / ((1u64 << 53) as f64)
272}
273
274fn uniform_symmetric(state: &mut u64, scale: f64) -> f64 {
275    (uniform_f64(state) - 0.5) * 2.0 * scale
276}
277
278#[cfg(test)]
279mod tests {
280    use super::*;
281    use ndarray::array;
282
283    #[test]
284    fn fast_ica_runs_on_small_input() {
285        // Mixed two sources; ICA should return finite output.
286        let x = array![
287            [1.0, 0.5],
288            [0.2, 0.9],
289            [-1.0, -0.4],
290            [0.3, -0.6],
291            [-0.7, 0.1],
292            [0.6, -0.9]
293        ];
294        let ica = FastIca::fit_with(x.view(), 2, IcaFun::LogCosh, 100, 1e-4, 7).unwrap();
295        let s = ica.transform(x.view());
296        assert_eq!(s.dim(), (6, 2));
297        for v in s.iter() {
298            assert!(v.is_finite());
299        }
300    }
301}