1use ferrolearn_core::error::FerroError;
79use ferrolearn_core::pipeline::{FittedPipelineTransformer, PipelineTransformer};
80use ferrolearn_core::traits::{Fit, Transform};
81use ndarray::{Array1, Array2};
82use num_traits::Float;
83use rand::SeedableRng;
84use rand_distr::{Distribution, Uniform};
85use std::any::TypeId;
86
87fn reject_non_finite<F: Float>(x: &Array2<F>) -> Result<(), FerroError> {
98 if x.iter().any(|v| !v.is_finite()) {
99 return Err(FerroError::InvalidParameter {
100 name: "X".into(),
101 reason: "Input X contains NaN or infinity.".into(),
102 });
103 }
104 Ok(())
105}
106
107#[derive(Debug, Clone, Copy, PartialEq, Eq)]
113pub enum NMFSolver {
114 MultiplicativeUpdate,
116 CoordinateDescent,
118}
119
120#[derive(Debug, Clone, Copy, PartialEq, Eq)]
122pub enum NMFInit {
123 Random,
125 Nndsvd,
127}
128
129#[derive(Debug, Clone)]
139pub struct NMF<F> {
140 n_components: usize,
142 max_iter: usize,
144 tol: f64,
146 solver: NMFSolver,
148 init: NMFInit,
150 random_state: Option<u64>,
152 _marker: std::marker::PhantomData<F>,
153}
154
155impl<F: Float + Send + Sync + 'static> NMF<F> {
156 #[must_use]
161 pub fn new(n_components: usize) -> Self {
162 Self {
163 n_components,
164 max_iter: 200,
165 tol: 1e-4,
166 solver: NMFSolver::MultiplicativeUpdate,
167 init: NMFInit::Random,
168 random_state: None,
169 _marker: std::marker::PhantomData,
170 }
171 }
172
173 #[must_use]
175 pub fn with_max_iter(mut self, max_iter: usize) -> Self {
176 self.max_iter = max_iter;
177 self
178 }
179
180 #[must_use]
182 pub fn with_tol(mut self, tol: f64) -> Self {
183 self.tol = tol;
184 self
185 }
186
187 #[must_use]
189 pub fn with_solver(mut self, solver: NMFSolver) -> Self {
190 self.solver = solver;
191 self
192 }
193
194 #[must_use]
196 pub fn with_init(mut self, init: NMFInit) -> Self {
197 self.init = init;
198 self
199 }
200
201 #[must_use]
203 pub fn with_random_state(mut self, seed: u64) -> Self {
204 self.random_state = Some(seed);
205 self
206 }
207
208 #[must_use]
210 pub fn n_components(&self) -> usize {
211 self.n_components
212 }
213
214 #[must_use]
216 pub fn max_iter(&self) -> usize {
217 self.max_iter
218 }
219
220 #[must_use]
222 pub fn tol(&self) -> f64 {
223 self.tol
224 }
225
226 #[must_use]
228 pub fn solver(&self) -> NMFSolver {
229 self.solver
230 }
231
232 #[must_use]
234 pub fn init(&self) -> NMFInit {
235 self.init
236 }
237
238 #[must_use]
240 pub fn random_state(&self) -> Option<u64> {
241 self.random_state
242 }
243}
244
245#[derive(Debug, Clone)]
254pub struct FittedNMF<F> {
255 components_: Array2<F>,
257 reconstruction_err_: F,
259 n_iter_: usize,
261}
262
263impl<F: Float + Send + Sync + 'static> FittedNMF<F> {
264 #[must_use]
266 pub fn components(&self) -> &Array2<F> {
267 &self.components_
268 }
269
270 #[must_use]
272 pub fn reconstruction_err(&self) -> F {
273 self.reconstruction_err_
274 }
275
276 #[must_use]
278 pub fn n_iter(&self) -> usize {
279 self.n_iter_
280 }
281
282 pub fn inverse_transform(&self, w: &Array2<F>) -> Result<Array2<F>, FerroError> {
291 let n_components = self.components_.nrows();
292 if w.ncols() != n_components {
293 return Err(FerroError::ShapeMismatch {
294 expected: vec![w.nrows(), n_components],
295 actual: vec![w.nrows(), w.ncols()],
296 context: "FittedNMF::inverse_transform".into(),
297 });
298 }
299 Ok(w.dot(&self.components_))
300 }
301}
302
303fn reconstruction_error<F: Float + 'static>(x: &Array2<F>, w: &Array2<F>, h: &Array2<F>) -> F {
309 let wh = w.dot(h);
310 let mut err = F::zero();
311 for (a, b) in x.iter().zip(wh.iter()) {
312 let diff = *a - *b;
313 err = err + diff * diff;
314 }
315 err.sqrt()
316}
317
318fn eps<F: Float>() -> F {
320 F::from(1e-12).unwrap_or_else(F::epsilon)
321}
322
323fn init_random<F: Float>(
325 n_samples: usize,
326 n_features: usize,
327 n_components: usize,
328 seed: u64,
329) -> (Array2<F>, Array2<F>) {
330 let mut rng: rand::rngs::StdRng = SeedableRng::seed_from_u64(seed);
331 let uniform = Uniform::new(0.0f64, 1.0f64).unwrap();
332
333 let mut w = Array2::<F>::zeros((n_samples, n_components));
334 for elem in &mut w {
335 *elem = F::from(uniform.sample(&mut rng)).unwrap_or_else(F::zero) + eps::<F>();
336 }
337
338 let mut h = Array2::<F>::zeros((n_components, n_features));
339 for elem in &mut h {
340 *elem = F::from(uniform.sample(&mut rng)).unwrap_or_else(F::zero) + eps::<F>();
341 }
342
343 (w, h)
344}
345
346fn ndarray_to_ferray_f64(a: &Array2<f64>) -> Result<ferray::Array<f64, ferray::Ix2>, FerroError> {
348 let (m, n) = a.dim();
349 let data: Vec<f64> = a.iter().copied().collect();
350 ferray::Array::<f64, ferray::Ix2>::from_vec(ferray::Ix2::new([m, n]), data).map_err(|e| {
351 FerroError::NumericalInstability {
352 message: format!("ferray array construction failed: {e}"),
353 }
354 })
355}
356
357fn ndarray_to_ferray_f32(a: &Array2<f32>) -> Result<ferray::Array<f32, ferray::Ix2>, FerroError> {
359 let (m, n) = a.dim();
360 let data: Vec<f32> = a.iter().copied().collect();
361 ferray::Array::<f32, ferray::Ix2>::from_vec(ferray::Ix2::new([m, n]), data).map_err(|e| {
362 FerroError::NumericalInstability {
363 message: format!("ferray array construction failed: {e}"),
364 }
365 })
366}
367
368#[allow(
375 clippy::type_complexity,
376 reason = "(U, s, Vt) is the standard thin-SVD triple, not worth a named struct"
377)]
378fn svd_full_f64(a: &Array2<f64>) -> Result<(Array2<f64>, Array1<f64>, Array2<f64>), FerroError> {
379 let fa = ndarray_to_ferray_f64(a)?;
380 let (u, s, vt) =
381 ferray::linalg::svd_lapack(&fa, false).map_err(|e| FerroError::NumericalInstability {
382 message: format!("ferray svd_lapack (gesdd) failed: {e}"),
383 })?;
384 Ok((u.into_ndarray(), s.into_ndarray(), vt.into_ndarray()))
385}
386
387#[allow(
389 clippy::type_complexity,
390 reason = "(U, s, Vt) is the standard thin-SVD triple, not worth a named struct"
391)]
392fn svd_full_f32(a: &Array2<f32>) -> Result<(Array2<f32>, Array1<f32>, Array2<f32>), FerroError> {
393 let fa = ndarray_to_ferray_f32(a)?;
394 let (u, s, vt) =
395 ferray::linalg::svd_lapack(&fa, false).map_err(|e| FerroError::NumericalInstability {
396 message: format!("ferray svd_lapack (gesdd) failed: {e}"),
397 })?;
398 Ok((u.into_ndarray(), s.into_ndarray(), vt.into_ndarray()))
399}
400
401fn svd_flip_u_based<F: Float>(u: &mut Array2<F>, vt: &mut Array2<F>) {
409 let (m, k) = u.dim();
410 let n = vt.ncols();
411 for col in 0..k {
412 let mut max_abs = F::neg_infinity();
413 let mut max_row = 0usize;
414 for row in 0..m {
415 let av = u[[row, col]].abs();
416 if av > max_abs {
417 max_abs = av;
418 max_row = row;
419 }
420 }
421 if u[[max_row, col]] < F::zero() {
422 for row in 0..m {
423 u[[row, col]] = -u[[row, col]];
424 }
425 for j in 0..n {
426 vt[[col, j]] = -vt[[col, j]];
427 }
428 }
429 }
430}
431
432#[allow(
438 clippy::type_complexity,
439 reason = "(U, s, Vt) is the standard thin-SVD triple, not worth a named struct"
440)]
441fn nndsvd_svd<F: Float + Send + Sync + 'static>(
442 x: &Array2<F>,
443 n_components: usize,
444) -> Result<(Array2<F>, Array1<F>, Array2<F>), FerroError> {
445 let (n_samples, n_features) = x.dim();
446 let k = n_components;
447
448 let full = if TypeId::of::<F>() == TypeId::of::<f64>() {
450 let x_f64: &Array2<f64> = unsafe { &*(std::ptr::from_ref(x).cast::<Array2<f64>>()) };
453 let (u, s, vt) = svd_full_f64(x_f64)?;
454 let u_f: Array2<F> = unsafe { std::mem::transmute_copy::<Array2<f64>, Array2<F>>(&u) };
455 let s_f: Array1<F> = unsafe { std::mem::transmute_copy::<Array1<f64>, Array1<F>>(&s) };
456 let vt_f: Array2<F> = unsafe { std::mem::transmute_copy::<Array2<f64>, Array2<F>>(&vt) };
457 std::mem::forget(u);
458 std::mem::forget(s);
459 std::mem::forget(vt);
460 Some((u_f, s_f, vt_f))
461 } else if TypeId::of::<F>() == TypeId::of::<f32>() {
462 let x_f32: &Array2<f32> = unsafe { &*(std::ptr::from_ref(x).cast::<Array2<f32>>()) };
465 let (u, s, vt) = svd_full_f32(x_f32)?;
466 let u_f: Array2<F> = unsafe { std::mem::transmute_copy::<Array2<f32>, Array2<F>>(&u) };
467 let s_f: Array1<F> = unsafe { std::mem::transmute_copy::<Array1<f32>, Array1<F>>(&s) };
468 let vt_f: Array2<F> = unsafe { std::mem::transmute_copy::<Array2<f32>, Array2<F>>(&vt) };
469 std::mem::forget(u);
470 std::mem::forget(s);
471 std::mem::forget(vt);
472 Some((u_f, s_f, vt_f))
473 } else {
474 None
475 };
476
477 let (mut u, s, mut vt) = match full {
478 Some((u_full, s_full, vt_full)) => {
479 let u_t = u_full.slice(ndarray::s![.., ..k]).to_owned();
481 let s_t = s_full.slice(ndarray::s![..k]).to_owned();
482 let vt_t = vt_full.slice(ndarray::s![..k, ..]).to_owned();
483 (u_t, s_t, vt_t)
484 }
485 None => {
486 let max_iter = n_features * n_features * 100 + 1000;
488 let xtx = x.t().dot(x);
489 let (eigenvalues, eigenvectors) = jacobi_eigen_symmetric(&xtx, max_iter)?;
490 let mut indices: Vec<usize> = (0..n_features).collect();
491 indices.sort_by(|&a, &b| {
492 eigenvalues[b]
493 .partial_cmp(&eigenvalues[a])
494 .unwrap_or(std::cmp::Ordering::Equal)
495 });
496 let mut s_t = Array1::<F>::zeros(k);
497 let mut vt_t = Array2::<F>::zeros((k, n_features));
498 for (row, &idx) in indices.iter().take(k).enumerate() {
499 let sv = eigenvalues[idx].max(F::zero()).sqrt();
500 s_t[row] = sv;
501 for j in 0..n_features {
502 vt_t[[row, j]] = eigenvectors[[j, idx]];
503 }
504 }
505 let mut u_t = Array2::<F>::zeros((n_samples, k));
507 for row in 0..k {
508 let sv = s_t[row];
509 if sv <= eps::<F>() {
510 continue;
511 }
512 for i in 0..n_samples {
513 let mut acc = F::zero();
514 for j in 0..n_features {
515 acc = acc + x[[i, j]] * vt_t[[row, j]];
516 }
517 u_t[[i, row]] = acc / sv;
518 }
519 }
520 (u_t, s_t, vt_t)
521 }
522 };
523
524 svd_flip_u_based(&mut u, &mut vt);
525 Ok((u, s, vt))
526}
527
528fn init_nndsvd<F: Float + Send + Sync + 'static>(
539 x: &Array2<F>,
540 n_components: usize,
541) -> Result<(Array2<F>, Array2<F>), FerroError> {
542 let (n_samples, n_features) = x.dim();
543 let (u, s, vt) = nndsvd_svd(x, n_components)?;
544
545 let machine_eps = if TypeId::of::<F>() == TypeId::of::<f32>() {
547 F::from(f32::EPSILON).unwrap_or_else(F::epsilon)
548 } else {
549 F::from(f64::EPSILON).unwrap_or_else(F::epsilon)
550 };
551
552 let mut w = Array2::<F>::zeros((n_samples, n_components));
553 let mut h = Array2::<F>::zeros((n_components, n_features));
554
555 let sqrt_s0 = s[0].max(F::zero()).sqrt();
557 for i in 0..n_samples {
558 w[[i, 0]] = sqrt_s0 * u[[i, 0]].abs();
559 }
560 for j in 0..n_features {
561 h[[0, j]] = sqrt_s0 * vt[[0, j]].abs();
562 }
563
564 for comp in 1..n_components {
565 let mut xp_nrm_sq = F::zero();
567 let mut xn_nrm_sq = F::zero();
568 for i in 0..n_samples {
569 let val = u[[i, comp]];
570 if val > F::zero() {
571 xp_nrm_sq = xp_nrm_sq + val * val;
572 } else {
573 xn_nrm_sq = xn_nrm_sq + val * val;
574 }
575 }
576 let mut yp_nrm_sq = F::zero();
577 let mut yn_nrm_sq = F::zero();
578 for j in 0..n_features {
579 let val = vt[[comp, j]];
580 if val > F::zero() {
581 yp_nrm_sq = yp_nrm_sq + val * val;
582 } else {
583 yn_nrm_sq = yn_nrm_sq + val * val;
584 }
585 }
586 let x_p_nrm = xp_nrm_sq.sqrt();
587 let y_p_nrm = yp_nrm_sq.sqrt();
588 let x_n_nrm = xn_nrm_sq.sqrt();
589 let y_n_nrm = yn_nrm_sq.sqrt();
590
591 let m_p = x_p_nrm * y_p_nrm;
592 let m_n = x_n_nrm * y_n_nrm;
593
594 let (use_positive, nrm_u, nrm_v, sigma) = if m_p > m_n {
596 (true, x_p_nrm, y_p_nrm, m_p)
597 } else {
598 (false, x_n_nrm, y_n_nrm, m_n)
599 };
600
601 let lbd = (s[comp] * sigma).max(F::zero()).sqrt();
602
603 for i in 0..n_samples {
605 let val = u[[i, comp]];
606 let part = if use_positive {
607 if val > F::zero() { val } else { F::zero() }
608 } else if val < F::zero() {
609 -val
610 } else {
611 F::zero()
612 };
613 w[[i, comp]] = if nrm_u > F::zero() {
614 lbd * (part / nrm_u)
615 } else {
616 F::zero()
617 };
618 }
619 for j in 0..n_features {
621 let val = vt[[comp, j]];
622 let part = if use_positive {
623 if val > F::zero() { val } else { F::zero() }
624 } else if val < F::zero() {
625 -val
626 } else {
627 F::zero()
628 };
629 h[[comp, j]] = if nrm_v > F::zero() {
630 lbd * (part / nrm_v)
631 } else {
632 F::zero()
633 };
634 }
635 }
636
637 for val in &mut w {
639 if *val < machine_eps {
640 *val = F::zero();
641 }
642 }
643 for val in &mut h {
644 if *val < machine_eps {
645 *val = F::zero();
646 }
647 }
648
649 Ok((w, h))
650}
651
652fn jacobi_eigen_symmetric<F: Float + Send + Sync + 'static>(
657 a: &Array2<F>,
658 max_iter: usize,
659) -> Result<(Array1<F>, Array2<F>), FerroError> {
660 let n = a.nrows();
661 if n == 0 {
662 return Ok((Array1::zeros(0), Array2::zeros((0, 0))));
663 }
664 if n == 1 {
665 let eigenvalues = Array1::from_vec(vec![a[[0, 0]]]);
666 let eigenvectors = Array2::from_shape_vec((1, 1), vec![F::one()]).unwrap();
667 return Ok((eigenvalues, eigenvectors));
668 }
669
670 let mut mat = a.to_owned();
671 let mut v = Array2::<F>::zeros((n, n));
672 for i in 0..n {
673 v[[i, i]] = F::one();
674 }
675
676 let tol = F::from(1e-12).unwrap_or_else(F::epsilon);
677
678 for _iteration in 0..max_iter {
679 let mut max_off = F::zero();
680 let mut p = 0;
681 let mut q = 1;
682 for i in 0..n {
683 for j in (i + 1)..n {
684 let val = mat[[i, j]].abs();
685 if val > max_off {
686 max_off = val;
687 p = i;
688 q = j;
689 }
690 }
691 }
692
693 if max_off < tol {
694 let eigenvalues = Array1::from_shape_fn(n, |i| mat[[i, i]]);
695 return Ok((eigenvalues, v));
696 }
697
698 let app = mat[[p, p]];
699 let aqq = mat[[q, q]];
700 let apq = mat[[p, q]];
701
702 let theta = if (app - aqq).abs() < tol {
703 F::from(std::f64::consts::FRAC_PI_4).unwrap_or_else(F::one)
704 } else {
705 let tau = (aqq - app) / (F::from(2.0).unwrap() * apq);
706 let t = if tau >= F::zero() {
707 F::one() / (tau.abs() + (F::one() + tau * tau).sqrt())
708 } else {
709 -F::one() / (tau.abs() + (F::one() + tau * tau).sqrt())
710 };
711 t.atan()
712 };
713
714 let c = theta.cos();
715 let s = theta.sin();
716
717 let mut new_mat = mat.clone();
718 for i in 0..n {
719 if i != p && i != q {
720 let mip = mat[[i, p]];
721 let miq = mat[[i, q]];
722 new_mat[[i, p]] = c * mip - s * miq;
723 new_mat[[p, i]] = new_mat[[i, p]];
724 new_mat[[i, q]] = s * mip + c * miq;
725 new_mat[[q, i]] = new_mat[[i, q]];
726 }
727 }
728
729 new_mat[[p, p]] = c * c * app - F::from(2.0).unwrap() * s * c * apq + s * s * aqq;
730 new_mat[[q, q]] = s * s * app + F::from(2.0).unwrap() * s * c * apq + c * c * aqq;
731 new_mat[[p, q]] = F::zero();
732 new_mat[[q, p]] = F::zero();
733
734 mat = new_mat;
735
736 for i in 0..n {
737 let vip = v[[i, p]];
738 let viq = v[[i, q]];
739 v[[i, p]] = c * vip - s * viq;
740 v[[i, q]] = s * vip + c * viq;
741 }
742 }
743
744 Err(FerroError::ConvergenceFailure {
745 iterations: max_iter,
746 message: "Jacobi eigendecomposition did not converge in NMF NNDSVD init".into(),
747 })
748}
749
750fn solve_multiplicative_update<F: Float + 'static>(
756 x: &Array2<F>,
757 w: &mut Array2<F>,
758 h: &mut Array2<F>,
759 max_iter: usize,
760 tol: f64,
761) -> usize {
762 let tol_f = F::from(tol).unwrap_or_else(F::epsilon);
763 let epsilon = eps::<F>();
764 let mut prev_err = reconstruction_error(x, w, h);
765
766 for iteration in 0..max_iter {
767 let wt = w.t();
769 let numerator_h = wt.dot(x);
770 let denominator_h = wt.dot(&*w).dot(&*h);
771
772 for (h_val, (num, den)) in h
773 .iter_mut()
774 .zip(numerator_h.iter().zip(denominator_h.iter()))
775 {
776 *h_val = *h_val * (*num / (*den + epsilon));
777 }
778
779 let ht = h.t();
781 let numerator_w = x.dot(&ht);
782 let denominator_w = w.dot(&*h).dot(&ht);
783
784 for (w_val, (num, den)) in w
785 .iter_mut()
786 .zip(numerator_w.iter().zip(denominator_w.iter()))
787 {
788 *w_val = *w_val * (*num / (*den + epsilon));
789 }
790
791 let err = reconstruction_error(x, w, h);
793 if (prev_err - err).abs() < tol_f {
794 return iteration + 1;
795 }
796 prev_err = err;
797 }
798
799 max_iter
800}
801
802fn update_cd_sweep<F: Float + 'static>(x: &Array2<F>, w: &mut Array2<F>, ht: &Array2<F>) -> F {
821 let n_components = ht.ncols();
822 let n_samples = w.nrows();
823
824 let hht = ht.t().dot(ht);
826 let xht = x.dot(ht);
828
829 let mut violation = F::zero();
830 for t in 0..n_components {
831 let hess = hht[[t, t]];
832 for i in 0..n_samples {
833 let mut grad = -xht[[i, t]];
835 for r in 0..n_components {
836 grad = grad + hht[[t, r]] * w[[i, r]];
837 }
838 let pg = if w[[i, t]] == F::zero() {
840 grad.min(F::zero())
841 } else {
842 grad
843 };
844 violation = violation + pg.abs();
845 if hess != F::zero() {
846 w[[i, t]] = (w[[i, t]] - grad / hess).max(F::zero());
847 }
848 }
849 }
850 violation
851}
852
853fn solve_coordinate_descent<F: Float + 'static>(
864 x: &Array2<F>,
865 w: &mut Array2<F>,
866 h: &mut Array2<F>,
867 max_iter: usize,
868 tol: f64,
869 update_h: bool,
870) -> usize {
871 let tol_f = F::from(tol).unwrap_or_else(F::epsilon);
872
873 let mut ht = h.t().to_owned();
875 let xt = x.t().to_owned();
876
877 let mut violation_init = F::zero();
878 let mut n_iter = if max_iter == 0 { 0 } else { 1 };
879
880 for iteration in 1..=max_iter {
881 n_iter = iteration;
882 let mut violation = F::zero();
883
884 violation = violation + update_cd_sweep(x, w, &ht);
886
887 if update_h {
889 violation = violation + update_cd_sweep(&xt, &mut ht, w);
890 }
891
892 if iteration == 1 {
893 violation_init = violation;
894 }
895 if violation_init == F::zero() {
896 break;
897 }
898 if violation / violation_init <= tol_f {
899 break;
900 }
901 }
902
903 if update_h {
905 for k in 0..h.nrows() {
906 for j in 0..h.ncols() {
907 h[[k, j]] = ht[[j, k]];
908 }
909 }
910 }
911
912 n_iter
913}
914
915impl<F: Float + Send + Sync + 'static> Fit<Array2<F>, ()> for NMF<F> {
920 type Fitted = FittedNMF<F>;
921 type Error = FerroError;
922
923 fn fit(&self, x: &Array2<F>, _y: &()) -> Result<FittedNMF<F>, FerroError> {
933 let (n_samples, n_features) = x.dim();
934
935 if self.n_components == 0 {
936 return Err(FerroError::InvalidParameter {
937 name: "n_components".into(),
938 reason: "must be at least 1".into(),
939 });
940 }
941 if n_samples == 0 {
942 return Err(FerroError::InsufficientSamples {
943 required: 1,
944 actual: 0,
945 context: "NMF::fit".into(),
946 });
947 }
948 if self.n_components > n_samples.min(n_features) {
949 return Err(FerroError::InvalidParameter {
950 name: "n_components".into(),
951 reason: format!(
952 "n_components ({}) exceeds min(n_samples, n_features) = {}",
953 self.n_components,
954 n_samples.min(n_features)
955 ),
956 });
957 }
958
959 reject_non_finite(x)?;
967
968 for &val in x {
970 if val < F::zero() {
971 return Err(FerroError::InvalidParameter {
972 name: "X".into(),
973 reason: "NMF requires all entries in X to be non-negative".into(),
974 });
975 }
976 }
977
978 let seed = self.random_state.unwrap_or(0);
979
980 let (mut w, mut h) = match self.init {
982 NMFInit::Random => init_random(n_samples, n_features, self.n_components, seed),
983 NMFInit::Nndsvd => init_nndsvd(x, self.n_components)?,
984 };
985
986 let n_iter = match self.solver {
988 NMFSolver::MultiplicativeUpdate => {
989 solve_multiplicative_update(x, &mut w, &mut h, self.max_iter, self.tol)
990 }
991 NMFSolver::CoordinateDescent => {
992 solve_coordinate_descent(x, &mut w, &mut h, self.max_iter, self.tol, true)
993 }
994 };
995
996 let reconstruction_err = reconstruction_error(x, &w, &h);
997
998 Ok(FittedNMF {
999 components_: h,
1000 reconstruction_err_: reconstruction_err,
1001 n_iter_: n_iter,
1002 })
1003 }
1004}
1005
1006impl<F: Float + Send + Sync + 'static> Transform<Array2<F>> for FittedNMF<F> {
1007 type Output = Array2<F>;
1008 type Error = FerroError;
1009
1010 fn transform(&self, x: &Array2<F>) -> Result<Array2<F>, FerroError> {
1030 let n_features = self.components_.ncols();
1031 if x.ncols() != n_features {
1032 return Err(FerroError::ShapeMismatch {
1033 expected: vec![x.nrows(), n_features],
1034 actual: vec![x.nrows(), x.ncols()],
1035 context: "FittedNMF::transform".into(),
1036 });
1037 }
1038
1039 reject_non_finite(x)?;
1045
1046 for &val in x {
1048 if val < F::zero() {
1049 return Err(FerroError::InvalidParameter {
1050 name: "X".into(),
1051 reason: "NMF requires all entries in X to be non-negative".into(),
1052 });
1053 }
1054 }
1055
1056 let n_samples = x.nrows();
1057 let n_components = self.components_.nrows();
1058
1059 let mut w = Array2::<F>::zeros((n_samples, n_components));
1062 let mut h = self.components_.clone();
1063
1064 solve_coordinate_descent(x, &mut w, &mut h, 200, 1e-4, false);
1066
1067 Ok(w)
1068 }
1069}
1070
1071impl<F: Float + Send + Sync + 'static> PipelineTransformer<F> for NMF<F> {
1076 fn fit_pipeline(
1084 &self,
1085 x: &Array2<F>,
1086 _y: &Array1<F>,
1087 ) -> Result<Box<dyn FittedPipelineTransformer<F>>, FerroError> {
1088 let fitted = self.fit(x, &())?;
1089 Ok(Box::new(fitted))
1090 }
1091}
1092
1093impl<F: Float + Send + Sync + 'static> FittedPipelineTransformer<F> for FittedNMF<F> {
1094 fn transform_pipeline(&self, x: &Array2<F>) -> Result<Array2<F>, FerroError> {
1100 self.transform(x)
1101 }
1102}
1103
1104#[cfg(test)]
1109mod tests {
1110 use super::*;
1111 use approx::assert_abs_diff_eq;
1112 use ndarray::array;
1113
1114 fn small_dataset() -> Array2<f64> {
1116 array![
1117 [1.0, 2.0, 3.0],
1118 [4.0, 5.0, 6.0],
1119 [7.0, 8.0, 9.0],
1120 [10.0, 11.0, 12.0],
1121 ]
1122 }
1123
1124 fn medium_dataset() -> Array2<f64> {
1126 array![
1127 [5.0, 3.0, 0.0, 1.0],
1128 [4.0, 0.0, 0.0, 1.0],
1129 [1.0, 1.0, 0.0, 5.0],
1130 [1.0, 0.0, 0.0, 4.0],
1131 [0.0, 1.0, 5.0, 4.0],
1132 [0.0, 0.0, 4.0, 3.0],
1133 ]
1134 }
1135
1136 #[test]
1137 fn test_nmf_basic_fit() {
1138 let nmf = NMF::<f64>::new(2).with_random_state(42);
1139 let x = small_dataset();
1140 let fitted = nmf.fit(&x, &()).unwrap();
1141 assert_eq!(fitted.components().dim(), (2, 3));
1142 }
1143
1144 #[test]
1145 fn test_nmf_components_non_negative() {
1146 let nmf = NMF::<f64>::new(2).with_random_state(42);
1147 let x = small_dataset();
1148 let fitted = nmf.fit(&x, &()).unwrap();
1149 for &val in fitted.components() {
1150 assert!(
1151 val >= 0.0,
1152 "component value should be non-negative, got {val}"
1153 );
1154 }
1155 }
1156
1157 #[test]
1158 fn test_nmf_transform_dimensions() {
1159 let nmf = NMF::<f64>::new(2).with_random_state(42);
1160 let x = small_dataset();
1161 let fitted = nmf.fit(&x, &()).unwrap();
1162 let projected = fitted.transform(&x).unwrap();
1163 assert_eq!(projected.dim(), (4, 2));
1164 }
1165
1166 #[test]
1167 fn test_nmf_transform_non_negative() {
1168 let nmf = NMF::<f64>::new(2).with_random_state(42);
1169 let x = small_dataset();
1170 let fitted = nmf.fit(&x, &()).unwrap();
1171 let projected = fitted.transform(&x).unwrap();
1172 for &val in &projected {
1173 assert!(val >= 0.0, "W value should be non-negative, got {val}");
1174 }
1175 }
1176
1177 #[test]
1178 fn test_nmf_reconstruction_error_decreases() {
1179 let nmf_few = NMF::<f64>::new(2).with_random_state(42).with_max_iter(10);
1180 let nmf_many = NMF::<f64>::new(2).with_random_state(42).with_max_iter(200);
1181 let x = small_dataset();
1182 let fitted_few = nmf_few.fit(&x, &()).unwrap();
1183 let fitted_many = nmf_many.fit(&x, &()).unwrap();
1184 assert!(
1185 fitted_many.reconstruction_err() <= fitted_few.reconstruction_err() + 1e-6,
1186 "more iterations should reduce error: few={}, many={}",
1187 fitted_few.reconstruction_err(),
1188 fitted_many.reconstruction_err()
1189 );
1190 }
1191
1192 #[test]
1193 fn test_nmf_reconstruction_error_positive() {
1194 let nmf = NMF::<f64>::new(2).with_random_state(42);
1195 let x = small_dataset();
1196 let fitted = nmf.fit(&x, &()).unwrap();
1197 assert!(fitted.reconstruction_err() >= 0.0);
1198 }
1199
1200 #[test]
1201 fn test_nmf_coordinate_descent_solver() {
1202 let nmf = NMF::<f64>::new(2)
1203 .with_solver(NMFSolver::CoordinateDescent)
1204 .with_random_state(42);
1205 let x = medium_dataset();
1206 let fitted = nmf.fit(&x, &()).unwrap();
1207 assert_eq!(fitted.components().dim(), (2, 4));
1208 for &val in fitted.components() {
1209 assert!(val >= 0.0, "CD component should be non-negative, got {val}");
1210 }
1211 }
1212
1213 #[test]
1214 fn test_nmf_nndsvd_init() {
1215 let nmf = NMF::<f64>::new(2)
1216 .with_init(NMFInit::Nndsvd)
1217 .with_random_state(42);
1218 let x = medium_dataset();
1219 let fitted = nmf.fit(&x, &()).unwrap();
1220 assert_eq!(fitted.components().dim(), (2, 4));
1221 for &val in fitted.components() {
1222 assert!(
1223 val >= 0.0,
1224 "NNDSVD component should be non-negative, got {val}"
1225 );
1226 }
1227 }
1228
1229 #[test]
1230 fn test_nmf_cd_with_nndsvd() {
1231 let nmf = NMF::<f64>::new(2)
1232 .with_solver(NMFSolver::CoordinateDescent)
1233 .with_init(NMFInit::Nndsvd)
1234 .with_random_state(42);
1235 let x = medium_dataset();
1236 let fitted = nmf.fit(&x, &()).unwrap();
1237 assert_eq!(fitted.components().dim(), (2, 4));
1238 }
1239
1240 #[test]
1241 fn test_nmf_invalid_n_components_zero() {
1242 let nmf = NMF::<f64>::new(0);
1243 let x = small_dataset();
1244 assert!(nmf.fit(&x, &()).is_err());
1245 }
1246
1247 #[test]
1248 fn test_nmf_invalid_n_components_too_large() {
1249 let nmf = NMF::<f64>::new(10);
1250 let x = small_dataset(); assert!(nmf.fit(&x, &()).is_err());
1252 }
1253
1254 #[test]
1255 fn test_nmf_negative_input_rejected() {
1256 let nmf = NMF::<f64>::new(1);
1257 let x = array![[1.0, -2.0], [3.0, 4.0]];
1258 assert!(nmf.fit(&x, &()).is_err());
1259 }
1260
1261 #[test]
1262 fn test_nmf_transform_shape_mismatch() {
1263 let nmf = NMF::<f64>::new(2).with_random_state(42);
1264 let x = small_dataset();
1265 let fitted = nmf.fit(&x, &()).unwrap();
1266 let x_bad = array![[1.0, 2.0]]; assert!(fitted.transform(&x_bad).is_err());
1268 }
1269
1270 #[test]
1271 fn test_nmf_transform_negative_rejected() {
1272 let nmf = NMF::<f64>::new(2).with_random_state(42);
1273 let x = small_dataset();
1274 let fitted = nmf.fit(&x, &()).unwrap();
1275 let x_neg = array![[1.0, -2.0, 3.0]];
1276 assert!(fitted.transform(&x_neg).is_err());
1277 }
1278
1279 #[test]
1280 fn test_nmf_reproducibility() {
1281 let nmf1 = NMF::<f64>::new(2).with_random_state(42);
1282 let nmf2 = NMF::<f64>::new(2).with_random_state(42);
1283 let x = small_dataset();
1284 let fitted1 = nmf1.fit(&x, &()).unwrap();
1285 let fitted2 = nmf2.fit(&x, &()).unwrap();
1286 for (a, b) in fitted1.components().iter().zip(fitted2.components().iter()) {
1287 assert_abs_diff_eq!(a, b, epsilon = 1e-10);
1288 }
1289 }
1290
1291 #[test]
1292 fn test_nmf_single_component() {
1293 let nmf = NMF::<f64>::new(1).with_random_state(42);
1294 let x = small_dataset();
1295 let fitted = nmf.fit(&x, &()).unwrap();
1296 assert_eq!(fitted.components().nrows(), 1);
1297 let projected = fitted.transform(&x).unwrap();
1298 assert_eq!(projected.ncols(), 1);
1299 }
1300
1301 #[test]
1302 fn test_nmf_n_iter_positive() {
1303 let nmf = NMF::<f64>::new(2).with_random_state(42);
1304 let x = small_dataset();
1305 let fitted = nmf.fit(&x, &()).unwrap();
1306 assert!(fitted.n_iter() > 0);
1307 }
1308
1309 #[test]
1310 fn test_nmf_getters() {
1311 let nmf = NMF::<f64>::new(3)
1312 .with_max_iter(100)
1313 .with_tol(1e-5)
1314 .with_solver(NMFSolver::CoordinateDescent)
1315 .with_init(NMFInit::Nndsvd)
1316 .with_random_state(99);
1317 assert_eq!(nmf.n_components(), 3);
1318 assert_eq!(nmf.max_iter(), 100);
1319 assert_abs_diff_eq!(nmf.tol(), 1e-5);
1320 assert_eq!(nmf.solver(), NMFSolver::CoordinateDescent);
1321 assert_eq!(nmf.init(), NMFInit::Nndsvd);
1322 assert_eq!(nmf.random_state(), Some(99));
1323 }
1324
1325 #[test]
1326 fn test_nmf_f32() {
1327 let nmf = NMF::<f32>::new(1).with_random_state(42);
1328 let x: Array2<f32> = array![[1.0f32, 2.0], [3.0, 4.0], [5.0, 6.0]];
1329 let fitted = nmf.fit(&x, &()).unwrap();
1330 let projected = fitted.transform(&x).unwrap();
1331 assert_eq!(projected.ncols(), 1);
1332 }
1333
1334 #[test]
1335 fn test_nmf_zero_entries() {
1336 let nmf = NMF::<f64>::new(2).with_random_state(42);
1337 let x = array![[0.0, 0.0, 1.0], [0.0, 1.0, 0.0], [1.0, 0.0, 0.0]];
1338 let fitted = nmf.fit(&x, &()).unwrap();
1339 assert_eq!(fitted.components().dim(), (2, 3));
1340 }
1341
1342 #[test]
1343 fn test_nmf_pipeline_integration() {
1344 use ferrolearn_core::pipeline::{FittedPipelineEstimator, Pipeline, PipelineEstimator};
1345 use ferrolearn_core::traits::Predict;
1346
1347 struct SumEstimator;
1348
1349 impl PipelineEstimator<f64> for SumEstimator {
1350 fn fit_pipeline(
1351 &self,
1352 _x: &Array2<f64>,
1353 _y: &Array1<f64>,
1354 ) -> Result<Box<dyn FittedPipelineEstimator<f64>>, FerroError> {
1355 Ok(Box::new(FittedSumEstimator))
1356 }
1357 }
1358
1359 struct FittedSumEstimator;
1360
1361 impl FittedPipelineEstimator<f64> for FittedSumEstimator {
1362 fn predict_pipeline(&self, x: &Array2<f64>) -> Result<Array1<f64>, FerroError> {
1363 let sums: Vec<f64> = x.rows().into_iter().map(|r| r.sum()).collect();
1364 Ok(Array1::from_vec(sums))
1365 }
1366 }
1367
1368 let pipeline = Pipeline::new()
1369 .transform_step("nmf", Box::new(NMF::<f64>::new(2).with_random_state(42)))
1370 .estimator_step("sum", Box::new(SumEstimator));
1371
1372 let x = small_dataset();
1373 let y = Array1::from_vec(vec![0.0, 1.0, 2.0, 3.0]);
1374
1375 let fitted = pipeline.fit(&x, &y).unwrap();
1376 let preds = fitted.predict(&x).unwrap();
1377 assert_eq!(preds.len(), 4);
1378 }
1379
1380 #[test]
1381 fn test_nmf_medium_dataset_mu() {
1382 let nmf = NMF::<f64>::new(3)
1383 .with_solver(NMFSolver::MultiplicativeUpdate)
1384 .with_random_state(42)
1385 .with_max_iter(500);
1386 let x = medium_dataset();
1387 let fitted = nmf.fit(&x, &()).unwrap();
1388 assert_eq!(fitted.components().dim(), (3, 4));
1389 assert!(
1391 fitted.reconstruction_err() < 10.0,
1392 "reconstruction error too large: {}",
1393 fitted.reconstruction_err()
1394 );
1395 }
1396
1397 #[test]
1398 fn test_nmf_insufficient_samples() {
1399 let nmf = NMF::<f64>::new(1);
1400 let x = Array2::<f64>::zeros((0, 3));
1401 assert!(nmf.fit(&x, &()).is_err());
1402 }
1403
1404 #[test]
1410 fn test_nmf_nndsvd_init_matches_sklearn() {
1411 let x: Array2<f64> = array![
1412 [2.744, 3.576, 3.014, 2.724, 2.118, 3.229],
1413 [2.188, 4.459, 4.818, 1.917, 3.959, 2.644],
1414 [2.84, 4.628, 0.355, 0.436, 0.101, 4.163],
1415 [3.891, 4.35, 4.893, 3.996, 2.307, 3.903],
1416 [0.591, 3.2, 0.717, 4.723, 2.609, 2.073],
1417 [1.323, 3.871, 2.281, 2.842, 0.094, 3.088],
1418 [3.06, 3.085, 4.719, 3.409, 1.798, 2.185],
1419 [3.488, 0.301, 3.334, 3.353, 1.052, 0.645],
1420 [1.577, 1.819, 2.851, 2.193, 4.942, 0.51],
1421 [1.044, 0.807, 3.266, 1.266, 2.332, 1.222],
1422 [0.795, 0.552, 3.282, 0.691, 0.983, 1.844],
1423 [4.105, 0.486, 4.19, 0.48, 4.882, 2.343],
1424 ];
1425 let result = init_nndsvd(&x, 3);
1426 assert!(
1427 result.is_ok(),
1428 "init_nndsvd should succeed on the 12x6 fixture"
1429 );
1430 let (w, h) = match result {
1431 Ok(pair) => pair,
1432 Err(_) => return,
1433 };
1434 assert_eq!(h.dim(), (3, 6));
1435 assert_eq!(w.dim(), (12, 3));
1436 let sk_h: [[f64; 6]; 3] = [
1438 [
1439 1.7741112747292909,
1440 2.0242809154746326,
1441 2.3973352370908354,
1442 1.7673562411021053,
1443 1.7161976904877514,
1444 1.750048077389621,
1445 ],
1446 [
1447 0.0,
1448 1.5897115406013913,
1449 0.0,
1450 0.41349076122979694,
1451 0.0,
1452 1.0036389906737029,
1453 ],
1454 [
1455 0.0,
1456 0.00561171974406682,
1457 0.0,
1458 1.6761257358883248,
1459 0.3407178035232749,
1460 0.0,
1461 ],
1462 ];
1463 for (k, row) in sk_h.iter().enumerate() {
1464 for (j, &expected) in row.iter().enumerate() {
1465 assert_abs_diff_eq!(h[[k, j]], expected, epsilon = 1e-9);
1466 }
1467 }
1468 let sk_w_col0_head = [1.5111518021294372, 1.7749072337282812, 1.0616211988216013];
1470 for (i, &expected) in sk_w_col0_head.iter().enumerate() {
1471 assert_abs_diff_eq!(w[[i, 0]], expected, epsilon = 1e-9);
1472 }
1473 }
1474
1475 #[test]
1476 fn test_nmf_more_components_lower_error() {
1477 let nmf1 = NMF::<f64>::new(1).with_random_state(42).with_max_iter(300);
1478 let nmf2 = NMF::<f64>::new(2).with_random_state(42).with_max_iter(300);
1479 let x = medium_dataset();
1480 let fitted1 = nmf1.fit(&x, &()).unwrap();
1481 let fitted2 = nmf2.fit(&x, &()).unwrap();
1482 assert!(
1483 fitted2.reconstruction_err() <= fitted1.reconstruction_err() + 1e-6,
1484 "more components should reduce error: 1comp={}, 2comp={}",
1485 fitted1.reconstruction_err(),
1486 fitted2.reconstruction_err()
1487 );
1488 }
1489}