mdarray_linalg_faer/
solve.rs1use faer::linalg::solvers::Solve as FaerSolve;
9use faer_traits::ComplexField;
10use mdarray::{Array, Dim, Layout, Shape, Slice};
11use mdarray_linalg::solve::{Solve, SolveError};
12use num_complex::ComplexFloat;
13
14use crate::{Faer, into_faer, into_faer_mut};
15
16impl<T, D: Dim> Solve<T, D> for Faer
17where
18 T: ComplexFloat
19 + ComplexField
20 + Default
21 + std::convert::From<<T as num_complex::ComplexFloat>::Real>,
22{
23 fn solve<R: Dim, La: Layout, Lb: Layout>(
25 &self,
26 a: &mut Slice<T, (D, D), La>,
27 b: &Slice<T, (D, R), Lb>,
28 ) -> Result<Array<T, (D, R)>, SolveError> {
29 let ash = *a.shape();
30 let (m, n) = (ash.dim(0), ash.dim(1));
31
32 let bsh = *b.shape();
33 let (b_m, b_n) = (bsh.dim(0), bsh.dim(1));
34
35 if m != n {
36 return Err(SolveError::InvalidDimensions);
37 }
38
39 if b_m != m {
40 return Err(SolveError::InvalidDimensions);
41 }
42
43 let a_faer = into_faer_mut(a);
44
45 let solver = a_faer.partial_piv_lu();
46
47 let b_faer = into_faer(b);
48 let x_faer = solver.solve(b_faer);
49
50 let mut x_mda =
51 Array::from_elem(<(D, R) as Shape>::from_dims(&[m, b_n]), T::default());
52
53 let mut x_faer_mut = into_faer_mut(&mut x_mda);
54 for i in 0..m {
55 for j in 0..b_n {
56 x_faer_mut[(i, j)] = x_faer[(i, j)];
57 }
58 }
59
60 Ok(x_mda)
61 }
62
63 fn solve_write<R: Dim, La: Layout, Lb: Layout>(
65 &self,
66 a: &mut Slice<T, (D, D), La>,
67 b: &mut Slice<T, (D, R), Lb>,
68 ) -> Result<(), SolveError> {
69 let ash = *a.shape();
70 let (m, n) = (ash.dim(0), ash.dim(1));
71
72 let bsh = *b.shape();
73 let (b_m, b_n) = (bsh.dim(0), bsh.dim(1));
74
75 if m != n {
76 return Err(SolveError::InvalidDimensions);
77 }
78
79 if b_m != m {
80 return Err(SolveError::InvalidDimensions);
81 }
82
83 let a_faer = into_faer(a);
84
85 let solver = a_faer.partial_piv_lu();
86
87 let b_faer = into_faer(b).to_owned();
88 let x_faer = solver.solve(b_faer);
89
90 let mut b_faer_mut = into_faer_mut(b);
91 for i in 0..m {
92 for j in 0..b_n {
93 b_faer_mut[(i, j)] = x_faer[(i, j)];
94 }
95 }
96
97 Ok(())
98 }
99}