Skip to main content

sparse_ir_core/
sampling.rs

1//! Sparse sampling in imaginary time
2//!
3//! This module provides `TauSampling` for transforming between IR basis coefficients
4//! and values at sparse sampling points in imaginary time.
5
6use crate::Matrix;
7use crate::error::Error;
8use crate::fitters::InplaceFitter;
9use crate::gemm::GemmBackendHandle;
10use crate::matrix::Mat;
11use crate::traits::StatisticsType;
12use num_complex::Complex;
13use tenferro_tensor::{TensorScalar, TypedTensor, TypedTensorView, TypedTensorViewMut};
14
15/// Copy a host matrix into the internal column-major container
16#[doc(hidden)]
17pub fn mat_from_matrix<T: TensorScalar + Copy>(m: &Matrix<T>) -> Result<Mat<T>, Error> {
18    Ok(Mat::from_typed(m)?)
19}
20
21/// Move axis from position `src` to position `dst`
22///
23/// This is equivalent to numpy.moveaxis or libsparseir's movedim. The other
24/// axes keep their order.
25///
26/// # Arguments
27/// * `arr` - Input tensor
28/// * `src` - Source axis position
29/// * `dst` - Destination axis position
30///
31/// # Returns
32/// A new tensor with the axes permuted
33///
34/// # Panics
35///
36/// Panics if `src` or `dst` is not an axis of `arr`.
37///
38/// # Example
39/// ```
40/// use sparse_ir::TypedTensor;
41/// use sparse_ir::sampling::movedim;
42///
43/// // A 4D tensor with shape (2, 3, 4, 5) and entries 1000 i + 100 j + 10 k + l
44/// let shape = [2usize, 3, 4, 5];
45/// let data: Vec<f64> = (0..120)
46///     .map(|lin| {
47///         let (i, j, k, l) = (lin % 2, lin / 2 % 3, lin / 6 % 4, lin / 24);
48///         (1000 * i + 100 * j + 10 * k + l) as f64
49///     })
50///     .collect();
51/// let arr = TypedTensor::from_vec_col_major(shape.to_vec(), data).unwrap();
52///
53/// // movedim(arr, 0, 2) moves axis 0 to position 2
54/// let moved = movedim(&arr, 0, 2);
55///
56/// // Result shape: (3, 4, 2, 5) with axes permuted as [1, 2, 0, 3]
57/// assert_eq!(moved.shape(), &[3, 4, 2, 5]);
58/// // Element [2, 3, 1, 4] of the result is element [1, 2, 3, 4] of arr
59/// let at = |t: &TypedTensor<f64>, idx: [usize; 4]| {
60///     let s = t.shape();
61///     t.host_data().unwrap()[idx[0] + s[0] * (idx[1] + s[1] * (idx[2] + s[2] * idx[3]))]
62/// };
63/// assert_eq!(at(&moved, [2, 3, 1, 4]), at(&arr, [1, 2, 3, 4]));
64/// ```
65pub fn movedim<T: TensorScalar + Copy>(
66    arr: &TypedTensor<T>,
67    src: usize,
68    dst: usize,
69) -> TypedTensor<T> {
70    let shape = arr.shape().to_vec();
71    let rank = shape.len();
72    assert!(
73        src < rank,
74        "src axis {} out of bounds for rank {}",
75        src,
76        rank
77    );
78    assert!(
79        dst < rank,
80        "dst axis {} out of bounds for rank {}",
81        dst,
82        rank
83    );
84    // Output axis k reads input axis perm[k].
85    let mut perm: Vec<usize> = (0..rank).collect();
86    perm.remove(src);
87    perm.insert(dst, src);
88    let out_shape: Vec<usize> = perm.iter().map(|&p| shape[p]).collect();
89    let data = arr
90        .host_data()
91        .expect("an owned tensor is compact host storage");
92    let len = data.len();
93    let mut out = Vec::with_capacity(len);
94    let mut out_idx = vec![0usize; rank];
95    for _ in 0..len {
96        // Column-major offset of the input element
97        let mut offset = 0;
98        let mut stride = 1;
99        for axis in 0..rank {
100            let k = perm.iter().position(|&p| p == axis).unwrap();
101            offset += out_idx[k] * stride;
102            stride *= shape[axis];
103        }
104        out.push(data[offset]);
105        for (k, i) in out_idx.iter_mut().enumerate() {
106            *i += 1;
107            if *i < out_shape[k] {
108                break;
109            }
110            *i = 0;
111        }
112    }
113    TypedTensor::from_vec_col_major(out_shape, out).expect("the shape matches the data")
114}
115
116/// Check the shape of a given sampling matrix against its points: some
117/// points, one row per point and at least one column
118///
119/// # Errors
120///
121/// * [`Error::EmptyInput`] named `sampling_points` if there are no points
122/// * [`Error::ShapeMismatch`] of the input if the matrix does not have one
123///   row per point
124/// * [`Error::EmptyInput`] named `matrix` if it has no columns: it describes
125///   no basis function
126pub(crate) fn check_sampling_matrix_shape(
127    n_points: usize,
128    (rows, cols): (usize, usize),
129) -> Result<(), Error> {
130    if n_points == 0 {
131        return Err(Error::EmptyInput {
132            name: "sampling_points",
133        });
134    }
135    if rows != n_points {
136        return Err(Error::ShapeMismatch {
137            which: crate::error::ArrayRole::Input,
138            expected: vec![n_points, cols],
139            actual: vec![rows, cols],
140        });
141    }
142    if cols == 0 {
143        return Err(Error::EmptyInput { name: "matrix" });
144    }
145    Ok(())
146}
147
148/// `Ok` if every entry of a given sampling matrix is finite (the fitter
149/// factorizes it)
150///
151/// # Errors
152///
153/// [`Error::NonFiniteInput`] named `matrix` at the first NaN or infinite
154/// entry in row-major order; for a complex entry, `value` is its real part
155/// if that is not finite, and its imaginary part otherwise
156pub(crate) fn check_finite_matrix<T: Copy>(
157    matrix: &Mat<T>,
158    non_finite_part: impl Fn(T) -> Option<f64>,
159) -> Result<(), Error> {
160    let (rows, cols) = *matrix.shape();
161    for i in 0..rows {
162        for j in 0..cols {
163            if let Some(value) = non_finite_part(matrix[[i, j]]) {
164                return Err(Error::NonFiniteInput {
165                    name: "matrix",
166                    index: vec![i, j],
167                    value,
168                });
169            }
170        }
171    }
172    Ok(())
173}
174
175/// Sparse sampling in imaginary time
176///
177/// Allows transformation between the IR basis and a set of sampling points
178/// in imaginary time (τ).
179pub struct TauSampling<S>
180where
181    S: StatisticsType,
182{
183    /// Sampling points in imaginary time, in the order given (τ ∈ [-β, β]
184    /// unless given with a matrix)
185    sampling_points: Vec<f64>,
186
187    /// Real matrix fitter for least-squares fitting
188    fitter: crate::fitters::RealMatrixFitter,
189
190    /// Marker for statistics type
191    _phantom: std::marker::PhantomData<S>,
192}
193
194impl<S> TauSampling<S>
195where
196    S: StatisticsType,
197{
198    /// Create a new TauSampling with default sampling points
199    ///
200    /// The default sampling points are the roots of the first discarded basis
201    /// function u_L (the extrema of u_{L-1} when u_L is not available), which
202    /// gives near-optimal conditioning.
203    /// SVD is computed lazily on first call to `fit` or `fit_nd`.
204    ///
205    /// # Arguments
206    /// * `basis` - Any basis implementing the `Basis` trait
207    ///
208    /// # Returns
209    /// A new TauSampling object
210    ///
211    /// # Errors
212    ///
213    /// The errors of [`Basis::default_tau_sampling_points`](crate::basis_trait::Basis::default_tau_sampling_points)
214    /// (e.g. NotSupported for a DLR, whose IR basis has the default points)
215    pub fn new(basis: &impl crate::basis_trait::Basis<S>) -> Result<Self, Error>
216    where
217        S: 'static,
218    {
219        let sampling_points = basis.default_tau_sampling_points()?;
220        Self::with_sampling_points(basis, sampling_points)
221    }
222
223    /// Create a new TauSampling with custom sampling points
224    ///
225    /// SVD is computed lazily on first call to `fit` or `fit_nd`.
226    ///
227    /// # Arguments
228    /// * `basis` - Any basis implementing the `Basis` trait
229    /// * `sampling_points` - Custom sampling points in τ ∈ [-β, β]
230    ///
231    /// # Returns
232    /// A new TauSampling object
233    ///
234    /// The points are kept in the given order, and duplicates are accepted;
235    /// they only raise the condition number.
236    ///
237    /// # Errors
238    ///
239    /// * [`Error::EmptyInput`] if `sampling_points` is empty
240    /// * [`Error::OutOfDomain`] if a point is outside [-β, β] or NaN (from
241    ///   [`Basis::evaluate_tau`](crate::basis_trait::Basis::evaluate_tau))
242    pub fn with_sampling_points(
243        basis: &impl crate::basis_trait::Basis<S>,
244        sampling_points: Vec<f64>,
245    ) -> Result<Self, Error>
246    where
247        S: 'static,
248    {
249        // With no points the sampling matrix would have no rows.
250        if sampling_points.is_empty() {
251            return Err(Error::EmptyInput {
252                name: "sampling_points",
253            });
254        }
255
256        // Compute sampling matrix: A[i, l] = u_l(τ_i); evaluate_tau checks
257        // that every τ is in [-β, β].
258        let matrix = mat_from_matrix(&basis.evaluate_tau(&sampling_points)?)?;
259        let fitter = crate::fitters::RealMatrixFitter::new(matrix);
260
261        Ok(Self {
262            sampling_points,
263            fitter,
264            _phantom: std::marker::PhantomData,
265        })
266    }
267
268    /// Create a new TauSampling with custom sampling points and pre-computed matrix
269    ///
270    /// This constructor is useful when the sampling matrix is already computed
271    /// (e.g., from external sources or for testing).
272    ///
273    /// # Arguments
274    /// * `sampling_points` - Imaginary times τ that label the rows of
275    ///   `matrix`, in any order. There is no β to check them against, so
276    ///   any finite value is accepted and kept as given.
277    /// * `matrix` - Pre-computed sampling matrix (n_points × basis_size); row i
278    ///   belongs to `sampling_points[i]`
279    ///
280    /// Duplicate points are accepted; they only raise the condition number.
281    ///
282    /// # Errors
283    ///
284    /// * [`Error::EmptyInput`] if `sampling_points` is empty, or `matrix`
285    ///   has no columns
286    /// * [`Error::ShapeMismatch`] of the input if `matrix` does not have one
287    ///   row per point
288    /// * [`Error::NonFiniteInput`] for the first NaN or infinite point, then
289    ///   for the first NaN or infinite entry of `matrix`
290    pub fn from_matrix(sampling_points: Vec<f64>, matrix: &Matrix<f64>) -> Result<Self, Error> {
291        let matrix = mat_from_matrix(matrix)?;
292        check_sampling_matrix_shape(sampling_points.len(), *matrix.shape())?;
293        if let Some((i, &tau)) = sampling_points
294            .iter()
295            .enumerate()
296            .find(|(_, tau)| !tau.is_finite())
297        {
298            return Err(Error::NonFiniteInput {
299                name: "sampling_points",
300                index: vec![i],
301                value: tau,
302            });
303        }
304        check_finite_matrix(&matrix, |x: f64| (!x.is_finite()).then_some(x))?;
305
306        let fitter = crate::fitters::RealMatrixFitter::new(matrix);
307
308        Ok(Self {
309            sampling_points,
310            fitter,
311            _phantom: std::marker::PhantomData,
312        })
313    }
314
315    /// Get the sampling points
316    pub fn sampling_points(&self) -> &[f64] {
317        &self.sampling_points
318    }
319
320    /// Get the number of sampling points
321    pub fn n_sampling_points(&self) -> usize {
322        self.fitter.n_points()
323    }
324
325    /// Get the basis size
326    pub fn basis_size(&self) -> usize {
327        self.fitter.basis_size()
328    }
329
330    /// Get the sampling matrix
331    pub fn matrix(&self) -> &Matrix<f64> {
332        self.fitter.matrix()
333    }
334
335    /// Condition number of the sampling matrix, which fitting solves with
336    ///
337    /// Returns `σ_max / σ_min`, the ratio of the largest to the smallest of the
338    /// `min(n_sampling_points, basis_size)` singular values of the real
339    /// `n_sampling_points × basis_size` matrix [`Self::matrix`]. It bounds how
340    /// much [`Self::fit`] can amplify relative errors in the values.
341    ///
342    /// Returns `f64::INFINITY` if the smallest singular value is below `1e-15`
343    /// (numerically singular matrix). The singular value decomposition is the
344    /// one fitting uses: it is computed by the first call to this method or to
345    /// a fit, then cached.
346    ///
347    /// # Errors
348    ///
349    /// [`Error::DecompositionFailed`] if the singular value decomposition
350    /// fails, which a matrix of finite entries does not cause in practice
351    /// (the constructors reject non-finite entries)
352    pub fn condition_number(&self) -> Result<f64, Error> {
353        self.fitter.condition_number()
354    }
355
356    // ========================================================================
357    // 1D functions (real and complex)
358    // ========================================================================
359
360    /// Evaluate basis coefficients at sampling points
361    ///
362    /// Computes g(τ_i) = Σ_l a_l * u_l(τ_i) for all sampling points
363    ///
364    /// # Arguments
365    /// * `coeffs` - Basis coefficients (length = basis_size)
366    ///
367    /// # Returns
368    /// Values at sampling points (length = n_sampling_points)
369    ///
370    /// # Errors
371    ///
372    /// * [`Error::ShapeMismatch`] of the input if `coeffs` does not have length
373    ///   `basis_size`
374    pub fn evaluate(&self, coeffs: &[f64]) -> Result<Vec<f64>, Error> {
375        self.fitter.evaluate(None, coeffs)
376    }
377
378    /// Evaluate basis coefficients at sampling points, writing to output slice
379    ///
380    /// # Errors
381    ///
382    /// * [`Error::ShapeMismatch`] of the input if `coeffs` does not have length
383    ///   `basis_size`
384    /// * [`Error::ShapeMismatch`] of the output if `out` does not have length
385    ///   `n_sampling_points`
386    ///
387    /// Nothing is written to `out` on an error.
388    pub fn evaluate_to(&self, coeffs: &[f64], out: &mut [f64]) -> Result<(), Error> {
389        self.fitter.evaluate_to(None, coeffs, out)
390    }
391
392    /// Fit values at sampling points to basis coefficients
393    ///
394    /// # Errors
395    ///
396    /// * [`Error::ShapeMismatch`] of the input if `values` does not have length
397    ///   `n_sampling_points`
398    /// * [`Error::DecompositionFailed`] if the singular value decomposition
399    ///   fails
400    pub fn fit(&self, values: &[f64]) -> Result<Vec<f64>, Error> {
401        self.fitter.fit(None, values)
402    }
403
404    /// Fit values at sampling points to basis coefficients, writing to output slice
405    ///
406    /// # Errors
407    ///
408    /// * [`Error::ShapeMismatch`] of the input if `values` does not have length
409    ///   `n_sampling_points`
410    /// * [`Error::ShapeMismatch`] of the output if `out` does not have length
411    ///   `basis_size`
412    /// * [`Error::DecompositionFailed`] if the singular value decomposition
413    ///   fails
414    ///
415    /// Nothing is written to `out` on an error.
416    pub fn fit_to(&self, values: &[f64], out: &mut [f64]) -> Result<(), Error> {
417        self.fitter.fit_to(None, values, out)
418    }
419
420    /// Evaluate complex basis coefficients at sampling points
421    ///
422    /// # Errors
423    ///
424    /// * [`Error::ShapeMismatch`] of the input if `coeffs` does not have length
425    ///   `basis_size`
426    pub fn evaluate_zz(&self, coeffs: &[Complex<f64>]) -> Result<Vec<Complex<f64>>, Error> {
427        self.fitter.evaluate(None, coeffs)
428    }
429
430    /// Evaluate complex basis coefficients, writing to output slice
431    ///
432    /// # Errors
433    ///
434    /// * [`Error::ShapeMismatch`] of the input if `coeffs` does not have length
435    ///   `basis_size`
436    /// * [`Error::ShapeMismatch`] of the output if `out` does not have length
437    ///   `n_sampling_points`
438    ///
439    /// Nothing is written to `out` on an error.
440    pub fn evaluate_zz_to(
441        &self,
442        coeffs: &[Complex<f64>],
443        out: &mut [Complex<f64>],
444    ) -> Result<(), Error> {
445        self.fitter.evaluate_to(None, coeffs, out)
446    }
447
448    /// Fit complex values at sampling points to basis coefficients
449    ///
450    /// # Errors
451    ///
452    /// * [`Error::ShapeMismatch`] of the input if `values` does not have length
453    ///   `n_sampling_points`
454    /// * [`Error::DecompositionFailed`] if the singular value decomposition
455    ///   fails
456    pub fn fit_zz(&self, values: &[Complex<f64>]) -> Result<Vec<Complex<f64>>, Error> {
457        self.fitter.fit(None, values)
458    }
459
460    /// Fit complex values, writing to output slice
461    ///
462    /// # Errors
463    ///
464    /// * [`Error::ShapeMismatch`] of the input if `values` does not have length
465    ///   `n_sampling_points`
466    /// * [`Error::ShapeMismatch`] of the output if `out` does not have length
467    ///   `basis_size`
468    /// * [`Error::DecompositionFailed`] if the singular value decomposition
469    ///   fails
470    ///
471    /// Nothing is written to `out` on an error.
472    pub fn fit_zz_to(
473        &self,
474        values: &[Complex<f64>],
475        out: &mut [Complex<f64>],
476    ) -> Result<(), Error> {
477        self.fitter.fit_to(None, values, out)
478    }
479
480    // ========================================================================
481    // N-D functions (real)
482    // ========================================================================
483
484    /// Evaluate N-D real coefficients at sampling points
485    ///
486    /// # Arguments
487    /// * `coeffs` - N-dimensional array with `coeffs.shape().dim(dim) == basis_size`
488    /// * `dim` - Dimension along which to evaluate (0-indexed)
489    ///
490    /// # Returns
491    /// N-dimensional array with `result.shape().dim(dim) == n_sampling_points`
492    ///
493    /// # Errors
494    ///
495    /// * [`Error::AxisOutOfRange`] if `dim` is not an axis of `coeffs`
496    /// * [`Error::ShapeMismatch`] of the input if `coeffs` does not have
497    ///   `basis_size` along `dim`
498    pub fn evaluate_nd(
499        &self,
500        backend: Option<&GemmBackendHandle>,
501        coeffs: &TypedTensor<f64>,
502        dim: usize,
503    ) -> Result<TypedTensor<f64>, Error> {
504        self.fitter.evaluate_nd(backend, coeffs, dim)
505    }
506
507    /// Evaluate N-D real coefficients, writing to a mutable view
508    ///
509    /// `out` must have the shape of `coeffs` with `n_sampling_points` along
510    /// `dim`.
511    ///
512    /// # Errors
513    ///
514    /// * [`Error::AxisOutOfRange`] if `dim` is not an axis of `coeffs`
515    /// * [`Error::ShapeMismatch`] of the input if `coeffs` does not have
516    ///   `basis_size` along `dim`, and of the output if `out` does not have
517    ///   the shape of `coeffs` with `n_sampling_points` along `dim`
518    ///
519    /// Nothing is written to `out` then.
520    pub fn evaluate_nd_to(
521        &self,
522        backend: Option<&GemmBackendHandle>,
523        coeffs: &TypedTensorView<'_, f64>,
524        dim: usize,
525        out: &mut TypedTensorViewMut<'_, f64>,
526    ) -> Result<(), Error> {
527        InplaceFitter::evaluate_nd_dd_to(self, backend, coeffs, dim, out)
528    }
529
530    /// Fit N-D real values at sampling points to basis coefficients
531    ///
532    /// # Arguments
533    /// * `values` - N-dimensional array with `values.shape().dim(dim) == n_sampling_points`
534    /// * `dim` - Dimension along which to fit (0-indexed)
535    ///
536    /// # Returns
537    /// N-dimensional array with `result.shape().dim(dim) == basis_size`
538    ///
539    /// # Errors
540    ///
541    /// * [`Error::AxisOutOfRange`] if `dim` is not an axis of `values`
542    /// * [`Error::ShapeMismatch`] of the input if `values` does not have
543    ///   `n_sampling_points` along `dim`
544    /// * [`Error::DecompositionFailed`] if the singular value decomposition
545    ///   fails
546    pub fn fit_nd(
547        &self,
548        backend: Option<&GemmBackendHandle>,
549        values: &TypedTensor<f64>,
550        dim: usize,
551    ) -> Result<TypedTensor<f64>, Error> {
552        self.fitter.fit_nd(backend, values, dim)
553    }
554
555    /// Fit N-D real values, writing to a mutable view
556    ///
557    /// `out` must have the shape of `values` with `basis_size` along `dim`.
558    ///
559    /// # Errors
560    ///
561    /// * [`Error::AxisOutOfRange`] if `dim` is not an axis of `values`
562    /// * [`Error::ShapeMismatch`] of the input if `values` does not have
563    ///   `n_sampling_points` along `dim`, and of the output if `out` does not have
564    ///   the shape of `values` with `basis_size` along `dim`
565    /// * [`Error::DecompositionFailed`] if the singular value decomposition
566    ///   fails
567    ///
568    /// Nothing is written to `out` then.
569    pub fn fit_nd_to(
570        &self,
571        backend: Option<&GemmBackendHandle>,
572        values: &TypedTensorView<'_, f64>,
573        dim: usize,
574        out: &mut TypedTensorViewMut<'_, f64>,
575    ) -> Result<(), Error> {
576        InplaceFitter::fit_nd_dd_to(self, backend, values, dim, out)
577    }
578
579    // ========================================================================
580    // N-D functions (complex)
581    // ========================================================================
582
583    /// Evaluate N-D complex coefficients at sampling points
584    ///
585    /// # Arguments
586    /// * `coeffs` - N-dimensional complex array with `coeffs.shape().dim(dim) == basis_size`
587    /// * `dim` - Dimension along which to evaluate (0-indexed)
588    ///
589    /// # Returns
590    /// N-dimensional complex array with `result.shape().dim(dim) == n_sampling_points`
591    ///
592    /// # Errors
593    ///
594    /// * [`Error::AxisOutOfRange`] if `dim` is not an axis of `coeffs`
595    /// * [`Error::ShapeMismatch`] of the input if `coeffs` does not have
596    ///   `basis_size` along `dim`
597    pub fn evaluate_nd_zz(
598        &self,
599        backend: Option<&GemmBackendHandle>,
600        coeffs: &TypedTensor<Complex<f64>>,
601        dim: usize,
602    ) -> Result<TypedTensor<Complex<f64>>, Error> {
603        self.fitter.evaluate_nd(backend, coeffs, dim)
604    }
605
606    /// Evaluate N-D complex coefficients, writing to a mutable view
607    ///
608    /// `out` must have the shape of `coeffs` with `n_sampling_points` along
609    /// `dim`.
610    ///
611    /// # Errors
612    ///
613    /// * [`Error::AxisOutOfRange`] if `dim` is not an axis of `coeffs`
614    /// * [`Error::ShapeMismatch`] of the input if `coeffs` does not have
615    ///   `basis_size` along `dim`, and of the output if `out` does not have
616    ///   the shape of `coeffs` with `n_sampling_points` along `dim`
617    ///
618    /// Nothing is written to `out` then.
619    pub fn evaluate_nd_zz_to(
620        &self,
621        backend: Option<&GemmBackendHandle>,
622        coeffs: &TypedTensorView<'_, Complex<f64>>,
623        dim: usize,
624        out: &mut TypedTensorViewMut<'_, Complex<f64>>,
625    ) -> Result<(), Error> {
626        InplaceFitter::evaluate_nd_zz_to(self, backend, coeffs, dim, out)
627    }
628
629    /// Fit N-D complex values at sampling points to basis coefficients
630    ///
631    /// # Arguments
632    /// * `values` - N-dimensional complex array with `values.shape().dim(dim) == n_sampling_points`
633    /// * `dim` - Dimension along which to fit (0-indexed)
634    ///
635    /// # Returns
636    /// N-dimensional complex array with `result.shape().dim(dim) == basis_size`
637    ///
638    /// # Errors
639    ///
640    /// * [`Error::AxisOutOfRange`] if `dim` is not an axis of `values`
641    /// * [`Error::ShapeMismatch`] of the input if `values` does not have
642    ///   `n_sampling_points` along `dim`
643    /// * [`Error::DecompositionFailed`] if the singular value decomposition
644    ///   fails
645    pub fn fit_nd_zz(
646        &self,
647        backend: Option<&GemmBackendHandle>,
648        values: &TypedTensor<Complex<f64>>,
649        dim: usize,
650    ) -> Result<TypedTensor<Complex<f64>>, Error> {
651        self.fitter.fit_nd(backend, values, dim)
652    }
653
654    /// Fit N-D complex values, writing to a mutable view
655    ///
656    /// `out` must have the shape of `values` with `basis_size` along `dim`.
657    ///
658    /// # Errors
659    ///
660    /// * [`Error::AxisOutOfRange`] if `dim` is not an axis of `values`
661    /// * [`Error::ShapeMismatch`] of the input if `values` does not have
662    ///   `n_sampling_points` along `dim`, and of the output if `out` does not have
663    ///   the shape of `values` with `basis_size` along `dim`
664    /// * [`Error::DecompositionFailed`] if the singular value decomposition
665    ///   fails
666    ///
667    /// Nothing is written to `out` then.
668    pub fn fit_nd_zz_to(
669        &self,
670        backend: Option<&GemmBackendHandle>,
671        values: &TypedTensorView<'_, Complex<f64>>,
672        dim: usize,
673        out: &mut TypedTensorViewMut<'_, Complex<f64>>,
674    ) -> Result<(), Error> {
675        InplaceFitter::fit_nd_zz_to(self, backend, values, dim, out)
676    }
677}
678
679/// InplaceFitter implementation for TauSampling
680///
681/// Delegates to RealMatrixFitter which supports dd and zz operations.
682impl<S: StatisticsType> InplaceFitter for TauSampling<S> {
683    fn n_points(&self) -> usize {
684        self.n_sampling_points()
685    }
686
687    fn basis_size(&self) -> usize {
688        self.basis_size()
689    }
690
691    fn evaluate_nd_dd_to(
692        &self,
693        backend: Option<&GemmBackendHandle>,
694        coeffs: &TypedTensorView<'_, f64>,
695        dim: usize,
696        out: &mut TypedTensorViewMut<'_, f64>,
697    ) -> Result<(), Error> {
698        self.fitter.evaluate_nd_dd_to(backend, coeffs, dim, out)
699    }
700
701    fn evaluate_nd_zz_to(
702        &self,
703        backend: Option<&GemmBackendHandle>,
704        coeffs: &TypedTensorView<'_, Complex<f64>>,
705        dim: usize,
706        out: &mut TypedTensorViewMut<'_, Complex<f64>>,
707    ) -> Result<(), Error> {
708        self.fitter.evaluate_nd_zz_to(backend, coeffs, dim, out)
709    }
710
711    fn fit_nd_dd_to(
712        &self,
713        backend: Option<&GemmBackendHandle>,
714        values: &TypedTensorView<'_, f64>,
715        dim: usize,
716        out: &mut TypedTensorViewMut<'_, f64>,
717    ) -> Result<(), Error> {
718        self.fitter.fit_nd_dd_to(backend, values, dim, out)
719    }
720
721    fn fit_nd_zz_to(
722        &self,
723        backend: Option<&GemmBackendHandle>,
724        values: &TypedTensorView<'_, Complex<f64>>,
725        dim: usize,
726        out: &mut TypedTensorViewMut<'_, Complex<f64>>,
727    ) -> Result<(), Error> {
728        self.fitter.fit_nd_zz_to(backend, values, dim, out)
729    }
730}