Skip to main content

tensor4all_core/
traits.rs

1//! Abstract traits for matrix cross interpolation.
2//!
3//! The [`AbstractMatrixCI`] trait provides a common interface for all matrix
4//! cross interpolation implementations ([`MatrixLUCI`](crate::MatrixLUCI),
5//! [`MatrixACA`](crate::MatrixACA)).
6
7use crate::error::Result;
8use crate::scalar::Scalar;
9use tensor4all_tensorbackend::{submatrix, Matrix};
10
11/// Common interface for matrix cross interpolation objects.
12///
13/// Implementors provide low-rank approximations of matrices via
14/// selected pivot rows and columns. This trait unifies the API for
15/// [`MatrixLUCI`](crate::MatrixLUCI) and [`MatrixACA`](crate::MatrixACA).
16///
17/// # Examples
18///
19/// ```
20/// use tensor4all_core::{AbstractMatrixCI, MatrixLUCI};
21/// use tensor4all_tensorbackend::from_vec2d;
22///
23/// let m = from_vec2d(vec![
24///     vec![1.0_f64, 2.0],
25///     vec![3.0, 4.0],
26/// ]);
27/// let ci = MatrixLUCI::from_matrix(&m, None).unwrap();
28///
29/// // All AbstractMatrixCI methods are available:
30/// assert_eq!(ci.nrows(), 2);
31/// assert_eq!(ci.ncols(), 2);
32/// assert!(ci.rank() >= 1);
33/// assert!(!ci.is_empty());
34///
35/// // Full reconstruction
36/// let full = ci.to_matrix();
37/// for i in 0..2 {
38///     for j in 0..2 {
39///         assert!((full[[i, j]] - m[[i, j]]).abs() < 1e-10);
40///     }
41/// }
42/// ```
43pub trait AbstractMatrixCI<T: Scalar>: Sized {
44    /// Number of rows in the approximated matrix
45    ///
46    /// # Examples
47    ///
48    /// ```
49    /// use tensor4all_core::{AbstractMatrixCI, MatrixLUCI};
50    /// use tensor4all_tensorbackend::from_vec2d;
51    ///
52    /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0], vec![5.0, 6.0]]);
53    /// let ci = MatrixLUCI::from_matrix(&m, None).unwrap();
54    /// assert_eq!(ci.nrows(), 3);
55    /// ```
56    fn nrows(&self) -> usize;
57
58    /// Number of columns in the approximated matrix
59    ///
60    /// # Examples
61    ///
62    /// ```
63    /// use tensor4all_core::{AbstractMatrixCI, MatrixLUCI};
64    /// use tensor4all_tensorbackend::from_vec2d;
65    ///
66    /// let m = from_vec2d(vec![vec![1.0_f64, 2.0, 3.0]]);
67    /// let ci = MatrixLUCI::from_matrix(&m, None).unwrap();
68    /// assert_eq!(ci.ncols(), 3);
69    /// ```
70    fn ncols(&self) -> usize;
71
72    /// Current rank of the approximation (number of pivots)
73    ///
74    /// # Examples
75    ///
76    /// ```
77    /// use tensor4all_core::{AbstractMatrixCI, MatrixLUCI};
78    /// use tensor4all_tensorbackend::from_vec2d;
79    ///
80    /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
81    /// let ci = MatrixLUCI::from_matrix(&m, None).unwrap();
82    /// assert!(ci.rank() >= 1);
83    /// assert!(ci.rank() <= 2);
84    /// ```
85    fn rank(&self) -> usize;
86
87    /// Row indices selected as pivots (I set)
88    ///
89    /// # Examples
90    ///
91    /// ```
92    /// use tensor4all_core::{AbstractMatrixCI, MatrixLUCI};
93    /// use tensor4all_tensorbackend::from_vec2d;
94    ///
95    /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
96    /// let ci = MatrixLUCI::from_matrix(&m, None).unwrap();
97    /// let rows = ci.row_indices();
98    /// assert_eq!(rows.len(), ci.rank());
99    /// for &r in rows {
100    ///     assert!(r < ci.nrows());
101    /// }
102    /// ```
103    fn row_indices(&self) -> &[usize];
104
105    /// Column indices selected as pivots (J set)
106    ///
107    /// # Examples
108    ///
109    /// ```
110    /// use tensor4all_core::{AbstractMatrixCI, MatrixLUCI};
111    /// use tensor4all_tensorbackend::from_vec2d;
112    ///
113    /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
114    /// let ci = MatrixLUCI::from_matrix(&m, None).unwrap();
115    /// let cols = ci.col_indices();
116    /// assert_eq!(cols.len(), ci.rank());
117    /// for &c in cols {
118    ///     assert!(c < ci.ncols());
119    /// }
120    /// ```
121    fn col_indices(&self) -> &[usize];
122
123    /// Check if the approximation is empty (no pivots)
124    ///
125    /// # Examples
126    ///
127    /// ```
128    /// use tensor4all_core::{AbstractMatrixCI, MatrixACA};
129    ///
130    /// let aca = MatrixACA::<f64>::new(2, 2);
131    /// assert!(aca.is_empty());
132    /// ```
133    fn is_empty(&self) -> bool {
134        self.rank() == 0
135    }
136
137    /// Evaluate the approximation at position (i, j)
138    ///
139    /// # Examples
140    ///
141    /// ```
142    /// use tensor4all_core::{AbstractMatrixCI, MatrixLUCI};
143    /// use tensor4all_tensorbackend::from_vec2d;
144    ///
145    /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
146    /// let ci = MatrixLUCI::from_matrix(&m, None).unwrap();
147    /// assert!((ci.evaluate(0, 0) - 1.0).abs() < 1e-10);
148    /// assert!((ci.evaluate(1, 1) - 4.0).abs() < 1e-10);
149    /// ```
150    fn evaluate(&self, i: usize, j: usize) -> T;
151
152    /// Get a submatrix of the approximation
153    ///
154    /// # Examples
155    ///
156    /// ```
157    /// use tensor4all_core::{AbstractMatrixCI, MatrixLUCI};
158    /// use tensor4all_tensorbackend::from_vec2d;
159    ///
160    /// let m = from_vec2d(vec![
161    ///     vec![1.0_f64, 2.0, 3.0],
162    ///     vec![4.0, 5.0, 6.0],
163    ///     vec![7.0, 8.0, 9.0],
164    /// ]);
165    /// let ci = MatrixLUCI::from_matrix(&m, None).unwrap();
166    /// let sub = ci.submatrix(&[0, 2], &[1]);
167    /// assert_eq!(sub.nrows(), 2);
168    /// assert_eq!(sub.ncols(), 1);
169    /// assert!((sub[[0, 0]] - 2.0).abs() < 1e-10);
170    /// assert!((sub[[1, 0]] - 8.0).abs() < 1e-10);
171    /// ```
172    fn submatrix(&self, rows: &[usize], cols: &[usize]) -> Matrix<T>;
173
174    /// Get a row of the approximation
175    ///
176    /// # Examples
177    ///
178    /// ```
179    /// use tensor4all_core::{AbstractMatrixCI, MatrixLUCI};
180    /// use tensor4all_tensorbackend::from_vec2d;
181    ///
182    /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
183    /// let ci = MatrixLUCI::from_matrix(&m, None).unwrap();
184    /// let row0 = ci.row(0);
185    /// assert_eq!(row0.len(), 2);
186    /// assert!((row0[0] - 1.0).abs() < 1e-10);
187    /// assert!((row0[1] - 2.0).abs() < 1e-10);
188    /// ```
189    fn row(&self, i: usize) -> Vec<T> {
190        let cols: Vec<usize> = (0..self.ncols()).collect();
191        let sub = self.submatrix(&[i], &cols);
192        (0..self.ncols()).map(|j| sub[[0, j]]).collect()
193    }
194
195    /// Get a column of the approximation
196    ///
197    /// # Examples
198    ///
199    /// ```
200    /// use tensor4all_core::{AbstractMatrixCI, MatrixLUCI};
201    /// use tensor4all_tensorbackend::from_vec2d;
202    ///
203    /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
204    /// let ci = MatrixLUCI::from_matrix(&m, None).unwrap();
205    /// let col1 = ci.col(1);
206    /// assert_eq!(col1.len(), 2);
207    /// assert!((col1[0] - 2.0).abs() < 1e-10);
208    /// assert!((col1[1] - 4.0).abs() < 1e-10);
209    /// ```
210    fn col(&self, j: usize) -> Vec<T> {
211        let rows: Vec<usize> = (0..self.nrows()).collect();
212        let sub = self.submatrix(&rows, &[j]);
213        (0..self.nrows()).map(|i| sub[[i, 0]]).collect()
214    }
215
216    /// Get the full approximated matrix
217    ///
218    /// # Examples
219    ///
220    /// ```
221    /// use tensor4all_core::{AbstractMatrixCI, MatrixLUCI};
222    /// use tensor4all_tensorbackend::from_vec2d;
223    ///
224    /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
225    /// let ci = MatrixLUCI::from_matrix(&m, None).unwrap();
226    /// let full = ci.to_matrix();
227    /// assert_eq!(full.nrows(), 2);
228    /// assert_eq!(full.ncols(), 2);
229    /// for i in 0..2 {
230    ///     for j in 0..2 {
231    ///         assert!((full[[i, j]] - m[[i, j]]).abs() < 1e-10);
232    ///     }
233    /// }
234    /// ```
235    fn to_matrix(&self) -> Matrix<T> {
236        let rows: Vec<usize> = (0..self.nrows()).collect();
237        let cols: Vec<usize> = (0..self.ncols()).collect();
238        self.submatrix(&rows, &cols)
239    }
240
241    /// Get available row indices (rows without pivots)
242    ///
243    /// # Examples
244    ///
245    /// ```
246    /// use tensor4all_core::{AbstractMatrixCI, MatrixACA};
247    /// use tensor4all_tensorbackend::from_vec2d;
248    ///
249    /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0], vec![5.0, 6.0]]);
250    /// let aca = MatrixACA::from_matrix_with_pivot(&m, (1, 0)).unwrap();
251    /// let avail = aca.available_rows();
252    /// // Row 1 was used as pivot, so 0 and 2 remain
253    /// assert_eq!(avail, vec![0, 2]);
254    /// ```
255    fn available_rows(&self) -> Vec<usize> {
256        let pivot_rows: std::collections::HashSet<usize> =
257            self.row_indices().iter().copied().collect();
258        (0..self.nrows())
259            .filter(|i| !pivot_rows.contains(i))
260            .collect()
261    }
262
263    /// Get available column indices (columns without pivots)
264    ///
265    /// # Examples
266    ///
267    /// ```
268    /// use tensor4all_core::{AbstractMatrixCI, MatrixACA};
269    /// use tensor4all_tensorbackend::from_vec2d;
270    ///
271    /// let m = from_vec2d(vec![vec![1.0_f64, 2.0, 3.0], vec![4.0, 5.0, 6.0]]);
272    /// let aca = MatrixACA::from_matrix_with_pivot(&m, (0, 1)).unwrap();
273    /// let avail = aca.available_cols();
274    /// // Column 1 was used as pivot, so 0 and 2 remain
275    /// assert_eq!(avail, vec![0, 2]);
276    /// ```
277    fn available_cols(&self) -> Vec<usize> {
278        let pivot_cols: std::collections::HashSet<usize> =
279            self.col_indices().iter().copied().collect();
280        (0..self.ncols())
281            .filter(|j| !pivot_cols.contains(j))
282            .collect()
283    }
284
285    /// Compute local error |A - CI| for given indices
286    ///
287    /// # Examples
288    ///
289    /// ```
290    /// use tensor4all_core::{AbstractMatrixCI, MatrixACA};
291    /// use tensor4all_tensorbackend::from_vec2d;
292    ///
293    /// let m = from_vec2d(vec![
294    ///     vec![1.0_f64, 2.0, 3.0],
295    ///     vec![4.0, 5.0, 6.0],
296    ///     vec![7.0, 8.0, 10.0],
297    /// ]);
298    /// let aca = MatrixACA::from_matrix_with_pivot(&m, (0, 0)).unwrap();
299    /// let err = aca.local_error(&m, &[1, 2], &[1, 2]);
300    /// // Error at pivot position (0,0) would be zero; off-pivot may be non-zero
301    /// assert_eq!(err.nrows(), 2);
302    /// assert_eq!(err.ncols(), 2);
303    /// ```
304    fn local_error(&self, a: &Matrix<T>, rows: &[usize], cols: &[usize]) -> Matrix<T>
305    where
306        T: std::ops::Sub<Output = T>,
307    {
308        let sub_a = submatrix(a, rows, cols);
309        let sub_ci = self.submatrix(rows, cols);
310
311        let mut result = Matrix::zeros(rows.len(), cols.len());
312        for i in 0..rows.len() {
313            for j in 0..cols.len() {
314                let diff = sub_a[[i, j]] - sub_ci[[i, j]];
315                result[[i, j]] = diff.abs();
316            }
317        }
318        result
319    }
320
321    /// Find a new pivot that maximizes the local error
322    ///
323    /// # Errors
324    ///
325    /// Returns an error when the operation fails (a shape or index mismatch, or
326    /// /// a backend failure).
327    ///
328    /// # Examples
329    ///
330    /// ```
331    /// use tensor4all_core::{AbstractMatrixCI, MatrixACA};
332    /// use tensor4all_tensorbackend::from_vec2d;
333    ///
334    /// let m = from_vec2d(vec![
335    ///     vec![1.0_f64, 2.0, 3.0],
336    ///     vec![4.0, 5.0, 6.0],
337    ///     vec![7.0, 8.0, 10.0],
338    /// ]);
339    /// let aca = MatrixACA::from_matrix_with_pivot(&m, (0, 0)).unwrap();
340    /// let ((r, c), err_val) = aca.find_new_pivot(&m).unwrap();
341    /// // New pivot must be in available rows/cols (not row 0 or col 0)
342    /// assert_ne!(r, 0);
343    /// assert_ne!(c, 0);
344    /// ```
345    fn find_new_pivot(&self, a: &Matrix<T>) -> Result<((usize, usize), T)>
346    where
347        T: std::ops::Sub<Output = T>,
348    {
349        let avail_rows = self.available_rows();
350        let avail_cols = self.available_cols();
351
352        self.find_new_pivot_in(a, &avail_rows, &avail_cols)
353    }
354
355    /// Find a new pivot in the given row/column subsets
356    ///
357    /// # Errors
358    ///
359    /// Returns an error when the operation fails (a shape or index mismatch, or
360    /// /// a backend failure).
361    ///
362    /// # Examples
363    ///
364    /// ```
365    /// use tensor4all_core::{AbstractMatrixCI, MatrixACA};
366    /// use tensor4all_tensorbackend::from_vec2d;
367    ///
368    /// let m = from_vec2d(vec![
369    ///     vec![1.0_f64, 2.0, 3.0],
370    ///     vec![4.0, 5.0, 6.0],
371    ///     vec![7.0, 8.0, 10.0],
372    /// ]);
373    /// let aca = MatrixACA::from_matrix_with_pivot(&m, (0, 0)).unwrap();
374    /// // Search only in rows [1,2] and cols [1,2]
375    /// let ((r, c), _) = aca.find_new_pivot_in(&m, &[1, 2], &[1, 2]).unwrap();
376    /// assert!(r == 1 || r == 2);
377    /// assert!(c == 1 || c == 2);
378    /// ```
379    fn find_new_pivot_in(
380        &self,
381        a: &Matrix<T>,
382        rows: &[usize],
383        cols: &[usize],
384    ) -> Result<((usize, usize), T)>
385    where
386        T: std::ops::Sub<Output = T>,
387    {
388        use crate::error::MatrixCIError;
389
390        if self.rank() == self.nrows().min(self.ncols()) {
391            return Err(MatrixCIError::FullRank);
392        }
393
394        if rows.is_empty() {
395            return Err(MatrixCIError::EmptyIndexSet {
396                dimension: "rows".to_string(),
397            });
398        }
399
400        if cols.is_empty() {
401            return Err(MatrixCIError::EmptyIndexSet {
402                dimension: "cols".to_string(),
403            });
404        }
405
406        let errors = self.local_error(a, rows, cols);
407
408        // Find maximum error position (comparing by abs_sq which returns f64)
409        let mut max_val_sq: f64 = errors[[0, 0]].abs_sq();
410        let mut max_i = 0;
411        let mut max_j = 0;
412
413        for i in 0..rows.len() {
414            for j in 0..cols.len() {
415                let val_sq: f64 = errors[[i, j]].abs_sq();
416                if val_sq > max_val_sq {
417                    max_val_sq = val_sq;
418                    max_i = i;
419                    max_j = j;
420                }
421            }
422        }
423
424        Ok(((rows[max_i], cols[max_j]), errors[[max_i, max_j]]))
425    }
426}
427
428#[cfg(test)]
429mod tests {
430    use super::*;
431    use crate::error::MatrixCIError;
432    use tensor4all_tensorbackend::from_vec2d;
433
434    struct ExactCi {
435        matrix: Matrix<f64>,
436        row_indices: Vec<usize>,
437        col_indices: Vec<usize>,
438    }
439
440    impl ExactCi {
441        fn new(matrix: Matrix<f64>, row_indices: Vec<usize>, col_indices: Vec<usize>) -> Self {
442            Self {
443                matrix,
444                row_indices,
445                col_indices,
446            }
447        }
448    }
449
450    impl AbstractMatrixCI<f64> for ExactCi {
451        fn nrows(&self) -> usize {
452            self.matrix.nrows()
453        }
454
455        fn ncols(&self) -> usize {
456            self.matrix.ncols()
457        }
458
459        fn rank(&self) -> usize {
460            self.row_indices.len()
461        }
462
463        fn row_indices(&self) -> &[usize] {
464            &self.row_indices
465        }
466
467        fn col_indices(&self) -> &[usize] {
468            &self.col_indices
469        }
470
471        fn evaluate(&self, i: usize, j: usize) -> f64 {
472            self.matrix[[i, j]]
473        }
474
475        fn submatrix(&self, rows: &[usize], cols: &[usize]) -> Matrix<f64> {
476            tensor4all_tensorbackend::submatrix(&self.matrix, rows, cols)
477        }
478    }
479
480    fn sample_matrix() -> Matrix<f64> {
481        from_vec2d(vec![
482            vec![1.0_f64, 2.0, 3.0],
483            vec![4.0, 5.0, 6.0],
484            vec![7.0, 8.0, 10.0],
485        ])
486    }
487
488    #[test]
489    fn default_helpers_return_expected_views_and_available_indices() {
490        let matrix = sample_matrix();
491        let ci = ExactCi::new(matrix.clone(), vec![1], vec![2]);
492
493        assert!(!ci.is_empty());
494        assert_eq!(ci.row(2), vec![7.0, 8.0, 10.0]);
495        assert_eq!(ci.col(1), vec![2.0, 5.0, 8.0]);
496        assert_eq!(ci.available_rows(), vec![0, 2]);
497        assert_eq!(ci.available_cols(), vec![0, 1]);
498
499        let full = ci.to_matrix();
500        assert_eq!(full.nrows(), matrix.nrows());
501        assert_eq!(full.ncols(), matrix.ncols());
502        for i in 0..matrix.nrows() {
503            for j in 0..matrix.ncols() {
504                assert_eq!(full[[i, j]], matrix[[i, j]]);
505            }
506        }
507
508        let local_error = ci.local_error(&matrix, &[0, 2], &[0, 1]);
509        for i in 0..local_error.nrows() {
510            for j in 0..local_error.ncols() {
511                assert_eq!(local_error[[i, j]], 0.0);
512            }
513        }
514    }
515
516    #[test]
517    fn find_new_pivot_covers_success_and_error_paths() {
518        let ci = ExactCi::new(sample_matrix(), vec![1], vec![2]);
519        let mut perturbed = sample_matrix();
520        perturbed[[2, 1]] += 10.0;
521
522        let ((row, col), error) = ci.find_new_pivot(&perturbed).unwrap();
523
524        assert_eq!((row, col), (2, 1));
525        assert_eq!(error, 10.0);
526
527        let ((row, col), error) = ci.find_new_pivot_in(&perturbed, &[0, 2], &[0, 1]).unwrap();
528
529        assert_eq!((row, col), (2, 1));
530        assert_eq!(error, 10.0);
531
532        assert!(matches!(
533            ci.find_new_pivot_in(&perturbed, &[], &[0]),
534            Err(MatrixCIError::EmptyIndexSet { dimension }) if dimension == "rows"
535        ));
536        assert!(matches!(
537            ci.find_new_pivot_in(&perturbed, &[0], &[]),
538            Err(MatrixCIError::EmptyIndexSet { dimension }) if dimension == "cols"
539        ));
540
541        let full_rank = ExactCi::new(sample_matrix(), vec![0, 1, 2], vec![0, 1, 2]);
542        assert!(matches!(
543            full_rank.find_new_pivot_in(&perturbed, &[0], &[0]),
544            Err(MatrixCIError::FullRank)
545        ));
546    }
547}