Skip to main content

solow_decomposition/
nmf.rs

1//! Non-negative Matrix Factorisation with multiplicative updates
2//! (Lee & Seung 2001) under the Frobenius objective.
3//!
4//! Approximates `X ≈ W · H` with `W ∈ ℝ⁺^{n × k}` and `H ∈ ℝ⁺^{k × d}`.
5//! Multiplicative updates preserve non-negativity by construction.
6
7use ndarray::{Array2, ArrayView2};
8use solow_core::{Error, Result};
9
10fn lcg_next(state: &mut u64) -> u64 {
11    *state = state
12        .wrapping_mul(6_364_136_223_846_793_005)
13        .wrapping_add(1_442_695_040_888_963_407);
14    *state
15}
16
17fn uniform_f64(state: &mut u64) -> f64 {
18    (lcg_next(state) >> 11) as f64 / ((1u64 << 53) as f64)
19}
20
21/// Fitted NMF factors.
22#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
23#[derive(Clone, Debug, PartialEq)]
24pub struct Nmf {
25    /// `W` — the sample-side factor `(n × n_components)`.
26    pub w: Array2<f64>,
27    /// `H` — the component-side factor `(n_components × d)`.
28    pub h: Array2<f64>,
29    /// Final Frobenius reconstruction error `‖X − W·H‖_F`.
30    pub reconstruction_err: f64,
31    /// Number of iterations run.
32    pub n_iter: usize,
33}
34
35impl Nmf {
36    /// Fit with defaults (`max_iter = 200`, `tol = 1e-4`).
37    pub fn fit(x: ArrayView2<'_, f64>, n_components: usize, seed: u64) -> Result<Self> {
38        Self::fit_with(x, n_components, 200, 1e-4, seed)
39    }
40
41    /// Full-configuration fit.
42    pub fn fit_with(
43        x: ArrayView2<'_, f64>,
44        n_components: usize,
45        max_iter: usize,
46        tol: f64,
47        seed: u64,
48    ) -> Result<Self> {
49        if x.nrows() == 0 || x.ncols() == 0 {
50            return Err(Error::Value("Nmf::fit_with: x must be non-empty".into()));
51        }
52        for &v in x.iter() {
53            if v < 0.0 || !v.is_finite() {
54                return Err(Error::Value(
55                    "Nmf::fit_with: X must be non-negative and finite".into(),
56                ));
57            }
58        }
59        if n_components == 0 || n_components > x.nrows().min(x.ncols()) {
60            return Err(Error::Value(format!(
61                "Nmf::fit_with: n_components must be in [1, min(n, d)] (got {n_components})"
62            )));
63        }
64        let (n, d) = (x.nrows(), x.ncols());
65        let k = n_components;
66        // Random non-negative init.
67        let mut state = seed.wrapping_add(0x1122_3344_5566_7788);
68        let mut w = Array2::<f64>::zeros((n, k));
69        let mut h = Array2::<f64>::zeros((k, d));
70        for i in 0..n {
71            for j in 0..k {
72                w[[i, j]] = uniform_f64(&mut state);
73            }
74        }
75        for i in 0..k {
76            for j in 0..d {
77                h[[i, j]] = uniform_f64(&mut state);
78            }
79        }
80        let mut prev_err = f64::INFINITY;
81        let mut n_iter_used = 0usize;
82        for it in 0..max_iter {
83            n_iter_used = it + 1;
84            // Update H: H ← H * (Wᵀ X) / (Wᵀ W H).
85            let wt_x = matmul(&transpose(&w), &array_view_to_array(x));
86            let wt_w = matmul(&transpose(&w), &w);
87            let wt_w_h = matmul(&wt_w, &h);
88            for i in 0..k {
89                for j in 0..d {
90                    let denom = wt_w_h[[i, j]] + 1e-12;
91                    h[[i, j]] *= wt_x[[i, j]] / denom;
92                }
93            }
94            // Update W: W ← W * (X Hᵀ) / (W H Hᵀ).
95            let x_ht = matmul(&array_view_to_array(x), &transpose(&h));
96            let h_ht = matmul(&h, &transpose(&h));
97            let w_h_ht = matmul(&w, &h_ht);
98            for i in 0..n {
99                for j in 0..k {
100                    let denom = w_h_ht[[i, j]] + 1e-12;
101                    w[[i, j]] *= x_ht[[i, j]] / denom;
102                }
103            }
104            // Frobenius error.
105            let reconstr = matmul(&w, &h);
106            let mut err = 0.0_f64;
107            for i in 0..n {
108                for j in 0..d {
109                    let dd = x[[i, j]] - reconstr[[i, j]];
110                    err += dd * dd;
111                }
112            }
113            err = err.sqrt();
114            if (prev_err - err).abs() < tol {
115                return Ok(Self {
116                    w,
117                    h,
118                    reconstruction_err: err,
119                    n_iter: n_iter_used,
120                });
121            }
122            prev_err = err;
123        }
124        // Final error.
125        let reconstr = matmul(&w, &h);
126        let mut err = 0.0_f64;
127        for i in 0..n {
128            for j in 0..d {
129                let dd = x[[i, j]] - reconstr[[i, j]];
130                err += dd * dd;
131            }
132        }
133        err = err.sqrt();
134        Ok(Self {
135            w,
136            h,
137            reconstruction_err: err,
138            n_iter: n_iter_used,
139        })
140    }
141}
142
143fn transpose(m: &Array2<f64>) -> Array2<f64> {
144    let (r, c) = m.dim();
145    let mut out = Array2::<f64>::zeros((c, r));
146    for i in 0..r {
147        for j in 0..c {
148            out[[j, i]] = m[[i, j]];
149        }
150    }
151    out
152}
153
154fn matmul(a: &Array2<f64>, b: &Array2<f64>) -> Array2<f64> {
155    let (r, mid) = a.dim();
156    let (_, c) = b.dim();
157    let mut out = Array2::<f64>::zeros((r, c));
158    for i in 0..r {
159        for k in 0..mid {
160            let aik = a[[i, k]];
161            if aik == 0.0 {
162                continue;
163            }
164            for j in 0..c {
165                out[[i, j]] += aik * b[[k, j]];
166            }
167        }
168    }
169    out
170}
171
172fn array_view_to_array(x: ArrayView2<'_, f64>) -> Array2<f64> {
173    x.to_owned()
174}
175
176#[cfg(test)]
177mod tests {
178    use super::*;
179    use ndarray::array;
180
181    #[test]
182    fn nmf_reduces_reconstruction_error() {
183        // A structured non-negative matrix — NMF with k=2 should reach low error.
184        let x = array![
185            [5.0, 4.0, 0.0, 0.0],
186            [4.0, 5.0, 0.0, 1.0],
187            [0.0, 1.0, 5.0, 4.0],
188            [0.0, 0.0, 4.0, 5.0],
189        ];
190        let nmf = Nmf::fit(x.view(), 2, 42).unwrap();
191        // Reconstruction should account for most of X's Frobenius norm.
192        let total: f64 = x.iter().map(|v| v * v).sum::<f64>().sqrt();
193        assert!(
194            nmf.reconstruction_err < 0.5 * total,
195            "err = {}",
196            nmf.reconstruction_err
197        );
198    }
199}