Skip to main content

mdarray_linalg_faer/
solve.rs

1// Linear system solver:
2//     A * X = B
3// where:
4//     - A is m × m         (square coefficient matrix)
5//     - B is m × n         (right-hand side matrix)
6//     - X is m × n         (solution matrix)
7
8use 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    /// Solves linear system AX = B with new allocated solution matrix.
24    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    /// Solves linear system AX = B, overwriting B with the solution X.
64    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}