1use ndarray::{Array1, Array2, ArrayView2};
4use solow_core::{Error, Result};
5use solow_manifold::isomap::jacobi_symmetric;
6
7#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
9#[derive(Copy, Clone, Debug, PartialEq, Eq)]
10pub enum IcaFun {
11 LogCosh,
13 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#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
34#[derive(Clone, Debug, PartialEq)]
35pub struct FastIca {
36 pub components: Array2<f64>,
39 pub whitening: Array2<f64>,
41 pub mean: Array1<f64>,
43 pub n_components: usize,
45 pub n_iter: usize,
47}
48
49impl FastIca {
50 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 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 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 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 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 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 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 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 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 pub fn transform(&self, x: ArrayView2<'_, f64>) -> Array2<f64> {
182 let (n, d) = (x.nrows(), x.ncols());
183 let mut out = Array2::<f64>::zeros((n, self.n_components));
185 for i in 0..n {
186 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 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 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 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 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 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 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 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}