use log::warn;
use nalgebra_sparse::na::RealField;
use crate::iteratives::amg::restrict::build_r;
use super::{*, coarsen::coarsen, graph::strength_graph, interpolate::build_p, rap};
pub struct Level<N> {
pub a: CsrMatrix<N>,
pub p: CsrMatrix<N>,
pub r: CsrMatrix<N>,
pub diag: CsrMatrix<N>,
}
pub struct Hierarchy<N> {
pub levels: Vec<Level<N>>,
}
pub fn setup<N>(a_fin: CsrMatrix<N>, theta: N, n_min: usize) -> Hierarchy<N>
where
N: RealField + Copy,
{
let mut levels = Vec::new();
let mut a = a_fin;
levels.reserve(10);
loop {
let (marks, coarse_of) = coarsen(&a, theta);
let s = strength_graph(&a, theta);
let p_candidate = build_p(&a, &marks, &coarse_of, &s);
let is_coarsest_level = p_candidate.ncols() >= a.nrows() || a.nrows() <= n_min;
if is_coarsest_level {
let r_coarsest = build_r(&a, &p_candidate);
let diag_coarsest = a.diagonal_as_csr();
levels.push(Level {
a,
p: p_candidate,
r: r_coarsest,
diag: diag_coarsest,
});
break;
}
let r_intermediate = build_r(&a, &p_candidate);
let diag_intermediate = a.diagonal_as_csr();
let a_coarse = rap(&r_intermediate, &a, &p_candidate);
levels.push(Level {
a,
p: p_candidate,
r: r_intermediate,
diag: diag_intermediate,
});
a = a_coarse;
if a.nrows() == 0 {
warn!("Matrix A became empty after RAP. Stopping.");
break;
}
if levels.len() > 20 { warn!("Too many levels generated ({}). Stopping.", levels.len());
break;
}
}
Hierarchy { levels }
}
impl<N: RealField + Copy> Hierarchy<N> {
pub fn vcycle(
&self,
l: usize,
b: &DVector<N>,
x: &mut DVector<N>,
residual_buffer: &mut DVector<N>,
tol: N,
nu_pre: usize,
nu_post: usize,
) {
let lev = &self.levels[l];
gauss_seidel::solve_with_initial_guess(&lev.a, b, x, nu_pre, tol);
if residual_buffer.len() != b.len() {
*residual_buffer = DVector::zeros(b.len());
}
residual_buffer.copy_from(b);
residual_buffer.axpy(-N::one(), &(&lev.a * &*x), N::one());
if l + 1 == self.levels.len() {
jacobi::solve_with_initial_guess(&lev.a, b, x, 50, tol);
return;
}
let residual_coarse = &lev.r * &*residual_buffer;
let mut error_coarse = DVector::<N>::zeros(residual_coarse.len());
let mut residual_buffer_coarse = DVector::<N>::zeros(residual_coarse.len());
self.vcycle(
l + 1,
&residual_coarse,
&mut error_coarse,
&mut residual_buffer_coarse,
tol,
nu_pre,
nu_post,
);
let correction = &lev.p * &error_coarse;
x.axpy(N::one(), &correction, N::one());
gauss_seidel::solve_with_initial_guess(&lev.a, b, x, nu_post, tol);
*residual_buffer = &lev.a * &*x - b;
}
}