1use 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#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
23#[derive(Clone, Debug, PartialEq)]
24pub struct Nmf {
25 pub w: Array2<f64>,
27 pub h: Array2<f64>,
29 pub reconstruction_err: f64,
31 pub n_iter: usize,
33}
34
35impl Nmf {
36 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 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 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 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 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 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 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 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 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}