1use 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
28pub struct LuSolver {
32 factor: Option<FullPivLu<Scalar>>,
34}
35
36impl LuSolver {
37 pub fn new() -> Self {
39 LuSolver { factor: None }
40 }
41
42 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 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; let _ = _pc_side;
97
98 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 let factor = FullPivLu::new(a.as_ref());
110 self.factor = Some(factor);
111
112 x.clone_from(b);
114
115 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 if let Some(monitors) = monitors {
125 for monitor in monitors {
126 monitor(1, 0.0, 0);
127 }
128 }
129
130 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
141pub struct QrSolver;
143
144impl QrSolver {
145 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 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; let _ = _pc_side;
186
187 if let Some(monitors) = monitors {
189 for monitor in monitors {
190 monitor(0, 0.0, 0);
191 }
192 }
193
194 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 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 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 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}