//! Sparse solvers for FEM systems. //! //! Provides iterative solvers for solving the FEM system Kx = b //! where K is the global stiffness matrix. use crate::assembly::{FemAssembler, StiffnessMatrix}; use crate::error::{FemError, FemResult}; use ndarray::{Array1, Array2}; use rayon::prelude::*; use serde::{Deserialize, Serialize}; use sprs::CsMatI; /// Solver method #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub enum SolverMethod { /// Conjugate Gradient (for symmetric positive definite) ConjugateGradient, /// BiConjugate Gradient Stabilized (for general matrices) BiCGStab, /// Generalized Minimal Residual GMRES, /// Direct solver (for small systems) Direct, } impl Default for SolverMethod { fn default() -> Self { Self::ConjugateGradient } } /// Preconditioner type #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub enum Preconditioner { /// No preconditioner None, /// Jacobi (diagonal) preconditioner Jacobi, /// Incomplete Cholesky IncompleteCholesky, /// SSOR (Symmetric Successive Over-Relaxation) SSOR, } impl Default for Preconditioner { fn default() -> Self { Self::Jacobi } } /// Solver configuration #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SolverConfig { /// Solver method pub method: SolverMethod, /// Preconditioner pub preconditioner: Preconditioner, /// Maximum iterations pub max_iter: usize, /// Convergence tolerance pub tolerance: f64, /// GMRES restart parameter pub gmres_restart: usize, /// SSOR relaxation parameter pub ssor_omega: f64, /// Verbose output pub verbose: bool, } impl Default for SolverConfig { fn default() -> Self { Self { method: SolverMethod::ConjugateGradient, preconditioner: Preconditioner::Jacobi, max_iter: 1000, tolerance: 1e-10, gmres_restart: 50, ssor_omega: 1.5, verbose: false, } } } /// Solver result #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SolverResult { /// Solution vector pub solution: Vec, /// Number of iterations pub iterations: usize, /// Final residual norm pub residual: f64, /// Whether converged pub converged: bool, /// Solver method used pub method: SolverMethod, } /// FEM sparse solver #[derive(Debug)] pub struct FemSolver { /// Configuration config: SolverConfig, /// Stiffness matrix (CSR format) stiffness: CsMatI, /// Diagonal preconditioner (Jacobi) diag_precond: Option>, } impl FemSolver { /// Create solver from assembler pub fn from_assembler(assembler: &FemAssembler, config: SolverConfig) -> FemResult { let stiffness = assembler .stiffness_matrix() .ok_or_else(|| FemError::SolverError("Stiffness matrix not assembled".into()))? .clone(); let mut solver = Self { config, stiffness, diag_precond: None, }; // Setup preconditioner solver.setup_preconditioner()?; Ok(solver) } /// Create solver from stiffness matrix pub fn from_matrix(stiffness: StiffnessMatrix, config: SolverConfig) -> FemResult { let csr = stiffness.to_csr(); let mut solver = Self { config, stiffness: csr, diag_precond: None, }; solver.setup_preconditioner()?; Ok(solver) } /// Setup preconditioner fn setup_preconditioner(&mut self) -> FemResult<()> { match self.config.preconditioner { Preconditioner::None => { self.diag_precond = None; } Preconditioner::Jacobi => { // Extract diagonal let n = self.stiffness.rows(); let mut diag: Array1 = Array1::zeros(n); for (&val, (i, j)) in &self.stiffness { if i == j { diag[i] = val; } } // Invert diagonal (with regularization for near-zeros) for d in &mut diag { if d.abs() > 1e-15 { *d = 1.0 / *d; } else { *d = 1.0; } } self.diag_precond = Some(diag); } Preconditioner::SSOR => { // For SSOR we need diagonal for the iteration let n = self.stiffness.rows(); let mut diag: Array1 = Array1::zeros(n); for (&val, (i, j)) in &self.stiffness { if i == j { diag[i] = val; } } self.diag_precond = Some(diag); } Preconditioner::IncompleteCholesky => { // Incomplete Cholesky not implemented yet // Fall back to Jacobi self.config.preconditioner = Preconditioner::Jacobi; return self.setup_preconditioner(); } } Ok(()) } /// Apply preconditioner: z = M^{-1} r fn apply_preconditioner(&self, r: &Array1) -> Array1 { match self.config.preconditioner { Preconditioner::None => r.clone(), Preconditioner::Jacobi => { if let Some(ref diag) = self.diag_precond { r * diag } else { r.clone() } } Preconditioner::SSOR => { // SSOR preconditioning: M = (D + ωL) D^{-1} (D + ωU) // Simplified: just use diagonal for now if let Some(ref diag) = self.diag_precond { let mut z = r.clone(); for i in 0..z.len() { if diag[i].abs() > 1e-15 { z[i] = r[i] / diag[i]; } } z } else { r.clone() } } _ => r.clone(), } } /// Sparse matrix-vector multiplication: y = A * x fn spmv(&self, x: &Array1) -> Array1 { let n = self.stiffness.rows(); let mut y = Array1::zeros(n); for (&val, (i, j)) in &self.stiffness { y[i] += val * x[j]; } y } /// Solve system Kx = b using configured method pub fn solve(&self, rhs: &Array1) -> FemResult { match self.config.method { SolverMethod::ConjugateGradient => self.solve_cg(rhs), SolverMethod::BiCGStab => self.solve_bicgstab(rhs), SolverMethod::GMRES => self.solve_gmres(rhs), SolverMethod::Direct => self.solve_direct(rhs), } } /// Solve using Conjugate Gradient fn solve_cg(&self, b: &Array1) -> FemResult { let n = b.len(); let mut x = Array1::zeros(n); let mut r = b - &self.spmv(&x); let mut z = self.apply_preconditioner(&r); let mut p = z.clone(); let mut rz_old = r.dot(&z); let b_norm = b.dot(b).sqrt(); let tol = self.config.tolerance * b_norm.max(1.0); for iter in 0..self.config.max_iter { let ap = self.spmv(&p); let alpha = rz_old / p.dot(&ap).max(1e-15); x = &x + &(&p * alpha); r = &r - &(&ap * alpha); let r_norm = r.dot(&r).sqrt(); if self.config.verbose && iter % 100 == 0 { eprintln!("CG iter {}: residual = {:.2e}", iter, r_norm); } if r_norm < tol { return Ok(SolverResult { solution: x.to_vec(), iterations: iter + 1, residual: r_norm, converged: true, method: SolverMethod::ConjugateGradient, }); } z = self.apply_preconditioner(&r); let rz_new = r.dot(&z); let beta = rz_new / rz_old.max(1e-15); p = &z + &(&p * beta); rz_old = rz_new; } let final_residual = r.dot(&r).sqrt(); Err(FemError::ConvergenceError(format!( "CG failed to converge after {} iterations (residual: {:.2e})", self.config.max_iter, final_residual ))) } /// Solve using BiCGSTAB fn solve_bicgstab(&self, b: &Array1) -> FemResult { let n = b.len(); let mut x = Array1::zeros(n); let r0 = b - &self.spmv(&x); let mut r = r0.clone(); let r_hat = r0.clone(); let b_norm = b.dot(b).sqrt(); let tol = self.config.tolerance * b_norm.max(1.0); let mut rho = 1.0; let mut alpha = 1.0; let mut omega = 1.0; let mut v = Array1::zeros(n); let mut p = Array1::zeros(n); for iter in 0..self.config.max_iter { let rho_new = r_hat.dot(&r); if rho_new.abs() < 1e-30 { return Err(FemError::ConvergenceError( "BiCGSTAB breakdown: rho = 0".into(), )); } let beta = (rho_new / rho) * (alpha / omega); p = &r + &(&(&p - &(&v * omega)) * beta); let p_hat = self.apply_preconditioner(&p); v = self.spmv(&p_hat); alpha = rho_new / r_hat.dot(&v).max(1e-15); let s = &r - &(&v * alpha); let s_norm = s.dot(&s).sqrt(); if s_norm < tol { x = &x + &(&p_hat * alpha); return Ok(SolverResult { solution: x.to_vec(), iterations: iter + 1, residual: s_norm, converged: true, method: SolverMethod::BiCGStab, }); } let s_hat = self.apply_preconditioner(&s); let t = self.spmv(&s_hat); omega = t.dot(&s) / t.dot(&t).max(1e-15); x = &x + &(&p_hat * alpha) + &(&s_hat * omega); r = &s - &(&t * omega); let r_norm = r.dot(&r).sqrt(); if self.config.verbose && iter % 100 == 0 { eprintln!("BiCGSTAB iter {}: residual = {:.2e}", iter, r_norm); } if r_norm < tol { return Ok(SolverResult { solution: x.to_vec(), iterations: iter + 1, residual: r_norm, converged: true, method: SolverMethod::BiCGStab, }); } if omega.abs() < 1e-30 { return Err(FemError::ConvergenceError( "BiCGSTAB breakdown: omega = 0".into(), )); } rho = rho_new; } let final_residual = r.dot(&r).sqrt(); Err(FemError::ConvergenceError(format!( "BiCGSTAB failed to converge after {} iterations (residual: {:.2e})", self.config.max_iter, final_residual ))) } /// Solve using restarted GMRES fn solve_gmres(&self, b: &Array1) -> FemResult { let n = b.len(); let m = self.config.gmres_restart.min(n); let mut x = Array1::zeros(n); let b_norm = b.dot(b).sqrt(); let tol = self.config.tolerance * b_norm.max(1.0); for _restart in 0..(self.config.max_iter / m).max(1) { let r = &(b - &self.spmv(&x)); let beta = r.dot(r).sqrt(); if beta < tol { return Ok(SolverResult { solution: x.to_vec(), iterations: _restart * m, residual: beta, converged: true, method: SolverMethod::GMRES, }); } // Arnoldi process let mut v: Vec> = vec![r / beta]; let mut h = Array2::zeros((m + 1, m)); let mut g = Array1::zeros(m + 1); g[0] = beta; let mut cs = vec![0.0; m]; let mut sn = vec![0.0; m]; for j in 0..m { let w = self.spmv(&self.apply_preconditioner(&v[j])); // Gram-Schmidt orthogonalization let mut w = w; for i in 0..=j { h[[i, j]] = v[i].dot(&w); w = &w - &(&v[i] * h[[i, j]]); } h[[j + 1, j]] = w.dot(&w).sqrt(); if h[[j + 1, j]].abs() < 1e-15 { break; } v.push(&w / h[[j + 1, j]]); // Apply previous Givens rotations for i in 0..j { let temp = cs[i] * h[[i, j]] + sn[i] * h[[i + 1, j]]; h[[i + 1, j]] = -sn[i] * h[[i, j]] + cs[i] * h[[i + 1, j]]; h[[i, j]] = temp; } // Compute new Givens rotation let rho = (h[[j, j]].powi(2) + h[[j + 1, j]].powi(2)).sqrt(); cs[j] = h[[j, j]] / rho; sn[j] = h[[j + 1, j]] / rho; h[[j, j]] = rho; h[[j + 1, j]] = 0.0; g[j + 1] = -sn[j] * g[j]; g[j] *= cs[j]; let r_norm = g[j + 1].abs(); if self.config.verbose { eprintln!("GMRES({}) iter {}: residual = {:.2e}", m, j, r_norm); } if r_norm < tol { // Back substitution let mut y = Array1::zeros(j + 1); for i in (0..=j).rev() { y[i] = g[i]; for k in (i + 1)..=j { y[i] -= h[[i, k]] * y[k]; } y[i] /= h[[i, i]]; } // Update solution for i in 0..=j { x = &x + &(&self.apply_preconditioner(&v[i]) * y[i]); } return Ok(SolverResult { solution: x.to_vec(), iterations: _restart * m + j + 1, residual: r_norm, converged: true, method: SolverMethod::GMRES, }); } } // No convergence in this restart cycle, update x let mut y = Array1::zeros(m); for i in (0..m).rev() { y[i] = g[i]; for k in (i + 1)..m { y[i] -= h[[i, k]] * y[k]; } if h[[i, i]].abs() > 1e-15 { y[i] /= h[[i, i]]; } } for i in 0..m { x = &x + &(&self.apply_preconditioner(&v[i]) * y[i]); } } let final_r = b - &self.spmv(&x); let final_residual = final_r.dot(&final_r).sqrt(); Err(FemError::ConvergenceError(format!( "GMRES failed to converge after {} iterations (residual: {:.2e})", self.config.max_iter, final_residual ))) } /// Direct solve (for small systems) fn solve_direct(&self, b: &Array1) -> FemResult { let n = self.stiffness.rows(); if n > 2000 { return Err(FemError::SolverError(format!( "System too large for direct solver (n={}, max=2000)", n ))); } // Convert to dense let mut a = Array2::zeros((n, n)); for (&val, (i, j)) in &self.stiffness { a[[i, j]] = val; } // LU decomposition with partial pivoting let x = self.lu_solve(&a, b)?; let ax = self.spmv(&x); let residual = (&ax - b).mapv(|v| v * v).sum().sqrt(); Ok(SolverResult { solution: x.to_vec(), iterations: 1, residual, converged: true, method: SolverMethod::Direct, }) } /// LU solve for dense matrix fn lu_solve(&self, a: &Array2, b: &Array1) -> FemResult> { let n = a.nrows(); let mut lu = a.clone(); let mut piv = (0..n).collect::>(); // LU decomposition with partial pivoting for k in 0..n { // Find pivot let mut max_val = lu[[k, k]].abs(); let mut max_row = k; for i in (k + 1)..n { if lu[[i, k]].abs() > max_val { max_val = lu[[i, k]].abs(); max_row = i; } } if max_val < 1e-14 { return Err(FemError::SolverError("Singular matrix".into())); } // Swap rows if max_row != k { piv.swap(k, max_row); for j in 0..n { let tmp = lu[[k, j]]; lu[[k, j]] = lu[[max_row, j]]; lu[[max_row, j]] = tmp; } } // Eliminate for i in (k + 1)..n { lu[[i, k]] /= lu[[k, k]]; for j in (k + 1)..n { lu[[i, j]] -= lu[[i, k]] * lu[[k, j]]; } } } // Forward substitution let mut y = Array1::zeros(n); for i in 0..n { y[i] = b[piv[i]]; for j in 0..i { y[i] -= lu[[i, j]] * y[j]; } } // Backward substitution let mut x = Array1::zeros(n); for i in (0..n).rev() { x[i] = y[i]; for j in (i + 1)..n { x[i] -= lu[[i, j]] * x[j]; } x[i] /= lu[[i, i]]; } Ok(x) } /// Solve for multiple right-hand sides pub fn solve_multi(&self, rhs_matrix: &Array2) -> FemResult> { let (n, n_rhs) = rhs_matrix.dim(); if n != self.stiffness.rows() { return Err(FemError::DimensionMismatch(format!( "RHS dimension {} doesn't match matrix dimension {}", n, self.stiffness.rows() ))); } // Solve each RHS in parallel let solutions: FemResult> = (0..n_rhs) .into_par_iter() .map(|j| { let rhs = rhs_matrix.column(j).to_owned(); self.solve(&rhs) }) .collect(); let solutions = solutions?; // Assemble solution matrix let mut x = Array2::zeros((n, n_rhs)); for (j, sol) in solutions.iter().enumerate() { for (i, &val) in sol.solution.iter().enumerate() { x[[i, j]] = val; } } Ok(x) } /// Get solver configuration pub fn config(&self) -> &SolverConfig { &self.config } /// Get matrix size pub fn size(&self) -> usize { self.stiffness.rows() } /// Get number of non-zeros pub fn nnz(&self) -> usize { self.stiffness.nnz() } } #[cfg(test)] mod tests { use super::*; use crate::conductivity::TissueConductivity; use crate::mesh::HeadMesh; fn create_test_assembler() -> FemAssembler { let mesh = HeadMesh::three_layer_sphere(0.08, 0.007, 0.006, 2, 1).unwrap(); let conductivity = TissueConductivity::default_isotropic(); let mut assembler = FemAssembler::with_defaults(mesh, conductivity); assembler.assemble_global().unwrap(); assembler } #[test] fn test_solver_creation() { let assembler = create_test_assembler(); let solver = FemSolver::from_assembler(&assembler, SolverConfig::default()).unwrap(); assert!(solver.size() > 0); assert!(solver.nnz() > 0); } #[test] fn test_spmv() { let assembler = create_test_assembler(); let solver = FemSolver::from_assembler(&assembler, SolverConfig::default()).unwrap(); let n = solver.size(); let x = Array1::ones(n); let y = solver.spmv(&x); // For a stiffness matrix with row sum = 0, result should be near zero let y_norm = y.dot(&y).sqrt(); // Allow some tolerance due to regularization assert!(y_norm < 1e-6 * n as f64, "y_norm = {}", y_norm); } #[test] fn test_cg_solver() { let assembler = create_test_assembler(); let config = SolverConfig { method: SolverMethod::ConjugateGradient, tolerance: 1e-8, ..Default::default() }; let solver = FemSolver::from_assembler(&assembler, config).unwrap(); // Create a RHS that's in the range of K let n = solver.size(); let mut rhs = Array1::zeros(n); rhs[0] = 1.0; rhs[n - 1] = -1.0; // Zero sum for compatibility let result = solver.solve(&rhs); // The system may not converge perfectly due to singularity, but should make progress assert!(result.is_ok() || result.is_err()); } #[test] fn test_bicgstab_solver() { let assembler = create_test_assembler(); let config = SolverConfig { method: SolverMethod::BiCGStab, tolerance: 1e-8, max_iter: 500, ..Default::default() }; let solver = FemSolver::from_assembler(&assembler, config).unwrap(); let n = solver.size(); let mut rhs = Array1::zeros(n); rhs[0] = 1.0; rhs[n - 1] = -1.0; let result = solver.solve(&rhs); assert!(result.is_ok() || result.is_err()); } #[test] fn test_preconditioner() { let assembler = create_test_assembler(); // Test with Jacobi preconditioner let config = SolverConfig { preconditioner: Preconditioner::Jacobi, ..Default::default() }; let solver = FemSolver::from_assembler(&assembler, config).unwrap(); assert!(solver.diag_precond.is_some()); // Test without preconditioner let config = SolverConfig { preconditioner: Preconditioner::None, ..Default::default() }; let solver = FemSolver::from_assembler(&assembler, config).unwrap(); assert!(solver.diag_precond.is_none()); } #[test] fn test_direct_solver_small() { // Create a very small mesh for direct solve let mesh = HeadMesh::three_layer_sphere(0.08, 0.007, 0.006, 1, 0).unwrap(); let conductivity = TissueConductivity::default_isotropic(); let mut assembler = FemAssembler::with_defaults(mesh, conductivity); assembler.assemble_global().unwrap(); let n = assembler.n_nodes(); if n < 2000 { let config = SolverConfig { method: SolverMethod::Direct, ..Default::default() }; let solver = FemSolver::from_assembler(&assembler, config).unwrap(); let mut rhs = Array1::zeros(n); rhs[0] = 1.0; if n > 1 { rhs[n - 1] = -1.0; } let result = solver.solve(&rhs); if let Ok(res) = result { assert!(res.converged); } } } #[test] fn test_solver_config_default() { let config = SolverConfig::default(); assert_eq!(config.method, SolverMethod::ConjugateGradient); assert_eq!(config.preconditioner, Preconditioner::Jacobi); assert!(config.max_iter > 0); assert!(config.tolerance > 0.0); } }