1use 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 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 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 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 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 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 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 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 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}