796 lines
24 KiB
Rust
796 lines
24 KiB
Rust
//! 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);
|
|
}
|
|
}
|