Skip to main content

sparse_ir_core/fitters/
common.rs

1//! Shared machinery for fitters.
2//!
3//! Every fitter operation is a contraction of a stored matrix with one axis of
4//! a column-major tensor. The tensor is addressed as `[pre, n, post]`, where
5//! `pre` is the product of the extents before the target axis and `post` the
6//! product of those after it, so no data movement is needed to bring the
7//! target axis to the front.
8//!
9//! This module provides the `InplaceFitter` trait, layout validation for
10//! tenferro tensors and views, and the SVD-based least-squares factors.
11
12use crate::error::{ArrayRole, Error, Result};
13use crate::fpu_check::FpuGuard;
14use crate::gemm::{GemmBackendHandle, GemmScalar, apply_along_axis};
15use num_complex::Complex;
16use tenferro_tensor::{TypedTensor, TypedTensorView, TypedTensorViewMut};
17
18// ============================================================================
19// InplaceFitter trait
20// ============================================================================
21
22/// Evaluation and fitting along one axis of an N-dimensional tensor, writing
23/// into a caller-provided output view.
24///
25/// Method suffixes follow BLAS naming (`d` = `f64`, `z` = `Complex<f64>`):
26/// `evaluate_nd_dz_to` maps real coefficients to complex values.
27///
28/// Inputs may have any strided host layout (non-contiguous inputs are copied
29/// once). Outputs must be host-resident compact column-major views whose
30/// shape equals the input shape with axis `dim` replaced by the output
31/// extent.
32///
33/// Each method returns `Ok(())` after writing the result to `out`. On an
34/// error nothing is written to `out`:
35/// - [`Error::NotSupported`] if the fitter does not support this pair of
36///   types (the default implementations), or if `out` is not a compact
37///   column-major view;
38/// - [`Error::AxisOutOfRange`] if `dim` is not an axis of the input;
39/// - [`Error::ShapeMismatch`] of the input if it does not have `basis_size`
40///   (evaluate) or `n_points` (fit) along `dim`, and of the output if `out`
41///   does not have the shape of the input with `n_points` (evaluate) or
42///   `basis_size` (fit) along `dim`;
43/// - [`Error::DecompositionFailed`] if a fit needs the singular value
44///   decomposition of the matrix and it fails.
45///
46/// An empty input of the right shape (a batch axis of extent 0) gives
47/// `Ok(())` without computing anything.
48pub trait InplaceFitter {
49    /// Number of sampling points.
50    fn n_points(&self) -> usize;
51
52    /// Number of basis functions.
53    fn basis_size(&self) -> usize;
54
55    /// Evaluate: `f64` coefficients to `f64` values.
56    fn evaluate_nd_dd_to(
57        &self,
58        backend: Option<&GemmBackendHandle>,
59        coeffs: &TypedTensorView<'_, f64>,
60        dim: usize,
61        out: &mut TypedTensorViewMut<'_, f64>,
62    ) -> Result<()> {
63        let _ = (backend, coeffs, dim, out);
64        Err(not_supported(
65            "evaluate_nd_dd_to (real coefficients to real values)",
66        ))
67    }
68
69    /// Evaluate: `f64` coefficients to `Complex<f64>` values.
70    fn evaluate_nd_dz_to(
71        &self,
72        backend: Option<&GemmBackendHandle>,
73        coeffs: &TypedTensorView<'_, f64>,
74        dim: usize,
75        out: &mut TypedTensorViewMut<'_, Complex<f64>>,
76    ) -> Result<()> {
77        let _ = (backend, coeffs, dim, out);
78        Err(not_supported(
79            "evaluate_nd_dz_to (real coefficients to complex values)",
80        ))
81    }
82
83    /// Evaluate: `Complex<f64>` coefficients to `f64` values.
84    fn evaluate_nd_zd_to(
85        &self,
86        backend: Option<&GemmBackendHandle>,
87        coeffs: &TypedTensorView<'_, Complex<f64>>,
88        dim: usize,
89        out: &mut TypedTensorViewMut<'_, f64>,
90    ) -> Result<()> {
91        let _ = (backend, coeffs, dim, out);
92        Err(not_supported(
93            "evaluate_nd_zd_to (complex coefficients to real values)",
94        ))
95    }
96
97    /// Evaluate: `Complex<f64>` coefficients to `Complex<f64>` values.
98    fn evaluate_nd_zz_to(
99        &self,
100        backend: Option<&GemmBackendHandle>,
101        coeffs: &TypedTensorView<'_, Complex<f64>>,
102        dim: usize,
103        out: &mut TypedTensorViewMut<'_, Complex<f64>>,
104    ) -> Result<()> {
105        let _ = (backend, coeffs, dim, out);
106        Err(not_supported(
107            "evaluate_nd_zz_to (complex coefficients to complex values)",
108        ))
109    }
110
111    /// Fit: `f64` values to `f64` coefficients.
112    fn fit_nd_dd_to(
113        &self,
114        backend: Option<&GemmBackendHandle>,
115        values: &TypedTensorView<'_, f64>,
116        dim: usize,
117        out: &mut TypedTensorViewMut<'_, f64>,
118    ) -> Result<()> {
119        let _ = (backend, values, dim, out);
120        Err(not_supported(
121            "fit_nd_dd_to (real values to real coefficients)",
122        ))
123    }
124
125    /// Fit: `f64` values to `Complex<f64>` coefficients.
126    fn fit_nd_dz_to(
127        &self,
128        backend: Option<&GemmBackendHandle>,
129        values: &TypedTensorView<'_, f64>,
130        dim: usize,
131        out: &mut TypedTensorViewMut<'_, Complex<f64>>,
132    ) -> Result<()> {
133        let _ = (backend, values, dim, out);
134        Err(not_supported(
135            "fit_nd_dz_to (real values to complex coefficients)",
136        ))
137    }
138
139    /// Fit: `Complex<f64>` values to `f64` coefficients.
140    fn fit_nd_zd_to(
141        &self,
142        backend: Option<&GemmBackendHandle>,
143        values: &TypedTensorView<'_, Complex<f64>>,
144        dim: usize,
145        out: &mut TypedTensorViewMut<'_, f64>,
146    ) -> Result<()> {
147        let _ = (backend, values, dim, out);
148        Err(not_supported(
149            "fit_nd_zd_to (complex values to real coefficients)",
150        ))
151    }
152
153    /// Fit: `Complex<f64>` values to `Complex<f64>` coefficients.
154    fn fit_nd_zz_to(
155        &self,
156        backend: Option<&GemmBackendHandle>,
157        values: &TypedTensorView<'_, Complex<f64>>,
158        dim: usize,
159        out: &mut TypedTensorViewMut<'_, Complex<f64>>,
160    ) -> Result<()> {
161        let _ = (backend, values, dim, out);
162        Err(not_supported(
163            "fit_nd_zz_to (complex values to complex coefficients)",
164        ))
165    }
166}
167
168/// The error of an [`InplaceFitter`] method that a fitter does not support
169fn not_supported(operation: &str) -> Error {
170    Error::NotSupported {
171        what: format!("{operation} for this sampling"),
172    }
173}
174
175/// The error of a tensor view that the fitters cannot address as compact
176/// column-major storage
177fn layout_error(what: impl Into<String>) -> Error {
178    Error::NotSupported { what: what.into() }
179}
180
181// ============================================================================
182// Scalar reinterpretation
183// ============================================================================
184
185/// `f64` or `Complex<f64>`, viewable as interleaved `f64` storage.
186///
187/// A column-major complex tensor of shape `[pre, n, post]` is a real tensor
188/// of shape `[2 * pre, n, post]`, so real-matrix contractions apply to it
189/// with `pre` doubled.
190pub trait FitScalar: GemmScalar {
191    /// Number of `f64` values per element (1 or 2).
192    const REAL_WIDTH: usize;
193    /// Reinterpret as `f64` storage.
194    fn as_f64(s: &[Self]) -> &[f64];
195    /// Reinterpret mutably as `f64` storage.
196    fn as_f64_mut(s: &mut [Self]) -> &mut [f64];
197}
198
199impl FitScalar for f64 {
200    const REAL_WIDTH: usize = 1;
201    fn as_f64(s: &[f64]) -> &[f64] {
202        s
203    }
204    fn as_f64_mut(s: &mut [f64]) -> &mut [f64] {
205        s
206    }
207}
208
209impl FitScalar for Complex<f64> {
210    const REAL_WIDTH: usize = 2;
211    fn as_f64(s: &[Self]) -> &[f64] {
212        // SAFETY: `Complex<f64>` is `#[repr(C)] { re: f64, im: f64 }`, so a
213        // slice of N complex values is a valid slice of 2N f64 values with
214        // the same lifetime and alignment.
215        unsafe { std::slice::from_raw_parts(s.as_ptr().cast::<f64>(), 2 * s.len()) }
216    }
217    fn as_f64_mut(s: &mut [Self]) -> &mut [f64] {
218        // SAFETY: as above; the unique borrow is transferred.
219        unsafe { std::slice::from_raw_parts_mut(s.as_mut_ptr().cast::<f64>(), 2 * s.len()) }
220    }
221}
222
223// ============================================================================
224// Axis geometry and layout validation
225// ============================================================================
226
227/// Extents of a tensor split around the target axis.
228#[derive(Clone, Copy, Debug, PartialEq, Eq)]
229pub(crate) struct AxisSplit {
230    pub pre: usize,
231    pub post: usize,
232}
233
234/// Check the input of an N-D evaluate or fit along axis `dim`
235///
236/// `Ok` if `dim` is an axis of the input and the input has extent `n_in`
237/// along it. Call this before reading `input_dims[dim]` anywhere else, and
238/// before allocating an output from the input shape.
239///
240/// # Errors
241///
242/// * [`Error::AxisOutOfRange`] if `dim` is not an axis of the input
243/// * [`Error::ShapeMismatch`] of the input, with the shape it should have,
244///   if its extent along `dim` is not `n_in`
245pub(crate) fn check_input_shape(input_dims: &[usize], dim: usize, n_in: usize) -> Result<()> {
246    let rank = input_dims.len();
247    if dim >= rank {
248        return Err(Error::AxisOutOfRange { axis: dim, rank });
249    }
250    if input_dims[dim] != n_in {
251        let mut expected = input_dims.to_vec();
252        expected[dim] = n_in;
253        return Err(Error::ShapeMismatch {
254            which: ArrayRole::Input,
255            expected,
256            actual: input_dims.to_vec(),
257        });
258    }
259    Ok(())
260}
261
262/// Check the shapes of an N-D evaluate or fit along axis `dim` that writes
263/// to `out`: the input must pass [`check_input_shape`], and `out` must have
264/// the shape of the input with `n_out` along `dim` (in particular its rank).
265///
266/// # Errors
267///
268/// The errors of [`check_input_shape`], then [`Error::ShapeMismatch`] of the
269/// output, with the shape it should have, if `out` has another shape
270pub(crate) fn check_nd_shapes(
271    input_dims: &[usize],
272    dim: usize,
273    n_in: usize,
274    out_dims: &[usize],
275    n_out: usize,
276) -> Result<()> {
277    check_input_shape(input_dims, dim, n_in)?;
278    let expected = replace_axis(input_dims, dim, n_out);
279    if out_dims != expected.as_slice() {
280        return Err(Error::ShapeMismatch {
281            which: ArrayRole::Output,
282            expected,
283            actual: out_dims.to_vec(),
284        });
285    }
286    Ok(())
287}
288
289/// `Ok` if a slice of `len` elements has the `expected` length
290///
291/// # Errors
292///
293/// [`Error::ShapeMismatch`] of `which`, with the lengths as one-entry shapes
294pub(crate) fn check_len(which: ArrayRole, len: usize, expected: usize) -> Result<()> {
295    if len == expected {
296        Ok(())
297    } else {
298        Err(Error::ShapeMismatch {
299            which,
300            expected: vec![expected],
301            actual: vec![len],
302        })
303    }
304}
305
306/// Split `shape` around `dim` after [`check_input_shape`].
307pub(crate) fn split_axis(shape: &[usize], dim: usize, expected: usize) -> Result<AxisSplit> {
308    check_input_shape(shape, dim, expected)?;
309    Ok(AxisSplit {
310        pre: shape[..dim].iter().product(),
311        post: shape[dim + 1..].iter().product(),
312    })
313}
314
315/// `shape` with axis `dim` replaced by `extent`.
316pub(crate) fn replace_axis(shape: &[usize], dim: usize, extent: usize) -> Vec<usize> {
317    let mut out = shape.to_vec();
318    out[dim] = extent;
319    out
320}
321
322/// Compact column-major host slice of a view, copying a strided view once.
323pub(crate) enum InputSlice<'a, T: tenferro_tensor::TensorScalar> {
324    Borrowed(&'a [T]),
325    Owned(Vec<T>),
326}
327
328impl<T: tenferro_tensor::TensorScalar> InputSlice<'_, T> {
329    pub(crate) fn get(&self) -> Result<&[T]> {
330        match self {
331            InputSlice::Borrowed(s) => Ok(s),
332            InputSlice::Owned(v) => Ok(v),
333        }
334    }
335}
336
337pub(crate) fn input_slice<'a, T: tenferro_tensor::TensorScalar>(
338    view: &TypedTensorView<'a, T>,
339) -> Result<InputSlice<'a, T>> {
340    if view.is_col_major_contiguous()? {
341        Ok(InputSlice::Borrowed(view.as_slice()?))
342    } else {
343        Ok(InputSlice::Owned(gather_col_major(view)?))
344    }
345}
346
347/// Copy a strided host view into a compact column-major buffer.
348fn gather_col_major<T: tenferro_tensor::TensorScalar>(
349    view: &TypedTensorView<'_, T>,
350) -> Result<Vec<T>> {
351    let shape = view.shape();
352    let strides = view.strides();
353    let storage = view.host_storage()?;
354    let len = view.n_elements();
355    if len == 0 {
356        return Ok(Vec::new());
357    }
358    // Validate the reachable offset range once, before allocating.
359    let (mut lo, mut hi) = (view.offset(), view.offset());
360    for (&n, &st) in shape.iter().zip(strides) {
361        let span = isize::try_from(n - 1)
362            .ok()
363            .and_then(|m| m.checked_mul(st))
364            .ok_or_else(|| layout_error("input view stride overflow"))?;
365        if span < 0 {
366            lo += span;
367        } else {
368            hi += span;
369        }
370    }
371    if lo < 0 || usize::try_from(hi).map_or(true, |h| h >= storage.len()) {
372        return Err(layout_error("input view exceeds its storage"));
373    }
374    let mut out = Vec::with_capacity(len);
375    let mut idx = vec![0usize; shape.len()];
376    let mut pos = view.offset();
377    for _ in 0..len {
378        // pos is within [lo, hi] by the range check above.
379        out.push(storage[pos as usize].clone());
380        for (k, i) in idx.iter_mut().enumerate() {
381            *i += 1;
382            pos += strides[k];
383            if *i < shape[k] {
384                break;
385            }
386            pos -= strides[k] * shape[k] as isize;
387            *i = 0;
388        }
389    }
390    Ok(out)
391}
392
393/// Compact column-major host slice of an output view.
394pub(crate) fn output_slice<'a, T: tenferro_tensor::TensorScalar>(
395    view: &'a mut TypedTensorViewMut<'_, T>,
396) -> Result<&'a mut [T]> {
397    if !view.is_col_major_contiguous()? {
398        return Err(layout_error(format!(
399            "output view must be compact column-major, got {}",
400            view.layout_summary()
401        )));
402    }
403    let offset = usize::try_from(view.offset())
404        .map_err(|_| layout_error("output view has a negative offset"))?;
405    let len = view.n_elements();
406    let storage = view
407        .host_storage_mut()
408        .map_err(|e| layout_error(e.to_string()))?;
409    storage
410        .get_mut(offset..offset + len)
411        .ok_or_else(|| layout_error("output view exceeds its storage"))
412}
413
414/// Run an axis operation `n_in -> n_out` from `input` into `out`.
415///
416/// `op` receives compact column-major `[pre, n_in, post]` and
417/// `[pre, n_out, post]` slices.
418pub(crate) fn run_to<Tin, Tout, F>(
419    input: &TypedTensorView<'_, Tin>,
420    dim: usize,
421    n_in: usize,
422    n_out: usize,
423    out: &mut TypedTensorViewMut<'_, Tout>,
424    op: F,
425) -> Result<()>
426where
427    Tin: tenferro_tensor::TensorScalar,
428    Tout: tenferro_tensor::TensorScalar,
429    F: FnOnce(&[Tin], AxisSplit, &mut [Tout]) -> Result<()>,
430{
431    check_nd_shapes(input.shape(), dim, n_in, out.shape(), n_out)?;
432    let split = split_axis(input.shape(), dim, n_in)?;
433    let x = input_slice(input)?;
434    let y = output_slice(out)?;
435    op(x.get()?, split, y)
436}
437
438/// Run an axis operation `n_in -> n_out` into a newly allocated tensor.
439pub(crate) fn run_alloc<Tin, Tout, F>(
440    input: &TypedTensor<Tin>,
441    dim: usize,
442    n_in: usize,
443    n_out: usize,
444    op: F,
445) -> Result<TypedTensor<Tout>>
446where
447    Tin: tenferro_tensor::TensorScalar,
448    Tout: tenferro_tensor::TensorScalar + num_traits::Zero,
449    F: FnOnce(&[Tin], AxisSplit, &mut [Tout]) -> Result<()>,
450{
451    let split = split_axis(input.shape(), dim, n_in)?;
452    let out_shape = replace_axis(input.shape(), dim, n_out);
453    // Owned tensors are compact column-major by tenferro's invariant, so the
454    // host buffer is the tensor in order (and avoids building a view).
455    let x = input.host_data()?;
456    debug_assert_eq!(x.len(), input.n_elements());
457    let mut y = vec![Tout::zero(); split.pre * n_out * split.post];
458    op(x, split, &mut y)?;
459    Ok(TypedTensor::from_vec_col_major(out_shape, y)?)
460}
461
462// ============================================================================
463// Complex <-> stacked-real layout helpers
464// ============================================================================
465
466/// Convert real `[pre, 2n, post]` (rows `0..n` real parts, `n..2n` imaginary
467/// parts) into complex `[pre, n, post]`.
468pub(crate) fn stacked_to_complex(
469    t: &[f64],
470    pre: usize,
471    n: usize,
472    post: usize,
473    y: &mut [Complex<f64>],
474) {
475    for q in 0..post {
476        for i in 0..n {
477            let re = &t[pre * (i + 2 * n * q)..][..pre];
478            let im = &t[pre * (n + i + 2 * n * q)..][..pre];
479            let dst = &mut y[pre * (i + n * q)..][..pre];
480            for p in 0..pre {
481                dst[p] = Complex::new(re[p], im[p]);
482            }
483        }
484    }
485}
486
487/// Convert complex `[pre, n, post]` into real `[pre, 2n, post]` with real
488/// parts in rows `0..n` and imaginary parts in rows `n..2n`.
489pub(crate) fn complex_to_stacked(
490    x: &[Complex<f64>],
491    pre: usize,
492    n: usize,
493    post: usize,
494) -> Vec<f64> {
495    let mut t = vec![0.0; 2 * pre * n * post];
496    for q in 0..post {
497        for i in 0..n {
498            let src = &x[pre * (i + n * q)..][..pre];
499            let base_re = pre * (i + 2 * n * q);
500            let base_im = pre * (n + i + 2 * n * q);
501            for p in 0..pre {
502                t[base_re + p] = src[p].re;
503                t[base_im + p] = src[p].im;
504            }
505        }
506    }
507    t
508}
509
510// ============================================================================
511// SVD-based least squares
512// ============================================================================
513
514/// Factors of the pseudo-inverse `A^+ = V diag(1/s) U^H` of an `n x m`
515/// matrix, stored for axis contractions.
516///
517/// Rank is `r = min(n, m)` without truncation: sampling matrices are well
518/// conditioned by construction.
519#[doc(hidden)]
520pub struct PinvFactors<T> {
521    /// `U^H`, contiguous `r x n`.
522    pub uh: Vec<T>,
523    /// `V diag(1/s)`, contiguous `m x r`.
524    pub v_scaled: Vec<T>,
525    /// Singular values `s`, descending (`min(n, m)` of them, before any
526    /// truncation).
527    pub s: Vec<f64>,
528    pub n: usize,
529    pub m: usize,
530    pub rank: usize,
531}
532
533impl<T: GemmScalar> PinvFactors<T> {
534    /// Least-squares solve along the middle axis of `[pre, n, post]` into
535    /// `[pre, m, post]`.
536    pub fn solve(
537        &self,
538        backend: Option<&GemmBackendHandle>,
539        values: &[T],
540        pre: usize,
541        post: usize,
542        out: &mut [T],
543    ) -> Result<()> {
544        let mut tmp = vec![T::zero(); pre * self.rank * post];
545        apply_along_axis(
546            backend, &self.uh, self.rank, self.n, values, pre, post, &mut tmp,
547        )?;
548        apply_along_axis(
549            backend,
550            &self.v_scaled,
551            self.m,
552            self.rank,
553            &tmp,
554            pre,
555            post,
556            out,
557        )?;
558        Ok(())
559    }
560}
561
562/// Scalars with an SVD through tenferro-linalg.
563pub trait SvdScalar: GemmScalar + tenferro_linalg::LinalgScalar {
564    fn conj_(self) -> Self;
565    fn scale(self, s: f64) -> Self;
566    fn real_part(v: <Self as tenferro_tensor::TensorScalar>::Real) -> f64;
567}
568
569impl SvdScalar for f64 {
570    fn conj_(self) -> Self {
571        self
572    }
573    fn scale(self, s: f64) -> Self {
574        self * s
575    }
576    fn real_part(v: f64) -> f64 {
577        v
578    }
579}
580
581impl SvdScalar for Complex<f64> {
582    fn conj_(self) -> Self {
583        self.conj()
584    }
585    fn scale(self, s: f64) -> Self {
586        self * s
587    }
588    fn real_part(v: f64) -> f64 {
589        v
590    }
591}
592
593/// Thin SVD of a contiguous column-major `n x m` matrix, returned as
594/// pseudo-inverse factors.
595///
596/// # Errors
597/// Propagates tenferro errors (for example a non-converging SVD).
598#[doc(hidden)]
599pub fn compute_pinv<T: SvdScalar>(a: &[T], n: usize, m: usize) -> Result<PinvFactors<T>> {
600    compute_pinv_impl(a, n, m, None, || sampling_matrix((n, m)))
601}
602
603/// Like [`compute_pinv`], but keeps only singular values `s_l > rtol * s_0`
604/// (truncated-SVD regularization).
605///
606/// # Errors
607/// [`Error::DecompositionFailed`] if the SVD fails.
608#[doc(hidden)]
609pub fn compute_pinv_truncated<T: SvdScalar>(
610    a: &[T],
611    n: usize,
612    m: usize,
613    rtol: f64,
614) -> Result<PinvFactors<T>> {
615    compute_pinv_impl(a, n, m, Some(rtol), || sampling_matrix((n, m)))
616}
617
618/// Name of a `rows × cols` sampling matrix in an error
619pub(crate) fn sampling_matrix((rows, cols): (usize, usize)) -> String {
620    format!("{rows} x {cols} sampling matrix")
621}
622
623/// The error of an SVD that failed; `matrix` names the matrix, e.g.
624/// "5 x 3 sampling matrix"
625fn svd_failed(matrix: String, err: impl std::fmt::Display) -> Error {
626    Error::DecompositionFailed {
627        reason: format!("the SVD of the {matrix} failed: {err}"),
628    }
629}
630
631/// [`compute_pinv`] of a matrix that `name` names in the error
632///
633/// # Errors
634/// [`Error::DecompositionFailed`] if the SVD does not converge or returns
635/// factors of the wrong shape.
636pub(crate) fn compute_pinv_of<T: SvdScalar>(
637    a: &[T],
638    n: usize,
639    m: usize,
640    name: impl FnOnce() -> String,
641) -> Result<PinvFactors<T>> {
642    compute_pinv_impl(a, n, m, None, name)
643}
644
645fn compute_pinv_impl<T: SvdScalar>(
646    a: &[T],
647    n: usize,
648    m: usize,
649    rtol: Option<f64>,
650    name: impl FnOnce() -> String,
651) -> Result<PinvFactors<T>> {
652    use tenferro_cpu::CpuBackend;
653    use tenferro_linalg::TypedTensorLinalgExt;
654    use tenferro_tensor::BackendSessionHost;
655
656    // Protect FPU state during SVD computation (required for Intel Fortran compatibility)
657    let _guard = FpuGuard::new_protect_computation();
658
659    let rank = n.min(m);
660    if rank == 0 {
661        return Ok(PinvFactors {
662            uh: Vec::new(),
663            v_scaled: vec![T::zero(); m * rank],
664            s: Vec::new(),
665            n,
666            m,
667            rank,
668        });
669    }
670    let tensor = TypedTensor::<T>::from_vec_col_major(vec![n, m], a.to_vec())?;
671    let mut host = CpuBackend::new();
672    let (u, s, vt) = match host.with_backend_session(|session| tensor.svd(session)) {
673        Ok(f) => f,
674        Err(e) => return Err(svd_failed(name(), e)),
675    };
676
677    let (u_shape, vt_shape) = (u.shape().to_vec(), vt.shape().to_vec());
678    let (u, s, vt) = (u.host_data()?, s.host_data()?, vt.host_data()?);
679    if u_shape[0] != n || u_shape[1] < rank || vt_shape[1] != m || vt_shape[0] < rank {
680        return Err(svd_failed(
681            name(),
682            format!("factors U {u_shape:?}, Vt {vt_shape:?}"),
683        ));
684    }
685    let (ldu, ldvt) = (u_shape[0], vt_shape[0]);
686    let singular_values: Vec<f64> = s[..rank].iter().map(|&v| T::real_part(v)).collect();
687    let rank = match rtol {
688        Some(rtol) => {
689            let s0 = T::real_part(s[0]);
690            (0..rank)
691                .take_while(|&l| T::real_part(s[l]) > rtol * s0)
692                .count()
693        }
694        None => rank,
695    };
696
697    // U^H: uh[l + r*i] = conj(U[i, l])
698    let mut uh = vec![T::zero(); rank * n];
699    for i in 0..n {
700        for l in 0..rank {
701            uh[l + rank * i] = u[i + ldu * l].conj_();
702        }
703    }
704    // V diag(1/s): V[j, l] = conj(Vt[l, j])
705    let mut v_scaled = vec![T::zero(); m * rank];
706    for l in 0..rank {
707        let inv_s = 1.0 / T::real_part(s[l]);
708        for j in 0..m {
709            v_scaled[j + m * l] = vt[l + ldvt * j].conj_().scale(inv_s);
710        }
711    }
712    Ok(PinvFactors {
713        uh,
714        v_scaled,
715        s: singular_values,
716        n,
717        m,
718        rank,
719    })
720}
721
722/// Singular values of a contiguous column-major `n x m` matrix, descending.
723///
724/// # Errors
725/// Propagates tenferro errors.
726pub fn singular_values<T: SvdScalar>(a: &[T], n: usize, m: usize) -> Result<Vec<f64>> {
727    use tenferro_cpu::CpuBackend;
728    use tenferro_linalg::TypedTensorLinalgExt;
729    use tenferro_tensor::BackendSessionHost;
730
731    let _guard = FpuGuard::new_protect_computation();
732    if n.min(m) == 0 {
733        return Ok(Vec::new());
734    }
735    let tensor = TypedTensor::<T>::from_vec_col_major(vec![n, m], a.to_vec())?;
736    let mut host = CpuBackend::new();
737    let s = host
738        .with_backend_session(|session| tensor.svdvals(session))
739        .map_err(|e| svd_failed(sampling_matrix((n, m)), e))?;
740    Ok(s.host_data()?.iter().map(|&v| T::real_part(v)).collect())
741}
742
743/// Condition number `σ_max / σ_min` from the singular values of a fitting matrix
744///
745/// Conventions shared by the `condition_number` methods of all samplings and
746/// by `spir_sampling_get_cond_num`:
747/// - `f64::INFINITY` if `σ_min < 1e-15` (numerically singular matrix);
748/// - `1.0` if there are no singular values (a matrix with a zero dimension);
749/// - `NaN` if any singular value is NaN, so a failed decomposition never
750///   yields a plausible finite value.
751pub(crate) fn condition_number_from_singular_values(s: &[f64]) -> f64 {
752    if s.is_empty() {
753        return 1.0;
754    }
755    if s.iter().any(|x| x.is_nan()) {
756        return f64::NAN;
757    }
758    let s_max = s.iter().copied().fold(f64::NEG_INFINITY, f64::max);
759    let s_min = s.iter().copied().fold(f64::INFINITY, f64::min);
760    if s_min.abs() < 1e-15 {
761        return f64::INFINITY;
762    }
763    s_max / s_min
764}