use crate::algorithms::decomposition::{
at, count_to_f64, descending_order, jacobi_eigen, mean_center,
};
use crate::rng::SplitMix64;
#[derive(Debug, Clone)]
pub struct IcaResult {
pub sources: Vec<Vec<f64>>,
}
#[must_use]
pub fn fast_ica(
mixed: &[Vec<f64>],
n_components: usize,
max_iter: usize,
tol: f64,
seed: u64,
) -> IcaResult {
let dim = mixed.first().map_or(0, Vec::len);
let k = n_components.min(dim);
let n = mixed.len();
if dim == 0 || k == 0 || n == 0 {
return IcaResult {
sources: Vec::new(),
};
}
let whitened = whiten(mixed, dim, k);
let unmixing = symmetric_fastica(&whitened, k, max_iter, tol, seed);
let sources = project(&whitened, &unmixing, k);
IcaResult { sources }
}
fn whiten(mixed: &[Vec<f64>], dim: usize, k: usize) -> Vec<Vec<f64>> {
let (centered, _means) = mean_center(mixed, dim);
let mut cov = vec![0.0_f64; dim * dim];
for row in ¢ered {
for i in 0..dim {
let ri = row.get(i).copied().unwrap_or(0.0);
for j in 0..dim {
let rj = row.get(j).copied().unwrap_or(0.0);
if let Some(slot) = cov.get_mut(i * dim + j) {
*slot = ri.mul_add(rj, *slot);
}
}
}
}
let denom = count_to_f64(centered.len());
if denom > 0.0 {
for c in &mut cov {
*c /= denom;
}
}
let (values, vectors) = jacobi_eigen(&cov, dim);
let order = descending_order(&values);
let top: Vec<usize> = order.into_iter().take(k).collect();
centered
.iter()
.map(|row| {
top.iter()
.map(|&col| {
let eig = values.get(col).copied().unwrap_or(0.0).max(1e-12);
let projection: f64 = (0..dim)
.map(|f| row.get(f).copied().unwrap_or(0.0) * at(&vectors, dim, f, col))
.sum();
projection / eig.sqrt()
})
.collect()
})
.collect()
}
#[allow(clippy::many_single_char_names)]
fn symmetric_fastica(
whitened: &[Vec<f64>],
k: usize,
max_iter: usize,
tol: f64,
seed: u64,
) -> Vec<Vec<f64>> {
let mut rng = SplitMix64::new(seed);
let mut w: Vec<Vec<f64>> = (0..k)
.map(|_| (0..k).map(|_| rng.standard_normal()).collect())
.collect();
symmetric_decorrelate(&mut w, k);
let n = count_to_f64(whitened.len()).max(1.0);
for _ in 0..max_iter {
let mut next = vec![vec![0.0_f64; k]; k];
for (c, wc) in w.iter().enumerate() {
let mut expectation = vec![0.0_f64; k];
let mut mean_gprime = 0.0_f64;
for row in whitened {
let u: f64 = row.iter().zip(wc).map(|(&x, &wi)| x * wi).sum();
let g = u.tanh();
let gprime = g.mul_add(-g, 1.0);
for (e, &x) in expectation.iter_mut().zip(row) {
*e = x.mul_add(g, *e);
}
mean_gprime += gprime;
}
mean_gprime /= n;
if let Some(target) = next.get_mut(c) {
for (slot, (&e, &wi)) in target.iter_mut().zip(expectation.iter().zip(wc)) {
*slot = (-mean_gprime).mul_add(wi, e / n);
}
}
}
symmetric_decorrelate(&mut next, k);
let delta = max_abs_change(&w, &next, k);
w = next;
if delta < tol {
break;
}
}
w
}
fn symmetric_decorrelate(w: &mut [Vec<f64>], k: usize) {
let mut gram = vec![0.0_f64; k * k];
for a in 0..k {
for b in 0..k {
let dot: f64 = w.get(a).zip(w.get(b)).map_or(0.0, |(wa, wb)| {
wa.iter().zip(wb).map(|(&x, &y)| x * y).sum()
});
if let Some(slot) = gram.get_mut(a * k + b) {
*slot = dot;
}
}
}
let (values, vectors) = jacobi_eigen(&gram, k);
let mut inv_sqrt = vec![0.0_f64; k * k];
for i in 0..k {
for j in 0..k {
let mut acc = 0.0_f64;
for (e, &lambda) in values.iter().enumerate().take(k) {
let safe = lambda.max(1e-12).sqrt();
acc += at(&vectors, k, i, e) * at(&vectors, k, j, e) / safe;
}
put(&mut inv_sqrt, k, i, j, acc);
}
}
let original: Vec<Vec<f64>> = w.iter().map(Clone::clone).collect();
for (a, row) in w.iter_mut().enumerate().take(k) {
for (b, slot) in row.iter_mut().enumerate().take(k) {
*slot = (0..k)
.map(|m| {
at(&inv_sqrt, k, a, m)
* original
.get(m)
.and_then(|r| r.get(b))
.copied()
.unwrap_or(0.0)
})
.sum();
}
}
}
fn put(matrix: &mut [f64], n: usize, i: usize, j: usize, value: f64) {
if let Some(slot) = matrix.get_mut(i * n + j) {
*slot = value;
}
}
fn max_abs_change(a: &[Vec<f64>], b: &[Vec<f64>], k: usize) -> f64 {
let mut worst = 0.0_f64;
for row in 0..k {
for col in 0..k {
let av = a.get(row).and_then(|r| r.get(col)).copied().unwrap_or(0.0);
let bv = b.get(row).and_then(|r| r.get(col)).copied().unwrap_or(0.0);
worst = worst.max((av.abs() - bv.abs()).abs());
}
}
worst
}
fn project(whitened: &[Vec<f64>], unmixing: &[Vec<f64>], k: usize) -> Vec<Vec<f64>> {
let mut sources: Vec<Vec<f64>> = whitened
.iter()
.map(|row| {
(0..k)
.map(|c| {
unmixing
.get(c)
.map_or(0.0, |w| row.iter().zip(w).map(|(&x, &wi)| x * wi).sum())
})
.collect()
})
.collect();
let n = count_to_f64(sources.len()).max(1.0);
for c in 0..k {
let mean: f64 = sources
.iter()
.map(|r| r.get(c).copied().unwrap_or(0.0))
.sum::<f64>()
/ n;
let var: f64 = sources
.iter()
.map(|r| {
let d = r.get(c).copied().unwrap_or(0.0) - mean;
d * d
})
.sum::<f64>()
/ n;
let std = var.sqrt().max(1e-12);
for row in &mut sources {
if let Some(slot) = row.get_mut(c) {
*slot = (*slot - mean) / std;
}
}
}
sources
}
#[cfg(test)]
mod tests {
use super::*;
fn mixture() -> Vec<Vec<f64>> {
(0..200)
.map(|i| {
let t = f64::from(i) * 0.2;
let s1 = (2.0 * t).sin();
let s2 = if (3.0 * t).sin() >= 0.0 { 1.0 } else { -1.0 };
vec![0.6_f64.mul_add(s2, s1), 0.4_f64.mul_add(s1, 1.2 * s2)]
})
.collect()
}
#[test]
fn recovers_two_unit_variance_sources() {
let r = fast_ica(&mixture(), 2, 300, 1e-5, 3);
assert_eq!(r.sources.len(), 200, "source row count");
let n = count_to_f64(r.sources.len());
for c in 0..2 {
let var: f64 = r
.sources
.iter()
.map(|row| {
let v = row.get(c).copied().unwrap_or(0.0);
v * v
})
.sum::<f64>()
/ n;
assert!((var - 1.0).abs() < 1e-6, "source {c} variance was {var}");
}
}
#[test]
fn deterministic_for_fixed_seed() {
let a = fast_ica(&mixture(), 2, 300, 1e-5, 9).sources;
let b = fast_ica(&mixture(), 2, 300, 1e-5, 9).sources;
let first_a = a.first().and_then(|r| r.first()).copied().unwrap_or(0.0);
let first_b = b.first().and_then(|r| r.first()).copied().unwrap_or(0.0);
assert_eq!(first_a.to_bits(), first_b.to_bits(), "non-deterministic");
}
#[test]
fn empty_input_is_empty() {
assert!(fast_ica(&[], 2, 10, 1e-4, 0).sources.is_empty());
}
}