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}