Skip to main content

kryst/solver/
direct_lu.rs

1//! Direct dense solvers using Faer: LU and QR factorizations.
2//!
3//! This module provides wrappers for direct dense linear solvers using the Faer library.
4//! It includes LU (with full pivoting) and QR solvers for square or rectangular systems.
5//! These solvers are suitable for small to medium-sized dense systems where direct methods are feasible.
6//!
7//! # Usage
8//! - Use `LuSolver` for general square systems (may be faster, but less stable for rank-deficient matrices).
9//! - Use `QrSolver` for square or rectangular systems (more stable for rank-deficient or nearly singular matrices).
10//!
11//! # References
12//! - Faer documentation: https://github.com/sarah-ek/faer-rs
13//! - Golub & Van Loan, Matrix Computations
14
15use crate::error::KError;
16use crate::solver::MonitorCallback;
17use crate::solver::legacy::LinearSolver;
18use crate::utils::convergence::{ConvergedReason, SolveStats};
19use crate::{parallel::UniverseComm, preconditioner::PcSide};
20use faer::linalg::solvers::{FullPivLu, Qr, SolveCore};
21use faer::{Conj, Mat, MatMut};
22
23type Scalar = f64;
24
25#[cfg(feature = "logging")]
26use crate::utils::profiling::StageGuard;
27
28/// LU solver using full pivoting from Faer.
29///
30/// Stores the LU factorization for reuse (if desired).
31pub struct LuSolver {
32    /// Cached LU factorization (if computed)
33    factor: Option<FullPivLu<Scalar>>,
34}
35
36impl LuSolver {
37    /// Create a new LU solver (no factorization yet).
38    pub fn new() -> Self {
39        LuSolver { factor: None }
40    }
41
42    /// Solve using the cached LU factorization.
43    ///
44    /// # Panics
45    /// Panics if called before any factorization has been performed.
46    ///
47    /// # Arguments
48    /// * `b` - Right-hand side vector
49    /// * `x` - Output vector (solution)
50    pub fn solve_cached(&self, b: &[Scalar], x: &mut [Scalar]) {
51        if let Some(factor) = &self.factor {
52            let n = b.len();
53            x.clone_from_slice(b);
54            let x_mat = MatMut::from_column_major_slice_mut(x, n, 1);
55            factor.solve_in_place_with_conj(Conj::No, x_mat);
56        } else {
57            panic!("LuSolver: solve_cached called before factorization");
58        }
59    }
60}
61
62impl LinearSolver<Mat<Scalar>, Vec<Scalar>> for LuSolver {
63    type Error = KError;
64    type Scalar = Scalar;
65
66    /// Solve Ax = b using LU factorization (full pivoting).
67    ///
68    /// # Arguments
69    /// * `a` - Matrix (Faer Mat)
70    /// * `pc` - (Unused) Preconditioner (not supported for direct solvers)
71    /// * `b` - Right-hand side vector
72    /// * `x` - On input: ignored; on output: solution vector
73    /// * `comm` - Communicator for parallel operations (unused for direct solvers)
74    /// * `monitors` - Optional callbacks to invoke at each iteration
75    /// * `work` - Optional workspace (unused for direct solvers)
76    ///
77    /// # Returns
78    /// * `Ok(SolveStats)` (always converged in 1 iteration)
79    fn solve(
80        &mut self,
81        a: &Mat<Scalar>,
82        pc: Option<
83            &(dyn crate::preconditioner::legacy::Preconditioner<Mat<Scalar>, Vec<Scalar>> + '_),
84        >,
85        b: &Vec<Scalar>,
86        x: &mut Vec<Scalar>,
87        _pc_side: PcSide,
88        _comm: &UniverseComm,
89        monitors: Option<&[Box<MonitorCallback<Self::Scalar>>]>,
90        _work: Option<&mut crate::context::ksp_context::Workspace>,
91    ) -> Result<SolveStats<Scalar>, KError> {
92        #[cfg(feature = "logging")]
93        let _guard = StageGuard::new("LuSolve");
94
95        let _ = pc; // Direct solvers do not use preconditioner
96        let _ = _pc_side;
97
98        // Call monitors at start if provided
99        if let Some(monitors) = monitors {
100            for monitor in monitors {
101                monitor(0, 0.0, 0);
102            }
103        }
104
105        #[cfg(feature = "logging")]
106        let _fact_guard = StageGuard::new("LuFactor");
107
108        // Compute LU factorization (overwrites any previous factor)
109        let factor = FullPivLu::new(a.as_ref());
110        self.factor = Some(factor);
111
112        // Copy b into x
113        x.clone_from(b);
114
115        // Solve in-place: x = A^{-1} b
116        let n = x.len();
117        let x_mat = MatMut::from_column_major_slice_mut(x, n, 1);
118        self.factor
119            .as_ref()
120            .unwrap()
121            .solve_in_place_with_conj(Conj::No, x_mat);
122
123        // Call monitors at end if provided
124        if let Some(monitors) = monitors {
125            for monitor in monitors {
126                monitor(1, 0.0, 0);
127            }
128        }
129
130        // For direct solvers, always converged in 1 iteration
131        Ok(SolveStats::new(1, 0.0, ConvergedReason::ConvergedAtol))
132    }
133}
134
135impl Default for LuSolver {
136    fn default() -> Self {
137        Self::new()
138    }
139}
140
141/// QR solver using Faer (for square or rectangular systems).
142pub struct QrSolver;
143
144impl QrSolver {
145    /// Create a new QR solver.
146    pub fn new() -> Self {
147        QrSolver
148    }
149}
150
151impl LinearSolver<Mat<Scalar>, Vec<Scalar>> for QrSolver {
152    type Error = KError;
153    type Scalar = Scalar;
154
155    /// Solve Ax = b using QR factorization.
156    ///
157    /// # Arguments
158    /// * `a` - Matrix (Faer Mat)
159    /// * `pc` - (Unused) Preconditioner (not supported for direct solvers)
160    /// * `b` - Right-hand side vector
161    /// * `x` - On input: ignored; on output: solution vector
162    /// * `comm` - Communicator for parallel operations (unused for direct solvers)
163    /// * `monitors` - Optional callbacks to invoke at each iteration
164    /// * `work` - Optional workspace (unused for direct solvers)
165    ///
166    /// # Returns
167    /// * `Ok(SolveStats)` (always converged in 1 iteration)
168    fn solve(
169        &mut self,
170        a: &Mat<Scalar>,
171        pc: Option<
172            &(dyn crate::preconditioner::legacy::Preconditioner<Mat<Scalar>, Vec<Scalar>> + '_),
173        >,
174        b: &Vec<Scalar>,
175        x: &mut Vec<Scalar>,
176        _pc_side: PcSide,
177        _comm: &UniverseComm,
178        monitors: Option<&[Box<MonitorCallback<Self::Scalar>>]>,
179        _work: Option<&mut crate::context::ksp_context::Workspace>,
180    ) -> Result<SolveStats<Scalar>, KError> {
181        #[cfg(feature = "logging")]
182        let _guard = StageGuard::new("QrSolve");
183
184        let _ = pc; // Direct solvers do not use preconditioner
185        let _ = _pc_side;
186
187        // Call monitors at start if provided
188        if let Some(monitors) = monitors {
189            for monitor in monitors {
190                monitor(0, 0.0, 0);
191            }
192        }
193
194        // Compute QR factorization
195        let factor = Qr::new(a.as_ref());
196        x.clone_from(b);
197        let n = x.len();
198        let x_mat = MatMut::from_column_major_slice_mut(x, n, 1);
199        factor.solve_in_place_with_conj(Conj::No, x_mat);
200
201        // Call monitors at end if provided
202        if let Some(monitors) = monitors {
203            for monitor in monitors {
204                monitor(1, 0.0, 0);
205            }
206        }
207
208        Ok(SolveStats::new(1, 0.0, ConvergedReason::ConvergedAtol))
209    }
210}
211
212impl Default for QrSolver {
213    fn default() -> Self {
214        Self::new()
215    }
216}
217
218#[cfg(test)]
219mod tests {
220    use super::*;
221    use crate::solver::legacy::LinearSolver;
222    use faer::Mat;
223
224    #[test]
225    fn lu_solver_solves_dense_system() {
226        // 3x3 system: [[2,1,1],[1,3,2],[1,0,0]] x = [4,5,6]
227        // True solution: [6,15,-23]
228        let a = Mat::from_fn(3, 3, |i, j| match (i, j) {
229            (0, 0) => 2.0,
230            (0, 1) => 1.0,
231            (0, 2) => 1.0,
232            (1, 0) => 1.0,
233            (1, 1) => 3.0,
234            (1, 2) => 2.0,
235            (2, 0) => 1.0,
236            (2, 1) => 0.0,
237            (2, 2) => 0.0,
238            _ => 0.0,
239        });
240        let b = vec![4.0, 5.0, 6.0];
241        let mut x = vec![0.0; 3];
242        let mut solver = LuSolver::new();
243        let stats = solver
244            .solve(
245                &a,
246                None,
247                &b,
248                &mut x,
249                PcSide::Left,
250                &UniverseComm::NoComm(crate::parallel::NoComm),
251                None,
252                None,
253            )
254            .unwrap();
255        let expected = vec![6.0, 15.0, -23.0];
256        let tol = 1e-10;
257        for (xi, ei) in x.iter().zip(expected.iter()) {
258            assert!((xi - ei).abs() < tol, "xi = {}, expected = {}", xi, ei);
259        }
260        assert!(
261            matches!(
262                stats.reason,
263                ConvergedReason::ConvergedAtol | ConvergedReason::ConvergedRtol
264            ),
265            "LU did not report Converged reason"
266        );
267    }
268
269    #[test]
270    fn qr_solver_solves_dense_system() {
271        // 3x3 system: [[2,1,1],[1,3,2],[1,0,0]] x = [4,5,6]
272        // True solution: [6,15,-23]
273        let a = Mat::from_fn(3, 3, |i, j| match (i, j) {
274            (0, 0) => 2.0,
275            (0, 1) => 1.0,
276            (0, 2) => 1.0,
277            (1, 0) => 1.0,
278            (1, 1) => 3.0,
279            (1, 2) => 2.0,
280            (2, 0) => 1.0,
281            (2, 1) => 0.0,
282            (2, 2) => 0.0,
283            _ => 0.0,
284        });
285        let b = vec![4.0, 5.0, 6.0];
286        let mut x = vec![0.0; 3];
287        let mut solver = QrSolver::new();
288        let stats = solver
289            .solve(
290                &a,
291                None,
292                &b,
293                &mut x,
294                PcSide::Left,
295                &UniverseComm::NoComm(crate::parallel::NoComm),
296                None,
297                None,
298            )
299            .unwrap();
300        let expected = vec![6.0, 15.0, -23.0];
301        let tol = 1e-10;
302        for (xi, ei) in x.iter().zip(expected.iter()) {
303            assert!((xi - ei).abs() < tol, "xi = {}, expected = {}", xi, ei);
304        }
305        assert!(
306            matches!(
307                stats.reason,
308                ConvergedReason::ConvergedAtol | ConvergedReason::ConvergedRtol
309            ),
310            "QR did not report Converged reason"
311        );
312    }
313}