1use 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
18pub trait InplaceFitter {
49 fn n_points(&self) -> usize;
51
52 fn basis_size(&self) -> usize;
54
55 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 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 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 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 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 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 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 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
168fn not_supported(operation: &str) -> Error {
170 Error::NotSupported {
171 what: format!("{operation} for this sampling"),
172 }
173}
174
175fn layout_error(what: impl Into<String>) -> Error {
178 Error::NotSupported { what: what.into() }
179}
180
181pub trait FitScalar: GemmScalar {
191 const REAL_WIDTH: usize;
193 fn as_f64(s: &[Self]) -> &[f64];
195 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 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 unsafe { std::slice::from_raw_parts_mut(s.as_mut_ptr().cast::<f64>(), 2 * s.len()) }
220 }
221}
222
223#[derive(Clone, Copy, Debug, PartialEq, Eq)]
229pub(crate) struct AxisSplit {
230 pub pre: usize,
231 pub post: usize,
232}
233
234pub(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
262pub(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
289pub(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
306pub(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
315pub(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
322pub(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
347fn 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 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 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
393pub(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
414pub(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
438pub(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 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
462pub(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
487pub(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#[doc(hidden)]
520pub struct PinvFactors<T> {
521 pub uh: Vec<T>,
523 pub v_scaled: Vec<T>,
525 pub s: Vec<f64>,
528 pub n: usize,
529 pub m: usize,
530 pub rank: usize,
531}
532
533impl<T: GemmScalar> PinvFactors<T> {
534 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
562pub 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#[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#[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
618pub(crate) fn sampling_matrix((rows, cols): (usize, usize)) -> String {
620 format!("{rows} x {cols} sampling matrix")
621}
622
623fn 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
631pub(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 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 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 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
722pub 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
743pub(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}