Skip to main content

mdarray_linalg_faer/lu/
context.rs

1// LU Decomposition with partial pivoting:
2//     P * A = L * U
3// where:
4//     - A is m × n         (input matrix)
5//     - P is m × m        (permutation matrix)
6//     - L is m × m        (lower triangular with ones on diagonal)
7//     - U is m × n         (upper triangular/trapezoidal matrix)
8
9use dyn_stack::{MemBuffer, MemStack};
10use faer_traits::ComplexField;
11use mdarray::{Array, Dim, Layout, Shape, Slice};
12use mdarray_linalg::lu::{InvError, LU};
13use num_complex::ComplexFloat;
14
15use super::simple::lu_faer;
16use crate::{Faer, into_faer_mut};
17
18fn map_cholesky_error(err: faer::linalg::cholesky::llt::factor::LltError) -> InvError {
19    match err {
20        faer::linalg::cholesky::llt::factor::LltError::NonPositivePivot { index } => {
21            InvError::NotPositiveDefinite {
22                lpm: index as i32 + 1,
23            }
24        }
25    }
26}
27
28impl<T, D0: Dim, D1: Dim> LU<T, D0, D1> for Faer
29where
30    T: ComplexFloat
31        + ComplexField
32        + Default
33        + std::convert::From<<T as num_complex::ComplexFloat>::Real>,
34{
35    /// Computes LU decomposition with new allocated matrices: L, U, P (permutation matrix)
36    fn lu<L: Layout>(
37        &self,
38        a: &mut Slice<T, (D0, D1), L>,
39    ) -> (Array<T, (D0, D0)>, Array<T, (D0, D1)>, Array<T, (D0, D0)>) {
40        let ash = *a.shape();
41        let (m, n) = (ash.dim(0), ash.dim(1));
42
43        let min_mn = m.min(n);
44
45        // Create shapes for L, U, and P matrices
46        let l_shape = <(D0, D0) as Shape>::from_dims(&[m, min_mn]);
47        let u_shape = <(D0, D1) as Shape>::from_dims(&[min_mn, n]);
48        let p_shape = <(D0, D0) as Shape>::from_dims(&[m, m]);
49
50        let mut l_mda = Array::from_elem(l_shape, T::default());
51        let mut u_mda = Array::from_elem(u_shape, T::default());
52        let mut p_mda = Array::from_elem(p_shape, T::default());
53
54        lu_faer(a, &mut l_mda, &mut u_mda, &mut p_mda);
55
56        (l_mda, u_mda, p_mda)
57    }
58
59    /// Computes LU decomposition overwriting existing matrices
60    fn lu_write<L: Layout, Ll: Layout, Lu: Layout, Lp: Layout>(
61        &self,
62        a: &mut Slice<T, (D0, D1), L>,
63        l: &mut Slice<T, (D0, D0), Ll>,
64        u: &mut Slice<T, (D0, D1), Lu>,
65        p: &mut Slice<T, (D0, D0), Lp>,
66    ) {
67        lu_faer::<T, D0, D1, L, Ll, Lu, Lp>(a, l, u, p);
68    }
69
70    /// Computes inverse with new allocated matrix
71    fn inv<L: Layout>(
72        &self,
73        a: &mut Slice<T, (D0, D1), L>,
74    ) -> Result<Array<T, (D0, D1)>, InvError> {
75        let ash = *a.shape();
76        let (m, n) = (ash.dim(0), ash.dim(1));
77
78        if m != n {
79            return Err(InvError::NotSquare {
80                rows: m as i32,
81                cols: n as i32,
82            });
83        }
84
85        let par = faer::get_global_parallelism();
86        let mut a_faer = into_faer_mut(a);
87
88        let mut row_perm_fwd = vec![0usize; m];
89        let mut row_perm_bwd = vec![0usize; m];
90
91        faer::linalg::lu::partial_pivoting::factor::lu_in_place(
92            a_faer.as_mut(),
93            &mut row_perm_fwd,
94            &mut row_perm_bwd,
95            par,
96            MemStack::new(&mut MemBuffer::new(
97                faer::linalg::lu::partial_pivoting::factor::lu_in_place_scratch::<usize, T>(
98                    m,
99                    n,
100                    par,
101                    faer::prelude::default(),
102                ),
103            )),
104            faer::prelude::default(),
105        );
106
107        let l_mat = a_faer.as_ref();
108        let u_mat = a_faer.as_ref();
109
110        let perm = unsafe {
111            faer::perm::Perm::new_unchecked(
112                row_perm_fwd.into_boxed_slice(),
113                row_perm_bwd.into_boxed_slice(),
114            )
115        };
116
117        let mut inv_mat = Array::<T, (D0, D1)>::from_elem(ash, T::zero());
118        let mut inv_mat_faer = into_faer_mut(&mut inv_mat);
119
120        faer::linalg::lu::partial_pivoting::inverse::inverse(
121            inv_mat_faer.as_mut(),
122            l_mat,
123            u_mat,
124            perm.as_ref(),
125            par,
126            MemStack::new(&mut MemBuffer::new(
127                faer::linalg::lu::partial_pivoting::inverse::inverse_scratch::<usize, T>(m, par),
128            )),
129        );
130        Ok(inv_mat)
131    }
132
133    /// Computes inverse overwriting the input matrix
134    fn inv_write<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> Result<(), InvError> {
135        let ash = *a.shape();
136        let (m, n) = (ash.dim(0), ash.dim(1));
137
138        if m != n {
139            return Err(InvError::NotSquare {
140                rows: m as i32,
141                cols: n as i32,
142            });
143        }
144
145        let par = faer::get_global_parallelism();
146        let mut a_faer = into_faer_mut(a);
147
148        let mut row_perm_fwd = vec![0usize; m];
149        let mut row_perm_bwd = vec![0usize; m];
150
151        faer::linalg::lu::partial_pivoting::factor::lu_in_place(
152            a_faer.as_mut(),
153            &mut row_perm_fwd,
154            &mut row_perm_bwd,
155            par,
156            MemStack::new(&mut MemBuffer::new(
157                faer::linalg::lu::partial_pivoting::factor::lu_in_place_scratch::<usize, T>(
158                    m,
159                    n,
160                    par,
161                    faer::prelude::default(),
162                ),
163            )),
164            faer::prelude::default(),
165        );
166
167        let l_mat = a_faer.as_ref();
168        let u_mat = a_faer.as_ref();
169
170        let perm = unsafe {
171            faer::perm::Perm::new_unchecked(
172                row_perm_fwd.into_boxed_slice(),
173                row_perm_bwd.into_boxed_slice(),
174            )
175        };
176
177        let mut inv_mat = faer::Mat::<T>::zeros(m, n);
178
179        faer::linalg::lu::partial_pivoting::inverse::inverse(
180            inv_mat.as_mut(),
181            l_mat,
182            u_mat,
183            perm.as_ref(),
184            par,
185            MemStack::new(&mut MemBuffer::new(
186                faer::linalg::lu::partial_pivoting::inverse::inverse_scratch::<usize, T>(m, par),
187            )),
188        );
189
190        for i in 0..m {
191            for j in 0..n {
192                a_faer[(i, j)] = inv_mat[(i, j)];
193            }
194        }
195
196        Ok(())
197    }
198
199    /// Computes the determinant of a square matrix. Panics if the matrix is non-square.
200    fn det<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> T {
201        let ash = *a.shape();
202        let (m, n) = (ash.dim(0), ash.dim(1));
203
204        assert_eq!(m, n, "determinant is only defined for square matrices");
205        let a_faer = into_faer_mut(a);
206        a_faer.determinant()
207    }
208
209    /// Computes the Cholesky decomposition, returning a lower-triangular matrix
210    fn cholesky<L: Layout>(
211        &self,
212        a: &mut Slice<T, (D0, D1), L>,
213    ) -> Result<Array<T, (D0, D1)>, InvError> {
214        let ash = *a.shape();
215        let (m, n) = (ash.dim(0), ash.dim(1));
216
217        if m != n {
218            return Err(InvError::NotSquare {
219                rows: m as i32,
220                cols: n as i32,
221            });
222        }
223
224        let mut l = a.to_tensor();
225        self.cholesky_write(&mut l)?;
226        Ok(l)
227    }
228
229    /// Computes the Cholesky decomposition in-place, overwriting the input matrix
230    fn cholesky_write<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> Result<(), InvError> {
231        let ash = *a.shape();
232        let (m, n) = (ash.dim(0), ash.dim(1));
233
234        if m != n {
235            return Err(InvError::NotSquare {
236                rows: m as i32,
237                cols: n as i32,
238            });
239        }
240
241        let par = faer::get_global_parallelism();
242
243        let result = {
244            let mut a_faer = into_faer_mut(a);
245            faer::linalg::cholesky::llt::factor::cholesky_in_place(
246                a_faer.as_mut(),
247                Default::default(),
248                par,
249                MemStack::new(&mut MemBuffer::new(
250                    faer::linalg::cholesky::llt::factor::cholesky_in_place_scratch::<T>(
251                        n,
252                        par,
253                        faer::prelude::default(),
254                    ),
255                )),
256                faer::prelude::default(),
257            )
258        };
259
260        match result {
261            Ok(_) => {
262                for i in 0..n {
263                    for j in i + 1..n {
264                        a[[i, j]] = T::zero();
265                    }
266                }
267                Ok(())
268            }
269            Err(err) => Err(map_cholesky_error(err)),
270        }
271    }
272}