use super::*;
pub fn solve<M, T>(a: &M, b: &DVector<T>, max_iter: usize, tol: T) -> Option<DVector<T>>
where
M: SpMatVecMul<T>,
T: SimdRealField + PartialOrd + Copy
{
let mut solver = ConjugateGradient {
x: DVector::<T>::zeros(a.nrows()),
r: DVector::<T>::zeros(a.nrows()),
p: DVector::<T>::zeros(a.nrows()),
ap: DVector::<T>::zeros(a.nrows()),
residual_dot: T::zero(),
tol,
max_iter,
iter: 0,
converged: false,
};
solver.init(a, b, None);
if solver.solve_iterations(a, b, max_iter) {
Some(solver.x.clone())
} else {
None
}
}
pub fn solve_with_initial_guess<M, T>(a: &M, b: &DVector<T>, x: &mut DVector<T>, max_iter: usize, tol: T) -> bool
where
M: SpMatVecMul<T>,
T: SimdRealField + PartialOrd + Copy
{
let mut solver = ConjugateGradient {
x: x.clone(),
r: DVector::<T>::zeros(a.nrows()),
p: DVector::<T>::zeros(a.nrows()),
ap: DVector::<T>::zeros(a.nrows()),
residual_dot: T::zero(),
tol,
max_iter,
iter: 0,
converged: false,
};
solver.init(a, b, Some(x));
let converged = solver.solve_iterations(a, b, max_iter);
*x = solver.x.clone();
converged
}
pub struct ConjugateGradient<T> {
pub x: DVector<T>,
pub r: DVector<T>,
pub p: DVector<T>,
pub ap: DVector<T>,
pub residual_dot: T,
pub tol: T,
pub max_iter: usize,
pub iter: usize,
pub converged: bool,
}
impl<M, T> IterativeSolver<M, DVector<T>, T> for ConjugateGradient<T>
where
M: SpMatVecMul<T>,
T: SimdRealField + PartialOrd + Copy,
{
fn init(&mut self, a: &M, b: &DVector<T>, x0: Option<&DVector<T>>) {
let n = a.nrows();
self.x = match x0 {
Some(x0) => x0.clone(),
None => DVector::<T>::zeros(n),
};
self.r = b - &a.mul_vec(&self.x);
self.p = self.r.clone();
self.residual_dot = self.r.dot(&self.r);
self.ap = DVector::<T>::zeros(n);
self.iter = 0;
self.converged = false;
}
fn step(&mut self, a: &M, _b: &DVector<T>) -> bool {
if self.converged {
return true;
}
let norm = self.r.magnitude();
if norm <= self.tol {
self.converged = true;
return true;
}
self.ap = a.mul_vec(&self.p);
let alpha = self.residual_dot / self.p.dot(&self.ap);
self.x.axpy(alpha, &self.p, T::one());
let new_r = &self.r - &self.ap * alpha;
let new_norm = new_r.magnitude();
if new_norm <= self.tol {
self.r = new_r;
self.converged = true;
return true;
}
let new_residual_dot = new_r.dot(&new_r);
let beta = new_residual_dot / self.residual_dot;
self.p = &new_r + &self.p * beta;
self.r = new_r;
self.residual_dot = new_residual_dot;
self.iter += 1;
false
}
fn reset(&mut self) {
self.x.fill(T::zero());
self.r.fill(T::zero());
self.p.fill(T::zero());
self.ap.fill(T::zero());
self.residual_dot = T::zero();
self.iter = 0;
self.converged = false;
}
fn hard_reset(&mut self) {
self.x = DVector::<T>::zeros(0);
self.r = DVector::<T>::zeros(0);
self.p = DVector::<T>::zeros(0);
self.ap = DVector::<T>::zeros(0);
self.residual_dot = T::zero();
self.iter = 0;
self.converged = false;
}
fn soft_reset(&mut self) {
self.r.fill(T::zero());
self.p.fill(T::zero());
self.ap.fill(T::zero());
self.residual_dot = T::zero();
self.iter = 0;
self.converged = false;
}
fn solution(&self) -> &DVector<T> {
&self.x
}
fn iterations(&self) -> usize {
self.iter
}
}