Skip to main content

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;