Skip to main content

sparse_ir_basis/
tsvd.rs

1//! High-precision truncated SVD implementation using nalgebra
2//!
3//! This module provides QR + SVD based truncated SVD decomposition
4//! with support for extended precision arithmetic.
5
6use 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/// Result of SVD decomposition
15#[derive(Debug, Clone)]
16pub struct SVDResult<T> {
17    /// Left singular vectors (m × rank)
18    pub u: DMatrix<T>,
19    /// Singular values (rank)
20    pub s: DVector<T>,
21    /// Right singular vectors (n × rank)
22    pub v: DMatrix<T>,
23    /// Effective rank
24    pub rank: usize,
25}
26
27/// Configuration for TSVD computation
28#[derive(Debug, Clone)]
29pub struct TSVDConfig<T> {
30    /// Relative tolerance for rank determination
31    pub rtol: T,
32}
33
34impl<T> TSVDConfig<T> {
35    pub fn new(rtol: T) -> Self {
36        Self { rtol }
37    }
38}
39
40/// Maximum number of implicit-shift QR sweeps per singular value in the SVD
41///
42/// nalgebra's `try_svd` iterates until convergence when its `max_niter` is 0,
43/// and a NaN or an infinity never converges: the SVD then loops forever. The
44/// limit turns non-convergence into [`Error::DecompositionFailed`]. The SVDs of
45/// the SVE matrices of the logistic and regularized Bose kernels (Λ from 1 to
46/// 1e5, f64 and Df64) converge within 1.4 sweeps per singular value; LAPACK's
47/// `dbdsqr` allows `MAXITR = 6` sweeps per singular value (the bound is
48/// convention-matched with it, no code is derived from it). 30 leaves a wide
49/// margin and still bounds the work on a matrix that does not converge.
50const SVD_MAX_SWEEPS_PER_SINGULAR_VALUE: usize = 30;
51
52/// Iteration limit of the SVD of an `nrows × ncols` matrix
53fn svd_max_sweeps(nrows: usize, ncols: usize) -> usize {
54    // Never 0, which nalgebra reads as "no limit"
55    SVD_MAX_SWEEPS_PER_SINGULAR_VALUE * nrows.min(ncols).max(1)
56}
57
58/// Reject a matrix with a NaN or infinite entry
59///
60/// Reports the first such entry in column-major (storage) order.
61fn 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/// Get appropriate epsilon value for SVD convergence based on type
81///
82/// Returns the machine epsilon (EPSILON constant) for the given type.
83/// This is preferred over approx::AbsDiffEq::default_epsilon() because
84/// the latter may return MIN_POSITIVE for Df64, which is too small and
85/// causes excessive iterations in SVD.
86#[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        // f64::EPSILON ≈ 2.22e-16
92        unsafe { std::ptr::read(&f64::EPSILON as *const f64 as *const T) }
93    } else if TypeId::of::<T>() == TypeId::of::<crate::Df64>() {
94        // Df64::EPSILON ≈ 2.465e-32
95        unsafe { std::ptr::read(&crate::Df64::EPSILON as *const crate::Df64 as *const T) }
96    } else {
97        // Fallback: use a reasonable default
98        T::from_f64(1e-15).unwrap_or(T::one() * T::from_f64(1e-15).unwrap_or(T::one()))
99    }
100}
101
102/// Full SVD by nalgebra's implicit-shift QR iteration, stopped after
103/// `max_niter` sweeps
104///
105/// The singular values are sorted in descending order. `max_niter` must be
106/// positive (nalgebra reads 0 as "no limit"). The matrix must not be empty.
107fn 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    // Use type-appropriate epsilon for SVD convergence
116    // For f64: f64::EPSILON (約 2.22e-16)
117    // For Df64: Df64::EPSILON (約 2.465e-32)
118    // Note: We use EPSILON (machine epsilon) instead of default_epsilon() (from approx trait)
119    // because default_epsilon() may return MIN_POSITIVE for Df64, which is too small.
120    let eps = get_epsilon_for_svd::<T>();
121
122    // Perform FULL SVD decomposition with explicit epsilon
123    // try_svd automatically sorts singular values in descending order
124    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
132/// [`svd_decompose`] returning an error instead of panicking
133fn 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    // A non-finite entry would exhaust the iteration limit; reject it early.
141    check_finite(matrix)?;
142    let svd = bounded_svd(matrix, svd_max_sweeps(matrix.nrows(), matrix.ncols()))?;
143
144    // Extract U, S, V matrices (already sorted by nalgebra)
145    let u_matrix = svd.u.unwrap();
146    let s_vector = svd.singular_values; // Sorted in descending order
147    let v_t_matrix = svd.v_t.unwrap();
148
149    // Calculate effective rank from sorted singular values
150    // Early termination is possible in rank calculation because values are sorted
151    let rank = calculate_rank_from_vector(&s_vector, rtol);
152
153    // Convert to thin SVD (truncate to effective rank)
154    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
161/// Perform SVD decomposition using nalgebra with sorted singular values
162///
163/// Note: This computes ALL singular values, not a truncated SVD.
164/// The truncation happens after SVD computation based on rtol.
165///
166/// # Arguments
167/// * `matrix` - Input matrix (m × n)
168/// * `rtol` - Relative tolerance for rank determination (used for rank calculation, not SVD convergence)
169///
170/// # Returns
171/// * `SVDResult` - Truncated SVD result with U, S, V matrices and rank
172///
173/// # Panics
174/// Panics if the matrix is empty, has a NaN or infinite entry, or the SVD
175/// iteration does not converge within its iteration limit
176/// ([`Error`] describes each case). [`tsvd`] reports these as errors.
177pub 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
184/// Calculate effective rank from sorted singular values
185///
186/// # Arguments
187/// * `singular_values` - Vector of singular values (sorted in descending order)
188/// * `rtol` - Relative tolerance for rank determination
189///
190/// # Returns
191/// * `usize` - Effective rank
192///
193/// # Note
194/// Since singular values are sorted in descending order by try_svd,
195/// this function can terminate early when a value below the threshold is found.
196fn 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    // First element is the maximum (sorted in descending order)
205    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            // Early termination: since values are sorted, all remaining values are also below threshold
214            break;
215        }
216    }
217
218    rank
219}
220
221/// Calculate rank from R matrix diagonal elements
222fn 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    // Find the maximum diagonal element
230    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 is zero, rank is zero
239    if max_diag_abs == Zero::zero() {
240        return 0;
241    }
242
243    // Check each diagonal element
244    for i in 0..dim {
245        let diag_abs = ComplexField::abs(r_matrix[(i, i)]);
246
247        // Check if the diagonal element is too small relative to the maximum diagonal element
248        if diag_abs < rtol * max_diag_abs {
249            rank = i;
250            break;
251        }
252    }
253
254    rank
255}
256
257/// Main TSVD function using QR + SVD approach
258///
259/// Computes the truncated SVD using the algorithm:
260/// 1. Apply QR decomposition to A to get Q and R
261/// 2. Compute SVD of R
262/// 3. Reconstruct final U and V matrices
263///
264/// # Arguments
265/// * `matrix` - Input matrix (m × n)
266/// * `config` - TSVD configuration
267///
268/// # Returns
269/// * `SVDResult` - Truncated SVD result
270///
271/// # Errors
272/// * [`Error::EmptyInput`] if the matrix has no rows or no columns
273/// * [`Error::InvalidParameter`] unless `0 < config.rtol < 1` (a NaN
274///   tolerance is rejected)
275/// * [`Error::NonFiniteInput`] if an entry of the matrix is NaN or infinite
276/// * [`Error::DecompositionFailed`] if the QR of the matrix overflows (its R
277///   factor has a non-finite entry) or the SVD iteration does not converge
278pub 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    // Written so that a NaN tolerance fails the check too
295    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    // The SVD iteration never converges on a NaN or an infinity (it looped
304    // forever before the iteration limit of `bounded_svd`).
305    check_finite(matrix)?;
306
307    // Step 1: Apply QR decomposition to A using nalgebra with early termination
308    // Convert config.rtol (T) to T::RealField for QR decomposition
309    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    // A finite matrix can still overflow in the QR (a column norm above
316    // f64::MAX). The input is valid, so a non-finite entry of R is a failure
317    // of the decomposition, not a non-finite input.
318    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        // check_finite reports only NonFiniteInput; pass anything else on.
328        Err(other) => return Err(other),
329    }
330
331    // Step 2: Apply QR-based rank estimation first
332    // Use type-specific epsilon for QR diagonal elements (more conservative than rtol)
333    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        // Matrix has zero rank
340        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    // Step 3: Truncate R to estimated rank and apply SVD
349    let r_truncated: DMatrix<T> = r_matrix.rows(0, qr_rank).into();
350    // Use rtol directly as T
351    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        // Matrix has zero rank
357        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    // Step 4: Reconstruct full SVD
366    // U = Q * U_R (Q is (m, qr_rank), U_R is (qr_rank, svd_result.rank))
367    let q_truncated: DMatrix<T> = q_matrix.columns(0, qr_rank).into();
368    let u_full = &q_truncated * &svd_result.u;
369
370    // V = P^T * V_R (apply inverse permutation matrix)
371    // Since A*P = Q*R, we have A = Q*R*P^T
372    // After SVD of R: A = Q*U_R*S_R*V_R^T*P^T = U*S*V^T
373    // where V^T = V_R^T*P^T, so V = P^(-1)*V_R = P^T*V_R
374    let mut v_full = svd_result.v.clone();
375    permutation.inv_permute_rows(&mut v_full);
376
377    // S_full = S_R (already correct size)
378    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
388/// Convenience function for f64 TSVD
389pub fn tsvd_f64(matrix: &DMatrix<f64>, rtol: f64) -> Result<SVDResult<f64>, Error> {
390    tsvd(matrix, TSVDConfig::new(rtol))
391}
392
393/// Convenience function for Df64 TSVD
394pub fn tsvd_df64(matrix: &DMatrix<Df64>, rtol: Df64) -> Result<SVDResult<Df64>, Error> {
395    tsvd(matrix, TSVDConfig::new(rtol))
396}
397
398/// Convenience function for Df64 TSVD from f64 matrix
399pub 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
407/// Compute SVD for DTensor using nalgebra-based TSVD
408///
409/// Supports both f64 and Df64 types. Uses nalgebra TSVD backend for both.
410///
411/// # Errors
412/// The errors of [`tsvd`]: [`Error::EmptyInput`], [`Error::NonFiniteInput`]
413/// and [`Error::DecompositionFailed`]
414///
415/// # Panics
416/// Panics if `T` is neither `f64` nor `Df64`
417pub 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    // Dispatch based on type: convert to appropriate DMatrix type
424    if TypeId::of::<T>() == TypeId::of::<f64>() {
425        // Convert to DMatrix<f64>
426        let matrix_f64 = DMatrix::from_fn(matrix.shape().0, matrix.shape().1, |i, j| {
427            CustomNumeric::to_f64(matrix[[i, j]])
428        });
429
430        // Use TSVD with appropriate tolerance for f64
431        let rtol = 2.0 * f64::EPSILON;
432        let result = tsvd(&matrix_f64, TSVDConfig::new(rtol))?;
433
434        // Convert back to DTensor<T>
435        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        // Convert to DMatrix<Df64> without going through f64 to preserve precision
450        // TypeId check ensures T == Df64 at runtime, so we can safely cast
451        let matrix_df64: DMatrix<Df64> =
452            DMatrix::from_fn(matrix.shape().0, matrix.shape().1, |i, j| {
453                // Safe: TypeId check guarantees T == Df64
454                unsafe { std::mem::transmute_copy(&matrix[[i, j]]) }
455            });
456
457        // Use TSVD with appropriate tolerance for Df64
458        let rtol = Df64::from(2.0) * Df64::epsilon();
459        let result = tsvd_df64(&matrix_df64, rtol)?;
460
461        // Convert back to DTensor<T> without going through f64 to preserve Df64 precision
462        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    /// A matrix with no rows or no columns is empty, not only 0 × 0; an empty
517    /// batch produces such shapes.
518    #[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    /// A 4 x 4 matrix of full rank with `bad` at (1, 2)
533    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    /// Before the fix, `tsvd` passed `max_niter = 0` (unbounded) to nalgebra's
556    /// `try_svd`, which never converges on a NaN or an infinity: these calls
557    /// did not return.
558    #[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    /// A NaN tolerance passed the `rtol <= 0 || rtol >= 1` check before the
570    /// fix and gave a silent rank-0 result. The message shows the value as a
571    /// number for Df64 too.
572    #[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    /// A finite matrix whose QR overflows: the norm of its column, about
618    /// 2.5e308, exceeds f64::MAX. The input is valid, so this is a failure of
619    /// the decomposition. Before the fix it was reported as a non-finite
620    /// entry of the input matrix, at the index of the infinite entry of R.
621    #[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    /// A tolerance out of range is shown in scientific notation, not with
634    /// hundreds of digits.
635    #[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    /// Non-convergence of the SVD iteration is reported as an error. A NaN
643    /// never converges, so it exhausts any iteration limit; `bounded_svd`
644    /// does not check finiteness itself.
645    #[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        // The limit used for real matrices is ample for a finite one.
655        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    /// Create Hilbert matrix of size n x n with generic type
668    /// H[i,j] = 1 / (i + j + 1)
669    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            // For high precision types like Df64, we need to do the division in type T
675            // to preserve precision, not in f64
676            T::one() / T::from((i + j + 1) as f64)
677        })
678    }
679
680    /// Reconstruct matrix from SVD with generic type: A = U * S * V^T
681    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        // A = U * S * V^T
690        // U: (m × k), S: (k), V: (n × k)
691        // Result: (m × n)
692        u * &DMatrix::from_diagonal(s) * &v.transpose()
693    }
694
695    /// Calculate Frobenius norm of matrix with generic type
696    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    /// Generic Hilbert matrix reconstruction test
711    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        // Compute TSVD with specified tolerance
723        let config = TSVDConfig::new(T::from(rtol));
724        let result = tsvd(&h, config).unwrap();
725
726        // Reconstruct matrix
727        let h_reconstructed = reconstruct_matrix_generic(&result.u, &result.s, &result.v);
728
729        // Calculate reconstruction error (in the same type T to preserve precision)
730        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        // Check that reconstruction error is within expected bounds
735        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        // Note: 10x10 Hilbert matrix has very large condition number (~1e13)
761        // Even with Df64, reconstruction is limited by nalgebra's matrix operations
762        // which may not fully utilize Df64's precision in intermediate calculations
763        test_hilbert_reconstruction_generic::<Df64>(10, 1e-28, 1e-30);
764    }
765
766    #[test]
767    fn test_hilbert_100x100_f64_reconstruction() {
768        // Large matrix test with f64 - expect reasonable performance
769        test_hilbert_reconstruction_generic::<f64>(100, 1e-12, 1e-12);
770    }
771
772    #[test]
773    fn test_hilbert_100x100_df64_reconstruction() {
774        // Large matrix test with Df64 - expect high precision but longer execution time
775        test_hilbert_reconstruction_generic::<Df64>(100, 1e-28, 1e-28);
776    }
777}