tensor4all_core/matrixlu.rs
1//! Rank-Revealing LU decomposition (rrLU) implementation.
2//!
3//! Provides [`RrLU`], a full-pivoting LU decomposition that reveals the
4//! numerical rank of a matrix. The decomposition is:
5//!
6//! ```text
7//! P_row * A * P_col = L * U
8//! ```
9//!
10//! where `P_row`, `P_col` are permutation matrices. The rank is determined
11//! by the number of pivots exceeding the tolerance thresholds in
12//! [`RrLUOptions`].
13//!
14//! # Examples
15//!
16//! ```
17//! use tensor4all_core::matrixlu::rrlu;
18//! use tensor4all_tensorbackend::from_vec2d;
19//!
20//! let m = from_vec2d(vec![
21//! vec![1.0_f64, 2.0],
22//! vec![3.0, 4.0],
23//! ]);
24//! let lu = rrlu(&m, None).unwrap();
25//! assert_eq!(lu.npivots(), 2);
26//! ```
27
28use crate::error::{MatrixCIError, Result};
29use crate::scalar::Scalar;
30use tensor4all_tensorbackend::{transpose, Matrix};
31
32/// Rank-Revealing LU decomposition.
33///
34/// Represents a matrix `A` as `P_row * A * P_col = L * U`, where `P_row`
35/// and `P_col` are permutation matrices, `L` is lower-triangular, and `U`
36/// is upper-triangular. One of `L` or `U` has unit diagonal, controlled by
37/// the `left_orthogonal` option.
38///
39/// # Examples
40///
41/// ```
42/// use tensor4all_core::matrixlu::rrlu;
43/// use tensor4all_tensorbackend::{from_vec2d, mat_mul};
44///
45/// let m = from_vec2d(vec![
46/// vec![1.0_f64, 2.0, 3.0],
47/// vec![4.0, 5.0, 6.0],
48/// vec![7.0, 8.0, 10.0],
49/// ]);
50///
51/// let lu = rrlu(&m, None).unwrap();
52/// assert_eq!(lu.npivots(), 3);
53///
54/// // Verify L * U reconstructs the permuted matrix
55/// let l = lu.left(false);
56/// let u = lu.right(false);
57/// let reconstructed = mat_mul(&l, &u).unwrap();
58///
59/// // Check reconstruction matches the permuted matrix
60/// for i in 0..3 {
61/// for j in 0..3 {
62/// let orig_row = lu.row_permutation()[i];
63/// let orig_col = lu.col_permutation()[j];
64/// assert!((reconstructed[[i, j]] - m[[orig_row, orig_col]]).abs() < 1e-10);
65/// }
66/// }
67/// ```
68#[derive(Debug, Clone)]
69pub struct RrLU<T: Scalar> {
70 /// Row permutation
71 row_permutation: Vec<usize>,
72 /// Column permutation
73 col_permutation: Vec<usize>,
74 /// Lower triangular matrix L
75 l: Matrix<T>,
76 /// Upper triangular matrix U
77 u: Matrix<T>,
78 /// Whether L is left-orthogonal (L has 1s on diagonal) or U is (U has 1s on diagonal)
79 left_orthogonal: bool,
80 /// Number of pivots
81 n_pivot: usize,
82 /// Last pivot error
83 error: f64,
84}
85
86impl<T: Scalar> RrLU<T> {
87 /// Create an empty rrLU for a matrix of given size.
88 ///
89 /// Used internally. Most users should call [`rrlu`] or [`rrlu_mut`]
90 /// instead.
91 ///
92 /// # Examples
93 ///
94 /// ```
95 /// use tensor4all_core::RrLU;
96 ///
97 /// let lu = RrLU::<f64>::new(3, 4, true);
98 /// assert_eq!(lu.nrows(), 3);
99 /// assert_eq!(lu.ncols(), 4);
100 /// assert_eq!(lu.npivots(), 0);
101 /// assert!(lu.is_left_orthogonal());
102 /// ```
103 pub fn new(nr: usize, nc: usize, left_orthogonal: bool) -> Self {
104 Self {
105 row_permutation: (0..nr).collect(),
106 col_permutation: (0..nc).collect(),
107 l: Matrix::zeros(nr, 0),
108 u: Matrix::zeros(0, nc),
109 left_orthogonal,
110 n_pivot: 0,
111 error: f64::NAN,
112 }
113 }
114
115 /// Number of rows
116 ///
117 /// # Examples
118 ///
119 /// ```
120 /// use tensor4all_core::matrixlu::rrlu;
121 /// use tensor4all_tensorbackend::from_vec2d;
122 ///
123 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0], vec![5.0, 6.0]]);
124 /// let lu = rrlu(&m, None).unwrap();
125 /// assert_eq!(lu.nrows(), 3);
126 /// ```
127 pub fn nrows(&self) -> usize {
128 self.l.nrows()
129 }
130
131 /// Number of columns
132 ///
133 /// # Examples
134 ///
135 /// ```
136 /// use tensor4all_core::matrixlu::rrlu;
137 /// use tensor4all_tensorbackend::from_vec2d;
138 ///
139 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0, 3.0], vec![4.0, 5.0, 6.0]]);
140 /// let lu = rrlu(&m, None).unwrap();
141 /// assert_eq!(lu.ncols(), 3);
142 /// ```
143 pub fn ncols(&self) -> usize {
144 self.u.ncols()
145 }
146
147 /// Number of pivots
148 ///
149 /// # Examples
150 ///
151 /// ```
152 /// use tensor4all_core::matrixlu::rrlu;
153 /// use tensor4all_tensorbackend::from_vec2d;
154 ///
155 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
156 /// let lu = rrlu(&m, None).unwrap();
157 /// assert_eq!(lu.npivots(), 2);
158 /// ```
159 pub fn npivots(&self) -> usize {
160 self.n_pivot
161 }
162
163 /// Row permutation
164 ///
165 /// # Examples
166 ///
167 /// ```
168 /// use tensor4all_core::matrixlu::rrlu;
169 /// use tensor4all_tensorbackend::from_vec2d;
170 ///
171 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
172 /// let lu = rrlu(&m, None).unwrap();
173 /// let perm = lu.row_permutation();
174 /// assert_eq!(perm.len(), 2);
175 /// // Permutation is a rearrangement of 0..nrows
176 /// let mut sorted = perm.to_vec();
177 /// sorted.sort();
178 /// assert_eq!(sorted, vec![0, 1]);
179 /// ```
180 pub fn row_permutation(&self) -> &[usize] {
181 &self.row_permutation
182 }
183
184 /// Column permutation
185 ///
186 /// # Examples
187 ///
188 /// ```
189 /// use tensor4all_core::matrixlu::rrlu;
190 /// use tensor4all_tensorbackend::from_vec2d;
191 ///
192 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
193 /// let lu = rrlu(&m, None).unwrap();
194 /// let perm = lu.col_permutation();
195 /// assert_eq!(perm.len(), 2);
196 /// let mut sorted = perm.to_vec();
197 /// sorted.sort();
198 /// assert_eq!(sorted, vec![0, 1]);
199 /// ```
200 pub fn col_permutation(&self) -> &[usize] {
201 &self.col_permutation
202 }
203
204 /// Get row indices (selected pivots)
205 ///
206 /// # Examples
207 ///
208 /// ```
209 /// use tensor4all_core::{matrixlu::rrlu, RrLUOptions};
210 /// use tensor4all_tensorbackend::from_vec2d;
211 ///
212 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
213 /// let lu = rrlu(&m, Some(RrLUOptions { max_bond_dim: 1, ..Default::default() })).unwrap();
214 /// let rows = lu.row_indices();
215 /// assert_eq!(rows.len(), 1);
216 /// assert!(rows[0] < 2);
217 /// ```
218 pub fn row_indices(&self) -> Vec<usize> {
219 self.row_permutation[0..self.n_pivot].to_vec()
220 }
221
222 /// Get column indices (selected pivots)
223 ///
224 /// # Examples
225 ///
226 /// ```
227 /// use tensor4all_core::{matrixlu::rrlu, RrLUOptions};
228 /// use tensor4all_tensorbackend::from_vec2d;
229 ///
230 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
231 /// let lu = rrlu(&m, Some(RrLUOptions { max_bond_dim: 1, ..Default::default() })).unwrap();
232 /// let cols = lu.col_indices();
233 /// assert_eq!(cols.len(), 1);
234 /// assert!(cols[0] < 2);
235 /// ```
236 pub fn col_indices(&self) -> Vec<usize> {
237 self.col_permutation[0..self.n_pivot].to_vec()
238 }
239
240 /// Get left matrix (optionally permuted)
241 ///
242 /// # Examples
243 ///
244 /// ```
245 /// use tensor4all_core::matrixlu::rrlu;
246 /// use tensor4all_tensorbackend::{from_vec2d, mat_mul};
247 ///
248 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
249 /// let lu = rrlu(&m, None).unwrap();
250 ///
251 /// // Unpermuted: L * U reconstructs the row/col-permuted matrix
252 /// let l = lu.left(false);
253 /// let u = lu.right(false);
254 /// let prod = mat_mul(&l, &u).unwrap();
255 /// for i in 0..2 {
256 /// for j in 0..2 {
257 /// let ri = lu.row_permutation()[i];
258 /// let cj = lu.col_permutation()[j];
259 /// assert!((prod[[i, j]] - m[[ri, cj]]).abs() < 1e-10);
260 /// }
261 /// }
262 /// ```
263 pub fn left(&self, permute: bool) -> Matrix<T> {
264 if permute {
265 let mut result = Matrix::zeros(self.l.nrows(), self.l.ncols());
266 for j in 0..self.l.ncols() {
267 for (new_i, &old_i) in self.row_permutation.iter().enumerate() {
268 result[[old_i, j]] = self.l[[new_i, j]];
269 }
270 }
271 result
272 } else {
273 self.l.clone()
274 }
275 }
276
277 /// Get right matrix (optionally permuted)
278 ///
279 /// # Examples
280 ///
281 /// ```
282 /// use tensor4all_core::matrixlu::rrlu;
283 /// use tensor4all_tensorbackend::{from_vec2d, mat_mul};
284 ///
285 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
286 /// let lu = rrlu(&m, None).unwrap();
287 /// let l = lu.left(false);
288 /// let u = lu.right(false);
289 /// assert_eq!(u.nrows(), lu.npivots());
290 /// assert_eq!(u.ncols(), lu.ncols());
291 /// // L * U reconstructs the permuted matrix
292 /// let prod = mat_mul(&l, &u).unwrap();
293 /// assert!((prod[[0, 0]] - m[[lu.row_permutation()[0], lu.col_permutation()[0]]]).abs() < 1e-10);
294 /// ```
295 pub fn right(&self, permute: bool) -> Matrix<T> {
296 if permute {
297 let mut result = Matrix::zeros(self.u.nrows(), self.u.ncols());
298 for (new_j, &old_j) in self.col_permutation.iter().enumerate() {
299 for i in 0..self.u.nrows() {
300 result[[i, old_j]] = self.u[[i, new_j]];
301 }
302 }
303 result
304 } else {
305 self.u.clone()
306 }
307 }
308
309 pub(crate) fn left_unpermuted(&self) -> &Matrix<T> {
310 &self.l
311 }
312
313 pub(crate) fn right_unpermuted(&self) -> &Matrix<T> {
314 &self.u
315 }
316
317 /// Get diagonal elements
318 ///
319 /// # Examples
320 ///
321 /// ```
322 /// use tensor4all_core::matrixlu::rrlu;
323 /// use tensor4all_tensorbackend::from_vec2d;
324 ///
325 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
326 /// let lu = rrlu(&m, None).unwrap();
327 /// let d = lu.diag();
328 /// assert_eq!(d.len(), lu.npivots());
329 /// // Diagonal elements are non-zero for full-rank matrices
330 /// for &val in &d {
331 /// assert!(val.abs() > 1e-14);
332 /// }
333 /// ```
334 pub fn diag(&self) -> Vec<T> {
335 let n = self.n_pivot;
336 if self.left_orthogonal {
337 (0..n).map(|i| self.u[[i, i]]).collect()
338 } else {
339 (0..n).map(|i| self.l[[i, i]]).collect()
340 }
341 }
342
343 /// Get pivot errors
344 ///
345 /// # Examples
346 ///
347 /// ```
348 /// use tensor4all_core::matrixlu::rrlu;
349 /// use tensor4all_tensorbackend::from_vec2d;
350 ///
351 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
352 /// let lu = rrlu(&m, None).unwrap();
353 /// let errs = lu.pivot_errors();
354 /// // One entry per pivot plus the final residual
355 /// assert_eq!(errs.len(), lu.npivots() + 1);
356 /// // Errors are non-negative
357 /// for &e in &errs {
358 /// assert!(e >= 0.0);
359 /// }
360 /// ```
361 pub fn pivot_errors(&self) -> Vec<f64> {
362 let mut errors: Vec<f64> = self.diag().iter().map(|d| f64::sqrt(d.abs_sq())).collect();
363 errors.push(self.error);
364 errors
365 }
366
367 /// Get last pivot error
368 ///
369 /// # Examples
370 ///
371 /// ```
372 /// use tensor4all_core::matrixlu::rrlu;
373 /// use tensor4all_tensorbackend::from_vec2d;
374 ///
375 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
376 /// let lu = rrlu(&m, None).unwrap();
377 /// // Full-rank decomposition has zero residual
378 /// assert_eq!(lu.last_pivot_error(), 0.0);
379 /// ```
380 pub fn last_pivot_error(&self) -> f64 {
381 self.error
382 }
383
384 /// Transpose the decomposition
385 ///
386 /// # Examples
387 ///
388 /// ```
389 /// use tensor4all_core::matrixlu::rrlu;
390 /// use tensor4all_tensorbackend::from_vec2d;
391 ///
392 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
393 /// let lu = rrlu(&m, None).unwrap();
394 /// let lu_t = lu.transpose();
395 /// assert_eq!(lu_t.nrows(), lu.ncols());
396 /// assert_eq!(lu_t.ncols(), lu.nrows());
397 /// assert_eq!(lu_t.npivots(), lu.npivots());
398 /// assert_eq!(lu_t.is_left_orthogonal(), !lu.is_left_orthogonal());
399 /// ```
400 pub fn transpose(&self) -> RrLU<T> {
401 RrLU {
402 row_permutation: self.col_permutation.clone(),
403 col_permutation: self.row_permutation.clone(),
404 l: transpose(&self.u),
405 u: transpose(&self.l),
406 left_orthogonal: !self.left_orthogonal,
407 n_pivot: self.n_pivot,
408 error: self.error,
409 }
410 }
411
412 /// Check if left-orthogonal (L has 1s on diagonal)
413 ///
414 /// # Examples
415 ///
416 /// ```
417 /// use tensor4all_core::{matrixlu::rrlu, RrLUOptions};
418 /// use tensor4all_tensorbackend::from_vec2d;
419 ///
420 /// let m = from_vec2d(vec![vec![1.0_f64, 2.0], vec![3.0, 4.0]]);
421 ///
422 /// let lu = rrlu(&m, None).unwrap();
423 /// assert!(lu.is_left_orthogonal()); // default
424 ///
425 /// let lu2 = rrlu(&m, Some(RrLUOptions {
426 /// left_orthogonal: false, ..Default::default()
427 /// })).unwrap();
428 /// assert!(!lu2.is_left_orthogonal());
429 /// ```
430 pub fn is_left_orthogonal(&self) -> bool {
431 self.left_orthogonal
432 }
433}
434
435fn validate_col_major_matrix_len(
436 nrows: usize,
437 ncols: usize,
438 actual_len: usize,
439) -> crate::Result<()> {
440 let expected = nrows
441 .checked_mul(ncols)
442 .ok_or_else(|| MatrixCIError::InvalidArgument {
443 message: format!("matrix shape {nrows} x {ncols} overflows usize"),
444 })?;
445 if actual_len != expected {
446 return Err(MatrixCIError::InvalidArgument {
447 message: format!(
448 "column-major matrix length mismatch: expected {expected}, got {actual_len}"
449 ),
450 });
451 }
452 Ok(())
453}
454
455#[inline]
456fn col_major_offset(nrows: usize, row: usize, col: usize) -> usize {
457 row + nrows * col
458}
459
460#[inline]
461fn col_major_get<T: Copy>(data: &[T], nrows: usize, row: usize, col: usize) -> T {
462 let offset = col_major_offset(nrows, row, col);
463 debug_assert!(row < nrows);
464 debug_assert!(offset < data.len());
465 // SAFETY: callers pass matrix-derived dimensions and prechecked ranges.
466 unsafe { *data.get_unchecked(offset) }
467}
468
469#[inline]
470fn col_major_set<T>(data: &mut [T], nrows: usize, row: usize, col: usize, value: T) {
471 let offset = col_major_offset(nrows, row, col);
472 debug_assert!(row < nrows);
473 debug_assert!(offset < data.len());
474 // SAFETY: callers pass matrix-derived dimensions and prechecked ranges.
475 unsafe {
476 *data.get_unchecked_mut(offset) = value;
477 }
478}
479
480fn submatrix_argmax_col_major<T: Scalar>(
481 data: &[T],
482 nrows: usize,
483 ncols: usize,
484 row_start: usize,
485 row_end: usize,
486 col_start: usize,
487 col_end: usize,
488) -> (usize, usize, T) {
489 debug_assert!(row_start < row_end);
490 debug_assert!(col_start < col_end);
491 debug_assert!(row_end <= nrows);
492 debug_assert!(col_end <= ncols);
493 debug_assert_eq!(data.len(), nrows * ncols);
494
495 let first = col_major_get(data, nrows, row_start, col_start);
496 let mut max_val = first.abs_sq();
497 let mut max_row = row_start;
498 let mut max_col = col_start;
499
500 for col in col_start..col_end {
501 let col_start_offset = col_major_offset(nrows, row_start, col);
502 for (offset, row) in (col_start_offset..).zip(row_start..row_end) {
503 // SAFETY: row and column loops stay within the prechecked region.
504 let value = unsafe { *data.get_unchecked(offset) };
505 let value_abs = value.abs_sq();
506 if value_abs > max_val {
507 max_val = value_abs;
508 max_row = row;
509 max_col = col;
510 }
511 }
512 }
513
514 (
515 max_row,
516 max_col,
517 col_major_get(data, nrows, max_row, max_col),
518 )
519}
520
521fn swap_rows_col_major<T>(data: &mut [T], nrows: usize, ncols: usize, row_a: usize, row_b: usize) {
522 debug_assert!(row_a < nrows);
523 debug_assert!(row_b < nrows);
524 debug_assert_eq!(data.len(), nrows * ncols);
525 if row_a == row_b {
526 return;
527 }
528
529 let ptr = data.as_mut_ptr();
530 for col in 0..ncols {
531 let offset_a = col_major_offset(nrows, row_a, col);
532 let offset_b = col_major_offset(nrows, row_b, col);
533 // SAFETY: row_a/row_b are in range, and each column offset is within
534 // the matrix-sized backing slice.
535 unsafe {
536 std::ptr::swap(ptr.add(offset_a), ptr.add(offset_b));
537 }
538 }
539}
540
541fn swap_cols_col_major<T>(data: &mut [T], nrows: usize, ncols: usize, col_a: usize, col_b: usize) {
542 debug_assert!(col_a < ncols);
543 debug_assert!(col_b < ncols);
544 debug_assert_eq!(data.len(), nrows * ncols);
545 if col_a == col_b || nrows == 0 {
546 return;
547 }
548
549 let start_a = col_major_offset(nrows, 0, col_a);
550 let start_b = col_major_offset(nrows, 0, col_b);
551 // SAFETY: distinct columns are non-overlapping contiguous ranges of
552 // length nrows in column-major storage.
553 unsafe {
554 std::ptr::swap_nonoverlapping(
555 data.as_mut_ptr().add(start_a),
556 data.as_mut_ptr().add(start_b),
557 nrows,
558 );
559 }
560}
561
562fn scale_column_tail<T: Scalar>(
563 data: &mut [T],
564 nrows: usize,
565 col: usize,
566 row_start: usize,
567 pivot: T,
568) {
569 if row_start >= nrows {
570 return;
571 }
572 let start = col_major_offset(nrows, row_start, col);
573 let end = col_major_offset(nrows, nrows - 1, col) + 1;
574 for value in &mut data[start..end] {
575 *value = *value / pivot;
576 }
577}
578
579fn scale_row_tail<T: Scalar>(
580 data: &mut [T],
581 nrows: usize,
582 ncols: usize,
583 row: usize,
584 col_start: usize,
585 pivot: T,
586) {
587 for col in col_start..ncols {
588 let value = col_major_get(data, nrows, row, col) / pivot;
589 col_major_set(data, nrows, row, col, value);
590 }
591}
592
593fn update_trailing_submatrix<T: Scalar>(data: &mut [T], nrows: usize, ncols: usize, pivot: usize) {
594 let tail_row_start = pivot + 1;
595 let tail_col_start = pivot + 1;
596 if tail_row_start >= nrows || tail_col_start >= ncols {
597 return;
598 }
599
600 let tail_len = nrows - tail_row_start;
601 let pivot_col_tail_start = col_major_offset(nrows, tail_row_start, pivot);
602 for col in tail_col_start..ncols {
603 let y = col_major_get(data, nrows, pivot, col);
604 let target_start = col_major_offset(nrows, tail_row_start, col);
605 let (before_target, target_and_after) = data.split_at_mut(target_start);
606 let pivot_col_tail = &before_target[pivot_col_tail_start..pivot_col_tail_start + tail_len];
607 let target_tail = &mut target_and_after[..tail_len];
608 for (target, &x) in target_tail.iter_mut().zip(pivot_col_tail.iter()) {
609 *target = *target - x * y;
610 }
611 }
612}
613
614fn extract_lu_from_factorized<T: Scalar>(
615 data: &[T],
616 nrows: usize,
617 ncols: usize,
618 rank: usize,
619 left_orthogonal: bool,
620) -> Result<(Matrix<T>, Matrix<T>)> {
621 debug_assert!(rank <= nrows.min(ncols));
622 debug_assert_eq!(data.len(), nrows * ncols);
623
624 let mut l_data = vec![T::zero(); nrows * rank];
625 for col in 0..rank {
626 let src_start = col_major_offset(nrows, col, col);
627 let src_end = col_major_offset(nrows, nrows - 1, col) + 1;
628 let dst_start = col_major_offset(nrows, col, col);
629 l_data[dst_start..dst_start + (nrows - col)].copy_from_slice(&data[src_start..src_end]);
630 }
631
632 let mut u_data = vec![T::zero(); rank * ncols];
633 for col in 0..ncols {
634 let rows_to_copy = rank.min(col + 1);
635 if rows_to_copy > 0 {
636 let src_start = col_major_offset(nrows, 0, col);
637 let dst_start = col_major_offset(rank, 0, col);
638 u_data[dst_start..dst_start + rows_to_copy]
639 .copy_from_slice(&data[src_start..src_start + rows_to_copy]);
640 }
641 }
642
643 if left_orthogonal {
644 for i in 0..rank {
645 l_data[col_major_offset(nrows, i, i)] = T::one();
646 }
647 } else {
648 for i in 0..rank {
649 u_data[col_major_offset(rank, i, i)] = T::one();
650 }
651 }
652
653 if l_data.iter().any(|&value| value.is_nan()) {
654 return Err(MatrixCIError::NaNEncountered {
655 matrix: "L".to_string(),
656 });
657 }
658 if u_data.iter().any(|&value| value.is_nan()) {
659 return Err(MatrixCIError::NaNEncountered {
660 matrix: "U".to_string(),
661 });
662 }
663
664 Ok((
665 Matrix::from_col_major_vec(nrows, rank, l_data),
666 Matrix::from_col_major_vec(rank, ncols, u_data),
667 ))
668}
669
670/// Options for rank-revealing LU decomposition.
671///
672/// # Examples
673///
674/// ```
675/// use tensor4all_core::RrLUOptions;
676///
677/// // Default: rel_tol = 1e-14, no absolute tolerance, no rank limit
678/// let opts = RrLUOptions::default();
679/// assert_eq!(opts.rel_tol, 1e-14);
680/// assert_eq!(opts.abs_tol, 0.0);
681/// assert!(opts.left_orthogonal);
682///
683/// // Limit rank to 5
684/// let opts = RrLUOptions { max_bond_dim: 5, ..Default::default() };
685/// assert_eq!(opts.max_bond_dim, 5);
686/// ```
687#[derive(Debug, Clone)]
688pub struct RrLUOptions {
689 /// Maximum rank
690 pub max_bond_dim: usize,
691 /// Relative tolerance
692 pub rel_tol: f64,
693 /// Absolute tolerance
694 pub abs_tol: f64,
695 /// Left orthogonal (L has 1s on diagonal) or right orthogonal (U has 1s)
696 pub left_orthogonal: bool,
697}
698
699impl Default for RrLUOptions {
700 fn default() -> Self {
701 Self {
702 max_bond_dim: usize::MAX,
703 rel_tol: 1e-14,
704 abs_tol: 0.0,
705 left_orthogonal: true,
706 }
707 }
708}
709
710/// Perform in-place rank-revealing LU decomposition.
711///
712/// The input matrix `a` is modified in place. Use [`rrlu`] for a
713/// non-destructive version.
714///
715/// # Errors
716///
717/// Returns [`MatrixCIError::InvalidArgument`] if the matrix shape product
718/// overflows `usize` or its backing storage length does not match the shape.
719/// Returns [`MatrixCIError::NaNEncountered`] if NaN values appear in the L or U
720/// factors.
721///
722/// # Examples
723///
724/// ```
725/// use tensor4all_core::{matrixlu::rrlu_mut, RrLUOptions};
726/// use tensor4all_tensorbackend::from_vec2d;
727///
728/// let mut m = from_vec2d(vec![
729/// vec![1.0_f64, 2.0],
730/// vec![3.0, 4.0],
731/// ]);
732/// let lu = rrlu_mut(&mut m, Some(RrLUOptions { max_bond_dim: 1, ..Default::default() })).unwrap();
733/// assert_eq!(lu.npivots(), 1);
734/// ```
735pub fn rrlu_mut<T: Scalar>(a: &mut Matrix<T>, options: Option<RrLUOptions>) -> Result<RrLU<T>> {
736 let opts = options.unwrap_or_default();
737 let nr = a.nrows();
738 let nc = a.ncols();
739 let data = a.as_col_major_mut_slice();
740 validate_col_major_matrix_len(nr, nc, data.len())?;
741 debug_assert_eq!(data.len(), nr * nc);
742
743 let mut lu = RrLU::new(nr, nc, opts.left_orthogonal);
744 let max_bond_dim = opts.max_bond_dim.min(nr).min(nc);
745 let mut max_error = 0.0f64;
746
747 while lu.n_pivot < max_bond_dim {
748 let k = lu.n_pivot;
749
750 if k >= nr || k >= nc {
751 break;
752 }
753
754 let (pivot_row, pivot_col, pivot_val) =
755 submatrix_argmax_col_major(data, nr, nc, k, nr, k, nc);
756
757 let pivot_abs = f64::sqrt(pivot_val.abs_sq());
758 lu.error = pivot_abs;
759
760 // Check stopping criteria (but add at least 1 pivot)
761 if lu.n_pivot > 0 && (pivot_abs < opts.rel_tol * max_error || pivot_abs < opts.abs_tol) {
762 break;
763 }
764
765 // Guard against tiny pivots to prevent NaN from division. A caller that
766 // sets both tolerances to zero is requesting a non-truncating
767 // decomposition, so only an exactly zero pivot stops the factorization.
768 let min_pivot_abs = if opts.rel_tol == 0.0 && opts.abs_tol == 0.0 {
769 0.0
770 } else {
771 f64::EPSILON
772 };
773 if pivot_abs <= min_pivot_abs {
774 if lu.n_pivot == 0 {
775 // First pivot is near-zero: the matrix is effectively zero
776 lu.error = pivot_abs;
777 }
778 break;
779 }
780
781 max_error = max_error.max(pivot_abs);
782
783 // Swap rows and columns
784 if pivot_row != k {
785 swap_rows_col_major(data, nr, nc, k, pivot_row);
786 lu.row_permutation.swap(k, pivot_row);
787 }
788 if pivot_col != k {
789 swap_cols_col_major(data, nr, nc, k, pivot_col);
790 lu.col_permutation.swap(k, pivot_col);
791 }
792
793 let pivot = col_major_get(data, nr, k, k);
794
795 // Eliminate
796 if opts.left_orthogonal {
797 scale_column_tail(data, nr, k, k + 1, pivot);
798 } else {
799 scale_row_tail(data, nr, nc, k, k + 1, pivot);
800 }
801
802 update_trailing_submatrix(data, nr, nc, k);
803
804 lu.n_pivot += 1;
805 }
806
807 let n = lu.n_pivot;
808 let (l, u) = extract_lu_from_factorized(data, nr, nc, n, opts.left_orthogonal)?;
809
810 // Set error to 0 if full rank
811 if n >= nr.min(nc) {
812 lu.error = 0.0;
813 }
814
815 lu.l = l;
816 lu.u = u;
817
818 Ok(lu)
819}
820
821/// Perform rank-revealing LU decomposition (non-destructive).
822///
823/// Clones the input matrix and calls [`rrlu_mut`].
824///
825/// # Errors
826///
827/// Returns [`MatrixCIError::InvalidArgument`] if the matrix shape product
828/// overflows `usize` or its backing storage length does not match the shape.
829/// Returns [`MatrixCIError::NaNEncountered`] if NaN values appear in the L or U
830/// factors.
831///
832/// # Examples
833///
834/// ```
835/// use tensor4all_core::matrixlu::rrlu;
836/// use tensor4all_tensorbackend::from_vec2d;
837///
838/// let m = from_vec2d(vec![
839/// vec![1.0_f64, 0.0],
840/// vec![0.0, 2.0],
841/// ]);
842/// let lu = rrlu(&m, None).unwrap();
843/// assert_eq!(lu.npivots(), 2);
844/// assert_eq!(lu.nrows(), 2);
845/// assert_eq!(lu.ncols(), 2);
846/// ```
847pub fn rrlu<T: Scalar>(a: &Matrix<T>, options: Option<RrLUOptions>) -> Result<RrLU<T>> {
848 let mut a_copy = a.clone();
849 rrlu_mut(&mut a_copy, options)
850}
851
852/// Convert L matrix to solve L * X = B given pivot matrix P
853///
854/// Modifies `c` in place so that the columns satisfy the triangular
855/// system defined by `p`. The matrix `c` must have at least `p.nrows()`
856/// columns.
857///
858/// # Errors
859/// Returns [`MatrixCIError::InvalidArgument`] when `p` is not square or
860/// `c` has too few columns, and [`MatrixCIError::SingularMatrix`] for a zero
861/// diagonal pivot.
862///
863/// # Examples
864///
865/// ```
866/// use tensor4all_core::matrixlu::cols_to_l_matrix;
867/// use tensor4all_tensorbackend::from_vec2d;
868///
869/// // Upper-triangular P (2x2)
870/// let p = from_vec2d(vec![
871/// vec![2.0_f64, 1.0],
872/// vec![0.0, 3.0],
873/// ]);
874/// // c has 3 rows, 2 columns (ncols >= p.nrows())
875/// let mut c = from_vec2d(vec![
876/// vec![4.0_f64, 5.0],
877/// vec![6.0, 9.0],
878/// vec![8.0, 7.0],
879/// ]);
880/// cols_to_l_matrix(&mut c, &p, true).unwrap();
881/// // After processing: c[:,0] was divided by p[0,0]=2
882/// assert!((c[[0, 0]] - 2.0).abs() < 1e-10);
883/// assert!((c[[1, 0]] - 3.0).abs() < 1e-10);
884/// assert!((c[[2, 0]] - 4.0).abs() < 1e-10);
885/// ```
886pub fn cols_to_l_matrix<T: Scalar>(
887 c: &mut Matrix<T>,
888 p: &Matrix<T>,
889 _left_orthogonal: bool,
890) -> Result<()> {
891 if p.nrows() != p.ncols() || c.ncols() < p.nrows() {
892 return Err(MatrixCIError::InvalidArgument {
893 message: format!(
894 "cols_to_l_matrix requires square p and c.ncols() >= p.nrows(), got p=({}, {}), c.ncols()={}",
895 p.nrows(),
896 p.ncols(),
897 c.ncols()
898 ),
899 });
900 }
901 let n = p.nrows();
902 for k in 0..n {
903 if p[[k, k]].abs_val() == 0.0 {
904 return Err(MatrixCIError::SingularMatrix);
905 }
906 }
907
908 for k in 0..n {
909 let pivot = p[[k, k]];
910 // c[:, k] /= pivot
911 for i in 0..c.nrows() {
912 let val = c[[i, k]] / pivot;
913 c[[i, k]] = val;
914 }
915
916 // c[:, k+1:] -= c[:, k] * p[k, k+1:]
917 for j in (k + 1)..c.ncols() {
918 let p_kj = p[[k, j]];
919 for i in 0..c.nrows() {
920 let c_ik = c[[i, k]];
921 let old = c[[i, j]];
922 c[[i, j]] = old - c_ik * p_kj;
923 }
924 }
925 }
926 Ok(())
927}
928
929/// Convert R matrix to solve X * U = B given pivot matrix P
930///
931/// Modifies `r` in place so that the rows satisfy the triangular
932/// system defined by `p`. The matrix `r` must have at least `p.nrows()`
933/// rows.
934///
935/// # Errors
936/// Returns [`MatrixCIError::InvalidArgument`] when `p` is not square or
937/// `r` has too few rows, and [`MatrixCIError::SingularMatrix`] for a zero
938/// diagonal pivot.
939///
940/// # Examples
941///
942/// ```
943/// use tensor4all_core::matrixlu::rows_to_u_matrix;
944/// use tensor4all_tensorbackend::from_vec2d;
945///
946/// // Lower-triangular P (2x2)
947/// let p = from_vec2d(vec![
948/// vec![2.0_f64, 0.0],
949/// vec![1.0, 3.0],
950/// ]);
951/// // r has 2 rows (nrows >= p.nrows()), 3 columns
952/// let mut r = from_vec2d(vec![
953/// vec![4.0_f64, 6.0, 8.0],
954/// vec![5.0, 9.0, 7.0],
955/// ]);
956/// rows_to_u_matrix(&mut r, &p, true).unwrap();
957/// // After processing: r[0,:] was divided by p[0,0]=2
958/// assert!((r[[0, 0]] - 2.0).abs() < 1e-10);
959/// assert!((r[[0, 1]] - 3.0).abs() < 1e-10);
960/// assert!((r[[0, 2]] - 4.0).abs() < 1e-10);
961/// ```
962pub fn rows_to_u_matrix<T: Scalar>(
963 r: &mut Matrix<T>,
964 p: &Matrix<T>,
965 _left_orthogonal: bool,
966) -> Result<()> {
967 if p.nrows() != p.ncols() || r.nrows() < p.nrows() {
968 return Err(MatrixCIError::InvalidArgument {
969 message: format!(
970 "rows_to_u_matrix requires square p and r.nrows() >= p.nrows(), got p=({}, {}), r.nrows()={}",
971 p.nrows(),
972 p.ncols(),
973 r.nrows()
974 ),
975 });
976 }
977 let n = p.nrows();
978 for k in 0..n {
979 if p[[k, k]].abs_val() == 0.0 {
980 return Err(MatrixCIError::SingularMatrix);
981 }
982 }
983
984 for k in 0..n {
985 let pivot = p[[k, k]];
986 // r[k, :] /= pivot
987 for j in 0..r.ncols() {
988 let val = r[[k, j]] / pivot;
989 r[[k, j]] = val;
990 }
991
992 // r[k+1:, :] -= p[k+1:, k] * r[k, :]
993 for i in (k + 1)..r.nrows() {
994 let p_ik = p[[i, k]];
995 for j in 0..r.ncols() {
996 let r_kj = r[[k, j]];
997 let old = r[[i, j]];
998 r[[i, j]] = old - p_ik * r_kj;
999 }
1000 }
1001 }
1002 Ok(())
1003}
1004
1005#[cfg(test)]
1006mod tests;