Initial commit
This commit is contained in:
@@ -0,0 +1,795 @@
|
||||
//! 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<f64>,
|
||||
/// 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<f64, usize>,
|
||||
/// Diagonal preconditioner (Jacobi)
|
||||
diag_precond: Option<Array1<f64>>,
|
||||
}
|
||||
|
||||
impl FemSolver {
|
||||
/// Create solver from assembler
|
||||
pub fn from_assembler(assembler: &FemAssembler, config: SolverConfig) -> FemResult<Self> {
|
||||
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<Self> {
|
||||
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<f64> = 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<f64> = 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<f64>) -> Array1<f64> {
|
||||
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<f64>) -> Array1<f64> {
|
||||
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<f64>) -> FemResult<SolverResult> {
|
||||
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<f64>) -> FemResult<SolverResult> {
|
||||
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<f64>) -> FemResult<SolverResult> {
|
||||
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<f64>) -> FemResult<SolverResult> {
|
||||
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<Array1<f64>> = 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<f64>) -> FemResult<SolverResult> {
|
||||
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<f64>, b: &Array1<f64>) -> FemResult<Array1<f64>> {
|
||||
let n = a.nrows();
|
||||
let mut lu = a.clone();
|
||||
let mut piv = (0..n).collect::<Vec<_>>();
|
||||
|
||||
// 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<f64>) -> FemResult<Array2<f64>> {
|
||||
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<Vec<SolverResult>> = (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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user