1use crate::Df64;
7use crate::col_piv_qr::ColPivQR;
8use crate::error::Error;
9use crate::matrix::Mat;
10use crate::numeric::CustomNumeric;
11use nalgebra::{ComplexField, DMatrix, DVector, RealField};
12use num_traits::{One, ToPrimitive, Zero};
13
14#[derive(Debug, Clone)]
16pub struct SVDResult<T> {
17 pub u: DMatrix<T>,
19 pub s: DVector<T>,
21 pub v: DMatrix<T>,
23 pub rank: usize,
25}
26
27#[derive(Debug, Clone)]
29pub struct TSVDConfig<T> {
30 pub rtol: T,
32}
33
34impl<T> TSVDConfig<T> {
35 pub fn new(rtol: T) -> Self {
36 Self { rtol }
37 }
38}
39
40const SVD_MAX_SWEEPS_PER_SINGULAR_VALUE: usize = 30;
51
52fn svd_max_sweeps(nrows: usize, ncols: usize) -> usize {
54 SVD_MAX_SWEEPS_PER_SINGULAR_VALUE * nrows.min(ncols).max(1)
56}
57
58fn check_finite<T>(matrix: &DMatrix<T>) -> Result<(), Error>
62where
63 T: ComplexField + ToPrimitive + Copy,
64{
65 for col in 0..matrix.ncols() {
66 for row in 0..matrix.nrows() {
67 let value = matrix[(row, col)];
68 if !value.is_finite() {
69 return Err(Error::NonFiniteInput {
70 name: "matrix",
71 index: vec![row, col],
72 value: ToPrimitive::to_f64(&value).unwrap_or(f64::NAN),
73 });
74 }
75 }
76 }
77 Ok(())
78}
79
80#[inline]
87fn get_epsilon_for_svd<T: RealField + Copy>() -> T {
88 use std::any::TypeId;
89
90 if TypeId::of::<T>() == TypeId::of::<f64>() {
91 unsafe { std::ptr::read(&f64::EPSILON as *const f64 as *const T) }
93 } else if TypeId::of::<T>() == TypeId::of::<crate::Df64>() {
94 unsafe { std::ptr::read(&crate::Df64::EPSILON as *const crate::Df64 as *const T) }
96 } else {
97 T::from_f64(1e-15).unwrap_or(T::one() * T::from_f64(1e-15).unwrap_or(T::one()))
99 }
100}
101
102fn bounded_svd<T>(
108 matrix: &DMatrix<T>,
109 max_niter: usize,
110) -> Result<nalgebra::SVD<T, nalgebra::Dyn, nalgebra::Dyn>, Error>
111where
112 T: ComplexField + RealField + Copy,
113{
114 debug_assert!(max_niter > 0, "max_niter = 0 would not bound the SVD");
115 let eps = get_epsilon_for_svd::<T>();
121
122 matrix
125 .clone()
126 .try_svd(true, true, eps, max_niter)
127 .ok_or_else(|| Error::DecompositionFailed {
128 reason: format!("the SVD did not converge within {max_niter} iterations"),
129 })
130}
131
132fn try_svd_decompose<T>(matrix: &DMatrix<T>, rtol: f64) -> Result<SVDResult<T>, Error>
134where
135 T: ComplexField + RealField + Copy + nalgebra::RealField + ToPrimitive,
136{
137 if matrix.is_empty() {
138 return Err(Error::EmptyInput { name: "matrix" });
139 }
140 check_finite(matrix)?;
142 let svd = bounded_svd(matrix, svd_max_sweeps(matrix.nrows(), matrix.ncols()))?;
143
144 let u_matrix = svd.u.unwrap();
146 let s_vector = svd.singular_values; let v_t_matrix = svd.v_t.unwrap();
148
149 let rank = calculate_rank_from_vector(&s_vector, rtol);
152
153 let u = DMatrix::from(u_matrix.columns(0, rank));
155 let s = DVector::from(s_vector.rows(0, rank));
156 let v = DMatrix::from(v_t_matrix.rows(0, rank).transpose());
157
158 Ok(SVDResult { u, s, v, rank })
159}
160
161pub fn svd_decompose<T>(matrix: &DMatrix<T>, rtol: f64) -> SVDResult<T>
178where
179 T: ComplexField + RealField + Copy + nalgebra::RealField + ToPrimitive,
180{
181 try_svd_decompose(matrix, rtol).unwrap_or_else(|err| panic!("SVD computation failed: {err}"))
182}
183
184fn calculate_rank_from_vector<T>(singular_values: &DVector<T>, rtol: f64) -> usize
197where
198 T: RealField + Copy + ToPrimitive,
199{
200 if singular_values.is_empty() {
201 return 0;
202 }
203
204 let max_sv = singular_values[0];
206 let threshold = max_sv * T::from_f64(rtol).unwrap_or(T::zero());
207
208 let mut rank = 0;
209 for &sv in singular_values.iter() {
210 if sv > threshold {
211 rank += 1;
212 } else {
213 break;
215 }
216 }
217
218 rank
219}
220
221fn calculate_rank_from_r<T: RealField>(r_matrix: &DMatrix<T>, rtol: T) -> usize
223where
224 T: ComplexField + RealField + Copy,
225{
226 let dim = r_matrix.nrows().min(r_matrix.ncols());
227 let mut rank = dim;
228
229 let mut max_diag_abs = Zero::zero();
231 for i in 0..dim {
232 let diag_abs = ComplexField::abs(r_matrix[(i, i)]);
233 if diag_abs > max_diag_abs {
234 max_diag_abs = diag_abs;
235 }
236 }
237
238 if max_diag_abs == Zero::zero() {
240 return 0;
241 }
242
243 for i in 0..dim {
245 let diag_abs = ComplexField::abs(r_matrix[(i, i)]);
246
247 if diag_abs < rtol * max_diag_abs {
249 rank = i;
250 break;
251 }
252 }
253
254 rank
255}
256
257pub fn tsvd<T>(matrix: &DMatrix<T>, config: TSVDConfig<T>) -> Result<SVDResult<T>, Error>
279where
280 T: ComplexField
281 + RealField
282 + Copy
283 + nalgebra::RealField
284 + std::fmt::Debug
285 + ToPrimitive
286 + CustomNumeric,
287{
288 let (m, n) = matrix.shape();
289
290 if m == 0 || n == 0 {
291 return Err(Error::EmptyInput { name: "matrix" });
292 }
293
294 if !(config.rtol > Zero::zero() && config.rtol < One::one()) {
296 return Err(Error::InvalidParameter {
297 name: "rtol",
298 value: format!("{:?}", CustomNumeric::to_f64(config.rtol)),
299 reason: "must be in (0, 1)".to_string(),
300 });
301 }
302
303 check_finite(matrix)?;
306
307 let qr_rtol = Some(config.rtol.clone().modulus());
310 let qr = ColPivQR::new_with_rtol(matrix.clone(), qr_rtol);
311 let q_matrix = qr.q();
312 let r_matrix = qr.r();
313 let permutation = qr.p();
314
315 match check_finite(&r_matrix) {
319 Ok(()) => {}
320 Err(Error::NonFiniteInput { index, value, .. }) => {
321 return Err(Error::DecompositionFailed {
322 reason: format!(
323 "the R factor of the QR decomposition has the non-finite entry {value} at index {index:?}"
324 ),
325 });
326 }
327 Err(other) => return Err(other),
329 }
330
331 let qr_rank = calculate_rank_from_r(
334 &r_matrix,
335 T::from_f64_unchecked(2.0) * get_epsilon_for_svd::<T>(),
336 );
337
338 if qr_rank == 0 {
339 return Ok(SVDResult {
341 u: DMatrix::zeros(m, 0),
342 s: DVector::zeros(0),
343 v: DMatrix::zeros(n, 0),
344 rank: 0,
345 });
346 }
347
348 let r_truncated: DMatrix<T> = r_matrix.rows(0, qr_rank).into();
350 let rtol_t = config.rtol;
352 let rtol_f64 = rtol_t.to_f64();
353 let svd_result = try_svd_decompose(&r_truncated, rtol_f64)?;
354
355 if svd_result.rank == 0 {
356 return Ok(SVDResult {
358 u: DMatrix::zeros(m, 0),
359 s: DVector::zeros(0),
360 v: DMatrix::zeros(n, 0),
361 rank: 0,
362 });
363 }
364
365 let q_truncated: DMatrix<T> = q_matrix.columns(0, qr_rank).into();
368 let u_full = &q_truncated * &svd_result.u;
369
370 let mut v_full = svd_result.v.clone();
375 permutation.inv_permute_rows(&mut v_full);
376
377 let s_full = svd_result.s.clone();
379
380 Ok(SVDResult {
381 u: u_full,
382 s: s_full,
383 v: v_full,
384 rank: svd_result.rank,
385 })
386}
387
388pub fn tsvd_f64(matrix: &DMatrix<f64>, rtol: f64) -> Result<SVDResult<f64>, Error> {
390 tsvd(matrix, TSVDConfig::new(rtol))
391}
392
393pub fn tsvd_df64(matrix: &DMatrix<Df64>, rtol: Df64) -> Result<SVDResult<Df64>, Error> {
395 tsvd(matrix, TSVDConfig::new(rtol))
396}
397
398pub fn tsvd_df64_from_f64(matrix: &DMatrix<f64>, rtol: f64) -> Result<SVDResult<Df64>, Error> {
400 let matrix_df64 = DMatrix::from_fn(matrix.nrows(), matrix.ncols(), |i, j| {
401 Df64::from(matrix[(i, j)])
402 });
403 let rtol_df64 = Df64::from(rtol);
404 tsvd(&matrix_df64, TSVDConfig::new(rtol_df64))
405}
406
407pub fn compute_svd_dtensor<T: CustomNumeric + 'static>(
418 matrix: &Mat<T>,
419) -> Result<(Mat<T>, Vec<T>, Mat<T>), Error> {
420 use nalgebra::DMatrix;
421 use std::any::TypeId;
422
423 if TypeId::of::<T>() == TypeId::of::<f64>() {
425 let matrix_f64 = DMatrix::from_fn(matrix.shape().0, matrix.shape().1, |i, j| {
427 CustomNumeric::to_f64(matrix[[i, j]])
428 });
429
430 let rtol = 2.0 * f64::EPSILON;
432 let result = tsvd(&matrix_f64, TSVDConfig::new(rtol))?;
433
434 let u = Mat::<T>::from_fn([result.u.nrows(), result.u.ncols()], |idx| {
436 let [i, j] = [idx[0], idx[1]];
437 T::from_f64_unchecked(result.u[(i, j)])
438 });
439
440 let s: Vec<T> = result.s.iter().map(|x| T::from_f64_unchecked(*x)).collect();
441
442 let v = Mat::<T>::from_fn([result.v.nrows(), result.v.ncols()], |idx| {
443 let [i, j] = [idx[0], idx[1]];
444 T::from_f64_unchecked(result.v[(i, j)])
445 });
446
447 Ok((u, s, v))
448 } else if TypeId::of::<T>() == TypeId::of::<Df64>() {
449 let matrix_df64: DMatrix<Df64> =
452 DMatrix::from_fn(matrix.shape().0, matrix.shape().1, |i, j| {
453 unsafe { std::mem::transmute_copy(&matrix[[i, j]]) }
455 });
456
457 let rtol = Df64::from(2.0) * Df64::epsilon();
459 let result = tsvd_df64(&matrix_df64, rtol)?;
460
461 let u = Mat::<T>::from_fn([result.u.nrows(), result.u.ncols()], |idx| {
463 let [i, j] = [idx[0], idx[1]];
464 T::convert_from(result.u[(i, j)])
465 });
466
467 let s: Vec<T> = result.s.iter().map(|x| T::convert_from(*x)).collect();
468
469 let v = Mat::<T>::from_fn([result.v.nrows(), result.v.ncols()], |idx| {
470 let [i, j] = [idx[0], idx[1]];
471 T::convert_from(result.v[(i, j)])
472 });
473
474 Ok((u, s, v))
475 } else {
476 panic!("SVD is only implemented for f64 and Df64");
477 }
478}
479
480#[cfg(test)]
481mod tests {
482 use super::*;
483 use nalgebra::DMatrix;
484 use num_traits::cast::ToPrimitive;
485
486 #[test]
487 fn test_svd_identity_matrix() {
488 let matrix = DMatrix::<f64>::identity(3, 3);
489 let result = svd_decompose(&matrix, 1e-12);
490
491 assert_eq!(result.rank, 3);
492 assert_eq!(result.s.len(), 3);
493 assert_eq!(result.u.nrows(), 3);
494 assert_eq!(result.u.ncols(), 3);
495 assert_eq!(result.v.nrows(), 3);
496 assert_eq!(result.v.ncols(), 3);
497 }
498
499 #[test]
500 fn test_tsvd_identity_matrix() {
501 let matrix = DMatrix::<f64>::identity(3, 3);
502 let result = tsvd_f64(&matrix, 1e-12).unwrap();
503
504 assert_eq!(result.rank, 3);
505 assert_eq!(result.s.len(), 3);
506 }
507
508 #[test]
509 fn test_tsvd_rank_one() {
510 let matrix = DMatrix::<f64>::from_fn(3, 3, |i, j| (i + 1) as f64 * (j + 1) as f64);
511 let result = tsvd_f64(&matrix, 1e-12).unwrap();
512
513 assert_eq!(result.rank, 1);
514 }
515
516 #[test]
519 fn test_tsvd_empty_matrix() {
520 for (rows, cols) in [(0, 0), (0, 3), (3, 0)] {
521 let matrix = DMatrix::<f64>::zeros(rows, cols);
522 assert!(
523 matches!(
524 tsvd_f64(&matrix, 1e-12),
525 Err(Error::EmptyInput { name: "matrix" })
526 ),
527 "{rows} x {cols}"
528 );
529 }
530 }
531
532 fn matrix_with_entry(bad: f64) -> DMatrix<f64> {
534 DMatrix::<f64>::from_fn(4, 4, |i, j| {
535 if (i, j) == (1, 2) {
536 bad
537 } else {
538 1.0 / (1.0 + i as f64 + j as f64) + if i == j { 1.0 } else { 0.0 }
539 }
540 })
541 }
542
543 fn assert_non_finite_at_1_2<T>(result: Result<SVDResult<T>, Error>, bad: f64) {
544 match result {
545 Err(Error::NonFiniteInput { name, index, value }) => {
546 assert_eq!(name, "matrix");
547 assert_eq!(index, vec![1, 2]);
548 assert!(value.is_nan() == bad.is_nan() && (bad.is_nan() || value == bad));
549 }
550 Err(other) => panic!("expected NonFiniteInput, got {other:?}"),
551 Ok(_) => panic!("expected NonFiniteInput, got Ok"),
552 }
553 }
554
555 #[test]
559 fn test_tsvd_rejects_non_finite_input() {
560 for bad in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
561 let matrix = matrix_with_entry(bad);
562 assert_non_finite_at_1_2(tsvd_f64(&matrix, 1e-12), bad);
563 assert_non_finite_at_1_2(tsvd_df64_from_f64(&matrix, 1e-28), bad);
564 let matrix_df64 = matrix.map(Df64::from);
565 assert_non_finite_at_1_2(tsvd_df64(&matrix_df64, Df64::from(1e-28)), bad);
566 }
567 }
568
569 #[test]
573 fn test_tsvd_rejects_nan_tolerance() {
574 let matrix = matrix_with_entry(0.5);
575 let expected = "invalid rtol = NaN: must be in (0, 1)";
576
577 let err = tsvd_f64(&matrix, f64::NAN).unwrap_err();
578 assert!(matches!(err, Error::InvalidParameter { name: "rtol", .. }));
579 assert_eq!(err.to_string(), expected);
580
581 let matrix_df64 = matrix.map(Df64::from);
582 let err = tsvd_df64(&matrix_df64, Df64::from(f64::NAN)).unwrap_err();
583 assert_eq!(err.to_string(), expected);
584 }
585
586 #[test]
587 fn test_compute_svd_dtensor_reports_errors() {
588 let empty = Mat::<f64>::zeros([0, 3]);
589 assert_eq!(
590 compute_svd_dtensor(&empty).unwrap_err(),
591 Error::EmptyInput { name: "matrix" }
592 );
593 let nan = Mat::<f64>::from_fn([2, 2], |idx| {
594 if idx[0] == 1 && idx[1] == 0 {
595 f64::NAN
596 } else {
597 1.0
598 }
599 });
600 assert!(matches!(
601 compute_svd_dtensor(&nan),
602 Err(Error::NonFiniteInput { name: "matrix", .. })
603 ));
604 let nan_df64 = Mat::<Df64>::from_fn([2, 2], |idx| {
605 Df64::from(if idx[0] == 1 && idx[1] == 0 {
606 f64::NAN
607 } else {
608 1.0
609 })
610 });
611 assert!(matches!(
612 compute_svd_dtensor(&nan_df64),
613 Err(Error::NonFiniteInput { name: "matrix", .. })
614 ));
615 }
616
617 #[test]
622 fn test_tsvd_reports_overflow_of_the_r_factor_as_decomposition_failure() {
623 let matrix = DMatrix::<f64>::from_column_slice(2, 1, &[f64::MAX, f64::MAX]);
624 match tsvd_f64(&matrix, 1e-12) {
625 Err(Error::DecompositionFailed { reason }) => assert!(
626 reason.starts_with("the R factor of the QR decomposition has the non-finite entry"),
627 "{reason}"
628 ),
629 other => panic!("expected DecompositionFailed, got {other:?}"),
630 }
631 }
632
633 #[test]
636 fn test_tsvd_shows_an_invalid_tolerance_compactly() {
637 let matrix = matrix_with_entry(0.5);
638 let err = tsvd_f64(&matrix, 1e300).unwrap_err();
639 assert_eq!(err.to_string(), "invalid rtol = 1e300: must be in (0, 1)");
640 }
641
642 #[test]
646 fn test_bounded_svd_reports_non_convergence() {
647 let matrix = matrix_with_entry(f64::NAN);
648 assert!(matches!(
649 bounded_svd(&matrix, 50),
650 Err(Error::DecompositionFailed { reason })
651 if reason == "the SVD did not converge within 50 iterations"
652 ));
653
654 let finite = matrix_with_entry(0.5);
656 let svd = bounded_svd(&finite, svd_max_sweeps(4, 4)).unwrap();
657 let reconstructed = svd.recompose().unwrap();
658 assert!((reconstructed - &finite).norm() < 1e-14 * finite.norm());
659 }
660
661 #[test]
662 #[should_panic(expected = "matrix has the non-finite entry inf at index [1, 2]")]
663 fn test_svd_decompose_panics_on_non_finite_input() {
664 svd_decompose(&matrix_with_entry(f64::INFINITY), 1e-12);
665 }
666
667 fn create_hilbert_matrix_generic<T>(n: usize) -> DMatrix<T>
670 where
671 T: nalgebra::RealField + From<f64> + Copy + std::ops::Div<Output = T>,
672 {
673 DMatrix::from_fn(n, n, |i, j| {
674 T::one() / T::from((i + j + 1) as f64)
677 })
678 }
679
680 fn reconstruct_matrix_generic<T>(
682 u: &DMatrix<T>,
683 s: &nalgebra::DVector<T>,
684 v: &DMatrix<T>,
685 ) -> DMatrix<T>
686 where
687 T: nalgebra::RealField + Copy,
688 {
689 u * &DMatrix::from_diagonal(s) * &v.transpose()
693 }
694
695 fn frobenius_norm_generic<T>(matrix: &DMatrix<T>) -> f64
697 where
698 T: nalgebra::RealField + Copy + ToPrimitive,
699 {
700 let mut sum = 0.0;
701 for i in 0..matrix.nrows() {
702 for j in 0..matrix.ncols() {
703 let val = matrix[(i, j)].to_f64().unwrap_or(0.0);
704 sum += val * val;
705 }
706 }
707 sum.sqrt()
708 }
709
710 fn test_hilbert_reconstruction_generic<T>(n: usize, rtol: f64, expected_max_error: f64)
712 where
713 T: nalgebra::RealField
714 + From<f64>
715 + Copy
716 + ToPrimitive
717 + std::fmt::Debug
718 + crate::numeric::CustomNumeric,
719 {
720 let h = create_hilbert_matrix_generic::<T>(n);
721
722 let config = TSVDConfig::new(T::from(rtol));
724 let result = tsvd(&h, config).unwrap();
725
726 let h_reconstructed = reconstruct_matrix_generic(&result.u, &result.s, &result.v);
728
729 let error_matrix = &h - &h_reconstructed;
731 let error_norm = frobenius_norm_generic(&error_matrix);
732 let relative_error = error_norm / frobenius_norm_generic(&h);
733
734 assert!(
736 relative_error <= expected_max_error,
737 "Relative reconstruction error {} exceeds expected maximum {}",
738 relative_error,
739 expected_max_error
740 );
741 }
742
743 #[test]
744 fn test_hilbert_5x5_f64_reconstruction() {
745 test_hilbert_reconstruction_generic::<f64>(5, 1e-12, 1e-14);
746 }
747
748 #[test]
749 fn test_hilbert_5x5_df64_reconstruction() {
750 test_hilbert_reconstruction_generic::<Df64>(5, 1e-28, 1e-28);
751 }
752
753 #[test]
754 fn test_hilbert_10x10_f64_reconstruction() {
755 test_hilbert_reconstruction_generic::<f64>(10, 1e-12, 1e-12);
756 }
757
758 #[test]
759 fn test_hilbert_10x10_df64_reconstruction() {
760 test_hilbert_reconstruction_generic::<Df64>(10, 1e-28, 1e-30);
764 }
765
766 #[test]
767 fn test_hilbert_100x100_f64_reconstruction() {
768 test_hilbert_reconstruction_generic::<f64>(100, 1e-12, 1e-12);
770 }
771
772 #[test]
773 fn test_hilbert_100x100_df64_reconstruction() {
774 test_hilbert_reconstruction_generic::<Df64>(100, 1e-28, 1e-28);
776 }
777}