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}